How Zero3 Enables Large Model Training on Limited GPU Memory: A Deep Dive into Sharded Data Parallelism
Zero3 reduces per-GPU memory usage from O(Ψ) to O(Ψ/Nd) by sharding parameters, gradients, and optimizer states across Nd devices, fetching only the active parameter shards on-demand during forward and backward passes.
Training models with hundreds of billions of parameters traditionally requires GPU clusters with massive individual memory capacity. According to the labmlai/annotated_deep_learning_paper_implementations repository, the Zero3 implementation (also referred to as Zero-DP) solves this limitation by partitioning every layer's tensors across all available devices and orchestrating precise memory cleanup between operations.
How Zero3 Shards Parameters Across GPUs
Partitioning Layer Parameters
When a Zero3Layer is instantiated in labml_nn/scaling/zero3/__init__.py, it immediately splits the layer's parameters into trainable and frozen groups. Each group is then divided into world_size shards—one per GPU process.
The initialization process follows this protocol:
- Rank 0 computes global chunk sizes and broadcasts the distribution plan to all ranks using
dist.broadcast. - Each rank allocates empty shards via
self._empty((s,)), sized exactly for its share of the parameters.
This occurs in Zero3Layer.__init__ (lines 81-130), ensuring that no single GPU ever holds the full parameter tensor.
Distributing Initial Weights with Padding
Before training begins, the full parameter tensors must be scattered to their respective shards. Rank 0 performs this by:
- Concatenating the full parameter tensors
- Padding the concatenated tensor so the total size is divisible by
world_size - Using
dist.scatterto distribute slices to each rank
Each rank stores only its assigned slice in self.chunk, as implemented in Zero3Layer.__init__ (lines 44-60). This initial distribution ensures that the aggregate memory across all devices holds the complete model state, while individual devices retain only a 1/Nd fraction.
On-Demand Parameter Fetching and Memory Cleanup
Dynamic Loading with fetch_params()
During the forward pass, layers do not keep their parameters resident in GPU memory. Instead, when computation reaches a Zero3Layer, it invokes fetch_params() (lines 31-63 in __init__.py). This method:
- Allocates a temporary buffer sized to
self.world_size * sum(self.chunk_size)on the current device - Executes
dist.all_gatherto collect all shards from every GPU into the buffer - Reconstructs the original parameter tensors from the gathered slices, restoring them to their original shapes for computation
This just-in-time loading ensures that parameters exist on the GPU only when actively needed for matrix operations.
Aggressive Memory Deallocation
Immediately after the forward computation completes, the layer invokes self._cleanup_params() to release memory. The cleanup routine (lines 6-14) performs a critical optimization:
- Records the current CUDA stream on each parameter tensor to ensure synchronization
- Resizes the underlying storage to 0 using
p.data.storage().resize_(0), which immediately frees the GPU memory without waiting for Python garbage collection
This aggressive deallocation strategy guarantees that memory usage remains bounded to the current layer's parameters rather than accumulating across the entire model depth.
Gradient Handling in Zero3
During the backward pass, Zero3 must reconcile gradients across the sharded parameters without materializing full tensors on any single device. The implementation achieves this through backward hooks registered on each trainable parameter:
- The hook captures gradients and writes them into a temporary buffer
- It then performs a reduce-scatter operation, distributing the aggregated gradients so each rank receives only the slice corresponding to its local parameter shard
- The reduced gradient slice is stored as the gradient of
self.chunk[TRAINING_PARAMS_IDX]
This mechanism, found in _backup_grads and related hooks (lines 104-149), allows the optimizer to update parameters using only local shard data, maintaining the O(Ψ/Nd) memory bound throughout the optimization step.
Pipeline Execution with Zero3Sequential
The Zero3Sequential wrapper (lines 49-96) orchestrates multiple Zero3Layer instances into a coordinated pipeline. It assigns each layer a dedicated CUDA fetch stream and a backup stream to overlap communication with computation.
During forward execution:
- The wrapper waits for the previous layer's gradient-backup to complete
- Triggers each layer's
fetch_params()just before the layer executes - Invokes
_cleanup_params()immediately after the forward pass returns
This sequential scheduling prevents parameter shards from multiple layers from coexisting simultaneously in GPU memory, further compressing the memory footprint during both training and inference.
Implementing Zero3 in Practice
Wrapping Individual Layers
To apply Zero3 to an existing model, wrap each layer in Zero3Layer and assemble them with Zero3Sequential:
from labml_nn.scaling.zero3 import Zero3Layer, Zero3Sequential
import torch.distributed as dist
modules = []
for layer in model_cfg.layers:
zl = Zero3Layer(
module=layer.to(device),
rank=dist.get_rank(),
world_size=dist.get_world_size(),
device=device,
dtype=torch.float16
)
modules.append(zl)
model = Zero3Sequential(modules)
The Zero3Layer constructor automatically handles parameter sharding, while Zero3Sequential manages the fetch and cleanup lifecycle.
Training with Sharded Optimizers
Use an optimizer designed for sharded parameters, such as Zero3Adam:
from labml_nn.optimizers import Zero3Adam
from labml import monit
optimizer = Zero3Adam(model.get_trainable_chunk(), lr=1e-4)
for batch in data_loader:
optimizer.zero_grad()
loss = model(batch['input']).mean()
loss.backward()
optimizer.step()
monit.step()
model.get_trainable_chunk() returns only the local trainable parameter shard required by the optimizer. The backward hooks inside Zero3Layer ensure that gradients are properly reduced and scattered before the optimizer step executes.
Using the Reference Implementation
The repository provides a complete fine-tuning script for GPT-NeoX:
python -m labml_nn.scaling.zero3.finetune_neox \
--model-config configs/gpt_neox.yaml \
--optimizer Zero3Adam \
--data-path /path/to/dataset
This script (labml_nn/scaling/zero3/finetune_neox.py) demonstrates the exact wiring of layers, streams, and optimizers required for large-scale training.
Summary
- Zero3 shards parameters, gradients, and optimizer states across
Nddevices, reducing per-GPU memory from O(Ψ) to O(Ψ/Nd). - On-demand fetching via
fetch_params()anddist.all_gatherreconstructs parameters only during active computation. - Aggressive cleanup using
storage().resize_(0)immediately frees memory after each layer completes. - Gradient reduce-scatter operations ensure backward passes maintain the sharded memory model without materializing full tensors.
- Zero3Sequential orchestrates layer execution with dedicated CUDA streams to prevent memory accumulation across layers.
Frequently Asked Questions
How much memory does Zero3 actually save compared to standard data parallelism?
Zero3 reduces per-device memory consumption proportional to the number of GPUs. While standard data parallelism replicates the full model (O(Ψ)) on every device, Zero3 maintains only O(Ψ/Nd) parameters per GPU, enabling models hundreds of billions of parameters large to train on modest clusters limited by aggregate memory rather than single-card capacity.
What is the performance overhead of fetching parameters on-demand?
The fetch_params() operation requires an all_gather communication round and temporary buffer allocation for each layer. However, Zero3Sequential mitigates this overhead by using dedicated CUDA fetch streams that overlap parameter gathering with computation from previous layers, keeping the critical path minimal.
Can Zero3 be used with mixed precision training?
Yes. The Zero3Layer constructor accepts a dtype parameter (commonly torch.float16 or torch.bfloat16), and the fetching mechanism reconstructs parameters in the specified precision. The sharding and cleanup logic operates independently of the tensor dtype, making it compatible with standard mixed precision training configurations.
How does Zero3 differ from ZeRO-2 (Zero-Offload)?
While ZeRO-2 typically shards only optimizer states and gradients, keeping full parameters on each GPU, Zero3 additionally shards the parameters themselves. This distinction allows Zero3 to train models larger than individual GPU memory, whereas ZeRO-2 is limited to models that fit in a single device's memory. The trade-off is increased communication volume due to the layer-wise all_gather operations required in Zero3.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →