Part of Language AI Handbook
Explains how ZeRO eliminates memory redundancy in distributed training by partitioning optimizer states, gradients, and parameters.
Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.
Article links
Make inline references clickable
ZeRO Optimization: Optimizer, Gradient, and Parameter Partitioning
Modern language models have crossed a threshold where no single GPU can hold all training state in memory. GPT-3 with 175 billion parameters requires roughly 700 GB of memory at 32-bit precision just to store its weights. Add the optimizer states and gradients that training requires, and you need several terabytes. A single A100 with 80 GB of HBM cannot hold even the parameters alone.
The standard response has been to distribute training across many GPUs. As we covered in the data parallelism chapter, each device holds a full model replica and processes a different data shard, then synchronizes gradients. This works well for moderate model sizes, but it scales memory requirements in a fundamentally wasteful way: 100 GPUs each holding a full copy of a large model means 99 copies of every tensor are pure redundancy. Adding more hardware does not reduce the per-device memory burden, it simply multiplies the number of devices carrying the same redundant data.
ZeRO (Zero Redundancy Optimizer) was introduced by Rajbhandari et al. at Microsoft Research in 2019 and implemented in the DeepSpeed library. Its central insight is that the redundancy in data-parallel training is not just wasteful, it is unnecessary. The parameters, gradients, and optimizer states held by each GPU are identical across all replicas because every device applies the same update to the same model. ZeRO eliminates this redundancy by partitioning each of these components across the fleet of GPUs, so each device owns a unique shard rather than a full copy.
This might seem to require each GPU to constantly request tensors from other devices, creating a communication bottleneck. ZeRO's elegance is that it achieves this partitioning while keeping communication volume equal to or only modestly above standard data parallelism. The three stages of ZeRO progressively partition more state while following a carefully designed communication schedule that hides most of the cost.
This chapter covers how ZeRO works mechanically through its three stages, the communication patterns each stage introduces, and the quantitative memory savings they achieve. We also examine ZeRO-Offload, ZeRO-Infinity, and ZeRO++, which extend the core idea to CPU memory, NVMe storage, and quantized communication respectively.
A family of memory optimization techniques for distributed training that eliminates redundancy by partitioning model states (optimizer states, gradients, and parameters) across data-parallel processes rather than replicating them. ZeRO is implemented in the DeepSpeed library and is mathematically equivalent to standard data parallelism.
The Memory Anatomy of Training
Before understanding ZeRO, you need to understand precisely what occupies GPU memory during training. The memory footprint has two main categories: model states and residual states.
Model States
Model states are the tensors directly associated with the model's learned parameters. These are the most important category because their size is fully determined by model architecture, and they grow linearly with parameter count.
The three components of model state are:
- Parameters (): The model weights themselves. At 16-bit (fp16 or bf16), each parameter takes 2 bytes. A 7B-parameter model needs 14 GB just for parameters.
- Gradients: One gradient tensor per parameter, matching the size of the parameters. Another 14 GB for a 7B model. These are computed during the backward pass and must persist until the optimizer step completes.
- Optimizer states: The most expensive component when using Adam. Adam maintains two additional tensors per parameter: the first moment (the exponential moving average of gradients) and the second moment (the exponential moving average of squared gradients). Both are stored in 32-bit precision to maintain numerical stability across many update steps. This is bytes per parameter, or 56 GB for a 7B model. Mixed-precision training also requires a 32-bit master copy of the parameters, adding another 4 bytes per parameter or 28 GB.
The reason Adam's moments are stored at 32-bit rather than 16-bit precision is subtle but important. The moment estimates accumulate small incremental updates over potentially millions of training steps. At 16-bit precision, the dynamic range is insufficient to represent the ratio between large and small gradient contributions without losing the small ones to rounding. The 32-bit master copy is similarly necessary because the optimizer step computes a small parameter delta, and adding a small 16-bit delta to a 16-bit parameter can lose precision at a critical scale where training dynamics matter.
Residual States
Residual states include activations, temporary buffers, and memory fragmentation. Activations are the intermediate outputs computed during the forward pass that must be retained for computing gradients during the backward pass. For a transformer with a sequence length of 2048 and a batch size of 4, the activation memory per layer can reach several gigabytes, and activations must be stored for all layers simultaneously. Gradient checkpointing (activation recomputation) trades activation memory for extra compute by discarding most activations and recomputing them during the backward pass. Since ZeRO targets model states specifically, it complements rather than replaces gradient checkpointing.
The Full Accounting
For mixed-precision training with Adam, the breakdown for a model with parameters is:
- 16-bit parameters: bytes
- 16-bit gradients: bytes
- 32-bit parameter copy for optimizer update: bytes
- 32-bit Adam first moment: bytes
- 32-bit Adam second moment: bytes
This totals to bytes per model. For a 7B model, that is 112 GB. For a 70B model, 1.12 TB. No single GPU can hold this, and in naive data parallelism, every GPU in the fleet holds this full amount. A training cluster with 256 A100s running a 7B model under standard data parallelism stores of model state, of which is pure redundancy.
def compute_training_memory(params_billions, precision_bits=16):
"""
Compute per-GPU memory requirements for mixed-precision Adam training.
Returns values in gigabytes.
"""
params = params_billions * 1e9
# Mixed-precision training state
fp16_params = params * 2 # 2 bytes per fp16 param
fp16_grads = params * 2 # 2 bytes per fp16 gradient
fp32_params = params * 4 # 32-bit master copy for optimizer
fp32_m = params * 4 # Adam first moment
fp32_v = params * 4 # Adam second moment
total_bytes = fp16_params + fp16_grads + fp32_params + fp32_m + fp32_v
total_gb = total_bytes / 1e9
return {
"fp16_params_gb": fp16_params / 1e9,
"fp16_grads_gb": fp16_grads / 1e9,
"fp32_master_gb": fp32_params / 1e9,
"fp32_m_gb": fp32_m / 1e9,
"fp32_v_gb": fp32_v / 1e9,
"total_gb": total_gb,
}
model_sizes = [1, 7, 13, 30, 70]
memory_profiles = {size: compute_training_memory(size) for size in model_sizes} Model Params Grads FP32 Master Adam m Adam v Total
------------------------------------------------------------------------
1B 2.0G 2.0G 4.0G 4.0G 4.0G 16.0G
7B 14.0G 14.0G 28.0G 28.0G 28.0G 112.0G
13B 26.0G 26.0G 52.0G 52.0G 52.0G 208.0G
30B 60.0G 60.0G 120.0G 120.0G 120.0G 480.0G
70B 140.0G 140.0G 280.0G 280.0G 280.0G 1120.0GThis table shows the per-GPU memory requirement under naive data parallelism, where each device holds a complete copy of all training state. Every row represents memory that is perfectly duplicated across every GPU in the training cluster, consuming resources that contribute zero additional model capacity. Notice that the Adam optimizer states (FP32 Master + Adam m + Adam v) account for bytes out of the total , making them the dominant memory component and the natural first target for ZeRO.
ZeRO Stage 1: Optimizer State Partitioning
ZeRO Stage 1 targets the largest single component of model state: the Adam optimizer states. In standard data parallelism, all GPUs hold identical copies of these tensors even though each GPU applies the optimizer update to the full set of parameters. Every GPU ends up doing the same work and storing the same state as all others. Stage 1 eliminates this by assigning each GPU responsibility for updating a shard of the parameters, so it only needs to store the optimizer states for that shard.
How the Partitioning Works
Before the training step begins, parameters are replicated across all GPUs as usual. The forward pass and backward pass proceed identically to standard data parallelism: each GPU processes its local data shard, computes a loss, and runs backpropagation to produce gradient tensors. Up to this point, Stage 1 is indistinguishable from baseline data parallelism.
After the backward pass, Stage 1 diverges. Instead of each GPU independently applying the optimizer update to the full parameter set (which would require holding the full optimizer state), Stage 1 performs a Reduce-Scatter operation:
- Each GPU computes gradients for its local data shard as normal.
- A Reduce-Scatter operation is performed: the gradients are averaged across all GPUs, and each GPU ends up with the averaged gradient for only its assigned parameter shard, rather than the full averaged gradient tensor.
- Each GPU applies the optimizer update to its owned shard using only the optimizer states for that shard. Because each GPU owns a different shard, the GPUs collectively perform one complete optimizer step across all parameters without any single GPU needing the full optimizer state.
- An All-Gather operation distributes the updated parameters back to all GPUs, restoring the full parameter replica everywhere.
The forward pass of the next iteration uses the complete, freshly-updated parameters that every GPU now holds. The key observation is that the communication operations in steps 2 and 4 together have the same total volume as the All-Reduce in standard data parallelism. A Reduce-Scatter followed by an All-Gather is mathematically equivalent to an All-Reduce, just decomposed into two phases.
Why This Works: Communication Equivalence
Standard data parallelism uses an All-Reduce to average gradients across all GPUs. An All-Reduce sends and receives bytes per GPU (in the ring-reduce formulation, it sends bytes and receives the same, which approximates to for large ). ZeRO Stage 1's Reduce-Scatter + All-Gather sends and receives the same total volume. The memory savings come for free in terms of communication cost.
This communication equivalence is what makes Stage 1 a nearly free win. You replace the optimizer's most expensive memory component with a sharded version that costs nothing extra in network traffic. The only overhead is that the Reduce-Scatter and All-Gather must be implemented correctly and potentially use different collective implementations than a standard All-Reduce, but modern communication libraries handle this transparently.
Memory Savings
The optimizer states (first moment, second moment, and 32-bit parameter copy) account for bytes in mixed-precision Adam. Stage 1 reduces this to bytes distributed across GPUs, so each GPU holds:
where:
- : 16-bit parameter replica (still a full copy, needed for forward and backward passes)
- : 16-bit gradient replica (still a full copy before Reduce-Scatter)
- : optimizer states for the owned shard only
For GPUs, the optimizer state contribution drops from 84 GB per GPU to about 1.31 GB per GPU for a 7B model. The total memory per GPU drops from 112 GB to roughly 29.31 GB. Stage 1 alone can turn an infeasible training job into a feasible one, at essentially no increase in communication cost compared to standard data parallelism. For many practical training runs on relatively modern hardware, Stage 1 is sufficient to fit the desired model.
def zero_stage1_memory(params_billions, num_gpus):
"""Memory per GPU under ZeRO Stage 1 (optimizer state partitioning)."""
params = params_billions * 1e9
# fp16 params + fp16 grads still fully replicated
replicated = (params * 2 + params * 2) / 1e9
# optimizer states partitioned across num_gpus
optimizer_per_gpu = (params * 4 + params * 4 + params * 4) / num_gpus / 1e9
return replicated + optimizer_per_gpu
def zero_stage2_memory(params_billions, num_gpus):
"""Memory per GPU under ZeRO Stage 2 (optimizer state + gradient partitioning)."""
params = params_billions * 1e9
# only fp16 params still fully replicated
replicated = (params * 2) / 1e9
# gradients + optimizer states partitioned
partitioned_per_gpu = (
(params * 2 + params * 4 + params * 4 + params * 4) / num_gpus / 1e9
)
return replicated + partitioned_per_gpu
def zero_stage3_memory(params_billions, num_gpus):
"""Memory per GPU under ZeRO Stage 3 (full partitioning of all model states)."""
params = params_billions * 1e9
# everything partitioned
partitioned_per_gpu = (
(params * 2 + params * 2 + params * 4 + params * 4 + params * 4)
/ num_gpus
/ 1e9
)
return partitioned_per_gpu
# Baseline: no ZeRO (standard data parallelism)
def baseline_memory(params_billions):
return memory_profiles[params_billions]["total_gb"]
gpu_counts = [1, 4, 8, 16, 32, 64, 128]
model_7b_memories = {
"Baseline": [baseline_memory(7)] * len(gpu_counts),
"Stage 1": [zero_stage1_memory(7, nd) for nd in gpu_counts],
"Stage 2": [zero_stage2_memory(7, nd) for nd in gpu_counts],
"Stage 3": [zero_stage3_memory(7, nd) for nd in gpu_counts],
}Per-GPU memory (GB) for a 7B parameter model
GPUs Baseline Stage 1 Stage 2 Stage 3
----------------------------------------------------
1 112.0G 112.0G 112.0G 112.0G
4 112.0G 49.0G 38.5G 28.0G
8 112.0G 38.5G 26.2G 14.0G
16 112.0G 33.2G 20.1G 7.0G
32 112.0G 30.6G 17.1G 3.5G
64 112.0G 29.3G 15.5G 1.8G
128 112.0G 28.7G 14.8G 0.9GThe progression shows how each stage supports larger models at the same hardware budget. Stage 3 with 64 GPUs reduces per-GPU memory from 112 GB to 1.75 GB for a 7B model, leaving ample headroom for activations and residuals.
ZeRO Stage 2: Gradient Partitioning
ZeRO Stage 2 extends Stage 1 by also partitioning the gradient tensors. In Stage 1, gradients are still replicated in full before the Reduce-Scatter operation: each GPU holds the complete averaged gradient tensor until after the optimizer step. Stage 2 observes that this full replication is also unnecessary. After the Reduce-Scatter averaging, each GPU only needs the gradients corresponding to its owned parameter shard to perform the optimizer update. There is no reason to keep the other shards' gradients around.
The Reduce-Scatter Pattern
Stage 2 uses a Reduce-Scatter as the fundamental communication primitive after the backward pass. The distinction from Stage 1 lies in when the gradient memory is freed:
- Each GPU starts with its local gradient tensors computed from its data shard.
- The Reduce-Scatter operation averages the gradients across all GPUs and simultaneously distributes them so that GPU ends up with the averaged gradient for parameter shard only.
- Each GPU runs the optimizer step on its shard using the received gradient shard and its stored optimizer state shard.
- An All-Gather restores the full updated parameters everywhere.
The key difference from Stage 1 is that between the Reduce-Scatter and the All-Gather, the full gradient tensor no longer exists on any single GPU. Each GPU holds only of the total gradient. This saves bytes per GPU compared to Stage 1. The communication volume is unchanged: the same Reduce-Scatter and All-Gather operations occur, just with the gradient memory freed earlier.
Why Gradients Can Be Sharded
The reason gradient sharding works correctly is the same reason Stage 1 works: after the optimizer step, each GPU reconstructs the full parameters via All-Gather. The gradient for shard is only ever needed by GPU (to update its owned parameters), so there is no reason any other GPU needs to retain it. The gradient is a means to an end, and once the optimizer step has consumed it, it can be discarded.
This also has a practical benefit for implementations: the gradient tensor for shard can be freed as soon as the Reduce-Scatter for that shard completes, even before all shards have finished their Reduce-Scatter. In practice, implementations bucket the Reduce-Scatter operations so that gradient memory is released progressively as each bucket's communication completes, reducing peak memory below the formula's worst case.
Memory Per GPU Under Stage 2
where:
- : 16-bit parameter replica (still a full copy, needed for forward and backward passes)
- : gradient shard plus optimizer state shard for the owned portion only
With GPUs and a 7B model, Stage 2 per-GPU memory is approximately 14 GB (parameters) plus 0.9 GB (sharded gradients and optimizer states), totaling roughly 14.9 GB. This fits on a single A100 with room to spare, compared to 112 GB for the baseline. The memory savings relative to Stage 1 come entirely from eliminating the full gradient replica.
Gradient Accumulation Compatibility
An important practical consideration: Stage 2 is fully compatible with gradient accumulation. When accumulating gradients across multiple micro-batches before the optimizer step, the partial gradients are accumulated locally on each GPU without any cross-GPU communication. The Reduce-Scatter only fires when the full accumulation is complete and the optimizer step is about to begin. This means the memory profile and communication pattern are unchanged when using gradient accumulation, which is commonly used to simulate large effective batch sizes on hardware with limited memory.
Stage 2 is also compatible with mixed-precision training in the standard way. The forward and backward passes use fp16 parameters and gradients, and the optimizer step converts the sharded gradient to fp32 before updating the fp32 optimizer states and master parameters. The fp16 parameter replica is then updated from the fp32 master, and the All-Gather broadcasts the updated fp16 parameters to all devices.
ZeRO Stage 3: Parameter Partitioning
ZeRO Stage 3 completes the partitioning by distributing the parameters themselves across the GPUs. In Stages 1 and 2, each GPU still holds a full 16-bit replica of all parameters for the forward and backward passes. This bytes is the last significant replicated component. Stage 3 eliminates it, reducing every component of model state to per GPU.
On-Demand Parameter Fetching
The challenge with partitioning parameters is that the forward and backward passes need parameters for every layer. If GPU only permanently stores a shard, how does it execute the forward pass through layers whose parameters live on other GPUs?
Stage 3 exploits the sequential structure of the forward pass: a transformer processes layers one at a time, and at any given moment, only the parameters for the current layer are actively needed. Stage 3 uses this property to fetch parameters on demand:
- Each GPU permanently owns only of the total parameters, distributed across all layers.
- Before layer 's computation begins, all GPUs participate in an All-Gather to collect the complete parameter set for layer from whichever GPUs own those shards.
- The computation for layer proceeds normally, using the fully gathered parameters.
- Once layer 's contribution to the activations (forward) or gradients (backward) has been computed, the gathered parameters are discarded. The GPU reverts to holding only its permanent shard.
The backward pass mirrors this: backpropagation traverses layers in reverse, fetching each layer's parameters via All-Gather before computing gradients, then immediately discarding the fetched parameters. A Reduce-Scatter ships the local gradient contribution to the GPU that owns each parameter shard, which then performs the optimizer update.
This "just-in-time" parameter fetching means that at any moment, a GPU holds:
- Its permanently-owned parameter shard, which is of all parameters
- The parameters for the layer currently being computed, fetched transiently via All-Gather
- Its permanently-owned gradient shard and optimizer state shard
The transient parameter buffer for the current layer is bounded by the size of the largest single layer, which is typically much smaller than . For a transformer model, the largest layer (usually the MLP FFN) has roughly parameters. For a 7B model with , that is about parameters, or 0.13 GB at fp16, which is negligible compared to the savings.
Memory Per GPU Under Stage 3
When all components are partitioned, the per-GPU memory for model state becomes:
where is the total model state (fp16 params + fp16 grads + fp32 master + fp32 Adam moments). Every component scales with . This is a perfect linear reduction in per-GPU memory with the number of GPUs.
For a 7B model with 64 GPUs:
Stage 3 makes it theoretically possible to train arbitrarily large models given enough GPUs, since the memory per GPU decreases linearly with . Doubling the GPU count halves the per-GPU model state memory. This is a fundamentally different scaling property than Stages 1 and 2, where the parameter replica creates a floor that does not decrease with more GPUs.
The Communication Cost
Stage 3's memory efficiency comes at a communication price. The total communication volume relative to standard data parallelism is:
- Baseline (Standard DP): 1 All-Reduce equivalent, total volume bytes per step
- Stage 1: 1 All-Reduce equivalent, total volume bytes (Reduce-Scatter + All-Gather for gradients/parameters)
- Stage 2: 1 All-Reduce equivalent, total volume bytes (same operations, gradient memory freed earlier)
- Stage 3: 3 All-Reduce equivalents, total volume bytes per step
Stage 3's forward pass requires one All-Gather to fetch parameters layer by layer (total volume ). The backward pass requires a second All-Gather for parameters (another ) plus a Reduce-Scatter for gradients (). The full step therefore transfers bytes versus for the baseline. This 3x communication overhead is manageable on high-bandwidth interconnects such as NVLink (600 GB/s bidirectional on A100 nodes) or InfiniBand (200-400 Gb/s between nodes), but can become a significant bottleneck on slower networks.
The communication-compute overlap is critical to Stage 3's practical performance. Modern implementations prefetch the parameters for the next layer while computing the current layer, so the All-Gather for layer proceeds in parallel with the computation for layer . If the network is fast enough to transfer a layer's parameters before the GPU finishes computing with the previous layer, the communication overhead is effectively hidden. This prefetching only works for sequential architectures where the execution order is predictable at scheduling time.
def communication_volume(params_billions, stage, num_gpus):
"""
Estimated total communication volume per training step in gigabytes.
Stage 1 and 2 match standard DP all-reduce volume.
Stage 3 triples it due to forward/backward parameter gathers.
"""
params = params_billions * 1e9
bytes_per_param = 2 # fp16
if stage in (1, 2):
# Equivalent to one all-reduce: 2 * params * bytes_per_param
volume = 2 * params * bytes_per_param
else:
# Stage 3: forward gather + backward gather + backward reduce-scatter
volume = 3 * 2 * params * bytes_per_param
return volume / 1e9
stages = [1, 2, 3]
gpu_counts_comm = [8, 16, 32, 64]
params_7b = 7
comm_table = {
stage: {
nd: communication_volume(params_7b, stage, nd) for nd in gpu_counts_comm
}
for stage in stages
}Communication volume per step (GB) for 7B parameters Stage 8GPUs 16GPUs 32GPUs 64GPUs ------------------------------------------------ Stage 1 28.0G 28.0G 28.0G 28.0G Stage 2 28.0G 28.0G 28.0G 28.0G Stage 3 84.0G 84.0G 84.0G 84.0G
Notice that the communication volume for Stages 1 and 2 is independent of the number of GPUs in this model, because the total data transferred in a ring-based all-reduce scales as for large . Stage 3's volume is also constant across GPU counts here because all three phases each transfer bytes in total, regardless of how many GPUs share the work. The communication volume for Stage 3 matches standard data parallelism practitioners already accept, just tripled. This is why it requires fast interconnects to achieve good hardware utilization.
Memory Savings Across Stages
Let's bring all three stages together and visualize their memory savings as a function of GPU count.

The A100 80 GB line marks the feasibility boundary for an important class of hardware. The baseline never dips below 112 GB, making a 7B model impossible to train on A100s under standard data parallelism regardless of how many GPUs you add. All three ZeRO stages cross below the feasibility line at 2 GPUs in this model, but they diverge sharply as the GPU count grows. The curve for Stage 3 is a true hyperbola (linear reduction), while Stages 1 and 2 flatten out as the unpartitioned components dominate at high GPU counts.

At 64 GPUs, Stage 3 fits a 70B model into 17.5 GB per GPU, well within even older GPU hardware with 24 or 32 GB. The baseline requires 1.12 TB per GPU for a 70B model, which is impossible on any current GPU. Stage 1 reduces that footprint to about 293 GB per GPU, so parameter and gradient replication still make it infeasible. This visualization illustrates why the choice of ZeRO stage has practical consequences: Stage 1 makes moderate models practical, while Stage 3 gives you access to a completely different tier of model sizes.
ZeRO-Offload and ZeRO-Infinity
The original three ZeRO stages partition state across GPU memory. Two extensions push this further by moving state off the GPU entirely, exploiting the memory hierarchy of modern servers.
ZeRO-Offload
ZeRO-Offload extends Stage 2 by offloading optimizer states and the optimizer step computation to CPU memory and CPU compute. The central observation is that the optimizer step, the Adam update computing new parameter values from gradients and moments, is a simple element-wise operation that does not require the GPU's massive parallelism. A CPU can execute it efficiently at lower cost. CPU RAM is also far more abundant: a modern server might have 1-2 TB of DRAM while the attached GPUs have 80-160 GB total.
The challenge with naively offloading to CPU is that the data transfer between GPU and CPU over PCIe is slow, typically 16-32 GB/s bidirectional, versus 2-4 TB/s for GPU HBM. ZeRO-Offload addresses this through pipelining: while the GPU runs the forward and backward passes for training step , the CPU runs the optimizer step for step . If the optimizer step takes longer than the GPU computation, there is a wait at the end of each step, but for large models where the optimizer state is large but the per-parameter update is cheap, the CPU work finishes before the GPU needs the updated parameters.
ZeRO-Offload reduces GPU memory requirements by the full optimizer state bytes, replacing them with CPU RAM. This allows training larger models on a single GPU or small cluster than would otherwise be possible. For a single A100 (80 GB), ZeRO-Offload can train models significantly larger than what fits in 80 GB of GPU HBM alone.
ZeRO-Infinity
ZeRO-Infinity extends the storage hierarchy further, to NVMe SSD storage. Modern NVMe drives offer 4-8 GB/s of sequential bandwidth and multiple terabytes of capacity. ZeRO-Infinity uses a three-tier heterogeneous memory hierarchy: NVMe for the largest tensors (optimizer states and inactive parameters), CPU RAM as a staging buffer for data moving between NVMe and GPU, and GPU HBM for active computation.
The key algorithmic challenge for ZeRO-Infinity is bandwidth management. GPU HBM operates at 2-4 TB/s; CPU RAM at 50-100 GB/s; NVMe at 4-8 GB/s. The system must carefully schedule which tensors live at which tier and when to prefetch data up the hierarchy to avoid stalling the GPU. ZeRO-Infinity uses a bandwidth-optimal data movement engine that analyzes the compute graph to schedule prefetches with enough lead time that the GPU is never blocked waiting for data from NVMe.
ZeRO-Infinity can train models with trillions of parameters on hardware configurations that would otherwise require a datacenter-scale GPU cluster. The original paper demonstrated training a 1 trillion-parameter model on a single DGX-2 node (16 V100 GPUs with 32 GB each), a configuration with roughly 0.5 TB of total GPU memory. Without ZeRO-Infinity, this would be impossible.
The tradeoff in both offloading approaches is latency and throughput. NVMe access is orders of magnitude slower than HBM, and the achievable training throughput depends critically on how well the prefetching hides the storage latency. In practice, ZeRO-Infinity achieves good throughput for compute-bound workloads (large models with large batch sizes) but can stall on memory-bound operations or when storage bandwidth becomes the bottleneck.
ZeRO++ and Communication Efficiency
ZeRO++ (published in 2023) addresses Stage 3's 3x communication overhead through three complementary improvements. Each targets a specific phase of the training step and uses different compression techniques.
Quantized Weights (qwZ)
The All-Gather during the forward pass communicates the parameters for each layer before computing with them. ZeRO++ compresses these parameters from 16-bit to 8-bit before transmission, halving the communication volume for this phase. After reception, each GPU dequantizes the parameters back to 16-bit (or bf16) before computation. The dequantization step takes negligible time compared to the actual layer computation.
The quantization uses a block-wise scheme where parameters are divided into small blocks, each with its own scale factor. This limits the quantization error relative to a global scale factor scheme, preserving model accuracy during the forward pass. The gradients computed from quantized parameters accumulate small errors, but these errors are comparable in magnitude to the noise already present from stochastic gradient descent and mini-batch sampling, so training stability is maintained.
Hierarchical Parameter Partitioning (hpZ)
Standard ZeRO Stage 3 distributes parameters uniformly across all GPUs, regardless of physical topology. This means that when a layer's All-Gather is performed, parameters must travel across slow inter-node InfiniBand links for any parameters owned by GPUs on other nodes.
ZeRO++ introduces hierarchical partitioning that exploits the two-level topology of modern clusters. Within a single node, NVLink provides 600 GB/s of bandwidth. Between nodes, InfiniBand provides 200-400 Gb/s. Under hpZ, parameters are partitioned only within each node: each GPU owns of the parameters within its node, but every node holds a complete copy. The forward-pass All-Gather is then a within-node operation using NVLink, avoiding the slow inter-node link entirely.
This does reintroduce some parameter replication (each node stores a complete copy), which somewhat reduces the memory savings of Stage 3. However, on multi-node training runs where inter-node bandwidth is the bottleneck, hpZ can dramatically improve throughput. The memory tradeoff is often acceptable because within-node NVLink bandwidth is used much more efficiently, and the remaining memory pressure can be addressed through other means.
Quantized Gradients (qgZ)
The Reduce-Scatter for gradients transfers bytes per step. ZeRO++ applies 1-bit quantization to these gradients before the Reduce-Scatter, reducing the gradient communication by up to 16x (from 16-bit to 1-bit). The 1-bit quantization uses the same error-feedback mechanism as 1-bit Adam: quantization errors are accumulated and added to the next gradient, so no information is permanently lost, just deferred.
This is the most aggressive of the three ZeRO++ optimizations and the one with the most potential for convergence impact. In practice, the gradient quantization tends to work well when the gradient signal is strong (early training) and can be problematic when gradients become very small (fine-grained fine-tuning). The qgZ component is typically used selectively, activated during phases where gradient magnitudes are large.
Together, qwZ, hpZ, and qgZ can reduce Stage 3's communication overhead by 4x or more, largely closing the gap with Stages 1 and 2 in terms of achieved hardware utilization. ZeRO++ makes Stage 3 practical on a wider range of cluster configurations, particularly those with slower inter-node interconnects.
Worked Example: Training State Partitioning
Let's trace through a concrete Stage 2 training step with GPUs and a toy model to make the mechanics concrete. This example uses a simplified SGD optimizer instead of Adam to keep the arithmetic clear, but the partitioning logic is identical.
import numpy as np
# Toy model: 8 parameters total, each GPU owns 2 parameters
N_d = 4
total_params = 8
params_per_gpu = total_params // N_d
# Initial parameters (same on all GPUs - full replica before forward pass)
np.random.seed(42)
params_global = np.array(
[0.1, -0.3, 0.5, 0.2, -0.1, 0.4, -0.2, 0.3], dtype=np.float32
)
# Each GPU processes different data, producing different local gradients
# (In reality, backprop on different mini-batches)
local_grads = {
0: np.array(
[0.04, -0.02, 0.06, 0.01, -0.03, 0.05, -0.01, 0.02], dtype=np.float32
),
1: np.array(
[0.02, -0.04, 0.03, 0.05, -0.01, 0.02, -0.04, 0.03], dtype=np.float32
),
2: np.array(
[0.06, -0.01, 0.04, 0.02, -0.02, 0.06, -0.02, 0.01], dtype=np.float32
),
3: np.array(
[0.02, -0.03, 0.05, 0.04, -0.04, 0.01, -0.03, 0.04], dtype=np.float32
),
}
# Step 1: Average gradients across GPUs (what Reduce-Scatter computes as average)
averaged_grads = np.mean([local_grads[i] for i in range(N_d)], axis=0)
# Step 2: Reduce-Scatter - each GPU receives the averaged gradient for its shard
owned_grad_shards = {
i: averaged_grads[i * params_per_gpu : (i + 1) * params_per_gpu]
for i in range(N_d)
}
# Step 3: Each GPU applies optimizer update to its owned parameter shard
# Simple SGD update (lr=0.1) for illustration
lr = 0.1
owned_param_shards = {
i: params_global[i * params_per_gpu : (i + 1) * params_per_gpu]
- lr * owned_grad_shards[i]
for i in range(N_d)
}
# Step 4: All-Gather - reconstruct full updated parameters on all GPUs
params_updated = np.concatenate([owned_param_shards[i] for i in range(N_d)])=== ZeRO Stage 2 Training Step (4 GPUs, 8 Parameters) === Initial parameters: [ 0.1 -0.3 0.5 0.2 -0.1 0.4 -0.2 0.3] Averaged gradients: [ 0.035 -0.025 0.045 0.03 -0.025 0.035 -0.025 0.025] After Reduce-Scatter - each GPU owns gradient shard: GPU 0 owns params [0, 1]: grad_shard = [ 0.035 -0.025] GPU 1 owns params [2, 3]: grad_shard = [0.045 0.03 ] GPU 2 owns params [4, 5]: grad_shard = [-0.025 0.035] GPU 3 owns params [6, 7]: grad_shard = [-0.025 0.025] After optimizer step on each shard: GPU 0: updated shard = [ 0.0965 -0.2975] GPU 1: updated shard = [0.4955 0.197 ] GPU 2: updated shard = [-0.0975 0.3965] GPU 3: updated shard = [-0.1975 0.2975] After All-Gather - full updated params: [ 0.0965 -0.2975 0.4955 0.197 -0.0975 0.3965 -0.1975 0.2975] Expected (naive full-param update): [ 0.0965 -0.2975 0.4955 0.197 -0.0975 0.3965 -0.1975 0.2975] Results match: True
This trace reveals something important: the final parameter values are identical to what a naive full-parameter update would produce. ZeRO is mathematically equivalent to standard data parallelism. The partitioning changes where computation happens and what tensors each GPU stores, but not the values that result. Each GPU performs the optimizer step only on its own shard, but because the averaged gradients are the same regardless of which GPU computes them, the resulting parameter updates are identical to a centralized update.
This mathematical equivalence is what makes ZeRO safe to use without any adjustment to hyperparameters, learning rate schedules, or convergence expectations. Training a model with ZeRO Stage 3 on 64 GPUs produces the same model as training it with Stage 1 on 4 GPUs or with standard data parallelism on a single GPU (given the same effective batch size and random seeds). ZeRO is purely an implementation optimization, not a change to the training algorithm.
Visualizing the Partitioning Structure
Let's visualize how the parameter shards, gradient shards, and optimizer state shards map to GPUs under each ZeRO stage, to build an intuitive picture of what each device is responsible for.

The diagram makes the partitioning structure concrete. Under Stage 1, every GPU holds the same full replica of parameters and gradients (shown in gray), but each GPU owns a distinct colored shard of the optimizer states. Stage 2 additionally shards the gradients. Stage 3 shards everything: each color appears in exactly one quarter of each row, and no two GPUs hold any overlapping permanent state. The All-Gather operations can be understood as temporarily merging the colors within a row before computation proceeds.
Practical Considerations
Choosing the Right Stage
The three stages present increasing memory savings at increasing communication cost. The right choice depends on your hardware configuration, model size, and network topology:
-
Stage 1: The recommended default for most training runs. It provides near-zero communication overhead increase over baseline data parallelism while delivering substantial optimizer state savings. For typical model sizes and reasonable GPU counts, Stage 1 often cuts the dominant memory component (optimizer states) by a factor of 8-64x, which is frequently enough to make a training job feasible. The implementation complexity is minimal because the communication pattern is nearly identical to standard data parallelism.
-
Stage 2: Suitable when optimizer state savings alone do not free enough memory, especially at small GPU counts where the term is still large. Stage 2 adds gradient partitioning with no increase in communication volume over Stage 1, making it a free upgrade in terms of communication cost. On a small cluster of 4-8 GPUs, Stage 2 may be the difference between fitting a model and not fitting it. The gradient bucketing implementation adds some complexity but is handled transparently by libraries like DeepSpeed and FSDP.
-
Stage 3: Use when Stage 2 memory is still insufficient, when training models that simply cannot fit in bytes regardless of how the optimizer states are handled, or when scaling to hundreds or thousands of GPUs where the communication overhead becomes proportionally smaller relative to compute time. At very large scale, the 3x communication volume of Stage 3 is spread over many GPUs, and the per-device communication fraction of total time decreases as the cluster grows.
Communication Topology Awareness
ZeRO's communication patterns are sensitive to network topology. Within a single node, NVLink provides 600 GB/s of bidirectional bandwidth on A100 DGX systems. Across nodes, InfiniBand typically offers 200-400 Gb/s per port. Stage 3 generates 3x more communication than Stages 1 and 2, which can become a bottleneck on inter-node links when the model is large and the nodes are connected by slower fabrics.
A practical heuristic: if you can fit the entire training run within a single node (using Stage 1 or Stage 2 and gradient checkpointing), the all-NVLink communication makes even Stage 3 fast. Once you go multi-node, the inter-node bandwidth becomes the governing constraint, and you should prefer the lowest ZeRO stage that fits the memory budget. ZeRO++ hierarchical partitioning (hpZ) is specifically designed to address this by keeping as much communication as possible within the fast NVLink domain.
Interaction with Gradient Checkpointing
ZeRO can be combined with gradient checkpointing (activation recomputation) and the two optimizations target entirely different memory categories. ZeRO reduces model state memory (parameters, gradients, optimizer states), while gradient checkpointing reduces activation memory by not storing intermediate activations and recomputing them during the backward pass at the cost of extra forward computation (roughly 33% more compute for transformers under full checkpointing).
Because they target different memory categories, they combine additively: using both together reduces both model state memory and activation memory simultaneously without any interference. For very large models with long sequences, both optimizations are typically needed at the same time. The combination allows training with batch sizes that would be impossible using either technique alone.
Interaction with Tensor Parallelism
Tensor parallelism, covered in the preceding chapter, splits individual layers across GPUs, so that each GPU holds a horizontal slice of each weight matrix. ZeRO Stage 3 instead splits parameters across the depth dimension (each GPU owns a contiguous set of parameters from all layers). Combining the two creates a more complex partitioning scheme where the same parameter may be split in two dimensions simultaneously.
In practice, tensor parallelism and ZeRO Stage 3 are rarely combined because they both target the same problem (parameter memory) and their combination creates awkward communication schedules. A more common pattern is to use tensor parallelism within a node (exploiting NVLink bandwidth) and ZeRO Stage 2 across nodes (low communication overhead), giving each level of the hierarchy a specialization. The Megatron-DeepSpeed integration supports this hierarchical combination.
Performance Benchmarks in Context
The original ZeRO paper reported training a 100B-parameter model on 400 V100 GPUs (32 GB each) with Stage 3, achieving over 38 teraflops per GPU, roughly 49% of the theoretical peak. This was the first demonstration of 100B-scale training outside of specialized supercomputing facilities. Throughput on these large-scale jobs was within 5% of the theoretical communication-compute overlap optimum, meaning ZeRO adds essentially no wall-clock overhead compared to baseline on well-provisioned hardware.
Subsequent work on LLaMA, Falcon, and Mistral training reported similar results: Stage 2 or Stage 3 with gradient checkpointing provides the right memory-compute tradeoff for models in the 7B-70B range on A100 clusters with NVLink within nodes and InfiniBand between nodes. The typical throughput efficiency is 45-55% of theoretical peak after accounting for communication, memory bandwidth constraints, and the overhead of mixed-precision operations.
Let's compare illustrative throughput scaling across stages by modeling compute and communication times:
def estimate_throughput(
params_billions,
num_gpus,
stage,
bandwidth_gbps=200,
exposed_communication_fraction=0.08,
):
"""
Simplified throughput model.
Models weak scaling: every GPU performs one unit of compute per step.
Communication time depends on volume, bandwidth, and overlap.
Returns relative throughput vs single-GPU baseline.
"""
compute_time_per_gpu = 1.0
# Ring collectives approach their asymptotic communication volume as N grows.
comm_gb = communication_volume(params_billions, stage, num_gpus)
ring_fraction = (num_gpus - 1) / num_gpus
bandwidth_effective = bandwidth_gbps / 8 # Gbps -> GB/s
raw_comm_time = ring_fraction * comm_gb / bandwidth_effective
# Most communication overlaps with layer computation. Only the exposed
# fraction extends the step time.
exposed_comm_time = raw_comm_time * exposed_communication_fraction
step_time = compute_time_per_gpu + exposed_comm_time
return num_gpus * compute_time_per_gpu / step_time
gpu_counts_perf = [8, 16, 32, 64, 128]
stages_perf = [1, 2, 3]
perf_data = {
stage: [estimate_throughput(7, nd, stage) for nd in gpu_counts_perf]
for stage in stages_perf
}
Stages 1 and 2 track near-ideal linear throughput scaling because their communication volume matches standard data parallelism, which practitioners already accept as overhead. Stage 3 falls below ideal due to 3x communication, but still provides substantial absolute speedup and enables training models that would otherwise require infeasible per-GPU memory. On faster interconnects, all three curves shift toward the ideal line, which is why newer clusters prioritize high network bandwidth alongside raw compute.
ZeRO in the PyTorch Ecosystem: FSDP
ZeRO's design philosophy was sufficiently influential that PyTorch incorporated a native implementation: Fully Sharded Data Parallel (FSDP), which implements ZeRO Stage 3 semantics directly in the PyTorch distributed training API. FSDP became part of PyTorch in version 1.11 and was substantially improved in subsequent releases.
FSDP's approach is slightly different from DeepSpeed's ZeRO Stage 3 in implementation, but the core idea is the same: each parameter is permanently owned by a subset of processes, and parameters are gathered on demand before computation. FSDP adds some features that make it easier to use within PyTorch's existing model definition patterns:
- Auto-wrapping: FSDP can automatically identify which submodules to shard based on a minimum parameter count threshold, without requiring manual annotation.
- Mixed sharding strategies: Different submodules can use different sharding strategies (full shard, hybrid shard, or no shard), allowing for hierarchical sharding that maps to the within-node/between-node topology.
- Activation checkpointing integration: FSDP integrates with PyTorch's native gradient checkpointing via
checkpoint_wrapper, making the combination straightforward.
FSDP is the standard tool for large-model training in PyTorch-native workflows, while DeepSpeed ZeRO remains dominant in environments that use the DeepSpeed infrastructure or require the ZeRO-Offload and ZeRO-Infinity extensions. The two systems have largely converged on similar feature sets, and the choice between them often comes down to ecosystem fit rather than fundamental capability differences.
Limitations and Impact
ZeRO has fundamentally changed what is achievable with commodity GPU clusters. Before DeepSpeed ZeRO, training models beyond a few billion parameters required specialized hardware (tensor and pipeline model parallelism requiring custom code, or proprietary clusters) that few organizations could access. ZeRO brought 100B+ model training within reach of researchers with access to a few hundred standard GPUs connected by InfiniBand, democratizing large-model research considerably.
Several practical limitations affect real-world deployments. The communication overhead of Stage 3 is manageable on NVLink clusters but becomes significant on Ethernet-interconnected nodes. A 100 Gb/s Ethernet cluster with Stage 3 may achieve only 50-60% hardware utilization where Stage 1 achieves 95%+. This communication-compute tradeoff is why many production training runs choose Stage 2 even when Stage 3 would technically fit the model into memory: the additional memory headroom from Stage 3 is sometimes not worth the throughput reduction.
ZeRO also interacts with other parallelism strategies in non-trivial ways. Pipeline parallelism partitions the model by layers, while ZeRO Stage 3 partitions parameters within every layer. Combining them requires careful bookkeeping to avoid double-partitioning or incomplete parameter shards during the pipeline's micro-batch execution. Tensor parallelism and ZeRO both address parameter memory, so using them together can lead to redundant communication. The DeepSpeed ZeRO++ extensions and Megatron-DeepSpeed integration provide working configurations for these combinations, but the interaction space is complex and benefits from careful tuning.
Another limitation is that ZeRO Stage 3's just-in-time parameter fetching assumes a predictable, sequential layer execution order. Non-sequential architectures, such as mixture-of-experts models where the routing mechanism may activate any expert in any order during forward pass execution, can generate communication patterns that are difficult to prefetch optimally. The ZeRO paper and DeepSpeed implementation include specialized handling for certain MoE configurations, but the general case remains more challenging than standard sequential architectures.
The memory savings from ZeRO also come at an implicit cost in debuggability. When a training run encounters a NaN gradient or a memory spike, understanding which GPU is responsible and what tensor caused the issue is harder when model state is spread across devices. Production training infrastructure typically augments ZeRO with careful logging of per-shard gradient norms and memory profiles to preserve visibility into training dynamics.
Despite these constraints, ZeRO's impact on the field has been decisive. GPT-NeoX, OPT, Falcon, Mistral, and many other open-source large models were trained using DeepSpeed ZeRO. The design philosophy, separating memory concerns from parallelism strategy and using communication-equivalent partitioning rather than model surgery, influenced subsequent systems including FSDP in PyTorch and the broader literature on memory-efficient distributed training. The central insight that redundancy is wasteful and unnecessary, and can be eliminated without any mathematical change to the training algorithm, remains one of the most practically impactful ideas in the infrastructure for large model training.
Summary
ZeRO eliminates the memory redundancy that makes naive data-parallel training wasteful at scale. Rather than replicating all training state across every GPU, it partitions the state so that each GPU permanently owns a shard and fetches the rest on demand. The total model state in mixed-precision Adam training is bytes, of which standard data parallelism replicates every byte across all devices. ZeRO progressively removes this redundancy layer by layer.
The three stages target progressively larger memory categories, each building on the previous:
- Stage 1 partitions optimizer states (Adam moments and 32-bit master parameters), reducing the largest memory component from to per GPU, at no extra communication cost over baseline data parallelism.
- Stage 2 additionally partitions gradients using Reduce-Scatter, saving another bytes per GPU with the same total communication volume as Stage 1 and baseline data parallelism.
- Stage 3 partitions the parameters themselves, achieving a perfect per-GPU memory footprint with 3x the communication of Stages 1 and 2. Each GPU fetches parameters on demand via All-Gather before each layer's computation.
Extensions like ZeRO-Offload and ZeRO-Infinity push state to CPU RAM and NVMe storage respectively, enabling models with trillions of parameters on practical hardware. ZeRO++ addresses Stage 3's communication overhead through quantized weight transmission, hierarchical partitioning that exploits within-node NVLink bandwidth, and quantized gradient communication.
The key practical insight is that ZeRO produces mathematically identical results to standard data parallelism. The partitioning is a pure implementation detail that changes where tensors live and what each GPU computes, without affecting the final parameter values or training dynamics. This correctness guarantee, combined with the linear memory reduction with GPU count under Stage 3, makes ZeRO the standard foundation for large-scale distributed training. FSDP in PyTorch brings these ideas natively into the PyTorch ecosystem, making ZeRO-style training accessible without requiring the full DeepSpeed infrastructure stack.
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about ZeRO Optimization.
ZeRO Optimization Quiz
Reference
Citation details
Cite or share this article.
Continue with the full handbook
This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.
Explore Language AI HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

Comments
No comments yet. Be the first to share your thoughts!