← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 8 · Chapters 16, 20

Data Parallelism: DDP, ZeRO and FSDP

DDP replicates everything, and that is its limit. The gradient all-reduce scales beautifully, but every rank still stores a full copy of the parameters, the gradients and the optimizer state. ZeRO observes that a replica which is identical everywhere is $P-1$ copies of wasted memory, and eliminates it one category at a time.

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.

1

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.

2

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.

3

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.

Bucket order must match on every rank. Ranks must issue collectives in the same sequence. If control flow differs between ranks — a conditional branch, a layer skipped on some inputs — some gradients never become ready, the bucket never fires, and the job hangs rather than erroring.

The memory ledger

Mixed-precision training with Adam keeps five things per parameter:

CategoryContentsBytes per parameter
ParametersBF16 weights used in the forward pass2
GradientsBF16 gradients2
Optimizer stateFP32 master copy of the weights4
Adam first moment4
Adam second moment4
Total16

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 ZeRO insight. A tensor that is bit-identical on all $P$ ranks and needed only briefly does not have to be replicated. Shard it, and gather or reconstruct just the piece each rank needs, when it needs it.

The three stages

StageShardedReplicatedBytes per parameterCommunication per step
DDPnothingparams, grads, optimizer$16$$2N$ (all-reduce)
ZeRO-1optimizer stateparams, grads$4+\dfrac{12}{P}$$2N$
ZeRO-2optimizer state, gradientsparams$2+\dfrac{14}{P}$$2N$
ZeRO-3 / FSDPeverythingnothing$\dfrac{16}{P}$$3N$
Stages 1 and 2 are free. They move the same number of bytes as DDP. The reason is the decomposition from the collectives note: DDP's all-reduce is a reduce-scatter followed by an all-gather. ZeRO-1 and ZeRO-2 simply stop after the reduce-scatter, update the shard each rank owns, and all-gather the updated parameters instead of the gradients. Same two halves, same volume, far less memory. If you are running plain DDP and fit comfortably, you are leaving this on the table for nothing.

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:

$$\underbrace{N}_{\text{all-gather (fwd)}}+\underbrace{N}_{\text{all-gather (bwd)}}+\underbrace{N}_{\text{reduce-scatter (grads)}}=3N,$$

which is $1.5\times$ DDP's volume. In exchange, memory falls by a factor of $P$ with no upper bound.

parameters gradients optimizer state exceeds 80 GB
DDP per rank
ZeRO-2 per rank
ZeRO-3 per rank
min stage that fits 80 GB
Model state only; activations are extra and are usually what consumes the remaining headroom. The 80 GB line marks a typical accelerator.

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.

One FSDP step
1forward, for each unit in order:
2all-gather this unit's parameter shards into a full set
3run the unit; save activations needed for backward
4free the gathered parameters immediately
5backward, for each unit in reverse:
6all-gather the parameters again (they were discarded)
7compute input and weight gradients
8reduce-scatter the weight gradients, so each rank keeps only its shard
9free the gathered parameters and the full gradient
10update: each rank steps only the shard it owns; no further communication

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.

Prefetching is what makes it viable. While unit $i$ computes, the gather for unit $i+1$ is already in flight. Without prefetch, every unit stalls on its own gather and FSDP is dramatically slower than DDP; with it, the extra $N$ of traffic largely disappears behind compute. If FSDP is much slower than expected, prefetching is the first thing to check.

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:

Shard within a node8 ranks over NVLink, all-gather stays local
→
Replicate across nodesplain gradient all-reduce between nodes
→
Resultmemory divided by 8, chatty traffic on the fast link

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

SituationUse
Model state fits comfortably per rankZeRO-1 or ZeRO-2 — same communication as DDP, strictly less memory
Parameters alone do not fitZeRO-3 / FSDP, with prefetching enabled and unit size tuned
Fast intra-node fabric, slow inter-nodehybrid sharding
A single layer does not fit on one devicedata parallelism is not the answer; go to tensor or pipeline parallelism
Takeaway

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