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.

One DDP step: backward fills buckets, each full bucket is all-reduced while backward continuesGPU computeforwardbwd L24-L17bwd L16-L9bwd L8-L1optimizerNCCL streamall-reduce B0all-reduce B1B2B0 readyB1 readyexposed communication: only the last bucketReducer (built once, at wrap time)1. broadcast rank 0 weights to all ranks2. group params into buckets, reverse order3. register an autograd hook per parameter4. hook marks grad ready; full bucket fires5. wait for all buckets, divide by world sizeInvariant after every stepevery rank holds identical averaged gradientsevery rank runs the same optimizer stepso weights stay bitwise identical, with noweight broadcast after initialisationbreak it and ranks silently drift apart
DDP overlaps gradient all-reduce with the rest of backward. Bucket order follows backward order, so the last bucket (the first layers) is the only communication that cannot hide behind compute.

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.

KnobWhat it doesWhen to change it
bucket_cap_mbTarget 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 buffersAlmost always: saves one gradient-sized copy of memory
find_unused_parameters=TrueWalks the graph each step to find params with no gradOnly when some parameters genuinely skip backward; it costs a graph traversal every step
static_graph=TruePromises the used-parameter set never changesFixed-graph models, including those using activation checkpointing; enables optimisations
register_comm_hookReplaces the default all-reducebf16 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 why

When 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. Without drop_last it 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 with all_reduce over an unpadded split.
  • BatchNorm. Each rank normalises with its local batch statistics. With small per-GPU batches that is noisy; SyncBatchNorm computes 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 break when one rank runs out of data. Use model.module for rank-local work, and wrap uneven input loops in torch.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=DETAIL names the offending parameters.
  • Rank 0 OOM at resume. torch.load without map_location puts every rank's copy on GPU 0. Load with map_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

  1. Launch with torchrun, call set_epoch every epoch, and save model.module from rank 0 followed by a barrier.
  2. Turn on gradient_as_bucket_view; set static_graph if your graph is fixed; leave find_unused_parameters off unless needed.
  3. Use no_sync() for every non-final micro-batch when accumulating.
  4. Measure tokens per second at one node and at N nodes, and compute scaling efficiency.
  5. If efficiency is low, run a no-sync step to separate input, straggler and network causes.
  6. Audit rank-conditional code paths and evaluation sampling before the first long run.
  7. Add a periodic parameter-checksum comparison across ranks to catch silent drift.
  8. When memory runs out, move to ZeRO or FSDP rather than shrinking the batch.
Key takeaway: Data parallelism keeps one invariant: identical weights on every rank after every step. DDP maintains it by averaging gradients in buckets that overlap backward. Most problems come from code around it: samplers, BatchNorm, rank-conditional collectives and uneven inputs. Measure scaling efficiency, and shard with ZeRO or FSDP once the model no longer fits.