Skip to content

dataloader

Module diagram

classDiagram
  class dataloader {
  }
  class memory_safe_dataloader {
  }
  dataloader --> memory_safe_dataloader

data.dataloader

MemorySafeDataLoader

Bases: DataLoader

WORKER_MEMORY_SAFETY_FACTOR class-attribute instance-attribute

WORKER_MEMORY_SAFETY_FACTOR = 2.0

WORKER_TIMEOUT_INTERVAL class-attribute instance-attribute

WORKER_TIMEOUT_INTERVAL = 1

_get_iterator

_get_iterator() -> _BaseDataLoaderIter
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def _get_iterator(self) -> "_BaseDataLoaderIter":
    if self.num_workers == 0:
        return _SingleProcessDataLoaderIter(self)
    else:
        self.check_worker_number_rationality()
        return _MemorySafeMultiProcessingDataLoaderIter(
            self,
            self.WORKER_MEMORY_SAFETY_FACTOR,
            self.WORKER_TIMEOUT_INTERVAL,
        )

memory_safe_dataloader

_WorkerStatus

_WorkerStatus(pid: Optional[int] = None)
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def __init__(self, pid: Optional[int] = None):
    self._pid = pid
    self._max_memory_usage = 0.0
    self._worker_ready_event = multiprocessing.Event()
_pid instance-attribute
_pid = pid
_max_memory_usage instance-attribute
_max_memory_usage = 0.0
_worker_ready_event instance-attribute
_worker_ready_event = Event()
update_pid
update_pid(pid: int) -> None
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def update_pid(self, pid: int) -> None:
    self._pid = pid
get_max_memory_usage
get_max_memory_usage() -> float
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def get_max_memory_usage(self) -> float:
    self._update_max_memory_usage()
    return self._max_memory_usage
_update_max_memory_usage
_update_max_memory_usage() -> None
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def _update_max_memory_usage(self) -> None:
    if self._pid is not None:
        try:
            _process = psutil.Process(self._pid)
            worker_memory_usage = _process.memory_info().rss / (1024**2)
            self._max_memory_usage = max(self._max_memory_usage, worker_memory_usage)
        except psutil.NoSuchProcess:
            raise DataLoaderWorkerInitializationError
is_worker_ready
is_worker_ready() -> bool
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def is_worker_ready(self) -> bool:
    return self._worker_ready_event.is_set()
set_worker_ready
set_worker_ready() -> None
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def set_worker_ready(self) -> None:
    self._worker_ready_event.set()
wait_worker_ready
wait_worker_ready(timeout: Optional[int] = None) -> None
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def wait_worker_ready(self, timeout: Optional[int] = None) -> None:
    self._worker_ready_event.wait(timeout=timeout)

_WorkerStatusManager

_WorkerStatusManager(worker_memory_safety_factor: float = 2.0, worker_status_list: Optional[List[_WorkerStatus]] = None)
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def __init__(
    self,
    worker_memory_safety_factor: float = 2.0,
    worker_status_list: Optional[List[_WorkerStatus]] = None,
):
    self.worker_memory_safety_factor = worker_memory_safety_factor
    if worker_status_list is None:
        worker_status_list = []
    self._worker_status_list = worker_status_list
worker_memory_safety_factor instance-attribute
worker_memory_safety_factor = worker_memory_safety_factor
_worker_status_list instance-attribute
_worker_status_list = worker_status_list
__getitem__
__getitem__(index: int) -> _WorkerStatus
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def __getitem__(self, index: int) -> _WorkerStatus:
    return self._worker_status_list[index]
_get_max_memory_usage
_get_max_memory_usage() -> None
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def _get_max_memory_usage(self) -> None:
    memory_usages = []
    for worker_status in self._worker_status_list:
        memory_usages.append(worker_status.get_max_memory_usage())
    return max(memory_usages)
is_memory_safe
is_memory_safe() -> bool
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def is_memory_safe(self) -> bool:
    max_memory_usage = self._get_max_memory_usage()
    memory_upper_bound = max_memory_usage * self.worker_memory_safety_factor
    available_memory = psutil.virtual_memory().available / 1024**2
    return available_memory > memory_upper_bound
is_workers_ready
is_workers_ready() -> bool
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def is_workers_ready(self) -> bool:
    return all(worker_status.is_worker_ready() for worker_status in self._worker_status_list)
wait_workers_ready
wait_workers_ready(timeout: Optional[int] = None) -> None
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def wait_workers_ready(self, timeout: Optional[int] = None) -> None:
    for worker_status in self._worker_status_list:
        if worker_status.is_worker_ready():
            continue
        worker_status.wait_worker_ready(timeout=timeout)
        break

_MemorySafeMultiProcessingDataLoaderIter

_MemorySafeMultiProcessingDataLoaderIter(loader: DataLoader, worker_memory_safety_factor: float, worker_timeout_interval: int = 1)

Bases: _MultiProcessingDataLoaderIter

Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def __init__(
    self,
    loader: DataLoader,
    worker_memory_safety_factor: float,
    worker_timeout_interval: int = 1,
):
    super(_MultiProcessingDataLoaderIter, self).__init__(loader)

    self._prefetch_factor = loader.prefetch_factor
    self._in_order = loader.in_order

    assert self._num_workers > 0
    assert self._prefetch_factor > 0

    if loader.multiprocessing_context is None:
        multiprocessing_context = multiprocessing
    else:
        multiprocessing_context = loader.multiprocessing_context

    self._worker_init_fn = loader.worker_init_fn

    # Adds forward compatibilities so classic DataLoader can work with DataPipes:
    #   Additional worker init function will take care of sharding in MP and Distributed
    if isinstance(self._dataset, (IterDataPipe, MapDataPipe)):
        self._worker_init_fn = functools.partial(
            _sharding_worker_init_fn,
            self._worker_init_fn,
            self._world_size,
            self._rank,
        )

    # No certainty which module multiprocessing_context is
    self._worker_result_queue = multiprocessing_context.Queue()  # type: ignore[var-annotated]
    self._worker_pids_set = False
    self._shutdown = False
    self._workers_done_event = multiprocessing_context.Event()

    self._index_queues = []
    self._workers = []

    # XXX: Added by SAIGE for memory safety
    worker_status_manager = _WorkerStatusManager(
        worker_memory_safety_factor=worker_memory_safety_factor,
        worker_status_list=[_WorkerStatus() for _ in range(self._num_workers)],
    )
    try:
        sample_index = next(iter(self._index_sampler))
    except Exception:
        # XXX: `self._index_sampler` StopIteration occurs occasionally.
        raise DataLoaderWorkerInitializationError(
            "length of dataloader is 0. check options[dataset, batch_size, drop_last]"
        )
    self._workers_status = [False for _ in range(self._num_workers)]
    # XXX: End

    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,
                # XXX: Added by SAIGE for memory safety
                worker_status_manager[i],
                sample_index,
                # XXX: End
            ),
        )
        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)

        # XXX: Added by SAIGE for memory safety
        self._workers_status[i] = True
        worker_status_manager[i].update_pid(w.pid)
        # NOTE: print the subprocess PID to the log for external termination handling.
        print(f"[AI module INFO] subprocess is started (PID: {w.pid})")
        # Wait for the data to be fetched and ready
        worker_status_manager.wait_workers_ready(timeout=worker_timeout_interval)
        if not worker_status_manager.is_memory_safe():
            raise CpuLowMemoryError
        # XXX: End

    # XXX: Added by SAIGE for memory safety
    while not worker_status_manager.is_workers_ready():
        worker_status_manager.wait_workers_ready(timeout=worker_timeout_interval)
        if not worker_status_manager.is_memory_safe():
            raise CpuLowMemoryError
    # XXX: End

    if self._pin_memory:
        self._pin_memory_thread_done_event = threading.Event()

        # Queue is not type-annotated
        self._data_queue = queue.Queue()  # type: ignore[var-annotated]
        if self._pin_memory_device == "xpu":
            current_device = torch.xpu.current_device()  # type: ignore[attr-defined]
        elif self._pin_memory_device == torch._C._get_privateuse1_backend_name():
            custom_device_mod = getattr(torch, torch._C._get_privateuse1_backend_name())
            current_device = custom_device_mod.current_device()
        else:
            current_device = torch.cuda.current_device()  # choose cuda for default
        pin_memory_thread = threading.Thread(
            target=_utils.pin_memory._pin_memory_loop,
            args=(
                self._worker_result_queue,
                self._data_queue,
                current_device,
                self._pin_memory_thread_done_event,
                self._pin_memory_device,
            ),
        )
        pin_memory_thread.daemon = True
        pin_memory_thread.start()
        # Similar to workers (see comment above), we only register
        # pin_memory_thread once it is started.
        self._pin_memory_thread = pin_memory_thread
    else:
        self._data_queue = self._worker_result_queue  # type: ignore[assignment]

    # In some rare cases, persistent workers (daemonic processes)
    # would be terminated before `__del__` of iterator is invoked
    # when main process exits
    # It would cause failure when pin_memory_thread tries to read
    # corrupted data from worker_result_queue
    # atexit is used to shutdown thread and child processes in the
    # right sequence before main process exits
    if self._persistent_workers and self._pin_memory:
        import atexit

        for w in self._workers:
            atexit.register(_MultiProcessingDataLoaderIter._clean_up_worker, w)

    # .pid can be None only before process is spawned (not the case, so ignore)
    _utils.signal_handling._set_worker_pids(id(self), tuple(w.pid for w in self._workers))  # type: ignore[misc]
    _utils.signal_handling._set_SIGCHLD_handler()
    self._worker_pids_set = True
    self._reset(loader, first_iter=True)
_prefetch_factor instance-attribute
_prefetch_factor = prefetch_factor
_in_order instance-attribute
_in_order = in_order
_worker_init_fn instance-attribute
_worker_init_fn = worker_init_fn
_worker_result_queue instance-attribute
_worker_result_queue = Queue()
_shutdown instance-attribute
_shutdown = False
_workers_done_event instance-attribute
_workers_done_event = Event()
_index_queues instance-attribute
_index_queues = []
_workers instance-attribute
_workers = []
_workers_status instance-attribute
_workers_status = [False for _ in (range(_num_workers))]
_pin_memory_thread_done_event instance-attribute
_pin_memory_thread_done_event = Event()
_data_queue instance-attribute
_data_queue = Queue()
_pin_memory_thread instance-attribute
_pin_memory_thread = pin_memory_thread
_worker_pids_set instance-attribute
_worker_pids_set = True

MemorySafeDataLoader

Bases: DataLoader

WORKER_MEMORY_SAFETY_FACTOR class-attribute instance-attribute
WORKER_MEMORY_SAFETY_FACTOR = 2.0
WORKER_TIMEOUT_INTERVAL class-attribute instance-attribute
WORKER_TIMEOUT_INTERVAL = 1
_get_iterator
_get_iterator() -> _BaseDataLoaderIter
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def _get_iterator(self) -> "_BaseDataLoaderIter":
    if self.num_workers == 0:
        return _SingleProcessDataLoaderIter(self)
    else:
        self.check_worker_number_rationality()
        return _MemorySafeMultiProcessingDataLoaderIter(
            self,
            self.WORKER_MEMORY_SAFETY_FACTOR,
            self.WORKER_TIMEOUT_INTERVAL,
        )

_worker_loop

_worker_loop(dataset_kind, dataset, index_queue, data_queue, done_event, auto_collation, collate_fn, drop_last, base_seed, init_fn, worker_id, num_workers, persistent_workers, shared_seed, worker_status: _WorkerStatus, sample_index)
Source code in SaigeToolkit/data/dataloader/memory_safe_dataloader.py
def _worker_loop(
    dataset_kind,
    dataset,
    index_queue,
    data_queue,
    done_event,
    auto_collation,
    collate_fn,
    drop_last,
    base_seed,
    init_fn,
    worker_id,
    num_workers,
    persistent_workers,
    shared_seed,
    # XXX: Added by SAIGE for memory safety
    worker_status: "_WorkerStatus",
    sample_index,
    # XXX: End
):
    # See NOTE [ Data Loader Multiprocessing Shutdown Logic ] for details on the
    # logic of this function.

    try:
        # Initialize C side signal handlers for SIGBUS and SIGSEGV. Python signal
        # module's handlers are executed after Python returns from C low-level
        # handlers, likely when the same fatal signal had already happened
        # again.
        # https://docs.python.org/3/library/signal.html#execution-of-python-signal-handlers
        signal_handling._set_worker_signal_handlers()

        torch.set_num_threads(1)
        seed = base_seed + worker_id
        random.seed(seed)
        torch.manual_seed(seed)
        if HAS_NUMPY:
            np_seed = _generate_state(base_seed, worker_id)
            import numpy as np

            np.random.seed(np_seed)

        from torch.utils.data import IterDataPipe
        from torch.utils.data.graph_settings import apply_random_seed

        shared_rng = torch.Generator()
        if isinstance(dataset, IterDataPipe):
            assert shared_seed is not None
            shared_rng.manual_seed(shared_seed)
            dataset = apply_random_seed(dataset, shared_rng)

        global _worker_info
        _worker_info = WorkerInfo(id=worker_id, num_workers=num_workers, seed=seed, dataset=dataset)

        from torch.utils.data import _DatasetKind

        init_exception = None

        try:
            if init_fn is not None:
                init_fn(worker_id)

            fetcher = _DatasetKind.create_fetcher(
                dataset_kind, dataset, auto_collation, collate_fn, drop_last
            )

            # XXX: Added by SAIGE for memory safety
            data = fetcher.fetch(sample_index)  # type: ignore[assignment]
            # XXX: End

        except Exception:
            init_exception = ExceptionWrapper(where=f"in DataLoader worker process {worker_id}")

        # XXX: Added by SAIGE for memory safety
        finally:
            worker_status.set_worker_ready()
        # XXX: End

        # When using Iterable mode, some worker can exit earlier than others due
        # to the IterableDataset behaving differently for different workers.
        # When such things happen, an `_IterableDatasetStopIteration` object is
        # sent over to the main process with the ID of this worker, so that the
        # main process won't send more tasks to this worker, and will send
        # `None` to this worker to properly exit it.
        #
        # Note that we cannot set `done_event` from a worker as it is shared
        # among all processes. Instead, we set the `iteration_end` flag to
        # signify that the iterator is exhausted. When either `done_event` or
        # `iteration_end` is set, we skip all processing step and just wait for
        # `None`.
        iteration_end = False

        watchdog = ManagerWatchdog()

        while watchdog.is_alive():
            try:
                r = index_queue.get(timeout=MP_STATUS_CHECK_INTERVAL)
            except queue.Empty:
                continue
            if isinstance(r, _ResumeIteration):
                # Acknowledge the main process
                data_queue.put((r, None))
                iteration_end = False

                if isinstance(dataset, IterDataPipe):
                    assert r.seed is not None
                    shared_rng.manual_seed(r.seed)
                    dataset = apply_random_seed(dataset, shared_rng)

                # Recreate the fetcher for worker-reuse policy
                fetcher = _DatasetKind.create_fetcher(
                    dataset_kind, dataset, auto_collation, collate_fn, drop_last
                )
                continue
            elif r is None:
                # Received the final signal
                assert done_event.is_set() or iteration_end
                break
            elif done_event.is_set() or iteration_end:
                # `done_event` is set. But I haven't received the final signal
                # (None) yet. I will keep continuing until get it, and skip the
                # processing steps.
                continue
            idx, index = r
            data: Union[_IterableDatasetStopIteration, ExceptionWrapper]
            if init_exception is not None:
                data = init_exception
                init_exception = None
            else:
                try:
                    data = fetcher.fetch(index)  # type: ignore[possibly-undefined]
                except Exception as e:
                    if isinstance(e, StopIteration) and dataset_kind == _DatasetKind.Iterable:
                        data = _IterableDatasetStopIteration(worker_id)
                        # Set `iteration_end`
                        #   (1) to save future `next(...)` calls, and
                        #   (2) to avoid sending multiple `_IterableDatasetStopIteration`s.
                        iteration_end = True
                    else:
                        # It is important that we don't store exc_info in a variable.
                        # `ExceptionWrapper` does the correct thing.
                        # See NOTE [ Python Traceback Reference Cycle Problem ]
                        data = ExceptionWrapper(where=f"in DataLoader worker process {worker_id}")
            data_queue.put((idx, data))
            del data, idx, index, r  # save memory
    except KeyboardInterrupt:
        # Main process will raise KeyboardInterrupt anyways.
        pass
    if done_event.is_set():
        data_queue.cancel_join_thread()
        data_queue.close()