ZeRO, the Zero Redundancy Optimizer behind DeepSpeed, removes the copies that plain data parallelism keeps on every GPU. Stage 1 shards the optimizer state, stage 2 also shards gradients, stage 3 also shards the weights. The memory side of that story, bytes per parameter and a worked 7B budget, is covered in ZeRO Optimizer, in depth. This article is about the decision you make repeatedly as a project grows: which stage to run now, what it will cost in communication, how gradient accumulation changes the answer, and how to move a run from one stage or cluster size to another without losing its optimizer state.

The short version: pick the lowest stage that fits, because each stage up the ladder buys memory with traffic. The long version needs a cost model, and building one takes about forty lines of Python.

The stages in one table

Use Psi for the parameter count and N for the number of data-parallel ranks. With bf16 mixed precision and Adam, each parameter costs 2 bytes of bf16 weight, 2 bytes of bf16 gradient and 12 bytes of fp32 state (master weight, first moment, second moment). DeepSpeed also has a stage 0, which is plain data parallelism inside the DeepSpeed engine and is the right baseline to measure against.

StageShardedBytes per parameter per rankCollectives per step
0nothing2 + 2 + 12 = 16all-reduce gradients
1optimizer state2 + 2 + 12/Nreduce-scatter gradients, all-gather updated weights
2+ gradients2 + (2 + 12)/Nreduce-scatter gradient buckets during backward, all-gather weights
3+ weights16/Nall-gather each layer in forward and again in backward, reduce-scatter gradients
What each ZeRO stage shards, and what it moves per optimizer stepbf16 paramsbf16 gradsfp32 optimizer statetraffic, k micro-batchesStage 0fullfullfull2 x Psi (boundary)Stage 1fullfull1/N2 x Psi (boundary)Stage 2full1/N1/N(k + 1) x PsiStage 31/N1/N1/N3k x PsiGreen is sharded across N data-parallel ranks; red is replicated on every rank.Psi = parameter count; traffic in elements sent per rank, ring collectives, (N-1)/N ignored.
Stages shard progressively more of the model state. The traffic column is per rank per optimizer step with k micro-batches of gradient accumulation, as DeepSpeed schedules it.

Counting the traffic

Traffic is where stages differ in cost, and it is easiest to count in elements sent per rank. A ring all-reduce of Psi elements sends about 2 Psi per rank, because it is a reduce-scatter (Psi) followed by an all-gather (Psi); the exact factor is 2(N-1)/N. Details of each collective are in GPU Collective Operations.

Without gradient accumulation the counts are simple. Stage 0 all-reduces gradients: 2 Psi. Stage 1 reduce-scatters gradients so each rank gets the slice it updates, then all-gathers the new weights: 2 Psi, the same as stage 0. Stage 2 does the same two collectives, just bucketed into backward: 2 Psi. Stage 3 gathers every layer's weights for forward (Psi), gathers them again for backward because they were freed (Psi) and reduce-scatters gradients (Psi): 3 Psi, one and a half times data parallelism.

Gradient accumulation changes the picture, and this is the part most guides skip. With k micro-batches per optimizer step, stages 0 and 1 keep a full-size gradient buffer on every rank, accumulate locally and communicate once at the accumulation boundary, so traffic stays at 2 Psi per step. Stage 2 does not keep full gradients; that is its whole saving. In DeepSpeed's implementation it reduce-scatters each micro-batch's gradients in backward hooks and accumulates into the local shard, so it sends k Psi of gradients plus Psi of weight all-gather: (k + 1) Psi per step. Stage 3 also has to re-gather weights for every micro-batch, so it sends about 3k Psi. This is implementation behaviour rather than something the ZeRO paper fixes; DeepSpeed's engine source says so in a comment (stage 2 and above communicate on non-boundary micro-steps as well), but confirm it in a profiler trace for the release you run.

The consequence: at k = 16, stage 2 moves about 8.5 times the bytes of stage 1 per optimizer step. On NVLink inside one node that rarely matters, because it overlaps with backward. Across nodes on Ethernet it can be the difference between compute-bound and network-bound. Stage 1 is the underrated stage for large accumulation counts on slow fabrics.

A cost model for a 13B job

Put the memory and the traffic side by side for a concrete job: a 13B decoder on two nodes of eight 80 GB GPUs (N = 16), bf16, AdamW, micro-batch 4 at 4,096 tokens, k = 8 accumulation steps. Assume an effective 40 GB/s per rank for collectives that cross nodes and about 400 TFLOP/s of achieved bf16 compute per GPU. Both are assumptions; measure yours with nccl-tests and a short training run.

PSI, N, K = 13e9, 16, 8
MICRO, SEQ = 4, 4096
BW = 40e9            # bytes/s per rank for cross-node collectives (assumed)
FLOPS = 400e12       # achieved bf16 FLOP/s per GPU (assumed)

def model_gb(stage):
    w, g, o = 2, 2, 12                          # bytes per parameter
    if stage >= 1: o /= N
    if stage >= 2: g /= N
    if stage >= 3: w /= N
    return PSI * (w + g + o) / 1e9

def traffic_elems(stage):
    return {0: 2, 1: 2, 2: K + 1, 3: 3 * K}[stage] * PSI

compute_s = 6 * PSI * MICRO * SEQ * K / FLOPS    # 6 FLOPs per param per token
for s in range(4):
    comm_s = traffic_elems(s) * 2 / BW          # bf16 = 2 bytes per element
    print(f"stage {s}: model state {model_gb(s):6.1f} GB/rank, "
          f"comm {comm_s:5.1f} s vs compute {compute_s:4.1f} s per step")

The output tells the story. Stage 0 needs 208 GB per rank and is out. Stage 1 needs 61.8 GB of model state, leaving under 20 GB for activations, buffers and allocator fragmentation: tight, workable only with activation checkpointing. Stage 2 needs 37.4 GB and stage 3 needs 13 GB. Communication per step is 1.3 s for stage 1, 5.8 s for stage 2 and 15.6 s for stage 3, against about 25.6 s of compute.

All three fit under the compute time if communication overlaps perfectly, but overlap is never perfect. Stage 1's traffic arrives in one burst at the boundary and is largely exposed, yet it is small. Stage 2's traffic overlaps well with backward. Stage 3's all-gathers sit on the critical path of every layer and need prefetching to hide. The practical pick here is stage 2, with stage 1 as the comparison run, and stage 3 only if you need a longer sequence or a larger micro-batch than stage 2's leftover memory allows.

The stage ladder

Turn that reasoning into a procedure you run whenever the model, sequence length or cluster changes. Each rung is one short run of a few hundred steps, recording tokens per second per GPU and peak memory.

for stage in (0, 1, 2, 3):
    run 200 steps with zero_optimization.stage = stage
    if OOM: continue
    record tokens/s/GPU, torch.cuda.max_memory_allocated(), step-time breakdown
    if peak memory < 90% of HBM: pick the fastest stage that fits with margin; stop
if nothing fits at stage 3:
    add activation checkpointing, then offload_optimizer to CPU, then offload_param
    or add GPUs / tensor parallelism rather than offloading weights

Switching stage is one config line, which is what makes the ladder cheap:

import argparse, deepspeed, torch

parser = argparse.ArgumentParser()
parser.add_argument("--stage", type=int, default=2)
parser.add_argument("--local_rank", type=int, default=-1)
args = parser.parse_args()

ds_config = {
    "train_micro_batch_size_per_gpu": 4,
    "gradient_accumulation_steps": 8,
    "bf16": {"enabled": True},
    "gradient_clipping": 1.0,
    "wall_clock_breakdown": True,              # forward/backward/step timings in the log
    "zero_optimization": {
        "stage": args.stage,
        "overlap_comm": True,
        "contiguous_gradients": True,
        "reduce_bucket_size": 5e8,
        "allgather_bucket_size": 5e8,
    },
}
model = build_model()                           # under deepspeed.zero.Init() for stage 3
optim = torch.optim.AdamW(model.parameters(), lr=2e-4)
engine, optim, _, _ = deepspeed.initialize(model=model, optimizer=optim, config=ds_config)

for batch in loader:
    loss = engine(batch)
    engine.backward(loss)
    engine.step()                               # steps the optimizer only at the boundary

Two notes on the config. The bucket sizes are in elements, so 5e8 in bf16 is a 1 GB buffer, and with overlap enabled more than one can be in flight; if a rung fails by a gigabyte or two, halve the buckets before moving up a stage. And DeepSpeed counts optimizer steps, not micro-batches, in its learning-rate schedule when you hand it the scheduler, so changing k changes the schedule.

Stage 3 across nodes

Stage 3 across nodes is where communication bites hardest, and DeepSpeed has three tools for it. hpZ, set with zero_hpz_partition_size equal to the GPUs per node, keeps a secondary copy of the weights sharded within each node so the backward all-gather stays on NVLink. qwZ (zero_quantized_weights) quantizes the forward weight all-gather and qgZ (zero_quantized_gradients) quantizes the gradient reduce-scatter. Together these are ZeRO++, whose paper reports up to 4x less cross-node communication volume. The quantized variants change numerics, so compare a loss curve before adopting them. Older guides also recommend MiCS, which sharded within a subgroup of ranks via mics_shard_size; current DeepSpeed has removed it and reports that key as no longer supported, so do not plan around it.

Offloading optimizer state to CPU or NVMe is the next rung after stage 3; its PCIe arithmetic is in Optimizer State Offloading.

Changing stage or cluster size mid-project

ZeRO checkpoints are sharded the way the run was sharded. Each rank writes its own slice of the optimizer state, and from stage 3 also its slice of the weights, so a checkpoint from stage 2 on 16 GPUs is not directly loadable on stage 3, or on 32 GPUs. Plan for this before the first long run.

For export, DeepSpeed writes a zero_to_fp32.py script into the checkpoint directory that consolidates the shards into a single fp32 state dict. That is enough for inference or for fine-tuning from scratch, but it drops the optimizer state, so a resumed run restarts Adam's moments. To resume with a different stage or world size, use Universal Checkpointing: save a normal ZeRO checkpoint, convert it with ds_to_universal.py, and load with the checkpoint config option set:

{
  "checkpoint": { "load_universal": true }
}

Test the round trip on a small model early: train 50 steps, save, convert, resume at a different N, and check that the loss continues from the same value rather than spiking.

Failure modes

  • Stage 2 slower than stage 1 at high accumulation. The (k + 1) Psi traffic above. Check the step breakdown and try stage 1 with activation checkpointing.
  • Stage 3 throughput collapses with a small micro-batch. Per-layer compute shrinks while the all-gather stays the same size, so communication is exposed. Raise the micro-batch, enlarge stage3_prefetch_bucket_size, or stay on stage 2.
  • Code that touches weights outside forward under stage 3. A sharded parameter is a placeholder; reading its shape or values directly gives an empty tensor. Wrap such access in deepspeed.zero.GatheredParameters.
  • Out of memory by a small margin after moving up a stage. Communication buckets and fragmentation, not model state. Reduce bucket sizes before changing anything else.
  • Loss spike after resume. Optimizer moments lost through zero_to_fp32.py or a world-size mismatch. Use a universal checkpoint.
  • Hangs at the first step on multi-node. Usually NCCL configuration rather than ZeRO; run the same job at stage 0 to separate the two.

Trade-offs against the alternatives

PyTorch FSDP exposes the same ideas under different names: NO_SHARD is data parallelism, SHARD_GRAD_OP shards gradients and optimizer state like stage 2, FULL_SHARD is stage 3, and HYBRID_SHARD shards within a node and replicates across nodes, the layout hpZ approximates for weights. There is no direct stage 1 equivalent. FSDP is the native choice inside a PyTorch-only stack and composes with tensor parallelism through DTensor; DeepSpeed offers more offload tiers, ZeRO++ and universal checkpoints. See PyTorch FSDP, in depth and DeepSpeed, in depth for the wider comparison.

Against tensor and pipeline parallelism, ZeRO keeps the model code unchanged, which is its main advantage, but it never reduces activation memory per GPU. When activations dominate, long sequences being the usual cause, ZeRO alone is the wrong lever.

What to do next

  1. Run the cost model above with your parameter count, N, accumulation steps and measured bandwidth.
  2. Measure bus bandwidth across nodes with nccl-tests before trusting any traffic estimate.
  3. Walk the ladder: 200 steps at each stage, recording tokens per second per GPU and peak memory.
  4. Pick the lowest stage that fits with 10% headroom; prefer stage 1 when accumulation is high and the fabric is slow.
  5. For stage 3 across nodes, try hpZ first, then the quantized options with a loss-curve comparison.
  6. Prove a universal-checkpoint round trip at a different world size before the first long run.
Key takeaway: Each ZeRO stage buys memory with communication: stages 0 to 2 move about 2 Psi per step without accumulation and stage 3 moves 3 Psi, but with k accumulation steps stage 2 and stage 3 scale with k while stage 1 does not. Model the cost, walk the ladder from the lowest stage up, choose the lowest stage that fits, and use universal checkpoints so the choice can change later.