Distributed Training: Data Parallel
Many GPUs, one model: copy it to every GPU, split the batch, average the gradients with a ring all-reduce, hide that behind the backward pass, and find out why adding GPUs eventually stops helping.
An interactive AI Infrastructure lesson: 23 steps, about 32 minutes, on a live simulation in your browser.
Your team wants to pre-train a 1.5-billion-parameter GPT (the GPT-2 XL shape) on 30 billion tokens. One training step takes a batch of 256 sequences of 1,024 tokens, 262,144 tokens in all, runs them forward and backward, and updates every weight once.
Start it on one H100 of dgx-1. It fits easily (29 GB of 80), and each step takes 5.95 s: 44,048 tokens a second at 42% MFU (model FLOPs utilisation, the share of the GPU's peak doing useful maths).
What you will learn
One GPU is too slow
- Eight days on one GPU: Training time is total FLOPs divided by the FLOPs you can deliver per second. When one GPU is already well used, the only lever left is more GPUs.
- Copy the model, split the batch: Data parallelism divides the batch, not the model: N copies each do 1/N of the maths, so the step shrinks by almost N as long as the copies can agree quickly.
Averaging gradients: the all-reduce
- Eight copies must stay identical: Data-parallel replicas stay identical because every step ends with the same averaged gradient on every GPU. The all-reduce of the full gradient is the cost of that agreement.
- The ring all-reduce: In a ring all-reduce every GPU sends and receives 2(N−1)/N of the buffer, just under twice its size, whether N is 8 or 8,000. Time depends on the slowest link, not on N.
- Reading all_reduce_perf: algbw and busbw: algbw = size / time describes the operation; busbw = algbw × 2(N−1)/N describes the wires. Compare busbw with the link speed, and look at large sizes: small buffers are latency-bound.
- Drill: how long is the all-reduce?: All-reduce time ≈ gradient bytes × 2(N−1)/N ÷ busbw. For data parallelism the bytes are 2 × parameters, whatever the batch size.
Hiding the communication
- Overlap: reduce while you compute: Communication only costs step time where it is exposed. DDP overlaps the gradient all-reduce with the backward pass, so a job slows down only when the all-reduce outlasts the backward pass.
- What DDP looks like in code: One process per GPU, one full model per process, one sampler shard per process: that is DDP. The communicator and the gradient hooks do the rest.
Across four servers
- The ring leaves the server: Across servers the ring runs at the slower of NVLink and the server's total NIC bandwidth. One 400 Gb/s NIC per GPU makes the network roughly as fast as NVLink for an all-reduce.
- 32 GPUs: six hours instead of eight days: Scaling efficiency = throughput on N GPUs ÷ (N/N₀ × throughput on N₀). Above 90% the job is compute-bound and more GPUs pay; the number falls as communication becomes exposed.
- Global batch size and N: Global batch = micro-batch × accumulation steps × data-parallel ranks. Hold it fixed and N is capped by the batch; grow it with N and you change the training run.
Where scaling stops
- A small batch on many GPUs: Adding GPUs shrinks each GPU's compute but not the gradient all-reduce. When the backward pass gets shorter than the all-reduce, the rest is exposed and every extra GPU buys less.
- Why adding GPUs stops helping: Data-parallel scaling ends when per-GPU compute is too small to hide the fixed gradient all-reduce, or when there are no more sequences to hand out. Both are set by the global batch.
- Break it: one cable with symbol errors: A ring runs at the speed of its slowest link, so one degraded cable slows every GPU in the job. A collective test across all servers finds it in a minute; per-node tests never will.
One slow GPU
- Break it: one hot GPU: A synchronous job runs at the pace of its slowest GPU. One straggler costs the whole job, so the fleet's worst GPU matters more than its average.
- Find the straggler, fix it: Diagnose a slow job in layers: the step breakdown says compute or communication, DCGM says which GPU, nvidia-smi's clock event reasons say why.
Launching it for real
- Four nodes: sbatch, srun and torchrun: Multi-node DDP is three layers: the scheduler allocates nodes, a launcher starts one process per GPU, and a rendezvous gives every process its rank.
- NCCL environment and NCCL_DEBUG=INFO: NCCL fails slow, not loud. NCCL_DEBUG=INFO shows the transport it chose; check for IB and GDRDMA before blaming the model.
- Drill: launch on one node: torchrun is the per-node launcher: it starts one process per GPU and hands each its LOCAL_RANK, RANK and WORLD_SIZE.
- Drill: make NCCL explain itself
- Break it: a model that does not fit: Data parallelism divides the work, never the memory: every GPU holds the full model, gradients and optimizer state. When that exceeds one GPU, you must shard or split the model.
Recap & playground
- Cheat sheet
- Playground: tune the pod