# How Zero3 Enables Large Model Training on Limited GPU Memory: A Deep Dive into Sharded Data Parallelism

> Discover how Zero3 uses sharded data parallelism to train large models on limited GPU memory by fetching active parameter shards on demand. Learn more today.

- Repository: [labml.ai/annotated_deep_learning_paper_implementations](https://github.com/labmlai/annotated_deep_learning_paper_implementations)
- Tags: deep-dive
- Published: 2026-03-04

---

**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`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/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`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/__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`:

```python
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`:

```python
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:

```bash
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`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/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.