Skip to content
Road to Intelligence

Concept · Chapter 9: How an LLM Is Actually Built

Sharded Training (ZeRO and FSDP)

Should knowUnderstand10 minDifficulty

Sharded training keeps data parallelism's simplicity but stores only a slice of the optimizer states, gradients and weights on each GPU, gathering each layer's full weights just in time to use them.

The problem

Plain data parallelism keeps an identical copy of 16 bytes per parameter on every GPU, so a model whose states don't fit on one GPU can't use it, however many GPUs you have.

The solution

Partition the model states across the data-parallel GPUs in three stages (optimizer states, then gradients, then the parameters themselves) and use collective operations to fetch what is needed when it is needed.

The consequence

Model-state memory per GPU falls in proportion to the number of GPUs for about 1.5 times the communication, making models with billions of parameters trainable on ordinary clusters; activations are untouched.

Remove the redundancy

With plain data parallelism, 64 GPUs hold 64 identical copies of the optimizer states. ZeRO ("Zero Redundancy Optimizer") stores each piece once.

  • Stage 1 splits the optimizer states (12 of the 16 bytes per parameter) across GPUs. Each GPU updates only its slice of the weights, then the updated weights are shared.
  • Stage 2 also splits the gradients: each GPU only needs the summed gradients for its slice, so the all-reduce becomes a reduce-scatter.
  • Stage 3 also splits the weights. Before computing a layer, the GPUs gather that layer's full weights, use them, and throw them away again.

In ZeRO's worked example, a 7.5-billion-parameter model needs 120 GB per GPU for model states with plain data parallelism; with 64 GPUs that falls to 31.4 GB with stage 1, 16.6 GB with stage 2 and 1.9 GB with stage 3 Established. Stage 3 increases the total communication of data parallelism to about 1.5 times Established.

Tiny example. LLaMA 7B's 107 GB of model states on 8 GPUs with stage 3: 107 ÷ 8 ≈ 13.4 GB each, leaving room for activations on an 80 GB GPU.

What it does not fix

Sharding divides the model states, not the activations. If one sequence's activations don't fit on a GPU, no number of GPUs helps; recomputation or model parallelism must. The memory calculator shows this: with LLaMA 7B, 4,096-token sequences and every activation stored, even ZeRO-3 on 1,024 GPUs doesn't fit.

PyTorch's FSDP (fully sharded data parallel) brings the same ideas into the framework itself Established, and Llama 3 used FSDP, sharding optimizer states and gradients Established.

What to remember

  • ZeRO-1 shards optimizer states, ZeRO-2 also gradients, ZeRO-3 also parameters.
  • ZeRO-3 on N GPUs: about 16 ÷ N bytes per parameter of model states per GPU.
  • ZeRO's example: a 7.5B model needs 120 GB per GPU with plain data parallelism, 1.9 GB with ZeRO-3 on 64 GPUs.
  • Cost: about 1.5× the communication of plain data parallelism. Activations are not sharded.
  • PyTorch's built-in version is FSDP (fully sharded data parallel).

Key papers

Essential

ZeRO: Memory Optimizations Toward Training Trillion Parameter Models

Samyam Rajbhandari, Jeff Rasley et al. · 2019

Explained where training memory goes (16 bytes per parameter with mixed-precision Adam) and how to remove the redundant copies. Its stages became DeepSpeed's ZeRO and PyTorch's FSDP.

How to read it: Section 3 ('Where did all the memory go?') and Figure 1 are the essentials; the rest is engineering detail.

~40 min readarXiv:1910.02054✓ verified 2026-10-04

Watch