I wanted to really understand the ZeRO paper (Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models, arXiv:1910.02054) — not just quote its tables. So I built the thing it describes: a cluster of 32 virtual GPUs (one Python thread each, each with its own real tensors and a byte-accurate memory tracker), a simulated collective interconnect that counts every byte moved, and from-scratch implementations of ZeRO-0 (baseline data parallelism), ZeRO-1, ZeRO-2 and ZeRO-3 running a real (tiny) transformer through real Adam steps. Then I measured what each stage does to memory, communication and computation.
The one-paragraph version of what I found: all four stages are the same algorithm — every stage in my runs produced identical losses and identical trained weights, because ZeRO changes where bytes live, not the math. Per-GPU model-state memory collapses exactly as the paper's formulas say (my measured/paper ratio was 1.0000 at every stage): 16Ψ → 4.4Ψ → 2.4Ψ → 0.5Ψ at N=32. Compute per GPU is unchanged (same forward/backward FLOPs in every stage). The only thing that grows is communication: 2Ψ elements per GPU per step for the baseline, still 2Ψ for ZeRO-2, and 3Ψ (1.5×) for ZeRO-3 — the price of partitioning the parameters themselves.
Per-GPU model states at N = 32, Ψ = 637,824 params (my toy GPT), mixed-precision layout:
| Stage | What is partitioned | Measured /GPU | Paper formula | Ratio | Comm /step /GPU | Collectives /step |
|---|---|---|---|---|---|---|
| ZeRO-0 (DP) | nothing | 10.21 MB | 16Ψ = 10.21 MB | 1.000 | 1.94Ψ | 10 |
| ZeRO-1 | + optimizer states | 2.79 MB | 4Ψ+12Ψ/N = 2.79 MB | 1.000 | 2.91Ψ † | 20 |
| ZeRO-2 | + gradients | 1.56 MB | 2Ψ+14Ψ/N = 1.56 MB | 1.000 | 1.94Ψ | 20 |
| ZeRO-3 | + parameters | 0.32 MB | 16Ψ/N = 0.32 MB | 1.000 | 2.89Ψ | 29 |
† my ZeRO-1 all-reduces gradients and re-syncs updated parameters (3·(N−1)/N·Ψ); the paper counts 2Ψ by using reduce-scatter + treating the parameter re-sync as the all-gather half — see the nuance I found.
And the correctness claim that matters most: ZeRO is exact data parallelism, not an approximation. Same seed, same data, all four stages:
| stage | loss @ step 3 | max Δloss vs ZeRO-0 | max Δmaster-weight vs ZeRO-0 |
|---|---|---|---|
| ZeRO-0 | 4.181497 | — | — |
| ZeRO-1 | 4.181497 | 0.0 | 0.00000 |
| ZeRO-2 | 4.181497 | 7e-8 | 0.00002 |
| ZeRO-3 | 4.181497 | 7e-8 | 0.00002 |
(The 7e-8 / 2e-5 diffs for stages 2/3 come from a different operator path, not from any
approximation — my ZeRO-2/3 compute gradients by hand inside torch.autograd.Function with
fp16 casts, so a couple of fp16 roundings land differently. Relative to weights of magnitude
~0.02, a 2e-5 absolute diff is fp16-rounding noise.)
The hatched block is activation memory — identical in all four stages. ZeRO-1/2/3 partition model states and nothing else; by ZeRO-3 the activations are ~10× the model states, which is exactly why the paper pairs ZeRO-1/2/3 with ZeRO-R (activation checkpointing/partitioning).
Mixed-precision training with Adam keeps, per parameter, per GPU:
| Buffer | dtype | Why it exists | Bytes |
|---|---|---|---|
| parameter | fp16 | what forward/backward actually uses | 2 |
| gradient | fp16 | what backward produces | 2 |
| master weights | fp32 | fp16 updates would vanish in rounding; the update happens here | 4 |
Adam m (1st moment) |
fp32 | momentum | 4 |
Adam v (2nd moment) |
fp32 | per-parameter adaptive step size | 4 |
| total | 16 |
The insight that makes ZeRO obvious in retrospect: 12 of those 16 bytes are optimizer bookkeeping,
and data parallelism never needed them replicated. In DP, each GPU averages gradients and steps
the whole optimizer — but GPU r could perfectly well own just 1/N of the optimizer states and
update only that slice. No GPU ever needs all of m/v/master. ZeRO is the observation that
params, grads and optimizer states can each be peeled off the "replicate everything" model one at
a time, as long as you add the communication to reassemble what a given step actually touches:
ZeRO-0 params ████ grads ████ optim ████████████ 16Ψ
ZeRO-1 params ████ grads ████ optim ▌ 4Ψ + 12Ψ/N (shard the fat 12)
ZeRO-2 params ████ grads ▌ optim ▌ 2Ψ + 14Ψ/N (never hold full grads)
ZeRO-3 params ▌ grads ▌ optim ▌ 16Ψ/N (params materialize on demand)
Every parameter tensor (weight + bias flattened together) is split into N near-equal contiguous
chunks; GPU r owns chunk r. This is what DeepSpeed's flat parameter partitioning does. In my
simulation, each GPU's fp32 master/m/v are the owned chunks as real tensors, and ZeRO-3's fp16
parameters are likewise only the owned chunks — a full parameter literally does not exist on the
device until an all-gather materializes it.
| during backward | gradient sync | optimizer step | parameter re-sync | |
|---|---|---|---|---|
| ZeRO-0 | full local grads (2Ψ) | all-reduce mean → full grads (2Ψ) | full Adam over Ψ params | none — everyone stepped everything |
| ZeRO-1 | full local grads (2Ψ) | all-reduce mean → full grads (2Ψ) | Adam on my 1/N slice, reading my slice of the mean grads | all-gather updated chunks (Ψ) |
| ZeRO-2 | a full grad never exists: each linear's grad is reduce-scattered inside backward | — (already done) | Adam on my shard | all-gather updated chunks (Ψ) |
| ZeRO-3 | the parameters themselves exist only as 1/N chunks: all-gathered before each use, released right after; grads reduce-scattered inside backward | — | Adam on my chunk | none — params stay sharded |
The engineering heart of ZeRO-2/3 is that full buffers are never materialized. In my code this
is done with custom torch.autograd.Functions:
- ZeRO-2 (
_PartitionedLinear): forward is a normal linear, butbackwardcomputes the gradient by hand, flattens it, reduce-scatters it, keeps only the owned 1/N shard — and returnsNonefor the parameter grads so autograd itself never allocates a full.grad. - ZeRO-3 (
_GatherLinear,_GatherEmbedding):forwardall-gathers the full parameter, uses it, and releases it before returning;backwardre-gathers it (the paper: parameters are needed twice per step per layer), computes gradients, and reduce-scatters them. The gathered copy shows up in my memory tracker as a transientgather.fp16buffer — ZeRO-3's peak memory is one layer's parameter, not the whole model's.
One subtlety I had to mirror from the real thing: since the full gradient is never stored in ZeRO-2/3, the reduce-scatter must happen during backward, per parameter. All 32 of my device threads execute the identical collective sequence, so the rendezvous aligns — real implementations give this the same shape with grad-hook buckets.
The paper's communication table counts ZeRO-1 as 2Ψ — same as baseline DP — by decomposing it as reduce-scatter(Ψ) + all-gather(Ψ). But which all-gather? If you reduce-scatter gradients, each rank ends up with only its shard of the mean gradient, and the fp32-master update plus the fp16 parameter re-sync needs an all-gather of updated params (Ψ) — that's the 2Ψ. My implementation instead all-reduces gradients (2Ψ) and all-gathers updated parameters (Ψ) = 3Ψ, because ZeRO-1's defining memory trait is that full gradients stay resident (4Ψ + 12Ψ/N total, matching the paper's memory table exactly). Both readings are "ZeRO-1"; the difference is whether the parameter re-sync is counted inside the 2Ψ or as an extra Ψ that real systems (DeepSpeed) pipeline under the next backward. Building both paths on a fake NIC is what made this legible to me — on paper the two tables look inconsistent until you realize they're counting different decompositions.
This was the clearest thing the simulation showed me. Per GPU per step:
- Forward ≈ 2Ψ·tokens FLOPs, backward ≈ 4Ψ·tokens — identical in all four stages. ZeRO partitions state, never the math. A ZeRO-3 GPU does the same matmuls as a ZeRO-0 GPU; it just fetches parameters over the network first. (That's why ZeRO is orthogonal to tensor/pipeline parallelism, which partition the math.)
- Optimizer step: elementwise Adam — Ψ elements at ZeRO-0, Ψ/N at ZeRO-1/2/3. It's the only compute that shrinks, and it's noise next to forward/backward.
- Communication volume is the thing that actually changes: 2Ψ → ~2Ψ → 2Ψ → 3Ψ (fp16 elements per GPU per step, ring-model asymptotes). ZeRO-3's extra comes from re-gathering every parameter once before forward use and once before backward use.
- Wall-clock in my 32-thread runs: 957 → 1055 → 995 → 1106 ms/step, with 670–789 ms of that inside collectives. On shared CPU threads this measures synchronization structure, not cluster throughput — the honest quantitative version is volume ÷ bandwidth (see the 7.5B table below). Note the collective count: 10 → 20 → 20 → 29 per step. Higher stages pay in sync points even when volume is equal.
For the toy model the absolute comm is a few MB; the cost becomes vivid at real scale:
| Stage | Model states /GPU | fits 32 GB? | Comm /step /GPU | @50 GB/s |
|---|---|---|---|---|
| ZeRO-0 | 120 GB | NO | 30 GB | 0.6 s |
| ZeRO-1 | 32.8 GB | NO | 30 GB | 0.6 s |
| ZeRO-2 | 18.3 GB | yes | 30 GB | 0.6 s |
| ZeRO-3 | 3.75 GB | yes, 32× less | 45 GB | 0.9 s |
ZeRO-3's +15 GB/step is the 1.5× communication tax; everything else on that row is what you buy with it. And activations are not in this table — for GPT-2-scale models the paper reports ~60 GB of activations at batch 32×1024 (cut to ~8 GB by activation checkpointing at ~33% forward-recompute overhead — a compute-for-memory trade on top of ZeRO).
- "ZeRO-1 is free memory" — almost. Partitioning optimizer states adds zero volume over baseline DP (both are ~2Ψ) — but my implementation still needed an extra parameter-resync all-gather, so how you count the re-sync changes the story. Memory tables and comm tables in papers are decompositions, not ground truth.
- ZeRO-2's memory win is really a when win. ZeRO-1 and ZeRO-2 have the same communication; the difference is ZeRO-1 reduces after backward (so full grads exist) and ZeRO-2 reduces during backward (so they never do). Same bytes on the wire, 2Ψ less resident.
- Compute does not shrink — at all. I half-expected ZeRO-3 GPUs to do 1/32nd of the work. No: every GPU runs the full model on its micro-batch. The 1/32 shows up only in memory and in the optimizer step. If you want compute partitioned, that's tensor/pipeline parallelism.
- Equivalence is testable, so test it. The single most convincing output of this whole exercise: four different memory layouts, bit-for-bit-same training. If your ZeRO changes the loss curve, you have a bug, not a tradeoff.
- By ZeRO-3, activations are the enemy. Model states at N=32 were 0.32 MB; activations were 3.12 MB for a tiny micro-batch. This is why the paper has ZeRO-R and why real runs stack ZeRO-3 + activation checkpointing + offload.
- Threads are not GPUs. Wall-clock shows sync structure only; the quantitative perf story is the volume ÷ bandwidth analysis. 32 threads share one GIL and one socket.
- Collectives are synchronous (rendezvous per op). Real systems bucket, pipeline and overlap reductions with compute — my timings are pessimistic for stages 2/3.
- Per-parameter (not bucketed) reductions, affine-free LayerNorm (so 100% of params are in sharding units — real ZeRO-3 shards LN affines too), no CPU/NVMe offload (ZeRO-Infinity), no activation checkpointing implemented.
- The toy model is fp16-storage/fp32-math on CPU: the per-op fp32 cast copies a real GPU wouldn't make are excluded from accounting (documented in the notebook).
pip install -r requirements.txt
jupyter lab zero_stages_simulation.ipynb # run all cells top-to-bottom (~2 min, CPU only)Or open the Colab badge above. zero_sim.py contains the same simulation code as an importable
module (the notebook was generated from it and is fully self-contained).
zero_stages_simulation.ipynb # the deliverable: concepts + 32-GPU simulation + 7 experiments
zero_sim.py # same code as an importable module
assets/ # figures used in this README
requirements.txt
Rajbhandari, S., Rasley, J., Ruwase, O., He, Y. ZeRO: Memory Optimizations Toward Training Trillion Parameter Models. arXiv:1910.02054 (2019). Memory/comm formulas cross-checked against the paper (Tables 1–3 and the communication analysis section).


