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:

  1. Rank 0 computes global chunk sizes and broadcasts the distribution plan to all ranks using dist.broadcast.
  2. 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.scatter to 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_gather to 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 Nd devices, reducing per-GPU memory from O(Ψ) to O(Ψ/Nd).
  • On-demand fetching via fetch_params() and dist.all_gather reconstructs 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →