What is FSDP (Fully Sharded Data Parallel)?
Abbreviated FSDP
FSDP (Fully Sharded Data Parallel) is PyTorch's way of training a model that is too big to replicate on every GPU. It is a form of data parallelism: each GPU still processes its own slice of the batch. The difference is that the model's parameters, gradients and optimizer state are split into shards, one per GPU, and a layer's full weights are rebuilt only for the moment that layer runs.
Plain data parallelism keeps a complete copy of everything on every GPU. FSDP keeps 1/N of it on each of N GPUs, which is how it trains models whose full training state would not fit on one card.
How it works
PyTorch's FSDP2 tutorial describes the cycle:
- Outside of forward and backward computation, parameters stay fully sharded.
- Before a layer's forward or backward pass, its sharded parameters are all-gathered into full parameters on every GPU.
- During backward, each GPU's local full gradients are reduce-scattered into sharded gradients.
- The optimizer updates the sharded parameters with the sharded gradients, so optimizer state stays sharded too.
The tutorial summarizes it as a decomposition of DDP's all-reduce into a reduce-scatter and an all-gather. Both are NCCL collectives. In FSDP2 you apply fully_shard to each transformer layer and then to the root model, so only one layer's full weights exist at a time while the rest remain sharded. The ideas come from the ZeRO paper (Rajbhandari et al., 2019), whose stage 3 shards parameters, gradients and optimizer state, as FSDP does.
Worked example: Llama 3.1 8B on 8 GPUs
Llama 3.1 8B has 8.03 billion parameters (computed from its published config.json). The ZeRO paper counts 16 bytes per parameter of model state for mixed-precision training with Adam: 2 for the weights, 2 for the gradients, and 12 for the fp32 weight copy and the two Adam moments.
| Setup | Model state per GPU |
|---|---|
| Data parallel, any N | 8.03B x 16 bytes = 128.5 GB |
| FSDP across 8 GPUs | 128.5 GB / 8 = 16.1 GB |
| FSDP across 16 GPUs | 128.5 GB / 16 = 8.0 GB |
Plain data parallelism does not fit on an 80GB H100 at all. With FSDP over 8 GPUs, the persistent state is 16.1 GB per GPU. On top of that each GPU briefly holds one layer's full weights (a single Llama 3.1 8B layer is about 218 million parameters, roughly 0.44 GB in BF16) plus activations, which depend on batch size and sequence length and are not reduced by sharding.
The price is communication. Per step, each GPU all-gathers the weights during the forward pass and again for the backward pass, and then reduce-scatters the gradients, so it moves about 1.5 times the volume of a plain all-reduce, per the ZeRO paper's accounting for its stage 3. In BF16 that is roughly 3 x 16 GB x 7/8 = 42 GB per GPU per step, against 28 GB for the all-reduce in data parallelism. FSDP overlaps these transfers with compute, but a slow link still shows up as idle GPUs.
How it differs from data parallelism
| Data parallelism | FSDP | |
|---|---|---|
| Weights per GPU | Full copy | 1/N shard, gathered per layer |
| Gradients | Full, all-reduced | Reduce-scattered to 1/N |
| Optimizer state | Full copy | 1/N shard |
| Communication | One all-reduce | All-gather and reduce-scatter |
| Largest trainable model | What fits on one GPU | Grows with GPU count |
FSDP also combines with other techniques. The PyTorch tutorial notes that FSDP2 can shard models with frozen and non-frozen parameters together, which is the setup for parameter-efficient fine-tuning, and that it supports NF4, the format QLoRA uses. For splitting layers rather than data, see tensor parallelism and pipeline parallelism.
What it means when you pick a GPU
Work out the model state first: parameters x 16 bytes, divided by the number of GPUs, plus activations and the one gathered layer. For a 70B model, the state is 70.55 billion x 16 = 1.13 TB, so 8 GPUs need 141 GB each for state alone, which does not fit even the 141GB of an H200 once activations are added. That is when you need more GPUs, or a smaller trainable footprint such as QLoRA.
Because FSDP all-gathers weights every step, interconnect speed matters more than for plain data parallelism. Prefer GPUs with NVLink inside a node, such as the H100 (900 GB/s) or A100 (600 GB/s), over PCIe-only cards, whose link is about 64 GB/s at Gen 4 (see NVLink vs PCIe). For BF16 weights and activations, see BF16. You can plan the memory with the VRAM calculator, and see Aquanode pricing for multi-GPU rentals by the hour.
Building on GPUs? Aquanode runs the workload.
Deploy on H100, H200, B200, A100 and MI300X across a multi-provider marketplace, without racking your own hardware or committing to one cloud's spec sheet.
See also
Data Parallelism
Data parallelism copies the whole model onto every GPU, splits each batch between them and averages gradients with an all-reduce. Simple, but memory-hungry.
NCCL (NVIDIA Collective Communications Library)
NCCL is NVIDIA's library for all-reduce and other multi-GPU communication. PyTorch uses it to sync GPUs, and NVLink vs PCIe decides how fast it runs.
QLoRA
QLoRA fine-tunes an LLM by freezing its weights in 4-bit and training small LoRA adapters on top, so a 65B model tunes on one 48GB GPU.
VRAM
VRAM is the memory attached to a GPU that holds the data it works on, and it caps which AI models fit. VRAM vs RAM, how to check yours, and how much AI needs.
BF16 (bfloat16)
BF16 is a 16-bit float with FP32's 8 exponent bits but only 7 mantissa bits. It is the default for training and costs 2 bytes per model parameter.
NVLink vs PCIe
NVLink and PCIe both move data in and out of a GPU, but at very different scales. Where each interconnect wins on bandwidth, latency, cost, and compatibility.