Training modern foundation models requires thousands of GPUs working together for weeks. Even university projects often benefit from spreading training across several GPUs. Distributed training is fundamentally about two constraints: time (we want training to finish sooner) and memory (the model and its training state may not fit on one device). Different parallelism strategies address each.
Data parallelism#
The simplest and most common approach: replicate the model on every device, give each device a different slice of the mini-batch, and average gradients before updating.
- Each of $N$ workers holds a full model copy.
- Each computes forward and backward passes on its local batch $\mathcal{B}_k$.
- Gradients are averaged across workers with an all-reduce operation:
- Every worker applies the identical update, so replicas stay synchronised.
The effective batch size is $N \times$ the per-device batch, so the learning rate and warm-up usually need adjusting.
All-reduce#
A naive approach sends all gradients to one parameter server โ a bandwidth bottleneck. Ring all-reduce arranges workers in a ring; each sends and receives chunks so that the per-worker communication volume is about $2\frac{N-1}{N}$ times the gradient size โ nearly independent of $N$. Libraries such as NCCL implement efficient all-reduce over NVLink and InfiniBand.
DDP in PyTorch#
DistributedDataParallel overlaps communication with computation: it all-reduces gradients in buckets as soon as they are computed during the backward pass.
# launch with: torchrun --nproc_per_node=4 train.py
import os, torch, torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data.distributed import DistributedSampler
dist.init_process_group("nccl")
rank = int(os.environ["LOCAL_RANK"]); torch.cuda.set_device(rank)
model = MyModel().cuda(rank)
model = DDP(model, device_ids=[rank])
sampler = DistributedSampler(train_ds, shuffle=True) # each rank gets a distinct shard
loader = torch.utils.data.DataLoader(train_ds, batch_size=64, sampler=sampler, num_workers=4)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
for epoch in range(epochs):
sampler.set_epoch(epoch) # reshuffle differently each epoch
for xb, yb in loader:
loss = torch.nn.functional.cross_entropy(model(xb.cuda(rank)), yb.cuda(rank))
opt.zero_grad(); loss.backward(); opt.step() # all-reduce happens in backward
if rank == 0:
torch.save(model.module.state_dict(), f"ckpt_{epoch}.pt") # save from one rank only
dist.destroy_process_group()Sharded data parallelism: ZeRO and FSDP#
Plain data parallelism replicates everything โ weights, gradients and optimiser states โ on every GPU. With Adam in mixed precision that is roughly 16 bytes per parameter per GPU: a 10-billion-parameter model needs about 160 GB per GPU, exceeding any single device.
ZeRO (Zero Redundancy Optimizer, Rajbhandari et al., 2020) removes the redundancy by sharding state across the $N$ data-parallel workers:
| Stage | Sharded | Memory per GPU (approx.) |
|---|---|---|
| ZeRO-1 | Optimiser states | ~4 + 12/N bytes/param |
| ZeRO-2 | + Gradients | ~2 + 14/N |
| ZeRO-3 / FSDP | + Parameters | ~16/N |
In ZeRO-3 / PyTorch FSDP (Fully Sharded Data Parallel), each layer's parameters are gathered from all shards just before they are needed (all-gather), used, and released; gradients are reduce-scattered back to their owning shards. Memory scales down with $N$ at the cost of extra communication. Offloading optimiser states or parameters to CPU memory or NVMe extends this further.
Model parallelism#
When a single layer or the activations are too large, or to reduce per-device compute, split the model itself.
Tensor parallelism#
Split individual weight matrices across devices. For a linear layer $\mathbf{Y} = \mathbf{X}\mathbf{W}$, partition $\mathbf{W}$ by columns so each device computes part of the output; Megatron-LM pairs a column-split first MLP layer with a row-split second layer so only one all-reduce is needed per MLP block. Attention heads split naturally across devices. Tensor parallelism requires very fast interconnects and is typically used within a server.
Pipeline parallelism#
Place consecutive groups of layers ("stages") on different devices. Naively, only one stage works at a time โ a large pipeline bubble. Micro-batching (GPipe) splits each batch into micro-batches that flow through the stages concurrently; schedules like 1F1B (one-forward-one-backward) further reduce idle time and memory.
Sequence / context parallelism#
Split long sequences across devices so attention over very long contexts fits in memory.
Expert parallelism#
In mixture-of-experts models, place different experts on different devices and route tokens to them.
3-D parallelism#
Frontier-scale training combines these: tensor parallelism inside a node, pipeline parallelism across nodes, and (sharded) data parallelism across replicas of the pipeline. Choosing the right combination depends on model size, sequence length, cluster topology and interconnect bandwidth.
| Strategy | Solves | Communication | Typical scope |
|---|---|---|---|
| Data parallel (DDP) | Time | Gradient all-reduce per step | Any |
| ZeRO / FSDP | Memory (states, params) | All-gather + reduce-scatter | Many GPUs |
| Tensor parallel | Memory + per-layer compute | Frequent, within layers | Within a node |
| Pipeline parallel | Memory (layers) | Activations between stages | Across nodes |
Practical concerns#
- Reproducibility and randomness: seed each rank; use
DistributedSampler; only rank 0 logs and saves. - Batch normalisation: statistics are per-GPU unless you use SyncBatchNorm.
- Fault tolerance: long runs on thousands of GPUs will see failures โ checkpoint frequently and support automatic restart.
- Communication efficiency: gradient compression, overlapping communication with computation, and topology-aware placement.
- Tools: PyTorch DDP/FSDP, DeepSpeed, Megatron-LM, Hugging Face Accelerate, JAX with
pjit/sharding annotations.