FSDP, ZeRO and DeepSpeed
Be able to explain the sharding of parameters, gradients and optimiser state.
Prerequisites
- FDistributed trainingrequired
Intuition
In data-parallel training every GPU has a full copy of the model, the gradients and the optimiser state. For a 7B model with Adam in bf16 that is ~84 GB per card — however many cards you have.
Sharding splits the three things across the cards instead:
| The stage | What is sharded | The memory per GPU (N cards) |
|---|---|---|
| DDP | nothing | 14 + 14 + 56 |
| ZeRO-1 | the optimiser state | 14 + 14 + 56/N |
| ZeRO-2 | + the gradients | 14 + (14+56)/N |
| ZeRO-3 / FSDP | + the parameters | (14+14+56)/N |
The price is communication: the parameters have to be gathered (all-gather) before each layer's computation and released afterwards. With a fast link (NVLink, InfiniBand) that is affordable; over slow Ethernet it dominates.
Formal
FSDP (PyTorch) and DeepSpeed ZeRO implement the same idea with different APIs. The flow in ZeRO-3 per layer:
- all-gather the layer's parameters from all the ranks.
- Compute the forward pass.
- Release the parameters (only your own shard is kept).
- The same backwards, plus a reduce-scatter of the gradients.
Complementary techniques:
- Activation checkpointing: do not save all the intermediate activations; recompute them in the backward pass. It saves a lot of memory for ~30 % more time. Nearly always the right choice when memory is short.
- CPU offload: move the optimiser state (and possibly the parameters) to the host memory. It makes very large models possible on little GPU memory, but is slow — the PCIe bandwidth becomes the bottleneck.
- Mixed precision: bf16 for the computation, fp32 for the master weights. The standard.
The decision order: start with DDP. Out of memory → activation checkpointing → ZeRO-2 → ZeRO-3/FSDP → offload. Every step costs speed, so only take the next one when the previous is not enough.
The most common practical mistake: saving checkpoints from all the ranks (it should be rank 0, or state_dict_type=FULL_STATE_DICT with the right configuration) and forgetting sampler.set_epoch() — then every epoch looks the same.
Code
import torch, functools
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, ShardingStrategy, MixedPrecision
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import apply_activation_checkpointing
dist.init_process_group("nccl")
rank = dist.get_rank(); torch.cuda.set_device(rank)
mp = MixedPrecision(param_dtype=torch.bfloat16, reduce_dtype=torch.bfloat16, buffer_dtype=torch.bfloat16)
wrap = functools.partial(transformer_auto_wrap_policy, transformer_layer_cls={TransformerBlock})
model = FSDP(model,
sharding_strategy=ShardingStrategy.FULL_SHARD, # = ZeRO-3
auto_wrap_policy=wrap, # shard per block, not the whole model
mixed_precision=mp,
device_id=rank)
apply_activation_checkpointing(model, check_fn=lambda m: isinstance(m, TransformerBlock))
auto_wrap_policy is decisive: without it the model is sharded as one unit, which means everything is gathered at once — and the memory gain does not materialise.
Mastery means
- Explains the sharding of parameters, gradients and optimiser state
- Chooses the ZeRO stage to suit the memory requirement
- Knows about offloading and activation checkpointing
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — arXiv (open access; licence per article)
- PyTorch — FSDP (BSD-3) — BSD-3-Clause