Concept · Chapter 9: How an LLM Is Actually Built
Sharded Training (ZeRO and FSDP)
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.
You should understand first
- Derivatives and Gradients
- Loss Functions
- Gradient Descent
- Probability and Distributions
- Expected Value and Variance
- Stochastic Gradient Descent (SGD)
- The Chain Rule
- Vectors
- Dot Product
- The Turing Test
- Symbolic AI
- Logic and Rules
- Expert Systems
- Knowledge Representation
- The Knowledge-Acquisition Bottleneck
- From Rules to Learning
- Supervised, Unsupervised and Self-Supervised Learning
- Features, Labels and Tasks
- Linear Regression
- Entropy
- Softmax
- Cross-Entropy Loss
- Logistic Regression
- The Perceptron
- Activation Functions
- The Artificial Neuron
- Matrix Multiplication
- Multilayer Perceptron (MLP)
- The Forward Pass
- Computational Graphs and Autodiff
- Backpropagation
- Momentum and Adam
- The Pretraining Loop
- Data Parallelism and All-Reduce
- Mixed-Precision Training (FP16 and BF16)
- Text as Data
- One-Hot Encoding
- Tokenization
- Conditional Probability and Bayes' Theorem
- Probability of Sequences
- Language Modeling
- Embeddings
- Attention
- Self-Attention
- Causal Masking
- Autoregressive Next-Token Prediction
- Pretraining at Scale
- Parameters, Tokens and Context Windows
- Where Training Memory Goes
- Sharded Training (ZeRO and FSDP)
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
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.
PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel
Yanli Zhao, Andrew Gu et al. · 2023 · PVLDB 16(12)
Describes PyTorch's built-in version of ZeRO-style sharding, the default way to train models that don't fit on one GPU.
Watch
Stanford Online
Stanford CS336 Language Modeling from Scratch | Spring 2025 | Lecture 7: Parallelism 1
A clear university lecture on how training is split across many GPUs, with the trade-offs between the methods.