spdl.dataloader.get_pytorch_dataloader¶
- get_pytorch_dataloader(dataset: Dataset[T], batch_size: int | None = 1, shuffle: bool = False, sampler: Sampler[K] | None = None, batch_sampler: Sampler[list[K]] | None = None, num_workers: int = 1, collate_fn: Callable[[list[T]], U] | None = None, pin_memory: bool = False, drop_last: bool = False, timeout: float | None = None, worker_init_fn: None = None, multiprocessing_context: str | BaseContext | None = None, generator: Generator | None = None, *, prefetch_factor: int = 2, persistent_workers: bool = False, pin_memory_device: str | None = None, in_order: bool = False, worker_init_concurrency: int | None = None) PyTorchDataLoader[U][source]¶
Build a process-backed data loader for a map-style dataset.
- Parameters:
dataset – Dataset from which samples are loaded.
batch_size – Number of samples per batch, or
Noneto disable batching.shuffle – Whether to sample indices randomly.
sampler – Optional sampler. Mutually exclusive with
batch_samplerandshuffle=True.batch_sampler – Optional sampler that yields batches of indices.
num_workers – Number of worker processes.
collate_fn – Optional function that combines samples into a batch.
pin_memory – Whether to move output tensors into pinned memory.
drop_last – Whether to discard the final incomplete batch.
timeout – Maximum time to wait for a batch, or
Nonefor no timeout.worker_init_fn – Unsupported PyTorch compatibility argument; must be
None.multiprocessing_context – Start method or multiprocessing context used to launch workers.
generator – Optional random generator used by the default sampler.
prefetch_factor – Number of batches buffered per worker.
persistent_workers – Unsupported PyTorch compatibility argument; must be
False.pin_memory_device – Unsupported PyTorch compatibility argument; must be
None.in_order – Whether to emit results in input order rather than completion order.
worker_init_concurrency –
Maximum number of workers that may deserialize the shared dataset concurrently. Useful when dataset construction contends for a shared resource such as filesystem bandwidth or database connections. This does not limit steady-state fetching.
Added in version 0.7.0: The
worker_init_concurrencyargument.
- Returns:
A reusable
PyTorchDataLoader.- Raises:
ValueError – If arguments are incompatible, unsupported, or out of range.
Example
Limit dataset deserialization to two workers while retaining eight workers for steady-state loading:
loader = get_pytorch_dataloader( dataset, num_workers=8, worker_init_concurrency=2, )