Data-parallel training copies the whole model onto every GPU. That is simple and fast until the model, its gradients and its optimizer state no longer fit on one device, which for mixed-precision Adam happens well before the weights alone look large. PyTorch's Fully Sharded Data Parallel (FSDP) keeps the data-parallel programming model, where every rank runs the same code on a different slice of the batch, but stores only a fraction of the model on each rank and assembles full layers just in time to use them.
The idea is the same partitioning that ZeRO sharding describes; this article is about PyTorch's implementation of it. It covers what happens inside one training step, the memory and communication arithmetic, how to set up the current API (FSDP2's fully_shard), the knobs that matter, how to checkpoint, and the failure modes that cost teams days. The code targets recent PyTorch releases; check the version you run, because this API has moved quickly.
The problem FSDP solves, in numbers
Count the bytes per parameter for mixed-precision training with Adam, keeping full-precision master weights. The fp32 parameter is 4 bytes, its fp32 gradient 4 bytes, and Adam's two moment estimates another 8 bytes. That is 16 bytes per parameter before a single activation is stored. A 7-billion-parameter model therefore needs about 112 GB of model state. An 80 GB GPU cannot hold it, so plain data parallelism, which replicates all of it on every rank, is not an option.
Shard that state across 8 GPUs and each rank holds about 14 GB. Nothing about the maths of training changed: each rank still computes gradients for its own micro-batch, and the averaged gradient is still applied to every parameter. What changed is where the bytes live and when they move.
One step, one rank: the unshard and reshard lifecycle
FSDP groups parameters into units, usually one unit per transformer block. Between uses, each unit exists only as shards. When the forward pass reaches block k, FSDP all-gathers the shards so every rank briefly holds block k's full weights, runs the block, and then frees the full weights again (resharding). Activations saved for backward stay; only the parameters are dropped.
Backward walks the blocks in reverse. Each block is all-gathered again, its gradients are computed, and a reduce-scatter both sums the gradients across ranks and leaves each rank holding only the slice that matches its parameter shard. The optimizer then updates the local shard with the local gradient slice. No rank ever holds the full optimizer state.
Two consequences follow directly. Peak parameter memory is the sharded state plus roughly one or two full blocks, the one computing and the one being prefetched. For the 7B example with 32 blocks, a block is about 220 million parameters, or about 0.44 GB in bf16, so the transient cost is small next to 14 GB of shards. The second consequence is traffic: two all-gathers and one reduce-scatter of the parameter volume per step, against DDP's single all-reduce, which is itself a reduce-scatter plus an all-gather. FSDP moves roughly 1.5 times the bytes DDP does, and the design only works if that communication hides behind compute.
FSDP1 and FSDP2: which one you are using
The original FSDP wrapped a module in a FullyShardedDataParallel object and flattened each unit's parameters into one large FlatParameter, which it then sharded. That made communication efficient but produced awkward side effects: parameter names and shapes were hidden behind views, mixing frozen and trainable parameters in one unit was painful, and memory release relied on CUDA stream bookkeeping that could make usage non-deterministic.
FSDP2, the fully_shard function in torch.distributed.fsdp, shards each parameter individually along dimension 0 and represents it as a DTensor, a tensor that knows its device mesh and placement. Your module keeps its own class and parameter names, frozen parameters can sit beside trainable ones, and PyTorch's documentation notes it avoids the recordStream path, which gives more predictable memory. New code should use FSDP2; the rest of this article does.
Setting it up: the minimal correct program
The order of operations matters more than any individual call. Build the model without allocating memory, apply fully_shard bottom-up (each block first, then the root), materialise the local shards, initialise them, and only then build the optimizer.
import os
import torch
import torch.distributed as dist
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
dist.init_process_group("nccl") # launched with torchrun
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
mesh = init_device_mesh("cuda", (dist.get_world_size(),))
with torch.device("meta"): # no memory allocated yet
model = Transformer(cfg)
mp = MixedPrecisionPolicy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32)
for block in model.layers: # one FSDP unit per transformer block
fully_shard(block, mesh=mesh, mp_policy=mp)
fully_shard(model, mesh=mesh, mp_policy=mp) # root: embeddings, final norm, head
model.to_empty(device="cuda") # materialise only the local shards
model.init_weights() # your init, applied to sharded DTensors
optim = torch.optim.AdamW(model.parameters(), lr=3e-4) # AFTER sharding: params are DTensors nowThe meta-device step is what makes large models loadable at all: a 7B model built normally would allocate every full tensor on one device before sharding it. to_empty allocates only each rank's shard, and your initialisation or checkpoint load then fills it. The optimizer must come last because fully_shard replaces parameters with DTensors; an optimizer built earlier holds references to the old tensors and silently updates nothing that matters.
Wrapping granularity: the decision that sets your peak
Each call to fully_shard defines one communication and memory unit. Calling it only on the root makes the whole model a single unit, so the forward pass all-gathers every parameter at once and you have spent the communication cost without saving any peak memory. Calling it on every linear layer produces hundreds of tiny collectives whose fixed latency dominates. Per transformer block is the standard answer because a block is large enough to use the interconnect efficiently and small enough that one or two unsharded blocks are cheap.
The root call still matters: it owns whatever is left over, typically the embedding table, the final norm and the output head. The documentation specifies that when reshard_after_forward is left unset, non-root units reshard after forward and the root does not, because the root's parameters are needed again immediately at the start of backward. Setting reshard_after_forward=False on other units trades memory for one fewer all-gather per unit, which is sometimes worth it on the last few blocks; an integer value reshards to a smaller group of ranks instead of all of them, a middle ground for multi-node jobs.
Mixed precision and offload
MixedPrecisionPolicy controls three dtypes. param_dtype is what the all-gathered weights are cast to for compute, so bf16 halves all-gather traffic and unsharded memory. reduce_dtype is the dtype of gradient reduction; keeping it at fp32 avoids losing small gradient contributions when many ranks are summed. The sharded master parameters stay in their original dtype, fp32 in the example, which is where the optimizer applies updates. cast_forward_inputs defaults to on and casts floating-point inputs to the parameter dtype.
CPUOffloadPolicy moves the sharded parameters, gradients and optimizer step to host memory, pinned by default for faster transfers. It lets a model train on fewer GPUs than it otherwise needs, at the price of PCIe traffic every step and a CPU-side optimizer step. Treat it as a capacity tool, not a speed tool, and see optimizer offload to CPU for how to judge whether host bandwidth can keep up.
Overlap and prefetch: where throughput is won or lost
FSDP issues all-gathers on a separate CUDA stream so the next unit's gather can run while the current unit computes. By default it prefetches one unit ahead, following the order modules ran in. When the interconnect is slow relative to compute, or a block is small, one unit ahead is not enough and the GPU stalls waiting for weights. set_modules_to_forward_prefetch and set_modules_to_backward_prefetch let you name which modules to gather early, at the cost of holding more unsharded blocks. Overlapping collectives with compute covers how to read the resulting profiler trace.
The quick diagnostic is a profiler timeline with communication and compute streams side by side. If NCCL kernels sit in gaps where no compute runs, you are exposed: raise the prefetch depth, increase the per-rank batch so compute per block grows, or move to HSDP so the gathers run over a faster link.
Hybrid sharding (HSDP) for multi-node jobs
Full sharding across 32 GPUs on 4 nodes means every all-gather crosses the slower inter-node network. Hybrid sharding passes a two-dimensional device mesh: parameters are sharded within a group, typically one node over NVLink, and replicated across groups. Gradients are reduce-scattered inside the node and then all-reduced across the replicas.
# 4 nodes x 8 GPUs: shard inside a node over NVLink, replicate across nodes.
mesh = init_device_mesh("cuda", (4, 8), mesh_dim_names=("replicate", "shard"))
for block in model.layers:
fully_shard(block, mesh=mesh, mp_policy=mp)
fully_shard(model, mesh=mesh, mp_policy=mp)Each rank now stores 1/8 of the model state rather than 1/32, so HSDP only works when the per-node share fits. When it does, it usually wins: the frequent gathers stay on the fast link and only one gradient all-reduce per step crosses nodes. NCCL collectives explains why those two link types behave so differently.
The training loop, gradient accumulation and clipping
Inside the loop FSDP is mostly invisible, with two exceptions. Gradient accumulation should skip the reduce-scatter on all but the last micro-batch, otherwise you pay the gradient communication once per micro-batch. The catch is memory: while synchronisation is off, each rank keeps unsharded gradients, so accumulation without sync spends memory to save bandwidth. Clipping must compute the norm over the whole model, not the local shard; the standard clipping utility handles DTensor parameters.
accum = 4
for step, batches in enumerate(loader): # each item: a list of `accum` micro-batches
for i, mb in enumerate(batches):
last = i == accum - 1
model.set_requires_gradient_sync(last) # skip reduce-scatter on early micro-batches
loss = model(mb["input_ids"], labels=mb["labels"]) / accum
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # norm computed over all shards
optim.step()
optim.zero_grad(set_to_none=True)
Checkpointing without gathering the model
Calling torch.save(model.state_dict()) on rank 0 saves only rank 0's shards. Gathering a full state dict to one rank works for small models but needs the whole model in one process's memory. PyTorch Distributed Checkpoint (DCP) has each rank write its own shards in parallel with metadata describing the global layout, so loading can reshard to a different number of GPUs. The anatomy of a training step walks through where memory goes in each phase of the step that these shards come from.
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
def save(step):
model_sd, optim_sd = get_state_dict(model, optim) # sharded, no gather to rank 0
dcp.save({"model": model_sd, "optim": optim_sd}, checkpoint_id=f"ckpt/step{step}")
def load(path):
model_sd, optim_sd = get_state_dict(model, optim) # templates with the right shapes
state = {"model": model_sd, "optim": optim_sd}
dcp.load(state, checkpoint_id=path) # works across a changed world size
set_state_dict(model, optim, model_state_dict=state["model"],
optim_state_dict=state["optim"])When you need a single-file export for inference, get_model_state_dict with StateDictOptions(full_state_dict=True, cpu_offload=True) gathers the full weights to CPU. Do that for export, not for the checkpoints you resume from.
Failure modes seen in practice
| Symptom | Likely cause | Fix |
|---|---|---|
| Loss never moves | Optimizer built before fully_shard | Build the optimizer after sharding |
| OOM at step 1 despite sharding | Only the root was wrapped, or the model was built on GPU before sharding | Wrap per block; build on the meta device |
| Job hangs with no error | Ranks took different code paths, so collectives were issued in different orders | Remove rank-dependent control flow around forward and backward; set an NCCL timeout |
| Resumed run diverges | Rank-0-only save, or optimizer state not saved | Use DCP for model and optimizer together |
| Low GPU utilisation | All-gathers exposed on a slow link | Deeper prefetch, larger micro-batch, HSDP |
| Memory grows during accumulation | Unsharded gradients held while sync is off | Accumulate fewer micro-batches, or sync every micro-batch |
The hang deserves emphasis. Every rank must issue the same collectives in the same order. A branch such as skipping a block for some inputs on one rank, or an exception caught on one rank but not others, leaves the remaining ranks waiting on an all-gather forever. If a step can raise, make every rank raise.
Trade-offs: FSDP against the alternatives
FSDP scales data parallelism when the model state is the problem. It does not reduce activation memory, which grows with sequence length and micro-batch size; pair it with activation checkpointing for that. It also does not split a single layer's compute, so when one block's matrix multiplies are too large or too slow for one GPU, tensor parallelism inside a node combined with FSDP across nodes is the usual next step, on a 2D mesh again. For models that fit comfortably with DDP, DDP stays faster because it moves fewer bytes.
What to do next
- Compute your model's bytes of state per parameter and divide by GPU count to predict the per-rank footprint before launching.
- Build the model on the meta device, apply
fully_shardper transformer block and then on the root, and build the optimizer last. - Set
MixedPrecisionPolicy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32)and confirm loss curves match a small full-precision baseline. - Profile one step; if collectives are exposed, tune prefetch or batch size before buying hardware.
- On more than one node, test HSDP against full sharding and keep whichever gives the higher tokens per second at equal memory headroom.
- Switch checkpointing to DCP, then test a resume on a different GPU count before you need one.