How far does splitting the batch get you?
Data parallelism is the axis nobody has to be talked into. Every GPU runs the ordinary single-device code on its own slice of the batch, the framework averages the gradients behind your back, and the result is mathematically identical to one enormous GPU with one enormous batch. Nothing about the model is cut.
Which is exactly why it runs out. The model is not cut, so every GPU still needs room for all of it, and the activations of its own tokens on top. ZeRO fixes the first half of that sentence and cannot touch the second.
This part follows the axis from a plain all-reduce to fully sharded data parallelism, counts the bytes at every step, and then works out the three walls it hits.
The 30-second version
Each rank runs forward and backward on its own tokens with no communication at all, then the gradients are averaged with an all-reduce so every rank applies the identical update. A ring all-reduce is a reduce-scatter followed by an all-gather, which moves two times the message minus a bit per rank, independent of how many ranks there are, and frameworks bucket the gradients so the reduction of the last layer overlaps the backward pass of the one before it. ZeRO observes that the replicas are redundant. Stage 1 shards the optimizer states, stage 2 shards the gradients too, stage 3 shards the weights as well and gathers each layer just before it runs. Stages 1 and 2 move the same 2 psi elements per step as plain data parallelism; stage 3 moves 3 psi, a 50% increase, in exchange for memory that falls linearly in the number of ranks. What none of them touch is the activations, because each rank still processes its own tokens, and that is the wall.What actually happens in a data-parallel step?
Four phases, and only one of them talks to anyone.
Each rank takes a different slice of the global batch, runs the forward pass, runs the backward pass, and ends up with a gradient for every parameter. Those gradients disagree, because each rank saw different tokens, and the optimizer wants the average over the whole batch.
One collective fixes that. After it, every rank holds the identical averaged gradient, applies the identical update to its identical weights, and the invariant holds for the next step.
For our running layout, where tensor parallelism has already cut the weights eight ways, each data-parallel rank is responsible for $\Psi = 8.81$ billion parameters. In BF16 that is 17.6 GB of gradient, and the all-reduce moves about twice that per rank.
Why does the all-reduce not get more expensive with more ranks?
This is the first thing interviewers check, and the answer is a two-line derivation.
A ring all-reduce is two halves. In the reduce-scatter, the message of $M$ bytes is cut into $d$ chunks, and each chunk walks around the ring being summed as it goes. After $d-1$ hops, rank $i$ owns the completely reduced version of one chunk. In the all-gather, those finished chunks walk around again until everyone has all of them.
Every hop moves one chunk, $M/d$ bytes, and there are $2(d-1)$ hops:
\[\text{bytes sent per rank} = \frac{2(d-1)M}{d} < 2M\]So bandwidth cost is flat in $d$. What is not flat is latency: there are $2(d-1)$ sequential hops, so a ring across 64 ranks pays 126 link latencies before the last byte lands. That is why large clusters use tree or hierarchical algorithms for small messages and rings for large ones, and why NCCL picks between them at runtime.
The second step of the figure animates the chunk grid and counts the bytes as they move. At $d = 8$ with a 17.6 GB gradient, each rank sends 30.8 GB, and it would send less than 35.2 GB no matter how large the ring got.
How does the collective hide behind the backward pass?
The reduction does not have to wait for the backward pass to finish, because gradients become final layer by layer, from the top of the network down.
So the framework buckets them. PyTorch’s DDP fills a bucket, 25 MB by default, and launches an asynchronous all-reduce the moment it is full. The reduction of the last layer overlaps the backward pass of the one before it, and only the tail sticks out past the end.
That works beautifully when there is a lot of compute per byte, and stops working when there is not. The collective is fixed by the parameter count. The backward pass shrinks with the tokens on the GPU.
At 8192 tokens per GPU the backward arithmetic is 292 ms and the all-reduce over 64 ranks is 694 ms, so more than half the step is exposed. The third step of the figure lets you drag the token count and watch the tail appear.
What is redundant about holding the optimizer state 512 times?
Nothing about a data-parallel step needs every rank to hold every optimizer moment. The ZeRO paper’s observation is that after the reduction, rank $i$ only ever updates its own slice of the parameters, so it only ever needs its own slice of the states.
Stage 1 shards the twelve bytes of FP32 master weight and Adam moments across the $d$ ranks. Stage 2 shards the two bytes of gradient too, which turns the all-reduce into a reduce-scatter. Stage 3 shards the two bytes of BF16 weight as well, which means gathering each layer’s parameters just before it runs and freeing them straight after.
\[\text{DDP: } 16\Psi \quad \text{ZeRO-1: } 4\Psi + \frac{12\Psi}{d} \quad \text{ZeRO-2: } 2\Psi + \frac{14\Psi}{d} \quad \text{ZeRO-3: } \frac{16\Psi}{d}\]The floors matter more than the slopes. Stage 1 never gets below $4\Psi$ and stage 2 never below $2\Psi$, whatever $d$ is, because the BF16 weight and then the gradient stay replicated on every rank. For our $\Psi$, that is 35.3 GB and 17.6 GB of permanent floor. Only stage 3 goes to zero.
What does each stage cost on the wire?
The ZeRO paper counts communication in elements, which is worth copying because it makes the comparison clean.
Plain data parallelism moves $2\Psi$ elements per rank per step: an all-reduce is a reduce-scatter of $\Psi$ followed by an all-gather of $\Psi$. Stage 1 moves the same $2\Psi$. Stage 2 moves the same $2\Psi$, because the reduce-scatter it uses is exactly the first half of the all-reduce and the all-gather of updated parameters is the second half.
Stage 3 moves $3\Psi$. The extra $\Psi$ is a second all-gather: the weights have to be gathered once for the forward pass and once again for the backward, because they were freed in between.
In BF16 those are four and six bytes per parameter. For our rank, 35.3 GB and 52.9 GB per step. A 50% increase in traffic for memory that falls linearly in $d$ is the trade, and at any reasonable scale it is a good one.
There is a middle option that shows up in real runs. Meta’s Llama 3 used FSDP to shard the optimizer states and the gradients, but deliberately did not reshard the parameters after the forward pass, so the backward pass reuses the gathered weights instead of paying for a second all-gather. That is stage 3’s memory for the forward and stage 2’s traffic overall, and it is the right choice when memory is not the binding constraint.
What does stage 3 actually do per layer?
Stage 3 is a promise that a layer’s weights will exist by the time the layer runs.
Before layer $i$ computes, its shard is all-gathered from the $d$ ranks into a full copy. The layer computes. The full copy is freed. In the backward pass the same gather happens again, the gradient is computed, a reduce-scatter sends each rank its own slice, and the copy is freed again.
Done naively, the GPU idles through every gather. Done with prefetch, the gather for layer $i+1$ is issued while layer $i$ is still computing, and the only exposed gather in the whole model is the first one. The fifth step of the figure has a prefetch toggle: turn it off and the compute row fills with holes.
The knob is wrapping granularity. Wrap too coarsely and the gathered chunk is enormous, which defeats the memory saving. Wrap too finely and each gather is too small to reach line rate on the network.
PyTorch’s FSDP2 shards each parameter on its own with DTensor rather than flattening a module’s parameters into one buffer, which is what turns wrapping into a per-module decision and makes it compose with tensor parallelism. In practice you wrap one transformer block per unit, and one block’s weights on our layout are 220 MB gathered.
What does gradient accumulation actually buy?
It buys arithmetic to hide the collective behind, and nothing else.
Run $k$ micro-batches, accumulate the gradients locally, reduce once. The collective is unchanged, because it is sized by the parameter count. The compute per optimizer step is $k$ times larger. So the exposed fraction falls as $1/k$ without changing the global batch’s learning dynamics at all, as long as you were going to run those tokens anyway.
The catch is numerical. Accumulating $k$ BF16 gradients loses the small contributions, so the accumulator has to be FP32. Meta reported doing exactly that for Llama 3: FP32 gradient accumulation across micro-batches, and an FP32 reduce-scatter across data-parallel workers.
The other thing to know is what accumulation does not buy. It does not reduce memory, because the accumulator is the same size as the gradient. And it does not help the pipeline, which wants micro-batches for a completely different reason that Part 4 gets to.
Where does the batch axis stop?
Three walls, and none of them is about bandwidth.
The activations never shard along this axis. After ZeRO-3 has taken the model states on our layout from 141 GB to 2.2 GB, the 182.5 GB of activations for a single 8192-token sequence is exactly where it was. That is the wall that sends you to tensor, pipeline or context parallelism, and it is the one that matters most.
The global batch has a ceiling. Global batch is tokens per rank times $d$, so at fixed tokens per rank, adding ranks grows the batch, and batch size is a hyperparameter that stops helping. Hold the batch fixed instead and the tokens per rank fall. Meta held Llama 3 405B at 16 million tokens per batch and halved the per-rank batch when the data-parallel degree went from 64 to 128; the reported MFU fell from 43% to 41%, and they attribute the drop to exactly that.
And falling tokens per rank means less arithmetic to hide the same fixed collective behind, which is the third wall and a consequence of the second.
All three say the same thing. Past a point you have to cut the model, not the batch.
The 30-second version
First the arithmetic: parameters times six bytes for stage 3 is the traffic per rank per step, and tokens per rank times six FLOPs per parameter is the compute it has to hide behind. If the second is smaller than the first, no amount of tuning fixes it and the answer is gradient accumulation or a larger micro-batch. If there is enough compute, then it is a prefetch or wrapping problem: check that the unit is a whole transformer block rather than the whole model or individual linears, check that forward prefetch and backward prefetch are enabled, and check that the first all-gather is not being issued inside the layer that needs it. Then check whether the job should be resharding after forward at all, because keeping the gathered parameters for the backward pass removes an entire all-gather at the cost of holding one unit longer. And confirm the data-parallel group is not spanning a slow tier of the network when a hybrid sharded arrangement would keep the all-gathers inside a node.Rapid fire: can you do these from memory?
- Derive the per-rank byte count of a ring all-reduce and say what does and does not grow with the ring size.
- Write the per-rank memory for DDP and the three ZeRO stages, and name the two floors.
- Give the communication volume in elements for each stage, and say where stage 3's extra volume comes from.
- Why does bucketing make the reduction overlap the backward pass, and when does it stop helping?
- Describe the lifecycle of one layer's weights under stage 3, in both passes.
- What does gradient accumulation change, what does it not change, and what precision does the accumulator need?
- Name the three reasons more data parallelism stops helping, and which of them is about memory.
Part 3 cuts the model for the first time, and starts with the smallest unit worth cutting: a single matrix multiply.