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.

Advertisement

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.

One FSDP training step, seen from a single rank (4 ranks, one shard each)Sharded state1/N of params, grads, Adamall-gather block kfull bf16 weightsforward block kkeep activationsreshard block kfree full weightsrepeat for every block; the next block's all-gather is prefetched while this one computesall-gather block kagain, in reverse orderbackward block kfull gradient for blockreduce-scattereach rank keeps 1/N gradoptimizer steplocal shard onlyPeak parameter memory = sharded state + roughly one or two unsharded blocks (current + prefetched)Traffic per step = all-gather (fwd) + all-gather (bwd) + reduce-scatter: about 1.5x what DDP's all-reduce moves
The per-block lifecycle. Forward all-gathers, computes and reshards each block in order; backward all-gathers again, computes gradients and reduce-scatters them, so the optimizer only ever touches a local shard.

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.

Advertisement

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 now

The 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

SymptomLikely causeFix
Loss never movesOptimizer built before fully_shardBuild the optimizer after sharding
OOM at step 1 despite shardingOnly the root was wrapped, or the model was built on GPU before shardingWrap per block; build on the meta device
Job hangs with no errorRanks took different code paths, so collectives were issued in different ordersRemove rank-dependent control flow around forward and backward; set an NCCL timeout
Resumed run divergesRank-0-only save, or optimizer state not savedUse DCP for model and optimizer together
Low GPU utilisationAll-gathers exposed on a slow linkDeeper prefetch, larger micro-batch, HSDP
Memory grows during accumulationUnsharded gradients held while sync is offAccumulate 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

  1. Compute your model's bytes of state per parameter and divide by GPU count to predict the per-rank footprint before launching.
  2. Build the model on the meta device, apply fully_shard per transformer block and then on the root, and build the optimizer last.
  3. Set MixedPrecisionPolicy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32) and confirm loss curves match a small full-precision baseline.
  4. Profile one step; if collectives are exposed, tune prefetch or batch size before buying hardware.
  5. On more than one node, test HSDP against full sharding and keep whichever gives the higher tokens per second at equal memory headroom.
  6. Switch checkpointing to DCP, then test a resume on a different GPU count before you need one.
Key takeaway: FSDP shards parameters, gradients and optimizer state across data-parallel ranks and rebuilds one block at a time: all-gather, compute, reshard in forward, and all-gather, compute, reduce-scatter in backward. That cuts model-state memory by the number of ranks for about 1.5 times DDP's communication, which pays off only when prefetch hides the gathers. Use FSDP2's fully_shard per block and then on the root, build on the meta device, create the optimizer after sharding, keep reductions in fp32, prefer HSDP across nodes, and checkpoint with DCP.