Skip to content

DDP in Depth

Overview

DistributedDataParallel (DDP) is the default way to scale training across GPUs: each rank holds a full copy of the model, processes its own shard of data, and averages gradients with an all-reduce every step. Correct usage is mostly about process setup and knowing what DDP synchronizes.

  • One process per GPU (world_size = total processes); data is sharded across ranks.
  • init_process_group("nccl") for GPU; "gloo" for CPU.
  • DistributedDataParallel(model) wraps the model; it synchronizes gradients before optimizer.step().
  • The backward pass triggers an all-reduce bucket per parameter group (bucket_cap_mb default 25 MB).
  • DDP averages gradients, not predictions.

DDP memory: each rank stores the full model + optimizer state. When the model alone doesn't fit, move to FSDP (Ch 05-02).


Launching vs Manual Spawning

torchrun --nproc_per_node=4 train_ddp.py

Manual spawn (from a driver script)

import torch.multiprocessing as mp
mp.spawn(main, args=(world_size,), nprocs=world_size, join=True)

Minimal DDP Training Loop

import os, torch, torch.nn as nn, torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

def main(rank, world_size):
 dist.init_process_group("nccl", rank=rank, world_size=world_size)
 torch.cuda.set_device(rank)

 model = nn.Sequential(nn.Linear(8, 8), nn.ReLU(), nn.Linear(8, 2)).cuda()
 ddp = DDP(model, device_ids=[rank]) # wraps & syncs init weights

 opt = torch.optim.SGD(ddp.parameters(), lr=0.1)
 data = torch.randn(64, 8).cuda() # per-rank shard

 for step in range(10):
 opt.zero_grad()
 loss = ddp(data).abs().mean() # forward on this rank's data
 loss.backward() # all-reduce of grads happens here
 opt.step()
 if rank == 0:
 print(f"step {step} loss {loss.item():.3f}")

 dist.destroy_process_group()

if __name__ == "__main__":
 import torch.multiprocessing as mp
 mp.spawn(main, args=(4,), nprocs=4, join=True)

-

How the All-Reduce Actually Works

Backward hooks in DDP bucket gradients by size (bucket_cap_mb), and as each bucket completes it all-reduces asynchronously— overlapping comm with the remaining compute.

backward() ──> reduce bucket 1 (async) ──> reduce bucket 2 ──>...
 └─ overlap with compute

optimizer.step() has to WAIT for all buckets to finish (sync_point).
  • bucket_cap_mb=25 is a good default; tune with profiler for large models.
  • NCCL uses ring all-reduce: bandwidth ∝ 2*(P-1)/P scaling per rank.
  • Gradient averaging requires dividing by world_size— DDP does this automatically.

What DDP Syncs (and Doesn't)

Item Behavior
Model weights at init synced once at wrap time
Gradients each step all-reduced (averaged)
Optimizer state NOT synced— make sure starting state matches
BatchNorm stats synced via SyncBatchNorm if used
Random seeds NOT synced— set them yourself
# reproducibility helper
def seed_everything(seed):
 import random
 random.seed(seed)
 torch.manual_seed(seed)
 torch.cuda.manual_seed_all(seed)

-

Common Mistakes

  1. Forgetting device_ids when wrapping on GPU.
  2. .cuda() before DDP— wrong device → crash; wrap then .to(device).
  3. Uneven data shards— use DistributedSampler with drop_last=True or pad.
  4. BatchNorm in DDP— stats aren't aggregated; use dist.SyncBatchNorm.convert_sync_batchnorm(model).
  5. Working with model instead of ddp.module— names change: ddp.module.fc1.
  6. Saving only rank 0's state_dict— all ranks should see identical grads/weights, but save on one rank to avoid races.
from torch.utils.data import DataLoader, DistributedSampler

N = 1000
sampler = DistributedSampler(
 torch.arange(N), num_replicas=world_size, rank=rank, shuffle=True)
loader = DataLoader(torch.arange(N), batch_size=16, sampler=sampler)
# IMPORTANT

SyncBatchNorm

ddp = DDP(dist.SyncBatchNorm.convert_sync_batchnorm(model), device_ids=[rank])
# BatchNorm stats are now all-reduced across ranks.

-

Key Takeaways

  • DDP = one full model per rank + gradient all-reduce. Simple, strong baseline.
  • Bucketed async all-reduce overlaps comm and compute— tune bucket_cap_mb.
  • Use torchrun; respect rank == 0 for logging/saving.
  • Watch seeds, samplers, and BatchNorm— DDP does not fix those for you.
  • Move to FSDP when a single GPU can't hold model + optimizer + activations.

-

  • Distributed Home
  • Fsdp
  • [Checkpointing](/06-pytorch/05-distributed-training/(04-checkpointing-fault-tolerance/)