Skip to content

About

ZeRO (0/1/2/3) data-parallel training simulated from scratch on 32 virtual GPUs - memory, communication and computation measured

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

1 Commit

Folders and files

Repository files navigation

ZeRO stages 0/1/2/3 — simulated from scratch on 32 virtual GPUs

Open In Colab

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.


Results in one table

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.)

Memory per GPU by stage

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).


Part 1 — The mental model that finally made this click

The 16 bytes hiding behind every parameter

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)

What "partitioning" means concretely

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.


Part 2 — What each stage actually does per step

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, but backward computes the gradient by hand, flattens it, reduce-scatters it, keeps only the owned 1/N shard — and returns None for the parameter grads so autograd itself never allocates a full .grad.
  • ZeRO-3 (_GatherLinear, _GatherEmbedding): forward all-gathers the full parameter, uses it, and releases it before returning; backward re-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 transient gather.fp16 buffer — 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.

A nuance I found: implementation vs paper accounting

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.

What changes in computation (spoiler: almost nothing)

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:

Scaling it up: Ψ = 7.5B on 32 GPUs (the paper's own example)

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).

Memory vs world size

7.5B projection


Things that surprised me / misconceptions I had to give up

  1. "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.
  2. 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.
  3. 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.
  4. 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.
  5. 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.

Honest limitations of my simulation

  • 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).

Run it

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).

Repo layout

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

Reference

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).

About

ZeRO (0/1/2/3) data-parallel training simulated from scratch on 32 virtual GPUs - memory, communication and computation measured

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages