Skip to content
AI-grafen
GFrontier LabModel training and fine-tuning· about 120 min· fast-moving, sources checked often· verified 2026-09-20· EN

FSDP, ZeRO and DeepSpeed

Be able to explain the sharding of parameters, gradients and optimiser state.

Prerequisites

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 stageWhat is shardedThe memory per GPU (N cards)
DDPnothing14 + 14 + 56
ZeRO-1the optimiser state14 + 14 + 56/N
ZeRO-2+ the gradients14 + (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:

  1. all-gather the layer's parameters from all the ranks.
  2. Compute the forward pass.
  3. Release the parameters (only your own shard is kept).
  4. 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

All the sources and licences