The memory model we are optimizing

Fix the accounting first, because ZeRO is entirely a story about bytes per parameter. Train a model of Ψ parameters with mixed-precision Adam and a standard optimizer keeps, per parameter: an fp16 weight (2 bytes), an fp16 gradient (2 bytes), and a bundle of fp32 optimizer states — a master fp32 copy of the weights (4), the Adam momentum (4), and the variance (4), so K = 12 bytes.

per-param bytes = 2 (fp16 weight)
                + 2 (fp16 gradient)
                + K (fp32 optimizer states),  K = 12
                = 16 bytes/param   →   total = 16Ψ

Plain distributed data parallel (DDP) replicates all 16Ψ bytes on every one of the N GPUs. That is the waste ZeRO attacks: the replicated state is identical across ranks, so storing N copies buys nothing. The question each ZeRO stage answers is simply which of those three buckets we are allowed to slice into 1/N-sized shards.

Advertisement

Where Stage 1 stopped

Stage 1 (P_os, partition optimizer states) shards only the K = 12 bytes. Each rank owns 1/N of the parameters’ optimizer states and is responsible for updating exactly that slice. But every rank still stores a full fp16 gradient buffer of 2Ψ bytes, because a normal all-reduce hands every rank the complete averaged gradient before the optimizer step.

Stage 1 memory/GPU = 2Ψ (weights)
                   + 2Ψ (full gradients)
                   + 12Ψ/N (sharded optimizer states)
                   = 4Ψ + 12Ψ/N

Notice the tension: a rank keeps the whole gradient but only ever uses 1/N of it — the slice matching its optimizer-state shard. The other (N-1)/N of the gradient is dead weight the instant the update runs. Stage 2 is the observation that we never needed to materialize it.

Advertisement

What Stage 2 adds: partition the gradients

Stage 2 (P_os+g) keeps everything Stage 1 does and additionally partitions the gradients. After the backward pass, rank i ends up holding only the reduced gradient for the parameters it owns — a 2Ψ/N-byte shard — not the full 2Ψ buffer.

Stage 2 memory/GPU = 2Ψ (weights, still replicated)
                   + 2Ψ/N (sharded gradients)
                   + 12Ψ/N (sharded optimizer states)
                   = 2Ψ + 14Ψ/N

Gradient memory has dropped from 2 bytes/param to 2/N bytes/param. As the world size grows the two /N terms vanish and each GPU asymptotes to 2Ψ — just the fp16 weights, which Stage 2 deliberately leaves replicated. That last replicated 2Ψ is the floor Stage 2 cannot cross, and it is precisely the wall Stage 3 knocks down.

Reduce-scatter, not all-reduce

The enabling move is swapping the collective. In DDP every micro-batch computes a local gradient and an all-reduce sums them and broadcasts the result, so all N ranks come away with the identical, full averaged gradient. That is exactly what forces each rank to own 2Ψ bytes.

A reduce-scatter does the summation but not the broadcast. Conceptually, cut the gradient into N equal chunks; reduce-scatter sums chunk i across all ranks and delivers the result only to rank i. Every rank contributes to reducing all chunks but walks away with just one fully-reduced shard — the shard whose optimizer states it also owns. Rank i then does its fp32 Adam update on that slice locally, with no other rank’s gradients needed. In practice this is overlapped with the backward pass: as each layer’s gradients finish, they are reduce-scattered immediately, and the off-shard portions can be freed instead of accumulated, which is what actually realizes the 2Ψ/N memory rather than merely renaming a full buffer.

Why the communication volume is unchanged

The counter-intuitive result is that Stage 2 sends no more bytes than DDP. The key identity: an all-reduce is a reduce-scatter followed by an all-gather. A bandwidth-optimal ring all-reduce of a Ψ-element buffer moves about 2Ψ of data per GPU — Ψ in the reduce-scatter half and Ψ in the all-gather half.

DDP:     all-reduce(grads)        ≈ 2Ψ moved/GPU
Stage 2: reduce-scatter(grads)   ≈  Ψ   (get my gradient shard)
       + all-gather(weights)     ≈  Ψ   (rebuild full fp16 weights)
                                 ≈ 2Ψ moved/GPU

Stage 2 uses only the first half of an all-reduce to collect gradients, then — because each rank updated a different slice of the weights — needs a matching all-gather after the optimizer step to reassemble the full fp16 weight tensor for the next forward pass. Reduce-scatter plus all-gather equals one all-reduce’s worth of traffic: 1x the DDP volume. Stage 2 saves gradient memory for free in communication terms.

A worked example: 7.5B parameters, 64 GPUs

Take Ψ = 7.5 billion parameters and N = 64 data-parallel GPUs, counting 1 GB = 10^9 bytes for round numbers.

Baseline DDP : 16Ψ            = 120 GB per GPU   (will not fit 32 GB)
Stage 1      : 4Ψ + 12Ψ/64  = 30 + 1.41   = 31.4 GB
Stage 2      : 2Ψ + 14Ψ/64  = 15 + 1.64   = 16.6 GB
Stage 3      : 16Ψ/64          = 1.9 GB

Read the columns. DDP is hopeless — 120 GB on a 32 GB card. Stage 1 already fits, dominated by the 4Ψ = 30 GB of replicated fp16 weights and the still-full fp16 gradients. Stage 2 halves that replicated block to 2Ψ = 15 GB by sharding the gradients, landing at 16.6 GB — comfortable headroom for activations. The whole difference between the Stage 1 and Stage 2 columns is one term: 2Ψ of gradient shrinking to 2Ψ/64, worth about 14.8 GB per GPU here, bought without moving a single extra byte over the network.