Skip to content

memory_safe_dataloader

data.dataloader.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()