Plain data parallelism keeps a full copy of the model, its gradients and its optimizer state on every GPU. With mixed-precision Adam that is about 16 bytes per parameter, so a 7-billion-parameter model needs roughly 112 GB per GPU before a single activation is stored, which no current single accelerator of the 80 GB class can hold. ZeRO (the Zero Redundancy Optimizer, from the DeepSpeed team) removes that redundancy by giving each of N data-parallel ranks one Nth of the state and moving the rest over the network only when it is needed.

The partitioning idea is covered in ZeRO sharding architecture and the per-stage algebra in ZeRO math. This page is about running it: what each stage changes inside one training step, how to size a job before launching it, which DeepSpeed settings actually move memory and throughput, how model code must change at stage 3, how checkpoints come out, and what breaks.

Advertisement

Where the 16 bytes per parameter go

Mixed-precision training with Adam keeps five things per parameter: the 16-bit weight (2 bytes), its 16-bit gradient (2 bytes), and in 32-bit precision the master weight and Adam's two moments (4 bytes each, 12 in total). The optimizer's 12 bytes are touched once per step, in the update; the weights are needed by every layer in forward and backward.

That access pattern is why ZeRO comes in three stages. Stage 1 shards the optimizer state, which is the largest piece and the least frequently used, so sharding it costs almost nothing extra in communication. Stage 2 also shards gradients, which is nearly free because data parallelism already has to reduce them across ranks. Stage 3 also shards the weights, which are needed constantly, so it is the only stage that adds real traffic. Activations are not touched by any stage; they depend on batch size, sequence length and activation checkpointing, and must be budgeted separately.

What each stage changes inside one step

It is easiest to reason about ZeRO as a change to the collectives in a data-parallel step. Plain data parallelism does forward, backward, an all-reduce of gradients, and a full optimizer step on every rank.

  • Stage 1. Gradients are still reduced across ranks, but each rank runs Adam only on the slice of parameters it owns, then the updated 16-bit weights are all-gathered so every rank has the full model for the next forward. Traffic is about the same as plain data parallelism, because an all-reduce is itself a reduce-scatter followed by an all-gather.
  • Stage 2. Instead of all-reducing gradients, ranks reduce-scatter them in buckets during backward, so each rank keeps only the gradient slice it will update. Full gradients never persist. Traffic is again about the same as plain data parallelism.
  • Stage 3. Weights are sharded too. Before each layer's forward, its parameters are all-gathered, used, and freed; the same happens again in backward, followed by a reduce-scatter of that layer's gradients. That is roughly one and a half times the communication volume of plain data parallelism, and it now sits on the critical path of every layer unless it is prefetched.
One ZeRO-3 training step on rank r of N (each rank owns 1/N of every tensor)Shard rfp16 params 1/NShard rfp32 master + m + v 1/Nall-gatherlayer k paramsforward kthen free full knext layerrepeatall-gatherlayer k againbackward kfull grads for kreduce-scatterkeep grad shard roptimizer stepAdam on shard r onlygrad shardupdate masterStage 1 shards only the green state; stage 2 adds the red reduce-scatter; stage 3 adds the yellow all-gathers.
Stage 3 per-layer lifecycle on one rank. Parameters are gathered just in time for forward and again for backward, gradients are reduce-scattered so each rank keeps its own slice, and Adam updates only the local shard of fp32 state.

The practical consequence: stages 1 and 2 are close to free on any interconnect that already handles data parallelism, so they are the default for anything larger than a few hundred million parameters. Stage 3 is the one to measure, because its cost depends on the ratio of per-layer compute to per-layer gather time, which gets worse as the micro-batch shrinks or the interconnect slows. The collectives themselves are explained in NCCL collectives.

Advertisement

Worked example: a 7B model on eight 80 GB GPUs

Take a 7-billion-parameter decoder, bf16 mixed precision, AdamW, and one node of eight 80 GB GPUs. Per-rank model-state memory follows directly from the 2 + 2 + 12 split, with N = 8:

ConfigurationBytes per parameter per rankModel state per rank
Plain data parallel2 + 2 + 12 = 16112 GB (does not fit)
ZeRO stage 12 + 2 + 12/8 = 5.538.5 GB
ZeRO stage 22 + (2 + 12)/8 = 3.7526.3 GB
ZeRO stage 316/8 = 214 GB

Now add what the table leaves out. Communication buckets hold temporary full-size buffers: with the documented default bucket sizes of 5e8 elements, a reduce bucket is on the order of a gigabyte in 16-bit, and there can be more than one in flight when communication overlaps compute. The CUDA caching allocator adds fragmentation, commonly several gigabytes. Activations for a micro-batch of four 4,096-token sequences on a 7B model are tens of gigabytes without activation checkpointing and a few gigabytes with it.

The decision falls out. Stage 1 at 38.5 GB leaves about 40 GB for activations, buffers and fragmentation, which is workable with activation checkpointing. Stage 2 at 26 GB is comfortable and costs nothing extra in communication, so it is the right starting point. Stage 3 is only worth its traffic here if you need a larger micro-batch or longer sequences than stage 2 allows. For a 70B model on the same node the numbers are ten times larger, stage 2 needs 263 GB per rank, and stage 3 across several nodes, or stage 3 plus offload, becomes mandatory.

Running it with DeepSpeed

DeepSpeed takes its settings from a JSON config. The block that controls ZeRO is zero_optimization. The config below is a sensible stage 3 starting point for the job above; each key is explained in the next section.

{
  "train_micro_batch_size_per_gpu": 4,
  "gradient_accumulation_steps": 8,
  "gradient_clipping": 1.0,
  "bf16": { "enabled": true },
  "optimizer": {
    "type": "AdamW",
    "params": { "lr": 2e-4, "betas": [0.9, 0.95], "weight_decay": 0.1 }
  },
  "zero_optimization": {
    "stage": 3,
    "overlap_comm": true,
    "contiguous_gradients": true,
    "reduce_bucket_size": 5e8,
    "stage3_prefetch_bucket_size": 5e8,
    "stage3_param_persistence_threshold": 1e5,
    "stage3_max_live_parameters": 1e9,
    "stage3_max_reuse_distance": 1e9,
    "stage3_gather_16bit_weights_on_model_save": true
  }
}

The training loop changes in three places. The model is built inside deepspeed.zero.Init so that at stage 3 no rank ever materialises the full model, deepspeed.initialize wraps it in an engine that owns the optimizer, and the loop calls engine.backward and engine.step instead of loss.backward() and optimizer.step(). The engine handles gradient accumulation: it counts micro-steps and only runs the optimizer on accumulation boundaries, so do not add your own accumulation logic on top.

import deepspeed
import torch

deepspeed.init_distributed()                      # NCCL process group from launcher env vars

with deepspeed.zero.Init(config_dict_or_path="ds_config.json"):
    model = build_model(cfg)                      # stage 3: each parameter is sharded as it is created

engine, optimizer, _, scheduler = deepspeed.initialize(
    model=model,
    model_parameters=model.parameters(),
    config="ds_config.json",
)
device = torch.device("cuda", engine.local_rank)

for step, batch in enumerate(loader):
    batch = {k: v.to(device) for k, v in batch.items()}
    loss = engine(**batch).loss                   # forward: all-gather layer by layer
    engine.backward(loss)                         # backward: gather again, reduce-scatter grads
    engine.step()                                 # steps only at accumulation boundaries
    if step % 2000 == 0:
        engine.save_checkpoint("ckpt", tag=f"step{step}")   # every rank writes its shard

The global batch is micro-batch times accumulation steps times data-parallel world size, here 4 x 8 x 8 = 256 sequences.

The settings that move memory and throughput

KeyDocumented defaultWhat it trades
stage0Which state is sharded. Pick the lowest stage that fits.
overlap_commset it explicitlyOverlaps gradient reduction with backward. Faster, but holds more bucket memory at once.
contiguous_gradientstrueCopies gradients into one contiguous buffer as they are produced, which reduces fragmentation.
reduce_bucket_size5e8Elements reduced per collective. Larger means fewer, more efficient calls but a bigger temporary buffer.
allgather_bucket_size5e8Same trade for the stage 1 and 2 weight all-gather.
stage3_prefetch_bucket_size5e8How many elements of upcoming layers to gather ahead of use. The main stage 3 throughput knob.
stage3_param_persistence_threshold1e5Parameters smaller than this stay unsharded on every rank; avoids many tiny gathers for biases and norms.
stage3_max_live_parameters1e9Upper bound on gathered parameters resident at once. Lower it if stage 3 runs out of memory mid-step.
stage3_max_reuse_distance1e9Keep a gathered parameter if it will be reused within this many parameters, instead of re-gathering.
sub_group_size1e9Parameters per optimizer sub-step; matters mostly with offload, where it bounds staging buffers.

Tune in this order. First get the job to fit, using stage and activation checkpointing. Then turn on overlap_comm and check that memory still fits. Then, for stage 3 only, raise stage3_prefetch_bucket_size until GPU utilisation stops rising or memory runs out.

Model code at stage 3: placeholders and gathered parameters

Stage 3 changes what a parameter is. Outside the forward and backward hooks DeepSpeed installs, each parameter tensor is an empty placeholder whose real data is spread across ranks. Code that reads weights directly, such as custom initialisation, weight tying, logging parameter norms or copying weights into another model, silently sees nothing or crashes on a shape mismatch.

The fix is to gather explicitly with deepspeed.zero.GatheredParameters. Passing modifier_rank tells DeepSpeed that one rank will change the tensor and that the change must be broadcast and re-partitioned when the context exits; without it, edits are discarded.

# Stage 3: outside forward/backward a parameter is a placeholder; its data lives in shards.
w = model.lm_head.weight
print(w.shape, w.ds_numel)        # e.g. torch.Size([0]) and the true element count

# Read or edit the full tensor: gather, act, and (with modifier_rank) re-partition the edit.
with deepspeed.zero.GatheredParameters(w, modifier_rank=0):
    if torch.distributed.get_rank() == 0:
        torch.nn.init.normal_(w, mean=0.0, std=0.02)

Two related rules. Parameters that are used outside the module that owns them, as in tied input and output embeddings, must be registered with DeepSpeed as external parameters or it will not gather them in time; recent versions detect common cases, but check the warning log. And zero.Init must wrap construction itself; otherwise every rank materialises the full model first.

Offload and ZeRO++

When sharding across all GPUs is still not enough, ZeRO can push state further down the memory hierarchy. offload_optimizer moves the fp32 master weights and Adam moments to CPU memory or NVMe and runs the update there; offload_param (stage 3 only) does the same for the 16-bit weights. Set device explicitly to cpu or nvme and keep pin_memory on, because unpinned host buffers make every transfer take a slow staging copy. Offload trades PCIe bandwidth for capacity; the budget is worked through in optimizer state offloading.

ZeRO++ attacks the other stage 3 cost, cross-node traffic. zero_quantized_weights quantises the weight all-gather, zero_quantized_gradients quantises the gradient reduce-scatter, and zero_hpz_partition_size keeps a secondary copy of the weights sharded only within a group of GPUs, usually set to the GPUs per node, so the backward all-gather stays on the fast intra-node links. All three default to off, help mainly on multi-node jobs, and should be validated against a baseline loss curve before a long run.

Checkpoints: sharded on the way in, consolidated on the way out

engine.save_checkpoint makes every rank write its own shard of model and optimizer state plus a small metadata file, so saving is parallel and needs no gather. The cost is that the checkpoint is laid out for a specific ZeRO configuration. Resuming on the same world size and stage just works with engine.load_checkpoint; changing the number of GPUs needs DeepSpeed's checkpoint conversion tooling, whose support depends on your version, so test a resume on the target shape before you need it.

For inference or sharing, export a plain state dict. With stage3_gather_16bit_weights_on_model_save enabled, engine.save_16bit_model gathers 16-bit weights at save time. For fp32 weights, DeepSpeed writes a zero_to_fp32.py script into the checkpoint directory, and the same function is importable:

# Offline, on one machine with enough CPU RAM, after training:
from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint

state = get_fp32_state_dict_from_zero_checkpoint("ckpt", tag="step40000")
torch.save(state, "model_fp32.pt")               # plain PyTorch state dict, no DeepSpeed needed

Failure modes seen in practice

  • Out of memory at the first step, not at load. Model state fit but buckets, prefetch and activations did not. Enable activation checkpointing, lower stage3_max_live_parameters and prefetch size, or reduce the micro-batch.
  • Stage 3 much slower than expected. Per-layer gathers are exposed. Check the micro-batch is not tiny, raise prefetch, confirm NCCL is using the fast interconnect, and compare against stage 2 at a smaller batch.
  • Hang at the first collective. Ranks disagree on model structure or one rank skipped a forward branch. Every rank must run the same sequence of gathered layers.
  • Weights look empty or zero in logging or evaluation code. Stage 3 placeholders; gather first.
  • Resume fails after changing GPU count. Sharded checkpoints are tied to layout; convert first.

Trade-offs against the alternatives

ZeRO and PyTorch FSDP implement the same idea; FSDP's full sharding corresponds to stage 3 and its gradient-and-optimizer sharding mode to stage 2. Choose DeepSpeed ZeRO when you want its offload tiers, ZeRO++ or an existing DeepSpeed stack; choose FSDP when you want to stay inside native PyTorch APIs. Both are data parallelism: every rank still processes whole layers, so a single layer too large for one GPU, or a batch too small to amortise communication, calls for tensor or pipeline parallelism, with ZeRO stage 1 or 2 layered across the data-parallel dimension.

What to do next

  1. Compute model-state bytes per rank for stages 1, 2 and 3 with the 2 + 2 + 12 rule, then add activations and 10 to 20 percent headroom.
  2. Start at the lowest stage that fits, usually stage 2, with bf16 and activation checkpointing.
  3. Write the config with every key you rely on set explicitly, including overlap_comm and any offload device.
  4. For stage 3, build the model inside zero.Init and audit code that touches weights directly for GatheredParameters.
  5. Measure tokens per second and peak memory for stage 2 and stage 3 on the real sequence length before committing to a long run.
  6. Save a checkpoint in the first hour, resume from it, and export it with zero_to_fp32 to prove the whole path works.
Key takeaway: ZeRO removes the copies of optimizer state, gradients and weights that plain data parallelism keeps on every GPU. Stage 1 shards Adam state and stage 2 adds gradients at essentially no extra communication; stage 3 shards weights too and adds per-layer all-gathers that must be prefetched to stay fast. Size the job with the 2 + 2 + 12 bytes rule plus activations and buffers, start at the lowest stage that fits, set every DeepSpeed key explicitly, gather parameters before touching them at stage 3, and prove checkpoint, resume and export early.