Data parallel training is the first scaling tool almost everyone reaches for and the one most often run slightly wrong. The idea fits in a sentence: put a full copy of the model on every GPU, give each copy a different slice of the batch, average the gradients, and take the same optimizer step everywhere. The details are where jobs go wrong: the sampler that feeds two ranks the same examples, the BatchNorm layer that never sees the global batch, the branch that runs on rank 0 only and hangs the other 63, the last gradient bucket that nobody overlapped.
This article is about PyTorch DistributedDataParallel (DDP) as an engine. The communication arithmetic is derived in DDP math and the collective itself in ring all-reduce; here we cover what the reducer does, a complete launchable script, the knobs that matter, how to measure scaling, and the traps that cost the most debugging time.
The one invariant data parallelism keeps
Data parallelism makes one promise: N workers each processing B examples produce the same update as one worker processing N times B examples. That holds because the gradient of a mean loss is the mean of per-example gradients, so averaging per-rank gradients is exactly the gradient of the union batch. Everything DDP does exists to keep a single invariant true: after every step, every rank holds identical weights. DDP never re-broadcasts weights after start-up; it relies on identical starting weights, identical averaged gradients and a deterministic optimizer.
The promise has fine print. Anything computed from the local batch only, rather than from gradients, does not get averaged: BatchNorm running statistics, per-rank loss values you log, data augmentation randomness. Anything that differs between ranks and feeds into the weights breaks the invariant without an error message. Keep that sentence in mind for every trap later in this article.
Inside DDP: the reducer, buckets and hooks
When you wrap a module, DDP builds a reducer. It broadcasts rank 0's parameters and buffers so every replica starts identical, then groups parameters into buckets of about bucket_cap_mb megabytes (25 by default) in roughly reverse registration order, because backward produces gradients for the last layers first. It registers an autograd hook on each parameter. During backward, each hook copies its gradient into the bucket's flat buffer and marks it ready; when every gradient in a bucket is ready, the reducer launches an asynchronous all-reduce for that bucket on NCCL's stream while autograd keeps computing earlier layers.
At the end of backward the reducer waits for all outstanding buckets, divides by the world size and copies averaged values back into each .grad. Your optimizer then runs locally, unaware anything distributed happened. The overlap is the whole point: if backward takes 300 ms and the all-reduces take 200 ms, a well-bucketed job hides most of the 200 ms, and only the final bucket, issued after the first layers finish, is exposed.
A complete training script
The script below is a complete, launchable DDP training loop. torchrun starts one process per GPU and sets RANK, LOCAL_RANK and WORLD_SIZE; the process group reads them from the environment.
# train_ddp.py launch: torchrun --nnodes=4 --nproc-per-node=8 \
# --rdzv-backend=c10d --rdzv-endpoint=$HEAD:29500 train_ddp.py
import os, datetime, torch, torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
def main():
dist.init_process_group("nccl", timeout=datetime.timedelta(minutes=20))
rank, local = dist.get_rank(), int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local)
torch.manual_seed(1234) # same init on every rank
model = build_model().cuda()
model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model) # only if it has BN
model = DDP(model, device_ids=[local], gradient_as_bucket_view=True)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
ds = build_dataset()
sampler = DistributedSampler(ds, shuffle=True, drop_last=True, seed=1234)
loader = DataLoader(ds, batch_size=32, sampler=sampler, num_workers=8,
pin_memory=True, drop_last=True)
for epoch in range(EPOCHS):
sampler.set_epoch(epoch) # new shuffle each epoch
for x, y in loader:
x, y = x.cuda(non_blocking=True), y.cuda(non_blocking=True)
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = model(x, y)
loss.backward() # all-reduce happens in here
opt.step(); opt.zero_grad(set_to_none=True)
if rank == 0: # one writer, unwrapped weights
torch.save({"model": model.module.state_dict(),
"opt": opt.state_dict(), "epoch": epoch}, f"ckpt_{epoch}.pt")
dist.barrier() # nobody races ahead of the save
dist.destroy_process_group()
if __name__ == "__main__":
main()Three lines carry most of the correctness. set_epoch changes the shuffle seed; without it every epoch replays the same order. Saving model.module keeps the module. prefix out of the checkpoint so it loads without DDP. The barrier stops fast ranks from starting the next epoch, and timing out, while rank 0 is still writing a large file.
Gradient accumulation without wasted all-reduces
Gradient accumulation runs several micro-batches before one optimizer step, to reach a target global batch without more GPUs. Under DDP the naive version all-reduces on every micro-batch, paying full communication for steps that do not update weights. no_sync() suspends the reducer; gradients accumulate locally and the final micro-batch, outside the context, triggers one all-reduce of the accumulated sum.
ACCUM = 4
for i, (x, y) in enumerate(loader):
last = (i + 1) % ACCUM == 0
ctx = contextlib.nullcontext() if last else model.no_sync()
with ctx:
loss = model(x.cuda(), y.cuda()) / ACCUM # keep the mean loss scale
loss.backward()
if last:
opt.step(); opt.zero_grad(set_to_none=True)Global batch is per-GPU batch times accumulation steps times world size: 32 x 4 x 32 = 4,096 here. Dividing the loss by ACCUM keeps gradient magnitude equal to a single large batch, so learning-rate settings transfer.
The knobs that matter
DDP has few constructor arguments, and each one trades memory, speed or flexibility.
| Knob | What it does | When to change it |
|---|---|---|
bucket_cap_mb | Target bucket size (default 25 MB) | Raise on fast fabrics to cut per-collective latency; lower to start overlap sooner on small models |
gradient_as_bucket_view=True | .grad tensors become views into bucket buffers | Almost always: saves one gradient-sized copy of memory |
find_unused_parameters=True | Walks the graph each step to find params with no grad | Only when some parameters genuinely skip backward; it costs a graph traversal every step |
static_graph=True | Promises the used-parameter set never changes | Fixed-graph models, including those using activation checkpointing; enables optimisations |
register_comm_hook | Replaces the default all-reduce | bf16 compression to halve bytes on slow links; custom schemes |
The compression hook is one line: model.register_comm_hook(None, default_hooks.bf16_compress_hook) from torch.distributed.algorithms.ddp_comm_hooks. It casts each bucket to bfloat16, all-reduces and casts back, halving traffic for fp32 gradients at the cost of rounding in the sum. Validate loss curves before adopting it for a long run.
Worked example: measuring scaling efficiency
Suppose a 350M-parameter model trains on one 8-GPU node at 410,000 tokens per second and you move to four nodes. The fp32 gradients are 1.4 GB. A ring all-reduce sends about twice the buffer per GPU, roughly 2.8 GB, so at a measured 45 GB/s bus bandwidth between nodes it needs about 62 ms per step. If backward takes 250 ms, almost all of it can overlap; the exposed part is the last bucket, around 25 MB, well under 2 ms. Expected throughput is close to 4 x 410,000. Measure it rather than trusting the estimate:
def tokens_per_sec(model, loader, steps=50, warmup=10):
it = iter(loader)
for i in range(warmup + steps):
if i == warmup:
torch.cuda.synchronize(); t0 = time.perf_counter()
x, y = next(it)
model(x.cuda(), y.cuda()).backward()
opt.step(); opt.zero_grad(set_to_none=True)
torch.cuda.synchronize()
local = steps * x.numel() / (time.perf_counter() - t0)
total = torch.tensor(local, device="cuda")
dist.all_reduce(total) # sum across ranks
return total.item()
eff = tokens_per_sec(model, loader) / (world_nodes * single_node_tps)
print(f"scaling efficiency {eff:.1%}") # under ~0.9: find out whyWhen efficiency falls well short, compare a step with no_sync() around everything (no communication) against a normal step. If the no-sync step is also slow, the problem is input loading or a straggler GPU, not the network. If only the normal step is slow, profile it and check whether NCCL kernels overlap backward or queue after it.
Silent correctness traps
These mistakes produce no error; they produce a model that trains worse.
- Missing
set_epoch. Every epoch sees the same order. Harmless on huge datasets, measurable on small ones. - Evaluation with
DistributedSampler. Withoutdrop_lastit pads by repeating samples so every rank gets equal counts, so a distributed validation score can count some examples twice. Evaluate on rank 0, or sum correct counts and totals withall_reduceover an unpadded split. - BatchNorm. Each rank normalises with its local batch statistics. With small per-GPU batches that is noisy;
SyncBatchNormcomputes them across ranks at the cost of an extra collective per layer. - Per-rank randomness. Model initialisation must use the same seed everywhere (DDP's broadcast covers you), but augmentation and dropout should differ by rank, or replicas do redundant work. Seed data workers with the rank mixed in.
- Logging rank 0's loss. It is one shard's loss, not the global one. All-reduce it before plotting if the curve drives decisions.
Failure modes
These are the failures that stop or hang a job.
- Hang from rank-conditional code. Any collective, including a forward through the DDP module, called by some ranks and not others blocks forever, then dies at the process-group timeout. Classic causes: evaluating the wrapped model on rank 0 only, and an early
breakwhen one rank runs out of data. Usemodel.modulefor rank-local work, and wrap uneven input loops intorch.distributed.algorithms.join.Join([model]). - Unused-parameter error. "Expected to have finished reduction in the prior iteration" means a parameter got no gradient. Fix the model or set
find_unused_parameters=True;TORCH_DISTRIBUTED_DEBUG=DETAILnames the offending parameters. - Rank 0 OOM at resume.
torch.loadwithoutmap_locationputs every rank's copy on GPU 0. Load withmap_location="cpu"or the local device. - Straggler. All-reduce runs at the speed of the slowest rank. A thermally throttled GPU or a host doing heavy decode drags all 32. Compare per-rank step times before blaming the fabric.
- Silent drift. Non-deterministic custom ops, or a buffer updated outside autograd, can make weights diverge. Check periodically by all-reducing a parameter checksum and comparing min and max across ranks.
Trade-offs
Plain DDP is right when the model, its gradients and its optimizer state fit on one GPU with room for activations. With Adam in mixed precision that is about 16 bytes per parameter before activations, so an 80 GB GPU holds roughly a 3 to 4 billion parameter model at best. Beyond that you shard: ZeRO partitions optimizer state and gradients across the same data-parallel ranks, and FSDP also shards parameters, trading extra all-gathers for memory. Both are still data parallelism: each rank sees different data.
Even when DDP fits, its costs rise with scale. Very large global batches can hurt convergence, and the learning-rate schedule must be retuned. Gradient accumulation trades throughput for batch size without more GPUs. And bitwise reproducibility across different world sizes is not available, because the reduction order changes; see ML reproducibility on GPUs.
What to do next
- Launch with
torchrun, callset_epochevery epoch, and savemodel.modulefrom rank 0 followed by a barrier. - Turn on
gradient_as_bucket_view; setstatic_graphif your graph is fixed; leavefind_unused_parametersoff unless needed. - Use
no_sync()for every non-final micro-batch when accumulating. - Measure tokens per second at one node and at N nodes, and compute scaling efficiency.
- If efficiency is low, run a no-sync step to separate input, straggler and network causes.
- Audit rank-conditional code paths and evaluation sampling before the first long run.
- Add a periodic parameter-checksum comparison across ranks to catch silent drift.
- When memory runs out, move to ZeRO or FSDP rather than shrinking the batch.