ZeRO, annotated
How to read this page
- Any dotted word explains itself on hover, focus or tap, and so does every symbol in every equation.
- The redrawn Figure 1 in §1 is the whole paper in one picture. The memory chart in §5.4 and the step-by-step schedule in §7 are live: change the model size, press Step, and watch the numbers.
Every idea climbs the same ladder: everyday picture, tiny example, diagram, the math, why it matters today. Two running examples carry the arithmetic:
The tiny model (illustrative, small enough to check by hand): Ψ = 8 parameters trained on 4 GPUs. The paper's model (its Figure 1): Ψ = 7.5 billion parameters on 64 GPUs. Both are trained with Adam in mixed precision. The pretraining lesson builds the memory bill and the three ZeRO stages in NumPy; its section 3 puts them next to tensor and pipeline parallelism, whose founding paper has its own Megatron-LM companion.
Abstract · original
“ZeRO eliminates memory redundancies in data- and model-parallel training while retaining low communication volume and high computational granularity, allowing us to scale the model size proportional to the number of devices with sustained high efficiency.”Rajbhandari et al. (2019), Abstract
Everyday picture
Sixty-four people are writing one encyclopedia together. Each keeps a complete copy of every volume, every draft and every margin note at their own desk, so the desks overflow long before the encyclopedia is finished. ZeRO says: keep one sixty-fourth of everything each, and borrow a volume from its keeper for the few minutes you need it.
What the paper claims
- A name for the waste: in data parallelism, every GPU stores the same weights, gradients and optimizer state.
- Three stages that remove it, sharding the optimizer state, then the gradients, then the weights, for 4×, 8× and Nd× less memory per GPU. The first two cost no extra communication; the third costs 1.5×.
- A second set of fixes (ZeRO-R) for everything else in memory: activations, temporary buffers and fragmentation.
- An implementation of the first two stages, ZeRO-100B, that trains models of over 100 billion parameters on 400 GPUs at 15 petaflops: 8× bigger and up to 10× faster than the best system of the day, and up to 13 billion parameters with no model changes at all.
Why it matters today
Sharding the training state became a standard ingredient of large training runs: the third stage is what PyTorch's FSDP does, and the paper's 16-bytes-per-parameter accounting is the first sum to do before any run.
1 Extended introduction · original
“Basic data parallelism (DP) does not reduce memory per device, and runs out of memory for models with more than 1.4B parameters on current generation of GPUs with 32 GB memory.”Rajbhandari et al. (2019), §1
Everyday picture
In 2019 there were two ways to train on many GPUs. Data parallelism is a set of identical kitchens, each cooking whole meals from its own share of the orders: fast and simple, but each kitchen needs the whole pantry. Model parallelism splits one recipe across cooks at one bench (the Megatron-LM approach): the pantry is shared, but the cooks must talk constantly, which only works when they stand side by side. The paper wants data parallelism's speed with model parallelism's small pantry.
Tiny example
Train the tiny model (Ψ = 8 parameters) on 4 GPUs. Mixed-precision Adam keeps 16 bytes per parameter (§3.1 explains them), so plain data parallelism puts 16 × 8 = 128 bytes on every GPU, 512 bytes in all, of which only 128 are distinct. Three quarters of the cluster's memory holds copies.
Hover or tap a stage's title or a coloured bar. Dashed outlines are memory a GPU no longer holds.
Reading it: each block of four bars is one way of training, and each bar is one GPU's memory, drawn to scale in bytes per parameter: purple for the 16-bit weights (2), orange for their gradients (2), blue for the optimizer state (12). In the top block every GPU holds all 16. Going down, the blue, then the orange, then the purple part is partitioned: each GPU keeps only its own quarter (the filled piece, a different quarter on each GPU) and the dashed outline marks what it dropped. The bars shrink from 16 to 7, 5.5 and finally 4 units: 128, 56, 44 and 32 bytes for the tiny model. Nothing is lost, because the four filled quarters in each column still add up to one complete copy.
What else the introduction sets out
- Model states (weights, gradients, optimizer state) are fixed by ZeRO-DP; residual states (activations, temporary buffers, fragments) by ZeRO-R.
- ZeRO combines with model parallelism: Nd-way ZeRO on top of Nm-way model parallelism divides the model states by up to Nd × Nm.
- Megatron-LM, the model-parallel state of the art, managed about 5 teraflops per GPU (under 5% of peak) on a 40-billion-parameter model split across two servers: talking inside every layer over the network between machines was the bottleneck.
Why it matters
The figure's promise is that memory per GPU falls as GPUs are added. With plain data parallelism, adding GPUs adds speed but never room; with ZeRO, it adds both. The pretraining lesson draws the same four stages for a 7B model.
2 Related work · original
Everyday picture
There were already four ways to squeeze a big model into small GPUs, and each had a catch. Splitting the model by layers (pipeline parallelism, §2.1) needs many micro-batches in flight to keep every stage busy. Throwing activations away and recomputing them (activation checkpointing, §2.2.1) saves memory at the price of extra arithmetic. Parking the model states in the CPU's memory (§2.2.2) makes every step wait on the slow link between CPU and GPU: up to half of training time in one earlier system. And optimizers that keep coarser statistics (§2.2.3) save memory but may change how the model learns.
Tiny example
GPipe, the pipeline system of the day, needs a batch that grows with the number of pipeline stages to hide its idle bubble. With 4 stages and 4 micro-batches, each stage idles for (4 − 1)/(4 + 4 − 1) = 3/7 of the time; with 32 micro-batches the idle share falls to 3/35, about 9%, but now 32 micro-batches' activations are in flight.
In Python:
# pipeline bubble: (p − 1) / (m + p − 1), for p = 4 stages
p = 4
round((p - 1) / (4 + p - 1), 3) # → 0.429
round((p - 1) / (32 + p - 1), 3) # → 0.086
Why it matters
ZeRO's claim is that it needs none of these trade-offs: the model, the optimizer and the batch size stay exactly as they were, and the result is the same model standard training would produce. Each alternative still has its place, and today's largest runs combine ZeRO-style sharding with pipeline and tensor parallelism; the pretraining lesson builds pipeline_schedule and bubble_fraction for the bubble above.
3 Where did all the memory go? · original
“For example, a 1.5B parameter GPT-2 model requires 3GB of memory for its weights (or parameters) in 16-bit precision, yet, it cannot be trained on a single GPU with 32GB memory using Tensorflow or PyTorch.”Rajbhandari et al. (2019), §3
Everyday picture
A recipe card fits in your pocket, but learning to cook from it fills a kitchen: the card, your notes on what went wrong, your running averages of how much to adjust, and every half-finished dish you keep to compare against. This section is the paper's audit of the kitchen.
3.1 Model states: optimizer states, gradients and parameters · original
Everyday picture
Adam keeps two running notes per weight: a running average of its gradients (momentum) and a running average of their squares (the variance). Mixed precision adds a third: the forward and backward passes use cheap 16-bit weights, but updates are applied to a full 32-bit master copy, because a tiny update added to a 16-bit number can round away to nothing.
Tiny example
For one parameter: 2 bytes for the fp16 weight, 2 for its fp16 gradient, then 4 + 4 + 4 for the fp32 master weight, momentum and variance. The paper lumps the last three together as the optimizer state, K = 12 bytes. Total: 2 + 2 + 12 = 16 bytes, eight times the 2 bytes the weight itself needs.
In words: “two bytes per parameter for the working weights, two for their gradients, and K = 12 for the optimizer's fp32 copies: sixteen bytes per parameter before a single activation is stored.”
With the numbers: the tiny model: 16 × 8 = 128 bytes. GPT-2 with 1.5 billion parameters: 16 × 1.5 × 109 = 24 GB, against the 3 GB its fp16 weights take. The paper's 7.5B model: 120 GB, nearly four 32 GB GPUs' worth before any activations.
In Python:
# bytes per parameter: fp16 weight, fp16 gradient, K for the optimizer state
K = 4 + 4 + 4
2 + 2 + K # → 16
# the tiny model: Ψ = 8
16 * 8 # → 128
# GPT-2, Ψ = 1.5 billion, in GB
16 * 1.5e9 / 1e9 # → 24.0
# its fp16 weights alone
2 * 1.5e9 / 1e9 # → 3.0
# the paper's Figure 1 model, Ψ = 7.5 billion
16 * 7.5e9 / 1e9 # → 120.0
Why it matters today
Sixteen bytes per parameter is still the rule of thumb for Adam in mixed precision (bf16 in place of fp16 changes nothing in the count). It is why a model you can run on one GPU usually cannot be trained on one. The pretraining lesson itemizes the same bill in training_memory, and the Adam companion explains the two running averages.
3.2 Residual memory consumption · original
Everyday picture
Beyond the model states, three things take memory. Activations: the forward pass's intermediate results, kept so the backward pass can compute gradients. Temporary buffers: scratch space, such as one long flat array that all the gradients are copied into so they can be summed across GPUs in one go. And fragments: free memory broken into pieces too small to use, like a car park with plenty of empty spaces but none two spaces wide for a van.
Tiny example
The paper's footnote counts a GPT-2-like model's activations as about 12 × hidden width × batch × sequence length × layers numbers. Take a toy of width 2, batch 1, sequence 3 and 1 layer: 12 × 2 × 1 × 3 × 1 = 72 numbers, or 144 bytes in fp16. Double the batch and it doubles: activations grow with the work in flight, not with the parameter count.
In words: “each layer keeps about twelve numbers for every hidden unit of every token in the batch.”
With the numbers: GPT-2 1.5B has h = 1,600 and L = 48 (the paper's Table 4). At b = 32 and s = 1,024: 12 × 1,600 × 32 × 1,024 × 48 = 30.2 billion numbers; at 2 bytes each that is 60.4 GB, the paper's “about 60 GB”. Activation checkpointing cuts it to about 8 GB, the paper says. The fp32 buffer that fuses all 1.5 billion gradients adds 1.5 × 109 × 4 bytes = 6 GB.
In Python:
# the toy: h = 2, b = 1, s = 3, L = 1
12 * 2 * 1 * 3 * 1 # → 72
# GPT-2 1.5B at batch 32, sequence 1,024
h, b, s, L = 1600, 32, 1024, 48
A = 12 * h * b * s * L
round(A / 1e9, 1) # → 30.2
# in GB at 2 bytes per fp16 number
round(A * 2 / 1e9, 1) # → 60.4
# one flat fp32 buffer holding every gradient, in GB
1.5e9 * 4 / 1e9 # → 6.0
Reading it: each bar is one item in the bill for training the 1.5-billion-parameter GPT-2, on a common scale, with the striped bar marking one 32 GB V100 GPU. The weights you would need to run the model are the shortest bar. The model states alone (24 GB) nearly fill the GPU; activations at batch 32 are more than twice the GPU; checkpointing brings them down to 8 GB, and the fused buffer adds 6 GB more. Model states plus checkpointed activations plus the buffer already exceed the 32 GB line, which is the paper's point: the weights are the small part.
Fragmentation
The paper reports running out of memory with over 30% still free in extreme cases: the free memory was there, but not in one contiguous piece. The PagedAttention companion meets the same problem, external fragmentation, in serving.
Why it matters today
Every large run still budgets these four items. Activations are the one that grows with batch and context length, which is why the pretraining lesson counts them per layer in activation_bytes, and why kernels like FlashAttention that never store the attention scores matter for training as well as speed.
4 ZeRO: insights and overview · original
“Both DP and MP keep all the model states needed over the entire training process, but not everything is required all the time.”Rajbhandari et al. (2019), §4.1
Everyday picture
A library does not buy every reader a copy of every book. It keeps one copy, and a book leaves the shelf only while someone is reading it. Training works the same way: layer 30's weights are needed only while layer 30 is computing, and the optimizer state for a weight only at the moment that weight is updated.
Tiny example: the three insights
- (a) Data parallelism scales better. Splitting a layer across GPUs shrinks each GPU's matrix multiplies and forces a conversation inside every layer. With the tiny model on 4 GPUs, model parallelism leaves each GPU 2 of the 8 weights to multiply by, then a round of talking; data parallelism gives each GPU all 8 weights and a quarter of the batch, and talks once per step.
- (b) Data parallelism wastes memory. 512 bytes of state across the cluster to hold 128 bytes of information.
- (c) Most state is idle most of the time. In the tiny model, if the 8 weights form 4 layers of 2, then at any moment of the forward pass only 2 of the 8 weights are in use.
The plan
ZeRO-DP keeps data parallelism's split of the batch (insight a), stores each model state once across the GPUs (b), and fetches each piece only when it is needed (c). ZeRO-R handles the residual states: partition the activations that model parallelism duplicates (§6.1), cap the temporary buffers at a constant size (§6.2), and pre-allocate contiguous space so memory does not fragment (§6.3).
Why it matters
Insight (c) is the one that aged best: “fetch just in time, then drop” is how stage 3 and FSDP handle weights, and it is the same trade activation checkpointing makes with activations: keep less, and rebuild what you need when you need it.
5 Deep dive into ZeRO-DP · original
Everyday picture
Three stages, each removing one kind of copy, cumulatively: first the optimizer state (Pos), then the gradients (Pos+g), then the parameters (Pos+g+p). Today they are usually called ZeRO stages 1, 2 and 3.
5.1 Pos: optimizer state partitioning · original
Everyday picture
Four accountants share one ledger of 8 accounts. Instead of all four updating all 8 accounts identically, each takes 2 accounts, keeps only those accounts' notes, updates them, and at the end of the day reads the other six fresh balances out to the others.
Tiny example
GPU 0 keeps the optimizer state for parameters 1 and 2 only, GPU 1 for 3 and 4, and so on. After the gradients are averaged, each GPU updates only its own 2 parameters, then an all-gather hands every GPU all 8 updated weights. Memory per GPU: 4 × 8 = 32 bytes of fp16 weights and gradients, plus 12 × 8 / 4 = 24 bytes of optimizer state: 56 bytes, down from 128.
In words: “every GPU still holds all the 16-bit weights and gradients, but only its own 1/Nd share of the optimizer state.”
With the numbers: tiny model: 32 + 24 = 56 bytes. The paper's model: 4 × 7.5 + 12 × 7.5 / 64 = 30 + 1.41 = 31.4 GB, against 120 GB. As Nd grows, the second term vanishes and memory approaches 4Ψ: the paper's “4× reduction”.
In Python:
K = 12
# M_os = 4Ψ + KΨ / N_d, tiny model on 4 GPUs
Psi, N_d = 8, 4
4 * Psi + K * Psi / N_d # → 56.0
# the paper's model on 64 GPUs, in GB
Psi, N_d = 7.5e9, 64
round((4 * Psi + K * Psi / N_d) / 1e9, 1) # → 31.4
# reduction against 16Ψ as N_d grows large
16 / 4 # → 4.0
Why it matters today
Stage 1 is nearly free: the gradient averaging and weight broadcast it needs cost the same as plain data parallelism (§7), and it removes three quarters of the memory. The pretraining lesson's zero_memory_per_gpu computes this and the next two formulas for any stage.
5.2 Pg: gradient partitioning · original
Everyday picture
If each accountant only updates their own 2 accounts, why receive the totals for all 8? Each needs only the totals for their own accounts. So instead of everyone learning every sum, each sum is delivered only to the accountant who owns it.
Tiny example
Four GPUs each computed a gradient for all 8 parameters. A reduce-scatter sums them so that GPU 0 ends up with the total for parameters 1 and 2 only, GPU 1 for 3 and 4, and so on; each GPU frees the rest. Gradient memory falls from 2 × 8 = 16 bytes to 16 / 4 = 4 bytes: per GPU, 16 + 28 = 44 bytes in all. A reduce-scatter is the first half of the ring all-reduce that plain data parallelism already runs.
In words: “every GPU keeps all the 16-bit weights, and a 1/Nd share of the gradients and the optimizer state.”
With the numbers: tiny model: 2 × 8 + 14 × 8 / 4 = 16 + 28 = 44 bytes. The paper's model: 15 + 1.64 = 16.6 GB. For large Nd memory approaches 2Ψ: the paper's “8× reduction”.
In Python:
K = 12
# M_os+g = 2Ψ + (2 + K)Ψ / N_d, tiny model on 4 GPUs
Psi, N_d = 8, 4
2 * Psi + (2 + K) * Psi / N_d # → 44.0
# the paper's model on 64 GPUs, in GB
Psi, N_d = 7.5e9, 64
round((2 * Psi + (2 + K) * Psi / N_d) / 1e9, 1) # → 16.6
# reduction against 16Ψ as N_d grows large
16 / 2 # → 8.0
To keep this fast, gradients are gathered into buckets, one per owner, and each bucket is reduced as soon as the backward pass has produced it, so communication overlaps with computation.
Why it matters today
Stage 2 is still communication-free relative to plain data parallelism, which is why it is a common default when the weights themselves fit. The lesson's ring_all_reduce runs the two halves, reduce-scatter and all-gather, separately so you can see each one.
5.3 Pp: parameter partitioning · original
“While this may seem to incur significant communication overhead at first glance, we show that this approach only increases the total communication volume of a baseline DP system to 1.5x, while enabling memory reduction proportional to Nd.”Rajbhandari et al. (2019), §5.3
Everyday picture
Now even the ledger itself is split: each accountant keeps only 2 accounts' balances. When the day's work reaches accounts 3 and 4, their keeper reads them out, everyone uses them, and everyone forgets them again.
Tiny example
Each GPU stores 2 of the 8 parameters, their 2 gradients and their optimizer state: 16 × 2 = 32 bytes, exactly a quarter of 128. The forward pass reaches layer 2 (parameters 3 and 4), so GPU 1 sends those two weights to the others; all four GPUs compute layer 2 on their own quarter of the batch and discard the weights. §7 animates the whole step.
In words: “every model state is stored exactly once across the cluster, so each GPU holds a 1/Nd share of all sixteen bytes per parameter.”
With the numbers: tiny model: 128 / 4 = 32 bytes. The paper's model: 120 / 64 = 1.9 GB (1.88 in the paper's Table 1). A trillion parameters on 1,024 GPUs: 16 × 1012 / 1,024 = 15.6 GB.
In Python:
K = 12
# M_os+g+p = (2 + 2 + K)Ψ / N_d, tiny model on 4 GPUs
Psi, N_d = 8, 4
(2 + 2 + K) * Psi / N_d # → 32.0
# the paper's model on 64 GPUs, in GB
round(16 * 7.5e9 / 64 / 1e9, 2) # → 1.88
# a trillion parameters on 1,024 GPUs, in GB
round(16 * 1e12 / 1024 / 1e9, 1) # → 15.6
Why it matters today
With stage 3 there is no model too big for data parallelism, only too few GPUs: memory per GPU falls in proportion to the number of GPUs. The paper describes it but did not implement it (its ZeRO-100B stops at stage 2). PyTorch's FSDP implements the same idea.
5.4 Implication on model size · original
Everyday picture
If each stage divides some of the bill by the number of GPUs, the question becomes: how many GPUs until the bill fits on one? The paper's Table 1 tabulates it for three model sizes. The chart below recomputes that table from the three formulas above.
Tiny example
A 32 GB GPU can hold 32 / 16 = 2 billion parameters' worth of model states under plain data parallelism, whatever the number of GPUs. With stage 3 on 64 GPUs it holds 32 × 64 / 16 = 128 billion. The paper's Table 1 marks exactly that: 7.5B fits at Nd = 64 with Pos, 128B with Pos+g+p.
Hover the chart, or focus it and use the arrow keys, to read the memory per GPU at each number of GPUs.
Reading it: the horizontal axis is the number of data-parallel GPUs, Nd, from 1 to 1,024, doubling at each step; the vertical axis is the model-state memory each GPU needs, on a log scale (each gridline is ten times the one below). The flat line at the top is plain data parallelism: 16Ψ no matter how many GPUs. Stages 1 and 2 (Pos and Pos+g) fall and then flatten at 4Ψ and 2Ψ, the parts they leave unsharded. Only stage 3 (Pos+g+p) keeps falling in a straight line. The purple dashed line is one 32 GB V100; wherever a curve drops beneath it, that configuration fits. Switch to 1T: only stage 3 ever crosses the line, at 1,024 GPUs (15.6 GB), which is the paper's “trillion parameters on 1,024 GPUs”. Hovering reproduces the paper's Table 1, for example 52.5, 41.3 and 30 GB for 7.5B on 4 GPUs.
Why it matters today
This chart is the planning tool for any large run: pick the model, pick the stage, read off the GPUs. Without ZeRO, the largest model data parallelism could train on those GPUs was under 1.5 billion parameters, the paper notes; with all three stages on 1,024 GPUs it is over a trillion.
6 Deep dive into ZeRO-R · original
Everyday picture
With the model states tidied away, the rest of the mess in the kitchen becomes the problem: half-finished dishes (activations), oversized mixing bowls (buffers) and counter space broken up by clutter (fragmentation). ZeRO-R has one fix for each.
6.1 Pa: partitioned activation checkpointing · original
Everyday picture
In model parallelism, the GPUs that share a layer all need the layer's whole input, so each keeps an identical copy of it. Sixteen identical copies of a saved checkpoint are sixteen times the memory for one copy's worth of information. Pa keeps one sixteenth each and reassembles the whole input with an all-gather only when the backward pass needs it.
Tiny example
A layer input of 16 numbers, shared by 4 model-parallel GPUs: normally each GPU saves all 16 as its checkpoint. With Pa, each saves 4, and before recomputing that layer the four GPUs all-gather their pieces back into 16.
In words: “keep one checkpoint per layer, the layer's input for every token in the batch, and split it evenly over the model-parallel GPUs.”
With the numbers: the paper's 100B model has L = 125 layers of width h = 8,192 (its Table 4 and appendix); at b = 32, s = 1,024 the checkpoints hold 125 × 32 × 1,024 × 8,192 = 33.6 billion numbers, the paper's “about 33 GB”. Split 16 ways: 2.1 billion, the paper's “about 2 GB”. (Stored as 2-byte fp16 values both figures would double; the 16× saving is the same either way.)
In Python:
# the toy: a 16-number checkpoint split over 4 GPUs
16 / 4 # → 4.0
# the 100B model: one checkpoint per layer
L, b, s, h = 125, 32, 1024, 8192
round(L * b * s * h / 1e9, 1) # → 33.6
# M_a with N_m = 16
round(L * b * s * h / 16 / 1e9, 1) # → 2.1
For very large models the partitioned checkpoints can go one step further, to the CPU's memory (Pa+cpu), leaving almost no activation memory on the GPU at the cost of moving them over the slow CPU link.
Why it matters today
Activation memory stayed the next wall once model states were sharded. Later work on Megatron-LM returned to it: Korthikanti et al. (2022), cited in the pretraining lesson, counts a tensor-parallel transformer's activations and cuts how many must be kept or recomputed.
6.2 CB: constant-size buffers · original
Everyday picture
Posting one big parcel is cheaper per item than posting many small letters, so libraries fuse all the gradients into one buffer before summing them across GPUs. But a buffer as big as the model grows with the model. ZeRO uses a fixed-size parcel instead: big enough to be efficient, and no bigger.
Tiny example
A 3-billion-parameter model in one fused fp32 buffer needs 3 × 109 × 4 bytes = 12 GB, the paper's example. A constant buffer holds the same gradients in several trips, each trip the same size whatever the model.
In Python:
# a fused fp32 buffer for 3 billion gradients, in GB
3e9 * 4 / 1e9 # → 12.0
Why it matters today
Buckets do double duty: a bucket can be sent as soon as its gradients exist, so communication overlaps with the rest of the backward pass, the trick the paper credits to NVIDIA's AMP in §5.2.
6.3 MD: memory defragmentation · original
Everyday picture
Fragmentation comes from mixing things that stay a long time with things that leave quickly. In the forward pass, checkpoints stay until the backward pass while the other activations are discarded; in the backward pass, weight gradients stay while activation gradients vanish. The gaps they leave are scattered. ZeRO copies the long-lived tensors into one pre-allocated contiguous block as they are produced, so the short-lived ones free up whole, reusable stretches.
Tiny example
Illustrative: 8 slots of memory, alternately filled by a checkpoint (kept) and a temporary (freed). After the temporaries go, 4 slots are free, but none are next to each other, so a request for 2 slots fails. Had the 4 checkpoints been placed in slots 1 to 4, slots 5 to 8 would be one free run of 4.
Why it matters today
Where a tensor lives, and for how long, is part of the design of a training system: the PyTorch FSDP paper describes being co-designed with PyTorch's caching memory allocator.
7 Communication analysis of ZeRO-DP · original
“In other words, we reschedule the parameter all-gather by spreading it across the entire forward propagation, and discarding the parameters once they have been used.”Rajbhandari et al. (2019), §7.2.2
Everyday picture
Saving memory is easy if you are willing to talk more. The surprise of this section is how little extra talking ZeRO needs: none for stages 1 and 2, half as much again for stage 3.
Tiny example: the baseline
Plain data parallelism averages gradients with an all-reduce, which the ring algorithm runs as a reduce-scatter followed by an all-gather. For the tiny model's 8 gradients on 4 GPUs, the reduce-scatter moves about 8 numbers through each GPU (6 exactly: three steps of two) and the all-gather another 8, so the paper counts 2Ψ per step (§7.1).
Stage 3, step by step
Reading it: the grid's rows are four GPUs and its columns the model's four shards of parameters, one per layer here; GPU i owns shard i (the purple diagonal) and nothing else. Press Step. In the forward pass the owner of each layer's shard sends it to everyone (the column lights up), all four compute that layer on their own quarter of the batch, and the borrowed copies are dropped before the next layer. The backward pass does the same in reverse order, and after each layer's gradients are computed they are summed onto the owner (the orange mark). The bars count what each GPU moves, in units of Ψ: one all-gather of every weight in the forward pass, one more in the backward pass, one reduce-scatter of the gradients. The total, 3Ψ, sits against plain data parallelism's 2Ψ: 1.5×. At no moment does any GPU hold more than its own shard plus one borrowed layer.
The math
In words: “plain data parallelism moves each gradient twice, once to sum it and once to share the sum; stage 3 moves each weight in twice, once for the forward pass and once for the backward, and each gradient once, to its owner.”
With the numbers: for the paper's 7.5-billion-parameter model, 15 billion numbers per GPU per step for plain data parallelism against 22.5 billion for stage 3: 1.5×. Stages 1 and 2 cost exactly 2Ψ: the reduce-scatter delivers each owner its gradients, and after the update an all-gather shares the new weights (§7.2.1).
In Python:
Psi = 7.5e9
# plain DP: reduce-scatter + all-gather
V_DP = Psi + Psi
# stage 3: forward all-gather + backward all-gather + gradient reduce-scatter
V_3 = Psi + Psi + Psi
V_DP / 1e9, V_3 / 1e9 # → (15.0, 22.5)
V_3 / V_DP # → 1.5
Why it matters today
This is why sharding won: for the price of half again the traffic, memory per GPU divides by the number of GPUs. In practice the next layer's weights are fetched while the current layer computes, so much of the extra traffic hides behind arithmetic. The pretraining lesson derives the ring's traffic in all_reduce_traffic, and the hardware lesson turns it into seconds over real links.
8 Communication analysis of ZeRO-R · original
Everyday picture
Partitioning activations (§6.1) adds one more conversation: before each layer is recomputed in the backward pass, its checkpoint must be gathered back together. Is that expensive? Only compared with nothing. Compared with the conversations model parallelism already has inside every layer, it is small.
Tiny example
In Megatron-LM with activation checkpointing, each transformer block runs 6 all-reduces (2 in the forward pass, 2 when the forward pass is recomputed, 2 in the backward pass), each of a message the size of the block's input, b × s × h numbers. An all-reduce moves twice its message, so a block moves 12 messages' worth. Pa adds one all-gather, which moves one message.
In words: “one extra all-gather per block, against six all-reduces that each cost two messages: about 8% more model-parallel traffic.”
With the numbers: for the 100B model's blocks (width 8,192) at batch 32 and 1,024 tokens, the message is 268 million numbers; Megatron-LM moves 3.2 billion per block and Pa adds 268 million, 8.3%, under the paper's “less than 10%”. (The paper writes the per-block total as 12 × seq_length × hidden_dim, leaving the batch size out; the ratio is the same.)
In Python:
b, s, h = 32, 1024, 8192
message = b * s * h
round(message / 1e6) # → 268
# Megatron-LM: 6 all-reduces, each moving 2 messages
V_MP = 6 * 2 * message
round(V_MP / 1e9, 1) # → 3.2
# P_a: one all-gather, one message
round(message / V_MP, 3) # → 0.083
The payoff can be larger than the cost. Partitioning activations frees memory for a bigger batch, and the data-parallel gradient traffic per sample falls as the batch grows. With 16-way model parallelism that is up to 16× the batch and an order of magnitude less data-parallel traffic per sample, the paper argues. Pa+cpu doubles the activation traffic again, to and from the CPU, and is worth it only when the batch would otherwise be tiny.
Why it matters today
The habit this section models, putting every communication on one ledger and comparing it against what it enables, is how parallelism strategies are still chosen. The Megatron-LM companion counts those 4 all-reduces per layer (without checkpointing) from the other side.
9 Step towards 1 trillion parameters · original
“It would require an exa-flop system to train a 1T parameter model in a reasonable time.”Rajbhandari et al. (2019), §9
Everyday picture
ZeRO solves the storage problem for a trillion parameters: the model fits. It does not solve the time problem: the cooking still takes as long as it takes. This section does both sums.
Tiny example
Fitting: 16 bytes × 1012 parameters = 16 TB of model states; over 1,024 GPUs that is about 16 GB each (15.6 exactly), inside a 32 GB V100. Megatron-LM alone topped out around 16 to 20 billion parameters within one 16-GPU DGX-2 server. Time: BERT-Large (330 million parameters) trained in 67 minutes on 1,024 GPUs. A trillion-parameter model does about 3,000 times the work per sample, so the same data would take 67 × 3,000 minutes: about 140 days, and more once data and sequence length grow with the model.
In Python:
# fitting: 16 bytes per parameter, a trillion parameters, 1,024 GPUs, in GB
round(16 * 1e12 / 1024 / 1e9, 1) # → 15.6
# time: work per sample relative to BERT-Large
round(1e12 / 330e6) # → 3030
# 67 minutes times the paper's 3,000, in days
round(67 * 3000 / 60 / 24, 1) # → 139.6
Why it matters today
The split the section draws, memory solved by systems and time bounded by compute, is the same one the scaling-laws companion formalizes: compute (about 6 × parameters × tokens) sets the price, and parallelism only decides whether you can pay it.
10 Implementation and evaluation · original
Everyday picture
The paper implements a subset, ZeRO-100B: stages 1 and 2 of ZeRO-DP plus all of ZeRO-R, aimed at models around 100 billion parameters, which fit and train in reasonable time on the hardware of the day. Stage 3 was left for later.
10.1 Implementation and methodology · original
Everyday picture
ZeRO-100B wraps an ordinary PyTorch model, like the standard data-parallel wrapper; the model is not changed. It runs on 400 V100 GPUs (25 DGX-2 servers of 16) with 800 Gbps between servers. Baselines: PyTorch's distributed data parallel without model parallelism, and Megatron-LM (as of September 2019) with it.
Tiny example
The models are GPT-2-like transformers whose size is set by depth and width. A transformer's weights are about 12 × layers × width², so the 100B configuration, 125 layers of width 8,192, has 12 × 125 × 8,1922 ≈ 100.7 billion parameters. (The rule of thumb is this page's check, not the paper's; the transformer lesson derives it.)
In Python:
# parameters ≈ 12 · layers · width², for 125 layers of width 8,192
round(12 * 125 * 8192 ** 2 / 1e9, 1) # → 100.7
Why it matters
Comparing per-GPU throughput keeps the comparison fair across different GPU counts; the paper's appendix notes that the baselines sometimes ran on fewer GPUs (256 or 384), which favours them.
10.2 Speed and model size · original
Everyday picture
Megatron-LM alone has to stretch its model parallelism across servers once a model outgrows one server's GPUs, and then every layer's conversation crawls over the network. ZeRO keeps model parallelism inside one server and lets sharded data parallelism span the servers.
Tiny example
Inside a DGX-2, GPUs talk at 300 GB/s per link; between servers, 12.5 GB/s per link: 24× slower. For 100B models ZeRO-100B sustains over 38 teraflops per GPU; across all 400 GPUs that is about 15 petaflops, over 30% of the hardware's peak.
In Python:
# link bandwidth inside a server against between servers, GB/s
300 / 12.5 # → 24.0
# aggregate throughput: 38 teraflops on each of 400 GPUs, in petaflops
38e12 * 400 / 1e15 # → 15.2
| Megatron-LM alone | ZeRO-100B with Megatron-LM | |
|---|---|---|
| Largest model run efficiently on 400 GPUs | under 40B (§1); 16 to 20B inside one 16-GPU server (§9) | 170B |
| Throughput on 100B-parameter models | far lower: the speedup is up to 10× | over 38 teraflops per GPU |
| A 40B model across two servers | about 5 teraflops per GPU |
Why it matters today
The layout ZeRO argued for, chatty model parallelism on the fast links inside a server and data parallelism across the slower network, is the one the hardware lesson derives from the wires themselves.
10.3 Super-linear scalability · original
Everyday picture
Normally doubling the GPUs at best doubles the speed. With ZeRO, doubling the GPUs also shrinks each GPU's share of the model states, which frees room for a bigger batch per GPU, and bigger batches keep each GPU busier. For a 60-billion-parameter model going from 64 to 400 GPUs, throughput per GPU rises: more than double the speed for double the GPUs.
Tiny example
The 60B model uses 16-way model parallelism (the paper's appendix), so each model-parallel slice holds 60 / 16 = 3.75 billion parameters. On 64 GPUs that is Nd = 4 copies; on 400 GPUs, Nd = 25. With stage 2, model states per GPU fall from 20.6 GB to 9.6 GB, and the appendix shows the batch per model-parallel group growing from 16 to 64.
In words: “model parallelism first divides the parameters by Nm; stage 2 then shares the gradients and optimizer state of each slice over the Nd data-parallel copies.”
With the numbers: (2 + 14/4) × 3.75 = 20.6 GB at 64 GPUs; (2 + 14/25) × 3.75 = 9.6 GB at 400. Eleven gigabytes per GPU freed for activations.
In Python:
K, Psi, N_m = 12, 60e9, 16
# data-parallel copies on 64 and on 400 GPUs
64 // N_m, 400 // N_m # → (4, 25)
# M_os+g per GPU, in GB
def M(N_d):
return (2 + (2 + K) / N_d) * Psi / N_m / 1e9
round(M(4), 1), round(M(25), 1) # → (20.6, 9.6)
Why it matters today
Memory and speed are not separate budgets: freed memory buys batch size, and batch size buys utilization. The paper's footnote adds the caveat that batches cannot grow forever before convergence suffers.
10.4 Democratizing large model training · original
Everyday picture
Model parallelism needs someone to rewrite the model. ZeRO does not: wrap the model and train. The paper's measure of that is the biggest model trainable with no model parallelism at all.
Tiny example
On 128 GPUs, plain data parallelism tops out at 1.4 billion parameters (16 × 1.4 = 22.4 GB of model states, plus activations, on a 32 GB GPU). ZeRO-100B trains up to 13 billion: with stage 2, (2 + 14/128) × 13 = 27.4 GB, still under 32. Throughput: over 40 teraflops per GPU on average, against under 20 for the 1.4B baseline.
In Python:
# plain data parallelism, 1.4 billion parameters, in GB
round(16 * 1.4, 1) # → 22.4
# ZeRO stage 2 on 128 GPUs, 13 billion parameters, in GB
round((2 + 14 / 128) * 13, 1) # → 27.4
Why it matters today
This is the result that made sharding a default: researchers without distributed-systems help could train models ten times bigger, larger than any published model at the time (T5 had 11 billion), on ordinary data-parallel code, and on clusters without the fastest links inside each server.
10.5 Memory and performance analysis · original
Everyday picture
An ablation: switch the optimizations on one at a time (the paper's configurations C1 to C5) and see which one buys what.
Tiny example
| Configuration | What changes | Largest model |
|---|---|---|
| C1 | Pos with constant buffers and defragmentation | 40B |
| C2 | + partitioned activations: activation memory ÷ 16 | 60B |
| C4 | Pos+g with partitioned activations: model states roughly halved | 140B |
| C5 | + activations moved to CPU memory | 150B |
Throughput follows memory: every configuration that frees memory allows a bigger batch and runs faster, with one exception. Moving activations to the CPU (C5) costs time on the slow link, so it helps only when the model cannot run without it, as for the 170B model; ZeRO turns it on only then.
Why it matters today
Offloading became its own line of work (ZeRO-Offload and ZeRO-Infinity, below), each designed around the finding here: the slow link to the CPU is worth using only for what it can move without stalling the GPU.
10.6 Turing-NLG · original
Everyday picture
The system result came with a model: Turing-NLG, 17 billion parameters, the largest language model published as of May 2020, trained end to end with ZeRO-100B at 41.4 teraflops per GPU.
Tiny example
Its perplexity of 10.21 (on the benchmark the paper names Webtext-103) beat the previous best, Megatron-LM's 8.3-billion-parameter model. Perplexity 10.21 means the model is, on average, about as unsure as a choice among 10 equally likely next tokens.
Why it matters today
A systems paper that trains the record model makes its own case. The Megatron-LM paper notes that Turing-NLG was trained using Megatron too: ZeRO's sharding across servers, Megatron's model parallelism inside each, the combination §1 proposed.
11 Concluding remarks · original
“Unlike existing approaches such as MP and PP, no model refactoring is necessary, and it is as easy to use as standard DP”Rajbhandari et al. (2019), §11
Everyday picture
The authors call the implementation “just a tip of the iceberg”: stages 1 and 2 gave 8× bigger models; stage 3 would give another order of magnitude.
Tiny example
On 400 GPUs with 16-way model parallelism (25 data-parallel copies), stage 2 leaves 2 + 14/25 = 2.56 bytes per parameter per slice; stage 3 leaves 16/25 = 0.64. The same memory then holds 4× the parameters, before counting activations.
In Python:
# bytes per parameter (per model-parallel slice) with N_d = 25
stage2 = 2 + 14 / 25
stage3 = 16 / 25
stage2, stage3 # → (2.56, 0.64)
stage2 / stage3 # → 4.0
Why it matters
The ease-of-use argument won as much as the memory argument: sharded data parallelism is a wrapper around an unchanged model, which is a large part of why it spread.
What happened next
| Development | What it does |
|---|---|
| DeepSpeed | The open-source library the paper released ZeRO in; the later ZeRO systems below are released there too |
| ZeRO-Offload (2021) | Moves data and computation to the CPU as well: over 13 billion parameters trained on a single GPU, with no model changes |
| ZeRO-Infinity (2021) | Uses GPU, CPU and NVMe memory together, enough for models of tens of trillions of parameters on existing clusters |
| PyTorch FSDP (Zhao et al., 2023) | Fully sharded data parallelism built into PyTorch: stage 3 as a standard library feature |
| Megatron-LM at scale (Narayanan et al., 2021) | Tensor, pipeline and data parallelism composed to train a trillion-parameter model on 3,072 GPUs; the Megatron-LM companion covers the tensor part |
To build the pieces yourself, the pretraining lesson walks from the 16-byte bill through the ZeRO stages to tensor and pipeline parallelism, and the hardware lesson explains the links underneath.
Glossary
Every term with hover guidance on this page, in one place.