Data Parallelism: DDP, ZeRO and FSDP
DDP, and why it works
The correctness argument was settled in the autograd note: parameters are replicated, replication is a fan-out, and the chain rule turns a fan-out into a sum. So a mean all-reduce of per-rank gradients gives exactly the full-batch gradient. Everything here is about making that all-reduce cheap.
Bucket the gradients. A large model has thousands of parameter tensors. One collective each would pay the per-message latency thousands of times. DDP groups gradients into buckets of roughly 25 MB and issues one all-reduce per bucket.
Fire on completion, not at the end. Backpropagation produces gradients in reverse layer order, so the last layer's bucket is ready long before the first layer's. A hook marks each gradient ready, and the bucket's collective launches the moment its last member arrives.
Overlap. Those collectives run on a separate stream while the remaining layers are still computing. In a healthy run almost all communication disappears behind the backward pass.
The memory ledger
Mixed-precision training with Adam keeps five things per parameter:
| Category | Contents | Bytes per parameter |
|---|---|---|
| Parameters | BF16 weights used in the forward pass | 2 |
| Gradients | BF16 gradients | 2 |
| Optimizer state | FP32 master copy of the weights | 4 |
| Adam first moment | 4 | |
| Adam second moment | 4 | |
| Total | 16 | |
Twelve of those sixteen bytes are optimizer state, and here is the key observation: the optimizer state is only touched during the update step, after the gradient all-reduce and before the next forward pass. For the rest of the step it sits idle, occupying identical bytes on every rank.
The three stages
| Stage | Sharded | Replicated | Bytes per parameter | Communication per step |
|---|---|---|---|---|
| DDP | nothing | params, grads, optimizer | $16$ | $2N$ (all-reduce) |
| ZeRO-1 | optimizer state | params, grads | $4+\dfrac{12}{P}$ | $2N$ |
| ZeRO-2 | optimizer state, gradients | params | $2+\dfrac{14}{P}$ | $2N$ |
| ZeRO-3 / FSDP | everything | nothing | $\dfrac{16}{P}$ | $3N$ |
ZeRO-3 is the one with a real cost. Parameters are no longer resident, so they must be gathered for the forward pass, discarded, and gathered again for the backward pass:
which is $1.5\times$ DDP's volume. In exchange, memory falls by a factor of $P$ with no upper bound.
How FSDP actually runs
FSDP is ZeRO-3 organized around units — typically one transformer block each. A unit is the granularity at which parameters are gathered and released.
Peak parameter memory is therefore one unit's worth of gathered weights plus $N/P$ of shards, not $N$. Unit size is the tuning knob:
Smaller units
Less transient memory, but more collectives, each too small to reach peak bandwidth. Latency starts to dominate.
Larger units
Efficient collectives, but a bigger transient spike and less opportunity to overlap, since a unit cannot start until its whole gather lands.
Hybrid sharding
Full ZeRO-3 across a 512-rank cluster shards very finely, but every all-gather then crosses the slow inter-node network. Hybrid sharding splits the difference by using two mesh axes:
Memory drops by the intra-node group size rather than the whole world size, which is usually enough, and the expensive parameter gathers never leave the node. This is the first genuinely two-dimensional parallelism scheme in the series, and it previews the device-mesh note.
Choosing a stage
| Situation | Use |
|---|---|
| Model state fits comfortably per rank | ZeRO-1 or ZeRO-2 — same communication as DDP, strictly less memory |
| Parameters alone do not fit | ZeRO-3 / FSDP, with prefetching enabled and unit size tuned |
| Fast intra-node fabric, slow inter-node | hybrid sharding |
| A single layer does not fit on one device | data parallelism is not the answer; go to tensor or pipeline parallelism |
Sixteen bytes per parameter, twelve of them idle most of the step. ZeRO-1 and ZeRO-2 remove that waste for free because DDP's all-reduce was already a reduce-scatter plus an all-gather. ZeRO-3 removes the rest for a 1.5x communication surcharge that prefetching mostly hides.
Check yourself
- Derive the bytes-per-parameter column for all four rows. derivation
- Explain why ZeRO-1 and ZeRO-2 move the same bytes as DDP, using the all-reduce decomposition. reasoning
- Show why ZeRO-3 needs two all-gathers rather than one. reasoning
- For a 70B model, find the smallest $P$ at which ZeRO-3 fits model state in 80 GB per rank. calculation
- Explain how a rank-dependent branch can hang a DDP job. debugging