What ZeRO-Infinity adds on top of stage 3
The companion pieces each own one rung. Stages 1–3 partition optimizer state, gradients and parameters across dp ranks, cutting the 16 bytes-per-parameter budget to 16Ψ/dp. CPU offload moves the fp32 optimizer bucket to host DRAM; NVMe offload pushes it onto flash. Each is a single move.
ZeRO-Infinity is the composition, and composing changes the objective. Once every tier is available simultaneously, the question is no longer ‘does this bucket fit?’ but ‘what is the cheapest placement of every bucket such that no tier stalls the GPU?’ Capacity becomes an aggregate — C_total = Σ_nodes (HBM + DRAM + NVMe) — and the binding constraint moves from bytes to bandwidth. Three mechanisms exist purely to make that constraint satisfiable, and none appear in a single-tier design: bandwidth-centric partitioning, a deep prefetch pipeline, and memory-centric tiling.
The FLOPs-per-byte test every tier must pass
Overlap is usually written t_transfer ≤ t_compute. Divide through and it becomes scale-free and far more useful. Let a GPU sustain R FLOP/s and a tier deliver BW bytes/s to it; that tier is invisible exactly when the work’s arithmetic intensity clears the machine’s balance point:
ait = FLOPs done / bytes moved
tier is hidden <=> ait >= R / BW (the machine balance)With activation checkpointing a training step costs forward (2P) + recomputed forward (2P) + backward (4P) = 8P FLOPs per token, so T tokens cost 8·P·T. The fp16 parameters must climb the ladder twice per step (once for forward, once for the recompute-plus-backward), moving 4P bytes. So ait_params = 8PT / 4P = 2T. Parameter streaming is hidden when 2T ≥ R/BW — a condition on batch size, not on the model. This single inequality decides every placement below.
Applying the test: which state can live where
Take R = 125 TFLOP/s sustained per GPU and read off what each tier demands:
| Tier | BW to one GPU | Required ait | Tokens needed per link |
|---|---|---|---|
| HBM | 2 TB/s | 63 FLOPs/byte | ~32 |
| DRAM over PCIe 4.0 x16 | 32 GB/s | 3,900 | ~2,000 |
| NVMe (one drive’s share) | 3 GB/s | 41,700 | ~20,900 |
Read that last row carefully: it is a per-link bar, the cost when one rank must pull everything it needs off its own drive. On those terms batch 16 × seq 2048 = 32,768 clears it and a lone 2048-token sequence misses by 10×. That is the problem statement, not the verdict — the next section divides the bar by dp. The Adam update, by contrast, fails the test at every tier and every batch size: it is elementwise, roughly 12 FLOPs against 24 bytes read and written, so ait ≈ 0.5. It can never be hidden behind GPU compute, which is why it does not run on the GPU at all — it runs on the CPU, beside the DRAM and flash it touches, where its absolute time is small next to a long step.
Bandwidth-centric partitioning: allgather, not broadcast
Classic stage 3 assigns each parameter tensor an owner. To materialize a layer, the owner reads the whole tensor from its own slow tier and broadcasts it. The aggregate bytes read per rank are the same either way — but the concurrency is not. Under owner-broadcast, the layer you need right now is gated by one rank’s link:
owner-broadcast: t_fetch = layer_bytes / bw_one
bandwidth-centric: t_fetch = layer_bytes / (dp * bw_one)ZeRO-Infinity instead stripes every tensor across all dp ranks and materializes it with an allgather, so all dp NVMe links and PCIe lanes read their slice at the same instant. For a 1T-parameter model with 128 layers, a layer is 2e12/128 = 15.7 GB of fp16 weights; at 3 GB/s one owner needs 5.2 s, while 64 ranks striped need 0.08 s. The flash tier’s required arithmetic intensity falls by a factor of dp — roughly 21,000 tokens per link becomes a few hundred — which is what makes NVMe-resident parameters practical at batch sizes a single-link design could never hide.
Prefetch depth: how far ahead the pipeline must run
Striping buys bandwidth, not timeliness: a fetch issued when compute reaches the layer is a stall however fast it runs. Let t_layer be one layer’s compute time and t_fetch the time to materialize the next. The lookahead depth d must satisfy d ≥ t_fetch / t_layer, and each unit of depth costs a resident buffer:
d >= t_fetch / t_layer buffer_bytes = d * layer_bytesWork it for the 1T model at a modest 4,096 tokens per GPU: step compute is 8 × 1e12 × 4096 / 1.25e14 ≈ 262 s, so t_layer ≈ 2.05 s. Under owner-broadcast, d ≥ 5.2/2.05 = 2.5 → d = 3, i.e. 47 GB of prefetch buffers — more than an 80 GB card can spare once activations are resident. Bandwidth-centric partitioning drops t_fetch to 0.08 s, so d = 1 suffices and the buffer is one layer. The two mechanisms are the same mechanism: striping is what makes shallow, affordable prefetch possible.
Memory-centric tiling: operators bigger than the device
Everything so far assumes a materialized layer fits. At extreme width it does not: for the 1T model, d_model = 25600 and the FFN’s first matrix alone is 25600 × 102400 = 2.6e9 parameters, 5.2 GB in fp16 — and there are two such matrices plus attention projections per layer. The historical fix was tensor parallelism, which drags in extra collectives and a fixed device group.
Memory-centric tiling avoids that. Split the operator along its output dimension into k tiles and execute them sequentially, gathering and releasing each tile in turn. Working memory becomes layer_bytes / k (doubled if you double-buffer), so with k = 8 the 5.2 GB matrix needs 655 MB resident, or 1.3 GB with a tile in flight. Because a column-split matmul produces a slice of the output, no cross-tile communication is needed. The cost is matmul shape: tile too finely and the kernels get small and utilization drops.