Finding out whether the GPU is waiting for data
A starved GPU looks exactly like a slow model unless you measure — split each step into "waiting for the batch" and "computing on it", and the split tells you which half to fix.
- 8 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.
"The GPU is waiting for data" means the expensive computing chip sits idle because batches are not being prepared as fast as they are consumed.
A tandoor cooks a naan in thirty seconds. But if the person rolling the dough takes two minutes per naan, the tandoor — the costliest thing in the kitchen — burns gas doing nothing, ninety seconds out of every cycle.
Buying a second tandoor fixes nothing. Speed here belongs to the dough station.
In training, the GPU — the fast chip built for exactly this arithmetic — is the tandoor, and the data pipeline is the dough station. Reading files, decoding photos, applying transforms: when that falls behind, the GPU waits, and training crawls no matter how powerful the chip.
Why it exists as a skill
The two diseases look identical from outside: "training is slow." But the cures point in opposite directions. Slow model → smaller model, better kernels, a faster GPU. Slow data → more workers, faster storage, cheaper transforms — and a faster GPU would change nothing.
People routinely guess wrong, and the wrong guess is expensive: renting a bigger GPU to starve it at higher hourly rates. One measurement replaces the guessing.
How it works
The trick is to time the two halves of every step separately.
one training step, timed in two pieces:
[ waiting for the next batch ] ← the dough station's share
[ computing on the batch ] ← the tandoor's share
waiting 1.4 s, computing 0.01 s → starving: fix the pipeline
waiting 0.01 s, computing 1.4 s → busy: the model is the costWhichever half dominates is where your effort — and money — should go.
A real example you have seen
Any download that crawls while your internet is fast: the bottleneck sits elsewhere — the server, the disk. Same diagnosis discipline: find the slow stage before paying to speed up a fast one.
Remember this
- A starved GPU and a slow model look identical until you time the halves.
- Waiting dominates → fix data. Computing dominates → fix model.
- Measure first; upgrade second. Never the reverse.
What to learn next
- num_workers, prefetching and the Windows spawn trap — the first rung of the ladder, in depth.
- Datasets that do not fit in memory — storage-layout fixes when workers are not enough.
- Image transforms with torchvision v2 — cheapening the per-sample work itself.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on Windows 11. Timings are from one machine — yours will differ; the split is what you are reading, and this technique needs no GPU to demonstrate.
The two-clock profiler
Twelve lines you can paste into any training loop:
import time
import torch
from torch import nn
from torch.utils.data import Dataset, DataLoader
class SlowDataset(Dataset):
"""Each sample costs 5 ms - a stand-in for JPEG decode + augmentation."""
def __len__(self):
return 256
def __getitem__(self, idx):
time.sleep(0.005)
return torch.rand(3, 32, 32), idx % 2
def profile_epoch(loader, model, opt, loss_fn):
wait, compute = 0.0, 0.0
mark = time.perf_counter()
for x, y in loader: # time spent on this line = waiting for data
wait += time.perf_counter() - mark
mark = time.perf_counter()
opt.zero_grad()
loss_fn(model(x.flatten(1)), y).backward()
opt.step()
compute += time.perf_counter() - mark
mark = time.perf_counter()
return wait, compute
if __name__ == "__main__":
model = nn.Sequential(nn.Linear(3 * 32 * 32, 128), nn.ReLU(), nn.Linear(128, 2))
opt = torch.optim.SGD(model.parameters(), lr=0.01)
loss_fn = nn.CrossEntropyLoss()
for workers in [0, 4]:
loader = DataLoader(SlowDataset(), batch_size=32, num_workers=workers,
persistent_workers=workers > 0)
profile_epoch(loader, model, opt, loss_fn) # warm-up: pays worker start-up once
wait, compute = profile_epoch(loader, model, opt, loss_fn)
print(f"num_workers={workers}: epoch {wait + compute:.2f} s "
f"(waiting {wait:.2f} s, computing {compute:.2f} s)")num_workers=0: epoch 1.44 s (waiting 1.43 s, computing 0.01 s) num_workers=4: epoch 0.36 s (waiting 0.36 s, computing 0.01 s)
Diagnosis at a glance: waiting dwarfs computing, so this run is starved — and four workers cut the epoch by 4x without touching the model. Still waiting-dominated, so on this synthetic job even more parallelism would help. On your real job, run this, read your own split, and act on that.
The warm-up epoch matters. The first pass pays worker start-up — the Windows spawn tax from the workers lesson. Profiling it would smear a one-time cost across your steady-state numbers. Warm up, then measure.
On a real GPU, one extra rule
GPU work is asynchronous — Python queues instructions and races ahead; the chip executes later. Timing with time.perf_counter() around GPU code measures queueing, not computing. Synchronise first:
torch.cuda.synchronize() # wait for the GPU to actually finish
mark = time.perf_counter()Call it before each timing mark (it exists on CPU-only installs too — a harmless no-op there, so the profiler above stays portable). The quick sanity check meanwhile is nvidia-smi while training: GPU utilisation bouncing near 30% with dips is the starvation signature; a steady 95%+ means the pipeline is keeping up.
The escalation ladder, cheapest first
- More workers — measure at 0, 2, 4, 8; keep the knee.
pin_memory=True+non_blocking=True. Pinned memory is RAM the operating system promises not to move, which the GPU can copy from directly; thenx.to(device, non_blocking=True)overlaps the copy with computation instead of stalling on it.- Cheapen the per-sample work — decode smaller images, move heavy transforms earlier (pre-resize the dataset once on disk).
- Faster storage layout — the memmap and streaming lessons.
- Only now, consider the model or the hardware.
Common mistakes
Profiling epoch zero. Start-up costs pollute it. Warm up first, as above.
Trusting task manager CPU% for the diagnosis. High CPU can be workers keeping up fine; low CPU can hide a disk bottleneck. Time the halves in the loop — the loop cannot lie about itself.
Fixing without a baseline number. "I added workers and it feels faster" is how folklore forms. Record wait/compute before and after; keep changes that move the number.
Timing GPU code without synchronising. The classic: compute appears free, waiting appears huge, and the diagnosis inverts. Synchronise, then trust.
Try it yourself
Fatten the model (nn.Linear(3072, 4096), an extra hidden layer or three) and rerun. Watch the split flip toward compute-dominated — then reason about which fixes on the ladder become pointless the moment it flips.
What to learn next
- num_workers, prefetching and the Windows spawn trap — the first rung of the ladder, in depth.
- Datasets that do not fit in memory — storage-layout fixes when workers are not enough.
- Image transforms with torchvision v2 — cheapening the per-sample work itself.
Researcher — Mathematics and papers.
The pipeline as a two-stage queue
Model the step as producer (pipeline, mean service time $t_d$ per batch, $W$ workers) feeding consumer (GPU, $t_c$ per batch): steady-state epoch time per batch is $\max(t_d / W, t_c)$ plus transfer, and utilisation of the consumer is $\min(1, t_c \cdot W / t_d)$. The two-clock profiler estimates exactly these: wait $\approx \max(0, t_d/W - t_c)$ per batch aggregated, compute $\approx t_c$. Prefetch depth (prefetch_factor × W batches) buffers variance in $t_d$ — long-tailed decode times, page-cache misses — not its mean; a persistently empty queue certifies mean-starvation, per Little's law, while occasional dips with a full-on-average queue indicate variance, fixable by deeper prefetch rather than more workers. This is the quantitative frame behind the escalation ladder.
Overlap mechanics: pinned memory and streams
cudaMemcpyAsync requires page-locked (pinned) host memory — pageable memory forces a synchronous staging copy. DataLoader(pin_memory=True) pins batches on a dedicated thread post-collate; tensor.to(device, non_blocking=True) then enqueues the H2D copy on a stream, overlapping with compute provided (1) the source is pinned and (2) no premature synchronisation intervenes. Full three-way overlap — load $b_{n+1}$, copy $b_{n+1}$, compute $b_n$ — is the standard prefetch-to-device pattern (a small wrapper advancing one batch ahead on a side stream); frameworks like DALI and Lightning's fabric implement it internally. Diagnostic tooling above the two-clock level: torch.profiler with schedule+TensorBoard traces shows H2D copies and kernel gaps on the timeline; Nsight Systems gives the same at driver granularity; the empty-gap-between-kernels signature is data starvation rendered visually.
torch.cuda.synchronize() inside measurement code is itself a perturbation — it drains all streams, serialising the overlap being measured. CUDA events (torch.cuda.Event(enable_timing=True)) bracket device work without global drains and are the precision instrument; the two-clock wall profiler remains the correct triage instrument because it attributes time to pipeline stages, which device-side events cannot see.
When the pipeline cannot be saved
At the limit — storage bandwidth exhausted, decode CPU-bound at all cores — remaining moves change the workload: GPU-side decode/augment (DALI; nvJPEG), storage co-design (FFCV; WebDataset sequential shards), caching decoded tensors across epochs (RAM or local NVMe tiers), and data echoing (Choi et al., 2020) — reusing each batch $e$ times, provably harmless to final quality for small $e$ on many workloads while multiplying effective pipeline throughput by $e$. Mohan et al. (2021), Analyzing and Mitigating Data Stalls in DNN Training, measured stall prevalence across production-representative jobs and found data stalls common enough to motivate all of the above — the paper to read before believing your job is special.
Reading
- Mohan et al. (2021), Analyzing and Mitigating Data Stalls in DNN Training — the empirical case.
- Choi et al. (2020), Faster Neural Network Training with Data Echoing.
- PyTorch docs: CUDA semantics (asynchronous execution, streams) and
torch.profiler— the normative tooling references.
What to learn next
- num_workers, prefetching and the Windows spawn trap — the first rung of the ladder, in depth.
- Datasets that do not fit in memory — storage-layout fixes when workers are not enough.
- Image transforms with torchvision v2 — cheapening the per-sample work itself.