skip to content

What does setting DataLoader's num_workers above 0 actually change in PyTorch?

level: middleimportance: must knowfreq 70%

answer

  1. separate processes, not threads
  2. a queue of batches ahead of you
  3. each copy pays memory
  4. they die at epoch end by default
  5. persistent_workers and prefetch_factor

basics

~20 s

It moves sample fetching and collation into that many separate processes, each prefetching batches ahead of the training loop. The costs are process startup per epoch, duplicated memory, a picklable dataset requirement, and worse tracebacks when something fails.

solid answer

~50 s

`num_workers=N` starts N subprocesses; each is assigned whole batches, runs `__getitem__` and `collate_fn` for them, and ships the finished batch back to the main process through shared memory. Each worker keeps `prefetch_factor` batches in flight (default 2 when workers are enabled), so up to `N * prefetch_factor` batches are queued ahead of the consumer — that buffer is what keeps the GPU from waiting. The costs are real. Workers are created and torn down every time you iterate the loader, unless `persistent_workers=True`. Each worker holds its own copy of the dataset object, so anything eagerly loaded in `__init__` is multiplied by N. The dataset and `collate_fn` must be picklable under the spawn start method, and spawn also requires the loader to be created under a `if __name__ == "__main__":` guard. Exceptions from workers arrive re-raised and harder to read, so the first debugging step is always `num_workers=0`.

code

python · 21 lines
python
import torch
from torch.utils.data import DataLoader, TensorDataset

dataset = TensorDataset(
    torch.randn(1024, 3, 32, 32),
    torch.randint(0, 10, (1024,)),
)

if __name__ == "__main__":  # required when the start method is spawn
    loader = DataLoader(
        dataset,
        batch_size=64,
        shuffle=True,
        num_workers=8,
        persistent_workers=True,  # survive across epochs
        prefetch_factor=4,        # batches queued per worker
        pin_memory=True,
    )
    for _ in range(2):
        for xb, yb in loader:
            pass

go deeper

for a junior

Know that num_workers loads data in parallel background processes so the GPU waits less, and that setting it to 0 is the way to get a readable traceback when something breaks.

for a middle

Explain that workers are processes with their own dataset copy, that prefetch_factor sets the queue depth, and that persistent_workers avoids per-epoch respawn. Name the picklability requirement.

for a senior

Reason about the ceiling: CPU quota, shared-memory size, copy-on-write erosion of a Python-object index, thread oversubscription, and measuring a data-only pass before adding workers.

for a principal

Set the platform defaults — container shm sizing, start method, worker/thread budgets per job class — and decide when the real answer is a different storage format or an offline preprocessing step rather than more processes per node.

## What actually happens With `num_workers=0`, everything runs inline: the training loop blocks while the next batch is read, decoded and collated. With `num_workers=N`, `DataLoader` spins up N processes through `torch.multiprocessing`. The main process assigns each worker a batch's worth of indices; the worker calls `__getitem__` for each index, runs `collate_fn`, writes the resulting tensors into shared memory, and puts a handle on a result queue. The main process reassembles batches **in order** by default, so shuffling semantics are unchanged by the worker count. `prefetch_factor` (default `None`, meaning 2 once `num_workers > 0`) controls how many batches each worker keeps in flight. The queue depth is roughly `num_workers * prefetch_factor` batches — that is the buffer absorbing variance in per-sample cost. Passing `prefetch_factor` with `num_workers=0` is an error. ## The costs **Startup per epoch.** Every `for batch in loader` builds a fresh iterator, which forks or spawns the workers again and tears them down at the end. For short epochs or small validation loaders, that overhead can dominate. `persistent_workers=True` (valid only with `num_workers > 0`) keeps them alive between epochs at the price of holding their memory the whole run. **Memory.** Each worker has its own copy of the dataset object. On Linux with the fork start method, that copy starts as copy-on-write, but Python's reference counting touches object headers and gradually un-shares pages — a dataset holding a large list of Python objects will steadily materialize N copies. Storing the index as a NumPy array or an Arrow table instead of a Python list is the standard mitigation. Batches themselves travel through `/dev/shm`, and in containers with a tiny default shared-memory size you get "bus error" or "unable to write to file" crashes until it is raised. **Picklability and start methods.** With spawn (macOS, Windows, and where the default has moved away from fork on newer Python versions), the dataset, `collate_fn` and `worker_init_fn` are pickled and sent, so lambdas, local closures and open file handles break, and the loader must be constructed under a `__main__` guard. With fork, they are inherited instead — which is why the same code can work on Linux and fail elsewhere. Write datasets that are picklable regardless: keep a path in `__init__`, open the handle lazily inside `__getitem__`. **CUDA in workers.** Initializing CUDA in a forked child is unsupported and will fail or hang. Keep worker code CPU-only and move batches to the device in the training loop; if you truly need CUDA in workers, use `multiprocessing_context="spawn"`. **Debuggability.** A worker exception is re-raised in the parent with the original traceback embedded, but breakpoints, `pdb` and profilers behave badly across processes. Reproduce with `num_workers=0` first. ## Choosing N Start around the number of physical cores available to the process, then measure — more workers is not monotonically better. Beyond some point they contend for CPU, thrash the page cache, and inflate memory. A frequent hidden cost is thread oversubscription: each worker may spin up its own BLAS/OpenMP thread pool, so N workers times T threads swamps the machine; setting `torch.set_num_threads(1)` inside `worker_init_fn`, or the equivalent environment variable, often speeds things up more than adding workers. In containers, remember the CPU quota may be far below the visible core count. The honest test is a data-only pass: iterate the loader with no model and time it. If that time is already below your step time, more workers buy nothing and you should look at the model instead. ## Ordering and determinism Because results are reassembled in index order, changing `num_workers` does not change which samples land in which batch for a map-style dataset with a fixed seed. It *can* change random augmentation, because per-worker RNG seeds derive from the worker id — a run with 4 workers and a run with 8 will not produce identical augmented pixels even with the same global seed.

  • Why can num_workers=32 be slower than num_workers=8?
    Because workers compete for the same finite resources. Past the CPU quota they context-switch rather than overlap, each holds a dataset copy so RAM and page-cache pressure rise, and every worker may start its own BLAS/OpenMP thread pool — 32 workers times several threads swamps the machine. Setting `torch.set_num_threads(1)` in `worker_init_fn` frequently recovers more throughput than adding workers.
  • What exactly does prefetch_factor control?
    How many batches each worker prepares ahead of what the training loop has consumed — the default is 2 once `num_workers > 0`, giving roughly `num_workers * prefetch_factor` batches in flight. Raising it smooths out occasional slow samples at the cost of memory held in the queue. It cannot be set when `num_workers=0`, since there is no worker to prefetch.
  • Why should a Dataset avoid opening an HDF5 or LMDB handle in __init__?
    Under spawn the handle is not picklable and the loader fails outright; under fork the handle is inherited by every worker, and concurrent reads through a shared file descriptor and offset produce corrupted or interleaved data. The fix is to store only the path in `__init__` and open the handle lazily on first access inside the worker, caching it on the instance.

saying these in an interview costs you the question

  • Says workers are threads inside one process
  • Assumes more workers always means more throughput
  • Forgets workers are recreated every epoch by default
  • Loads the whole dataset in __init__ then adds eight workers
  • Initializes CUDA tensors inside worker processes

context