skip to content

Why do PyTorch DataLoader workers sometimes produce identical random augmentations?

level: seniorimportance: nice to knowfreq 33%

answer

  1. state copied into every child
  2. the framework seeds what it can name
  3. your own generator is invisible to it
  4. build it in the worker, not the parent
  5. torch.initial_seed inside __getitem__

basics

~20 s

Because a random generator created in the parent process is copied unchanged into every worker. PyTorch seeds each worker's torch, Python random and NumPy global generators from base_seed plus worker id, but an RNG object your Dataset holds, or a third-party library's state, is duplicated.

solid answer

~50 s

When workers start, PyTorch derives a per-worker seed from a `base_seed` and the worker id and applies it to the **global** generators it knows about: `torch`, Python's `random`, and NumPy's global RNG. So augmentations built on `torch.rand` or `random.random` differ correctly across workers. What is not covered is state the framework cannot see. If your `Dataset.__init__` does `self.rng = numpy.random.default_rng(0)`, that object is copied or pickled into every worker with identical state, and all four workers then generate the same noise, the same crops, the same mixup pairs — one sample in four is really distinct. Third-party libraries that seed their own global state at import time have the same problem. The fix: never create the generator in the parent. Build it lazily inside the worker, seeded from `torch.initial_seed()` or `get_worker_info().seed`, or set it in a `worker_init_fn`.

code

python · 23 lines
python
import numpy as np
import torch
from torch.utils.data import DataLoader, Dataset


class NoisyDataset(Dataset):
    def __init__(self, n):
        self.n = n
        self.rng = None  # never build the generator in the parent process

    def __len__(self):
        return self.n

    def __getitem__(self, idx):
        if self.rng is None:
            # torch.initial_seed() is already per-worker and per-epoch
            self.rng = np.random.default_rng(torch.initial_seed() % 2**32)
        return torch.from_numpy(self.rng.normal(size=4))


if __name__ == "__main__":
    loader = DataLoader(NoisyDataset(8), batch_size=2, num_workers=2)
    print(torch.cat(list(loader)))  # distinct rows across workers

go deeper

for a junior

Know that DataLoader workers are separate processes and that anything random created before they start is copied, so randomness must be established inside the worker.

for a middle

Explain the mechanism: a base_seed plus worker id seeds the torch, random and NumPy global generators per worker, while an RNG object held by the dataset is duplicated instead.

for a senior

Show how you would catch it — print draws per worker id — and pick between lazy per-worker construction, worker_init_fn, and index-derived seeds depending on the reproducibility you need.

for a principal

Own the reproducibility contract for the team: what a run's seed is required to pin, whether augmentation must be invariant to worker count, and how a silent diversity loss like this would be detected before it shows up as a soft regression across many models.

## What PyTorch seeds for you Each time you create an iterator over a `DataLoader` with workers, the main process draws a `base_seed` from the loader's generator (the global RNG unless you passed `generator=`). Every worker then computes its own seed from `base_seed` and its worker id, and PyTorch applies that seed to the global generators it controls: `torch.manual_seed`, Python's `random.seed`, and — since PyTorch 1.9 — NumPy's global `numpy.random.seed`. Inside a worker, `torch.initial_seed()` returns that value, and `torch.utils.data.get_worker_info().seed` exposes it too. Two useful consequences follow. First, augmentations written against `torch.rand`, `random.randint` or `numpy.random.rand` genuinely differ per worker with no work from you. Second, because `base_seed` is re-drawn every time an iterator is created, the augmentations also differ every epoch — which is what you want, and which is why a run's augmentation stream is not reproducible from the epoch number alone. ## What it cannot seed The mechanism reaches global generator state, by name, in three libraries. It cannot reach: **A generator object your dataset owns.** `self.rng = np.random.default_rng(0)` in `__init__` creates a concrete object holding concrete state. Under fork the child inherits it byte-for-byte; under spawn it is pickled, state included. Every worker's `self.rng` is now at the same position in the same stream, so the *k*-th call in worker 0 and the *k*-th call in worker 3 return the same numbers. With four workers you get roughly a quarter of the augmentation diversity you think you have — and identically for `torch.Generator()` objects, `random.Random()` instances, and library handles created eagerly. **Other libraries' global state.** Anything with its own RNG seeded at import in the parent — a C-extension augmentation library, a simulator, a hashing salt — is inherited unchanged. If it exposes a seeding call, seed it in `worker_init_fn`. **State reset across epochs.** A generator stored on the dataset is not reset when a new iterator starts either, so under `persistent_workers=True` the streams continue, while with respawned workers they restart — a subtle difference in behaviour caused by a flag that supposedly only concerns performance. ## The two fixes **Lazy construction inside the worker.** Set `self.rng = None` in `__init__`, and on first use in `__getitem__` build it from `torch.initial_seed()`. Because that value is already per-worker and per-epoch, the derived generator inherits both properties. This keeps everything in one file and works identically at `num_workers=0`, where `torch.initial_seed()` simply returns the main process's seed. **`worker_init_fn`.** DataLoader calls it in each worker with the worker id before iteration starts. Use it to seed third-party libraries, attach a fresh generator to `get_worker_info().dataset`, or set `torch.set_num_threads(1)`. Note that under spawn it must be picklable — a module-level function, not a lambda. A third, stronger option exists when you want reproducibility independent of the worker count: derive the seed **from the sample index** inside `__getitem__`. Sample 4,712 then gets the same augmentation whether you run one worker or sixteen, which makes bug reports reproducible and makes a run's data stream a pure function of its seed. The cost is a hash and a fresh generator per call, and that augmentations repeat across epochs unless you fold the epoch number in too. ## How you would notice You usually would not, which is why it is worth knowing. Training still converges, just with less effective augmentation, so the symptom is a validation gap slightly worse than expected. The cheap check is to print a few random draws per worker: iterate a loader whose `__getitem__` returns its own random number along with `get_worker_info().id`, and look at whether the values repeat across ids. Ten lines of diagnostic beat a week of wondering.

  • How do you make shuffling reproducible across runs?
    Pass a seeded `torch.Generator` as DataLoader's `generator=` argument; the sampler draws from it, and the per-iterator `base_seed` handed to workers is drawn from it too. Combined with a deterministic `worker_init_fn` and a fixed worker count, that pins both the batch order and the per-worker augmentation seeds. Change the worker count and the augmentation stream changes, even though the batch composition does not.
  • Should augmentations repeat from one epoch to the next?
    No — you want fresh randomness each epoch, and PyTorch gives it by re-drawing base_seed every time an iterator is created, so worker seeds differ per epoch. If you deliberately want per-sample reproducibility instead, derive the seed from the sample index inside `__getitem__`, and fold in the epoch number if you still want variation across epochs.
  • What else is worker_init_fn good for besides seeding?
    Setting `torch.set_num_threads(1)` so N workers do not each launch a BLAS thread pool and oversubscribe the CPU; opening per-worker file or database handles that must not be shared across processes; configuring per-worker logging; and sharding an IterableDataset by mutating the worker's own dataset copy. Under the spawn start method it must be a picklable module-level function.

saying these in an interview costs you the question

  • Assumes a seed set in the parent covers every worker
  • Creates the dataset's RNG object in __init__
  • Thinks identical augmentations would raise an error
  • Believes augmentations repeat identically every epoch
  • Confuses seeding the sampler with seeding augmentation

context