num_workers, prefetching and the Windows spawn trap
num_workers puts extra processes on data preparation, prefetching keeps batches ready ahead of time — and on Windows, forgetting one guard line crashes training before it starts.
- 9 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
num_workers hires extra helper processes to prepare batches, so the trainer never stands idle waiting for data.
One cook in a kitchen chops, fries, and plates — and thirty diners wait. Hire four cooks and dishes flow out in parallel while the head chef only assembles and serves.
num_workers is the number of hired cooks. Each is a separate worker process — an independent copy of your program — preparing samples at the same time as the others.
But there is a catch, and on Windows it bites hard. Hiring a cook on Windows means training a brand-new person from scratch, every time. And if the hiring is done carelessly, each new cook starts by trying to hire more cooks — and the kitchen collapses before lunch.
Why it exists
Preparing one sample is real work: read a file from disk, decode a photo, apply random edits. Meanwhile the trainer — often an expensive GPU — computes on the previous batch.
With zero workers, these two jobs take turns: prepare, train, prepare, train. The expensive half sits idle half the time. Workers overlap the two — batches are prepared while training happens, and the next tray is ready the moment the current one finishes. That ready-ahead-of-time habit is called prefetching.
How it works, and the Windows trap
num_workers=0: prepare → train → prepare → train (turns)
num_workers=4: prepare prepare prepare prepare
↓ (queue of ready batches)
train → train → train → train (no waiting)Every new worker on Windows is started by re-reading your script from the top. If the script's top level starts training, each worker starts training, which starts more workers. The guard line you will see below exists to break that loop.
A real example you have seen
Any long download alongside a video call: your phone overlaps the two instead of freezing one for the other. Same principle — slow input work runs beside the main job instead of in front of it.
Remember this
- Workers prepare data in parallel with training; prefetching keeps batches ready early.
- More workers is not always faster — hiring has a cost, especially on Windows.
- On Windows, the
if __name__ == "__main__":guard is not optional.
What to learn next
- Finding out whether the GPU is waiting for data — measuring starvation instead of guessing.
- Samplers and weighted sampling — controlling which samples the workers fetch.
- IterableDataset for streams and huge files — when random access itself is the bottleneck.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on Windows 11, Python 3.10. All timings are from one machine and will differ on yours — the ratios are the lesson, not the digits.
The trap first, because it crashes first
import torch
from torch.utils.data import DataLoader, TensorDataset
ds = TensorDataset(torch.rand(8, 2))
loader = DataLoader(ds, num_workers=2)
for _ in loader: # this line starts worker processes
pass
print("finished")On Linux this runs. On Windows (and macOS), the workers are started by spawn — a fresh Python re-imports your script — and each import re-runs the loop, which spawns again. PyTorch detects the recursion and the run dies in a wall of tracebacks containing:
RuntimeError:
An attempt has been made to start a new process before the
current process has finished its bootstrapping phase.
...
RuntimeError: DataLoader worker (pid(s) 44584, 49376) exited unexpectedlyThe process ids in the last line change every run. The fix is structural, and it is the single most Windows-relevant line in PyTorch:
if __name__ == "__main__": # workers re-import this file; imports skip this block
main()Everything that does work — building loaders, training — goes under the guard. Class and function definitions stay at top level, because workers legitimately need to import those. Two related spawn rules: the dataset and collate function must be picklable (no lambdas, no locally-defined classes), which is why the collate lesson insisted on module-level def.
Measuring workers honestly
import time
import torch
from torch.utils.data import Dataset, DataLoader
class SlowDataset(Dataset):
"""Pretends each sample costs 10 ms of disk reading and decoding."""
def __len__(self):
return 512
def __getitem__(self, idx):
time.sleep(0.010)
return torch.rand(3, 64, 64), idx % 2
def time_one_epoch(num_workers):
loader = DataLoader(SlowDataset(), batch_size=32, num_workers=num_workers)
start = time.perf_counter()
for _ in loader:
pass
return time.perf_counter() - start
if __name__ == "__main__":
for n in [0, 2, 4]:
print(f"num_workers={n}: {time_one_epoch(n):.1f} s")num_workers=0: 5.4 s num_workers=2: 4.8 s num_workers=4: 3.7 s
Honest reading: four workers only won 1.7 seconds, on work that is 5.1 seconds of pure sleeping. Why so modest? On Windows, each spawned worker pays a start-up tax — a fresh Python importing torch takes a second or two — and this epoch pays it before the first batch flows.
persistent_workers: pay the tax once
By default, workers are dismissed at each epoch's end and re-hired at the next. On Windows that means paying the spawn tax every epoch. persistent_workers=True keeps them alive:
import time
import torch
from torch.utils.data import Dataset, DataLoader
class SlowDataset(Dataset):
def __len__(self):
return 512
def __getitem__(self, idx):
time.sleep(0.010)
return torch.rand(3, 64, 64), idx % 2
def three_epochs(persistent):
loader = DataLoader(SlowDataset(), batch_size=32, num_workers=4,
persistent_workers=persistent)
for epoch in range(3):
start = time.perf_counter()
for _ in loader:
pass
print(f" epoch {epoch}: {time.perf_counter() - start:.1f} s")
if __name__ == "__main__":
print("persistent_workers=False (workers rebuilt every epoch):")
three_epochs(False)
print("persistent_workers=True (workers built once):")
three_epochs(True)persistent_workers=False (workers rebuilt every epoch): epoch 0: 3.7 s epoch 1: 3.7 s epoch 2: 3.7 s persistent_workers=True (workers built once): epoch 0: 3.5 s epoch 1: 1.4 s epoch 2: 1.4 s
From epoch 1 onward, 3.7 becomes 1.4 — the spawn tax is gone and the parallelism finally shows its true value. On Windows, persistent_workers=True with num_workers > 0 should be your reflex.
Two companion settings, briefly: prefetch_factor (default 2) is how many batches each worker keeps ready ahead of demand — raise it for bursty storage. pin_memory=True stages batches in pinned memory — RAM the operating system will not shuffle around — enabling faster, overlappable copies to a GPU; it costs a little RAM and does nothing useful on CPU-only training. The full GPU-feeding story is the last lesson of this section.
Common mistakes
Cranking workers to 16 and trusting vibes. Each worker duplicates the dataset object and buffers prefetch_factor batches — RAM scales with workers. Past the point where the trainer stops waiting, more workers buy pure overhead. Measure at 0, 2, 4, 8; keep the knee.
The guard missing from imported code. The guard must protect the entry script — the file actually run. A guarded helper module does not save an unguarded train.py.
Unpicklable things on the dataset. Open file handles, database connections, lambdas stored on self — spawn serialises the dataset and dies. Store paths and configuration; open handles lazily inside __getitem__ (first use per worker).
Debugging with workers on. Breakpoints and print debugging inside __getitem__ misbehave across processes. Set num_workers=0 while debugging; restore after.
Try it yourself
Change the sleep to 0.001 and re-run the timing script. Watch workers lose to num_workers=0 — then explain it: when per-sample work is trivial, the queueing overhead exceeds the work. That inversion is real, and knowing it saves you from cargo-culting num_workers=8 everywhere.
What to learn next
- Finding out whether the GPU is waiting for data — measuring starvation instead of guessing.
- Samplers and weighted sampling — controlling which samples the workers fetch.
- IterableDataset for streams and huge files — when random access itself is the bottleneck.
Researcher — Mathematics and papers.
fork versus spawn, precisely
num_workers > 0 builds worker processes via Python multiprocessing with the platform default start method: fork on Linux, spawn on Windows and macOS (macOS switched for framework-safety reasons in Python 3.8; CUDA is additionally fork-hostile, so even Linux GPU workloads increasingly choose spawn or forkserver). Fork clones the parent's address space copy-on-write: near-zero start-up, dataset shared read-only until written — but it silently duplicates any inherited state, including RNGs (the classic identical-augmentation bug: NumPy generators forked with identical state across workers; torch re-seeds each worker as base_seed + worker_id, but NumPy/random require worker_init_fn intervention). Spawn imports the main module afresh and reconstructs the dataset from a pickle — hence the import-guard requirement, picklability constraints, and the measured start-up tax (a fresh interpreter importing torch costs seconds), amortised by persistent_workers.
The transport differs too: workers write sample tensors into shared-memory segments (file-backed on Linux /dev/shm, with its container-quota failure mode; section-object based on Windows), and the parent maps them zero-copy. pin_memory then re-stages into page-locked memory on a dedicated thread in the parent, enabling cudaMemcpyAsync overlap — the mechanics quantified in the starvation lesson.
A queueing model for choosing num_workers
Steady-state throughput is min(producer rate, consumer rate): W workers each producing a batch every t_data seconds feed a trainer consuming one every t_step. Starvation vanishes when W ≥ t_data / t_step — the practical sizing formula, with prefetch_factor × W buffered batches absorbing variance in t_data (long-tailed storage latencies, JPEG size skew). Beyond that W, added workers contribute memory pressure, context switching, and on shared filesystems, self-inflicted I/O contention — the empirically observed throughput decline past the knee. Little's law gives the buffer-occupancy view: queued batches = arrival rate × wait, so a persistently full prefetch queue certifies the trainer as bottleneck (healthy), while a persistently empty one certifies data starvation.
When __getitem__ is GIL-bound Python, processes (not threads) are the only parallelism; the 2024-25 free-threaded CPython work and torch's own pin_memory/decode thread pools are slowly shifting this landscape — as of torch 2.5, process workers remain the standard mechanism.
Alternatives when workers are not enough
CPU-side decode/augment ceilings motivated bypasses: NVIDIA DALI moves JPEG decode and augmentation to GPU; FFCV (Leclerc et al., 2023) restructures storage for sequential reads plus JIT-compiled transforms; WebDataset streams sharded tars sequentially, converting random access into bulk reads (IterableDataset lesson). Data echoing (Choi et al., 2020, Faster Neural Network Training with Data Echoing) attacks from the consumer side: reuse each fetched batch e times when the pipeline, not the optimiser, is the bottleneck — trading gradient freshness for throughput.
Reading
torch.utils.datadocs, "Platform-specific behaviors" — the normative spawn/fork statement.- Leclerc et al. (2023), FFCV: Accelerating Training by Removing Data Bottlenecks.
- Choi et al. (2020), Faster Neural Network Training with Data Echoing.
What to learn next
- Finding out whether the GPU is waiting for data — measuring starvation instead of guessing.
- Samplers and weighted sampling — controlling which samples the workers fetch.
- IterableDataset for streams and huge files — when random access itself is the bottleneck.