diff --git a/test/stateful_dataloader/test_worker_start.py b/test/stateful_dataloader/test_worker_start.py new file mode 100644 index 000000000..2fb369acf --- /dev/null +++ b/test/stateful_dataloader/test_worker_start.py @@ -0,0 +1,217 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import threading +import unittest +from multiprocessing.context import SpawnContext +from typing import Any, List, Optional, Tuple +from unittest.mock import patch + +import torch +from torch.utils.data import Dataset +from torchdata.stateful_dataloader import StatefulDataLoader + + +class _MapDataset(Dataset): + def __init__(self, length: int) -> None: + self.length = length + + def __len__(self) -> int: + return self.length + + def __getitem__(self, index: int): + return {"idx": index} + + +class _WorkerSeedDataset(Dataset): + def __len__(self) -> int: + return 8 + + def __getitem__(self, index: int) -> int: + return torch.initial_seed() + + +class _StartTracker: + def __init__(self, concurrent_starts: int) -> None: + self._lock = threading.Lock() + self._barrier = None + if concurrent_starts > 1: + self._barrier = threading.Barrier(concurrent_starts) + self.active_starts = 0 + self.max_active_starts = 0 + + def start_entered(self) -> None: + with self._lock: + self.active_starts += 1 + self.max_active_starts = max(self.max_active_starts, self.active_starts) + if self._barrier is not None: + self._barrier.wait(timeout=30) + + def start_exited(self) -> None: + with self._lock: + self.active_starts -= 1 + + +class _ProcessProxy: + def __init__( + self, + process: Any, + tracker: _StartTracker, + fail_start: bool = False, + ) -> None: + self._process = process + self._tracker = tracker + self._fail_start = fail_start + + @property + def daemon(self) -> bool: + return self._process.daemon + + @daemon.setter + def daemon(self, value: bool) -> None: + self._process.daemon = value + + @property + def pid(self) -> Optional[int]: + return self._process.pid + + def start(self) -> None: + self._tracker.start_entered() + try: + if self._fail_start: + raise RuntimeError("worker start failed") + self._process.start() + finally: + self._tracker.start_exited() + + def is_alive(self) -> bool: + return self._process.is_alive() + + def terminate(self) -> None: + self._process.terminate() + + def kill(self) -> None: + self._process.kill() + + def join(self, timeout: Optional[float] = None) -> None: + self._process.join(timeout) + + +class WorkerStartTest(unittest.TestCase): + def test_invalid_parallelism_is_rejected(self) -> None: + with self.assertRaisesRegex(ValueError, "must be positive"): + StatefulDataLoader( + _MapDataset(1), + spawn_worker_start_parallelism=0, + ) + + def test_spawn_starts_respect_parallelism_and_preserve_order(self) -> None: + for parallelism in (1, 2, 4): + with self.subTest(parallelism=parallelism): + rows, max_active_starts = self._run_with_tracked_starts(parallelism) + + self.assertEqual(rows, list(range(20))) + self.assertEqual(max_active_starts, parallelism) + + def test_parallel_spawn_preserves_worker_seeds(self) -> None: + def worker_seeds(parallelism: int) -> List[int]: + loader = StatefulDataLoader( + _WorkerSeedDataset(), + batch_size=1, + num_workers=4, + multiprocessing_context=SpawnContext(), + generator=torch.Generator().manual_seed(42), + spawn_worker_start_parallelism=parallelism, + ) + return [int(seed) for seed in loader] + + self.assertEqual(worker_seeds(1), worker_seeds(4)) + + def test_parallel_spawn_preserves_resume_and_persistent_workers(self) -> None: + def create_loader() -> StatefulDataLoader: + return StatefulDataLoader( + _MapDataset(20), + batch_size=2, + num_workers=4, + multiprocessing_context=SpawnContext(), + persistent_workers=True, + spawn_worker_start_parallelism=4, + ) + + loader = create_loader() + loader_iter = iter(loader) + rows = [] + for _ in range(3): + rows.extend(next(loader_iter)["idx"].tolist()) + state_dict = loader.state_dict() + + resumed_loader = create_loader() + resumed_loader.load_state_dict(state_dict) + rows.extend(row for batch in resumed_loader for row in batch["idx"].tolist()) + + self.assertEqual(rows, list(range(20))) + self.assertEqual( + [row for batch in resumed_loader for row in batch["idx"].tolist()], + list(range(20)), + ) + + def test_parallel_start_failure_cleans_up_started_workers(self) -> None: + self._assert_start_failure_cleans_up_workers(parallelism=2) + + def test_serial_start_failure_cleans_up_started_workers(self) -> None: + self._assert_start_failure_cleans_up_workers(parallelism=1) + + def _assert_start_failure_cleans_up_workers(self, parallelism: int) -> None: + context = SpawnContext() + original_process = context.Process + tracker = _StartTracker(parallelism) + processes = [] + + def create_process(*args, **kwargs): + process = _ProcessProxy( + original_process(*args, **kwargs), + tracker, + fail_start=len(processes) == 1, + ) + processes.append(process) + return process + + with patch.object(SpawnContext, "Process", side_effect=create_process): + loader = StatefulDataLoader( + _MapDataset(20), + batch_size=2, + num_workers=4, + multiprocessing_context=context, + spawn_worker_start_parallelism=parallelism, + ) + + with self.assertRaisesRegex(RuntimeError, "worker start failed"): + next(iter(loader)) + + self.assertEqual(len(processes), 4 if parallelism > 1 else 2) + self.assertEqual(tracker.max_active_starts, parallelism) + for process in processes: + if process.pid is not None: + self.assertFalse(process.is_alive()) + + def _run_with_tracked_starts(self, parallelism: int) -> Tuple[List[int], int]: + context = SpawnContext() + original_process = context.Process + tracker = _StartTracker(parallelism) + + def create_process(*args, **kwargs): + return _ProcessProxy(original_process(*args, **kwargs), tracker) + + with patch.object(SpawnContext, "Process", side_effect=create_process): + loader = StatefulDataLoader( + _MapDataset(20), + batch_size=2, + num_workers=4, + multiprocessing_context=context, + spawn_worker_start_parallelism=parallelism, + ) + rows = [row for batch in loader for row in batch["idx"].tolist()] + return rows, tracker.max_active_starts diff --git a/torchdata/stateful_dataloader/stateful_dataloader.py b/torchdata/stateful_dataloader/stateful_dataloader.py index c0adec0bc..a2407bdf8 100644 --- a/torchdata/stateful_dataloader/stateful_dataloader.py +++ b/torchdata/stateful_dataloader/stateful_dataloader.py @@ -24,6 +24,7 @@ import logging import queue import threading +from concurrent.futures import FIRST_EXCEPTION, ThreadPoolExecutor, wait from typing import Any, Dict, Iterable, List, Optional, TypeVar, Union @@ -82,6 +83,76 @@ logger = logging.getLogger(__name__) + +def _start_worker_processes(workers, max_parallelism): + if max_parallelism == 1: + for worker in workers: + worker.start() + return + + stop_starting = threading.Event() + + def start_worker(worker): + if stop_starting.is_set(): + return + try: + worker.start() + except BaseException: + stop_starting.set() + raise + + with ThreadPoolExecutor( + max_workers=min(max_parallelism, len(workers)), + thread_name_prefix="StatefulDataLoaderWorkerStart", + ) as executor: + futures = [executor.submit(start_worker, worker) for worker in workers] + done, pending = wait(futures, return_when=FIRST_EXCEPTION) + failed_future = next( + ( + future + for future in futures + if future in done + if not future.cancelled() + if future.exception() is not None + ), + None, + ) + if failed_future is not None: + for future in pending: + future.cancel() + + if failed_future is not None: + for future in futures: + if future is failed_future or future.cancelled(): + continue + secondary_exception = future.exception() + if secondary_exception is not None: + logger.error( + "An additional worker process failed to start", + exc_info=( + type(secondary_exception), + secondary_exception, + secondary_exception.__traceback__, + ), + ) + failed_future.result() + + +def _clean_up_failed_workers(worker_launches, done_event, result_queue): + done_event.set() + for _, worker in worker_launches: + if worker.pid is not None and worker.is_alive(): + worker.terminate() + for index_queue, worker in worker_launches: + if worker.pid is not None: + worker.join(timeout=_utils.MP_STATUS_CHECK_INTERVAL) + if worker.is_alive(): + worker.kill() + worker.join(timeout=_utils.MP_STATUS_CHECK_INTERVAL) + index_queue.close() + result_queue.close() + + _INDEX_SAMPLER_STATE = "_index_sampler_state" _SAMPLER_ITER_STATE = "_sampler_iter_state" _SAMPLER_ITER_YIELDED = "_sampler_iter_yielded" @@ -97,7 +168,8 @@ class StatefulDataLoader(DataLoader[_T_co]): checkpointing. All arguments are identical to ``torch.utils.data.DataLoader``, with - a new kwarg: ``snapshot_every_n_steps``. + additional ``snapshot_every_n_steps`` and + ``spawn_worker_start_parallelism`` keyword arguments. Args: dataset (Dataset): dataset from which to load the data. @@ -151,6 +223,10 @@ class StatefulDataLoader(DataLoader[_T_co]): are returned in a first-in, first-out order. Only applies when ``num_workers > 0``. (default: ``True``) snapshot_every_n_steps (int, optional): Defines how often the state is transferred from the dataloader workers to the dataloader. By default, it is set to ``1``, i.e., state is transferred every step. If the state is large, this value can be increased (and ideally set to the frequency of training checkpointing) to reduce the overhead of transferring state every step. + spawn_worker_start_parallelism (int, optional): Maximum number of worker + processes to start concurrently when using the ``spawn`` multiprocessing + context. Other multiprocessing contexts always start workers serially. + (default: ``1``) .. warning:: If the ``spawn`` start method is used, :attr:`worker_init_fn` @@ -210,6 +286,7 @@ def __init__( pin_memory_device: str = "", in_order: bool = True, snapshot_every_n_steps: Optional[int] = 1, + spawn_worker_start_parallelism: int = 1, ): torch._C._log_api_usage_once("python.stateful_data_loader") @@ -221,6 +298,9 @@ def __init__( if timeout < 0: raise ValueError("timeout option should be non-negative") + if spawn_worker_start_parallelism < 1: + raise ValueError("spawn_worker_start_parallelism must be positive") + if num_workers == 0 and prefetch_factor is not None: raise ValueError( "prefetch_factor option could only be specified in multiprocessing." @@ -250,6 +330,7 @@ def __init__( self.worker_init_fn = worker_init_fn self.multiprocessing_context = multiprocessing_context self.in_order = in_order + self.spawn_worker_start_parallelism = spawn_worker_start_parallelism # Adds forward compatibilities so classic DataLoader can work with DataPipes: # _DataPipeSerializationWrapper container makes it easier to serialize without redefining pickler @@ -948,43 +1029,62 @@ def __init__(self, loader, next_iter_state): _SHARED_SEED, self._shared_seed ) - for i in range(self._num_workers): - # No certainty which module multiprocessing_context is - index_queue = multiprocessing_context.Queue() # type: ignore[var-annotated] - # Need to `cancel_join_thread` here! - # See sections (2) and (3b) above. - index_queue.cancel_join_thread() - - w = multiprocessing_context.Process( - target=_worker_loop, - args=( - self._dataset_kind, - self._dataset, - index_queue, - self._worker_result_queue, - self._workers_done_event, - self._auto_collation, - self._collate_fn, - self._drop_last, - self._base_seed, - self._worker_init_fn, - i, - self._num_workers, - self._persistent_workers, - self._shared_seed, - worker_states[self._worker_key(i)], - ), + worker_start_parallelism = loader.spawn_worker_start_parallelism + uses_spawn_context = multiprocessing_context.get_start_method() == "spawn" + start_workers_in_parallel = worker_start_parallelism > 1 and uses_spawn_context + worker_launches = [] + try: + for i in range(self._num_workers): + # No certainty which module multiprocessing_context is + index_queue = multiprocessing_context.Queue() # type: ignore[var-annotated] + # Need to `cancel_join_thread` here! + # See sections (2) and (3b) above. + index_queue.cancel_join_thread() + + w = multiprocessing_context.Process( + target=_worker_loop, + args=( + self._dataset_kind, + self._dataset, + index_queue, + self._worker_result_queue, + self._workers_done_event, + self._auto_collation, + self._collate_fn, + self._drop_last, + self._base_seed, + self._worker_init_fn, + i, + self._num_workers, + self._persistent_workers, + self._shared_seed, + worker_states[self._worker_key(i)], + ), + ) + w.daemon = True + worker_launches.append((index_queue, w)) + if not start_workers_in_parallel: + w.start() + self._index_queues.append(index_queue) + self._workers.append(w) + + if start_workers_in_parallel: + _start_worker_processes( + [worker for _, worker in worker_launches], + worker_start_parallelism, + ) + index_queues = [index_queue for index_queue, _ in worker_launches] + self._index_queues.extend(index_queues) + self._workers.extend(worker for _, worker in worker_launches) + except BaseException: + _clean_up_failed_workers( + worker_launches, + self._workers_done_event, + self._worker_result_queue, ) - w.daemon = True - # NB: Process.start() actually take some time as it needs to - # start a process and pass the arguments over via a pipe. - # Therefore, we only add a worker to self._workers list after - # it started, so that we do not call .join() if program dies - # before it starts, and __del__ tries to join but will get: - # AssertionError: can only join a started process. - w.start() - self._index_queues.append(index_queue) - self._workers.append(w) + self._index_queues.clear() + self._workers.clear() + raise if self._pin_memory: self._pin_memory_thread_done_event = threading.Event()