The ladder: from stage 1 to full partitioning
ZeRO (Zero Redundancy Optimizer) attacks the redundancy in plain data parallelism, where every GPU holds an identical, complete copy of the parameters, gradients, and optimizer states. That replication is pure waste: N GPUs store the same numbers N times. ZeRO removes it in three stages, each partitioning one more category of state across the N data-parallel ranks.
Stage 1 shards the optimizer states — the largest consumer — so each rank owns 1/N of them. Stage 2 additionally shards the gradients: after backward, each rank keeps only the gradient slice it needs to update its optimizer shard. Stage 3 takes the final step and shards the parameters themselves. After Stage 3, nothing is fully replicated: every rank permanently holds 1/N of the parameters, 1/N of the gradients, and 1/N of the optimizer states. The model no longer has to fit on one device — only 1/N of it does. Everything in this article is about what that last step buys (near-linear memory) and what it costs (extra communication to reconstruct weights on demand).
The memory model: 16 bytes per parameter
To see what partitioning saves, you first need the per-parameter memory bill of mixed-precision Adam training. For a model with Ψ parameters, the standard accounting stores, per parameter:
fp16 parameters : 2 bytes
fp16 gradients : 2 bytes
optimizer states (K) : 12 bytes
- fp32 master copy : 4
- fp32 momentum : 4
- fp32 variance : 4
-----------------------------
total = 2 + 2 + K = 16 bytes / parameter (K = 12)The two easy-to-miss facts live in K. First, Adam keeps two running statistics — momentum and variance — each a full fp32 tensor. Second, mixed precision keeps an fp32 master copy of the weights (the fp16 copy is only for the forward/backward math), and that master copy lives inside K, not inside the leading 2. So the true cost is 16Ψ bytes, and the optimizer states — the K = 12 term — are three-quarters of it. That is why ZeRO shards them first.
Stage 3: divide everything by N
The three stages differ only in which terms get the /N divisor. Writing the per-parameter cost as a formula makes the progression exact:
Baseline (DDP) : 2 + 2 + K = 16 bytes/param
ZeRO-1 : 2 + 2 + K/N = 4 + 12/N bytes/param
ZeRO-2 : 2 + (2 + K)/N = 2 + 14/N bytes/param
ZeRO-3 : (2 + 2 + K)/N = 16/N bytes/paramStages 1 and 2 still carry constant terms — the 2 + 2 or the 2 — because a full-precision copy of the parameters (and, in Stage 1, the gradients) stays resident on every GPU. Those constants are what cap the model size: no matter how many GPUs you add, ZeRO-2 never drops below 2 bytes per parameter for the resident weights. Stage 3 puts the last constant under the divisor. The per-parameter cost becomes 16/N, with no floor — add GPUs and the per-device footprint keeps shrinking. This is what ‘near-linear’ memory scaling means, and it is the whole reason Stage 3 exists.
A worked memory example
Take a Ψ = 10 billion parameter model on GPUs with 40 GB of memory each. Under plain data parallelism the state alone is:
16 bytes × 10e9 params = 160 GB per GPUThat does not fit — not on a 40 GB card, not on an 80 GB one, and adding more GPUs in plain DDP does nothing, because every GPU still needs the whole 160 GB. Now shard with Stage 3 across N = 64 GPUs:
160 GB / 64 = 2.5 GB per GPU (persistent shard)The same 160 GB of state now spreads to 2.5 GB on each device, leaving the rest of the 40 GB for activations and the transient buffers described below. The canonical figure from the ZeRO paper is the same shape: a 7.5B model at N = 64 drops from about 120 GB to roughly 1.9 GB per GPU. The lesson is that Stage 3 turns ‘too big for any GPU’ into ‘trivially small per GPU,’ and the more GPUs you pool, the smaller each one’s share.
Just-in-time all-gather: how the weights reappear
If each rank only stores 1/N of every weight, how does it run a forward pass that needs the whole weight? The answer is the mechanism that defines Stage 3: a just-in-time all-gather. Parameters are organized into shardable units — typically one transformer layer’s worth. Immediately before a layer’s forward compute, the N ranks perform an all-gather so that every rank temporarily reconstructs that layer’s full parameters. The layer runs. Then the gathered full copy is freed, and each rank falls back to holding just its 1/N shard again.
The same dance repeats in the backward pass. Because the full weights were discarded after the forward, they must be all-gathered a second time before the layer’s backward compute, used to compute gradients, and freed again. Finally the freshly computed gradients are reduce-scattered: summed across ranks and split so each rank receives only the 1/N gradient slice matching its parameter shard. At no instant does any GPU hold more than one layer’s full weights beyond its permanent shards — the model is materialized in slices, on demand.
Peak transient memory
Just-in-time gathering means the memory formula 16Ψ/N describes the persistent footprint, not the peak. During a layer’s forward or backward, that layer’s full parameters are briefly resident on every rank. So the peak working-set is approximately:
peak ≈ (persistent 16Ψ/N shard) + (one layer's full params) + activationsThis is a deliberate and favorable trade. A single layer is a tiny fraction of a deep model — a 60-layer network materializes roughly 1/60 of the parameters at a time — so the transient bump is small compared to a full replica. It also explains a practical knob: the granularity of the shard unit. Gathering bigger chunks (several layers at once) overlaps communication with compute better and can raise throughput, but it raises the transient peak, since more full weights are resident at once. Smaller units keep peak memory down at the cost of more, smaller collectives. This tension — peak memory versus communication efficiency — is the central thing you tune when running Stage 3.