Multi-GPU and Distributed Training

FSDP: sharding a model across GPUs

DDP gives every GPU a full copy of the model, which caps model size at one card; FSDP cuts the parameters, gradients and optimizer state into shards and rebuilds each layer only for the moment it is needed.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

FSDP lets several GPUs hold one model between them, instead of each holding a whole copy.

Think of a long recipe book that will not fit on one small shelf. Four friends each take a quarter of the pages. When chapter three is needed, whoever holds those pages photocopies them for everyone. All four cook that chapter. Then the copies go in the bin.

Only one chapter is ever duplicated. The rest of the book stays split.

That is FSDP. Sharding means cutting something into pieces and giving one piece to each worker. FSDP shards the model's numbers, and briefly reassembles one layer at a time.

Why it exists

DDP puts a complete copy of the model on every GPU. Eight GPUs, eight identical copies. Splendid for speed, useless for size — the biggest model you can train is the biggest that fits on one card.

And the model weights are the small part. Training also stores the corrections and the optimizer's memory of past steps. With the common Adam optimizer, those add up to several times the weights themselves.

FSDP splits all three. The cost is conversation: every layer must be reassembled before use and thrown away after.

How it works

DDP:    gpu0 [WHOLE MODEL]  gpu1 [WHOLE MODEL]  gpu2 [WHOLE MODEL]

FSDP:   gpu0 [piece 1]  gpu1 [piece 2]  gpu2 [piece 3]

  running layer 4:
      everyone shares their piece of layer 4  ->  all have layer 4
      compute layer 4
      throw the borrowed pieces away          ->  back to one piece each
      move to layer 5

A real example you have seen

A group project where the report is too big to email. Each person keeps their own section on their laptop. When the group edits section three, whoever owns it shares the screen. Everyone works on it. Then it goes back to being one person's file.

Remember this

  • DDP copies the whole model to every GPU; FSDP gives each GPU a piece.
  • Pieces are gathered into a full layer, used, then thrown away at once.
  • You trade more talking between GPUs for fitting a bigger model.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Be warned before you plan around this lesson: FSDP is a GPU technique. The sharding and the forward pass below do run on CPU with the gloo backend, and that is how the outputs were captured. A full training step does not — reduce_scatter, which FSDP needs for gradients, is not implemented in gloo. Everything past the forward pass here needs real CUDA GPUs.

The modern API is fully_shard, sometimes called FSDP2. It replaced the older FullyShardedDataParallel wrapper class, which is still present and still works.

One import note, since this API moved recently. from torch.distributed.fsdp import fully_shard is the public path from PyTorch 2.6 onward. On 2.5 the same function lives at torch.distributed._composable.fsdp, and that is the path used to capture the output below. If the public import raises ImportError on your machine, you are on 2.5 or older.

Watching a model get cut up

fsdp_shard.py
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.distributed.device_mesh import DeviceMesh
from torch.distributed.fsdp import fully_shard

dist.init_process_group(backend="gloo")
rank, world = dist.get_rank(), dist.get_world_size()
mesh = DeviceMesh.from_group(dist.group.WORLD, "cpu")   # "cuda" on GPUs

torch.manual_seed(0)
model = nn.Sequential(nn.Linear(64, 64), nn.ReLU(), nn.Linear(64, 64))
total = sum(p.numel() for p in model.parameters())

for layer in model:                      # shard each block, then the whole model
    if isinstance(layer, nn.Linear):
        fully_shard(layer, mesh=mesh)
fully_shard(model, mesh=mesh)

w = model[0].weight                      # now a DTensor: a global view over shards
print(f"rank {rank}: 0.weight logical shape {tuple(w.shape)}, "
      f"my shard {tuple(w.to_local().shape)}, placement {w.placements}")
kept = sum(p.to_local().numel() for p in model.parameters())
print(f"rank {rank}: {total} parameters in the model, {kept} stored on this rank")

out = model(torch.randn(2, 64))          # gathers each block just in time
print(f"rank {rank}: forward output {tuple(out.shape)}")

dist.destroy_process_group()
bash
torchrun --nproc_per_node=2 fsdp_shard.py
Output
rank 0: 0.weight logical shape (64, 64), my shard (32, 64), placement (Shard(dim=0),)
rank 0: 8320 parameters in the model, 4160 stored on this rank
rank 0: forward output (2, 64)
rank 1: 0.weight logical shape (64, 64), my shard (32, 64), placement (Shard(dim=0),)
rank 1: 8320 parameters in the model, 4160 stored on this rank
rank 1: forward output (2, 64)

Lines are grouped by rank for reading. Three things to take from that output.

The weight still says it is 64x64. model[0].weight is now a DTensor — a tensor that knows it is a view over pieces living on several ranks. Your code, and print(model), still see the logical shape. w.to_local() reveals the real 32x64 slice this rank owns.

placement=Shard(dim=0) says the cut runs along the first dimension. Rows 0–31 on rank 0, rows 32–63 on rank 1.

8320 parameters, 4160 per rank. Exactly half each, with two ranks. That halving is the entire point.

The forward pass still returns the right shape, because FSDP gathered each block back to full size the instant it was needed, then released it.

Caveat, stated plainly: on this CPU run, calling .backward() on that output raises an autograd in-place-modification error. FSDP's backward path is written against CUDA collectives. Treat the CPU version as a way to see the sharding, not as a way to train.

A real GPU training loop

This block needs at least two CUDA GPUs, so no output is shown for it. Presenting invented multi-GPU numbers would be worse than showing none.

fsdp_train.py — needs 2 or more CUDA GPUs
import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.fsdp import fully_shard

local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
mesh = init_device_mesh("cuda", (dist.get_world_size(),))

model = nn.Sequential(*[nn.Linear(1024, 1024) for _ in range(8)]).cuda()

for block in model:                       # shard per block, not per parameter
    fully_shard(block, mesh=mesh)
fully_shard(model, mesh=mesh)

opt = torch.optim.AdamW(model.parameters(), lr=1e-4)   # AFTER sharding
x = torch.randn(8, 1024, device="cuda")

for step in range(3):
    opt.zero_grad(set_to_none=True)
    model(x).sum().backward()
    opt.step()
    if dist.get_rank() == 0:
        print(f"step {step} done")

dist.destroy_process_group()

Two lines carry most of the meaning. Shard per block, then the root: each fully_shard call creates one gather-and-release unit, so its size controls the memory/communication trade. Shard nothing and you are back to DDP; shard every single parameter and you drown in tiny collectives. One transformer block per call is the standard choice.

Build the optimizer after sharding. model.parameters() must return the sharded DTensor objects, or the optimizer allocates state for full-size parameters and you lose the memory saving you came for.

What FSDP actually saves

For a model with P parameters trained in float32 with Adam, the per-GPU fixed cost is roughly:

held on each GPUDDPFSDP over R ranks
parameters4P bytes4P/R
gradients4P4P/R
Adam moments8P8P/R
total fixed16P16P/R
transient gathernoneone block, full size
activationssame both wayssame both ways

At 8 ranks, a 1-billion-parameter model drops from about 16 GB of fixed state per GPU to about 2 GB. Activations are not sharded by this — they scale with batch size, and gradient checkpointing is the lever for those.

Communication rises by roughly 50% against DDP: an all-gather in the forward pass, an all-gather in the backward pass, and a reduce-scatter for gradients, instead of DDP's single all-reduce.

Common mistakes

Building the optimizer before fully_shard. No error, no warning, no memory saved.

Sharding every parameter individually. Thousands of tiny collectives, each paying full network latency. Group at block level.

Saving with plain torch.save(model.state_dict()). Each rank writes its own shard, so you get R incomplete files. Use torch.distributed.checkpoint — see saving from one rank.

Reaching for FSDP when DDP fits. FSDP is strictly slower per step. If the model, gradients and optimizer state fit on one card, DDP is the right answer. Try mixed precision and checkpointing before you reach for sharding.

Calling .to_local() and forgetting it. Once you unwrap a DTensor, the result is an ordinary tensor holding one rank's slice. Arithmetic on it is silently per-rank.

Try it yourself

Run fsdp_shard.py with --nproc_per_node=4. The stored-parameter count should fall to a quarter. Then remove the per-layer fully_shard calls, keep the root one, and see how the reported shard changes.

What to learn next

Researcher — Mathematics and papers.

ZeRO stages, and which one this is

FSDP implements ZeRO stage 3 (Rajbhandari et al., 2020). The stages partition, in order: stage 1 the optimizer states, stage 2 also the gradients, stage 3 also the parameters. Per-device fixed memory for $\Psi$ parameters with mixed-precision Adam is often quoted as $2\Psi + 2\Psi + 12\Psi$ bytes (fp16 weights, fp16 gradients, fp32 master weights plus two Adam moments); stage 3 divides all of it by the shard count $N_d$, giving $\frac{16\Psi}{N_d}$ against DDP's $16\Psi$.

The runtime is a sequence of communication units, each corresponding to one fully_shard call. For unit $i$ in the forward pass: all_gather its parameter shards, run the module, free the gathered buffer. In the backward pass: all_gather again (the parameters were discarded), compute gradients, reduce_scatter them so each rank retains only its own shard, free. Prefetching overlaps unit $i+1$'s gather with unit $i$'s compute, which is why unit granularity is the dominant tuning knob — too small and latency dominates, too large and the transient gathered buffer defeats the memory saving.

Volume per step per rank: $2 \times S\frac{N-1}{N}$ for the two all-gathers plus $S\frac{N-1}{N}$ for the reduce-scatter, against DDP's $2S\frac{N-1}{N}$ all-reduce — the $1.5\times$ figure. Latency is hidden only when compute per unit exceeds transfer per unit, which is why FSDP scales well on transformer blocks and poorly on shallow models.

FSDP2 and DTensor

The fully_shard API is a rewrite of the original FullyShardedDataParallel wrapper. The differences that matter in practice: parameters become per-parameter DTensor objects sharded on dim 0 rather than being flattened into one opaque FlatParameter, which makes state_dict readable, makes per-parameter optimizer settings and parameter groups work naturally, and composes with tensor parallelism through a 2-D DeviceMesh. It is a function applied to a module, not a wrapper class, so model keeps its type and attribute paths — no module. prefix appears in checkpoints.

HYBRID_SHARD shards within a node and replicates across nodes, which trades memory for keeping the expensive collectives on NVLink rather than the network. In FSDP2 that is expressed as a 2-D mesh with (replicate, shard) dimensions rather than a flag.

Choosing between the tools

Sharding attacks state memory, which is fixed in the parameter count. Activation checkpointing attacks activation memory, which scales with batch and sequence length. Tensor and pipeline parallelism attack the case where a single layer or a single activation does not fit at all. Large-scale training composes all four; the ordering that usually holds is: mixed precision, then checkpointing, then FSDP, then model parallelism — cheapest and least invasive first.

References

  • Rajbhandari et al. (2020), ZeRO: Memory Optimizations Toward Training Trillion Parameter Models, SC20 — the stage decomposition and the memory arithmetic above.
  • Zhao et al. (2023), PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel, VLDB 16(12) — the PyTorch implementation, unit granularity and prefetching measurements.
  • Ren et al. (2021), ZeRO-Offload: Democratizing Billion-Scale Model Training — pushing optimizer state to CPU when even sharding is not enough.
  • PyTorch documentation, FSDP2 and DTensor — the current API surface and the sharding placements.

What to learn next