FSDP: Fully Sharded Data Parallel Training at Scale

Michael BrenndoerferJanuary 22, 202647 min read

Part of Language AI Handbook

Explains how FSDP shards model parameters, gradients, and optimizer states across GPUs to train billion-parameter models.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

FSDP: Fully Sharded Data Parallel

Training large language models requires distributing both computation and memory across many GPUs. As we explored in earlier chapters on data parallelism, the standard approach replicates model weights on every GPU and splits only the data batches. This works well when the model fits in GPU memory, but modern LLMs with billions of parameters make that assumption impossible. A single GPT-3-scale model requires hundreds of gigabytes of memory just for weights, optimizer states, and gradients combined, while even the most powerful GPUs today offer only 40-80 GB of high-bandwidth memory.

Fully Sharded Data Parallel (FSDP) solves this by sharding everything across GPUs: model parameters, optimizer states, and gradients are each split into equal fragments, so every GPU holds only a 1/N1/N slice of each tensor. When a layer needs its full weights for a forward or backward pass, FSDP orchestrates an all-gather communication to reconstruct them temporarily, uses them, and then discards the reconstructed copy to reclaim memory. This approach makes it possible to train models that are NN times larger than what any single GPU could accommodate, while still achieving near-linear scaling in compute efficiency.

FSDP was introduced in PyTorch 1.12 as a production-ready implementation of the ideas pioneered by Microsoft's ZeRO optimizer. Understanding FSDP gives you the ability to train frontier-scale models using standard PyTorch, without requiring custom infrastructure or proprietary frameworks.

The Memory Problem at Scale

Before examining how FSDP works, it helps to understand precisely where memory goes during training. The memory breakdown directly determines which GPU configurations can run which models, and misunderstanding it leads to mysterious out-of-memory crashes in production training runs.

Four Components of Training Memory

For a model with PP parameters stored in float32, memory consumption breaks down across four components:

  • Parameters: 4P4P bytes (4 bytes per float32 value)
  • Gradients: 4P4P bytes (same shape as parameters, one gradient per parameter)
  • Optimizer states: 8P8P bytes for Adam (first and second moment estimates, each 4 bytes per parameter)
  • Activations: variable, roughly proportional to batch size and sequence length

The total for a training step without any memory optimization is approximately 16P16P bytes. For GPT-3 with P=175P = 175 billion parameters, that amounts to:

16×175×109 bytes≈2.8 terabytes16 \times 175 \times 10^9 \text{ bytes} \approx 2.8 \text{ terabytes}

That figure represents nearly 70 A100 GPUs worth of memory, just for the model state, before counting activations. Even the humble 7B-parameter models that became common in the open-source community after the Llama release require:

16×7×109 bytes≈112 GB16 \times 7 \times 10^9 \text{ bytes} \approx 112 \text{ GB}

A single A100 with 80 GB of memory cannot hold the full training state for even a 7B model, let alone its activations.

Why Optimizer States Are the Biggest Culprit

Many people new to large model training expect parameters to dominate memory. In practice, the Adam optimizer states are the largest single component at 8P8P bytes, accounting for half of the 16P16P total. This happens because Adam maintains two running averages per parameter: the first moment (exponential moving average of gradients) and the second moment (exponential moving average of squared gradients). Both are kept in float32 for numerical stability even when training in lower precision, because accumulated rounding errors in these statistics degrade convergence.

When using mixed precision training, the picture shifts somewhat. Parameters and gradients may be stored in float16 or bfloat16 (2 bytes each), but the optimizer states and a master copy of the parameters remain in float32. A common mixed precision accounting gives approximately 2P2P (fp16 params) +2P+ 2P (fp16 grads) +4P+ 4P (fp32 master params) +8P+ 8P (fp32 optimizer states) =16P= 16P bytes total, the same as pure float32. The mixed precision savings come from smaller activation memory (since intermediate activations are in fp16) and from faster matrix multiplications on hardware with Tensor Cores, not from reduced parameter or optimizer state memory.

Why Data Parallelism Does Not Help

Standard DistributedDataParallel (DDP) replicates the full model on every GPU. With 8 GPUs, you have 8 complete copies of the model, 8 complete sets of gradients, and 8 complete sets of optimizer states. The total memory consumption is exactly the same per GPU as running on a single device. The benefit of DDP is purely throughput: you can process 8 times as many data samples per step, but you cannot train a model that does not fit on one GPU.

This is the fundamental limitation that sharding addresses. Rather than replicating the model across GPUs, sharding divides it. Each GPU is responsible for a fraction of the total parameter space, and the GPUs communicate to temporarily share parameters when computation requires them.

Memory vs. Compute Scaling

Memory requirements scale linearly with parameter count, while compute (FLOPs) scales proportionally with parameters and sequence length. As models grow, memory becomes the binding constraint first. FSDP directly targets this constraint.

FSDP Architecture

FSDP works by sharding the three memory-intensive quantities, namely parameters, gradients, and optimizer states, across all participating GPUs. The elegance of the design is that the sharding is largely transparent to the model code: from the perspective of each individual layer, it receives its full parameter tensor when it needs it, computes normally, and returns its result. The orchestration of gathering and scattering tensors happens behind the scenes.

What Sharding Means Mathematically

Consider a linear layer with a weight matrix WW of shape (dout,din)(d_{\text{out}}, d_{\text{in}}), containing PW=dout×dinP_W = d_{\text{out}} \times d_{\text{in}} total parameters. With NN GPUs, FSDP flattens WW into a 1D tensor of length PWP_W and assigns an equal contiguous slice to each GPU:

  • GPU 0 holds elements 00 through ⌊PW/N⌋−1\lfloor P_W/N \rfloor - 1
  • GPU 1 holds elements ⌊PW/N⌋\lfloor P_W/N \rfloor through ⌊2PW/N⌋−1\lfloor 2P_W/N \rfloor - 1
  • GPU kk holds elements ⌊kPW/N⌋\lfloor kP_W/N \rfloor through ⌊(k+1)PW/N⌋−1\lfloor (k+1)P_W/N \rfloor - 1

Each GPU also holds a corresponding 1/N1/N shard of the gradient tensor for those same parameters, and a 1/N1/N shard of the Adam optimizer states (first moment mm and second moment vv) for those same parameter positions. This three-way alignment is what makes the local optimizer step possible: when it is time to update parameters, GPU kk has all the information it needs to perform the Adam update for its slice without any communication.

The critical constraint is that computing the forward pass through a linear layer requires the complete weight matrix WW, not a shard of it. An input xx of shape (B,din)(B, d_{\text{in}}) needs to be multiplied by the full WW to produce output of shape (B,dout)(B, d_{\text{out}}). FSDP handles this requirement with explicit communication phases.

The All-Gather and Reduce-Scatter Cycle

FSDP wraps each module (or group of modules) in a unit called an FSDP unit. These units are the atoms of the sharding system: all-gather and reduce-scatter operations occur at FSDP unit boundaries. Understanding the communication pattern for one unit clarifies the entire system.

During the forward pass, the following sequence occurs for each FSDP unit:

  1. Every GPU broadcasts its local parameter shard to all other GPUs via an all-gather. After this operation, all GPUs temporarily hold the complete parameter tensor for this unit.
  2. Each GPU uses the full parameter tensor to compute its portion of the forward pass (processing its slice of the data batch).
  3. Immediately after the computation completes, FSDP discards the reconstructed parameter tensor on each GPU, leaving only the original local shard. This is the key memory reclamation step.

During the backward pass, the cycle repeats with one addition:

  1. An all-gather reconstructs the full parameter tensor again (necessary for computing parameter gradients).
  2. The backward computation runs, producing full-parameter gradient tensors on each GPU.
  3. A reduce-scatter aggregates gradients across all GPUs and distributes the result. Unlike all-reduce (which gives every GPU the full summed gradient), reduce-scatter performs the sum and simultaneously shards it: GPU kk receives only the portion of the gradient corresponding to its parameter shard.
  4. The reconstructed parameter tensor is discarded again.

The optimizer step then runs locally on each GPU: each GPU updates its own parameter shard using its own gradient shard and its own optimizer state shard. No communication is required at this step, because the data for each parameter's update resides entirely on the GPU responsible for that parameter.

This complete cycle, all-gather then compute then discard (forward), all-gather then compute then reduce-scatter (backward), local optimizer step, is what defines FSDP's operation. The name "fully sharded" reflects the fact that all three memory-intensive components remain sharded at rest; they are only temporarily gathered when computation demands it.

All-Gather vs. All-Reduce

Standard data parallelism uses all-reduce at the end of backward to sum gradients across all GPUs, giving each GPU the full summed gradient. FSDP uses reduce-scatter instead, which sums gradients but also shards the result. Each GPU receives only the portion it owns. This saves memory because no GPU ever holds a complete gradient tensor.

Memory Profile Comparison

Under FSDP with NN GPUs, the memory breakdown per GPU changes dramatically compared to DDP:

  • Parameters (steady state): 4P/N4P/N bytes (only the local shard persists)
  • Gradients (steady state): 4P/N4P/N bytes (only the local gradient shard persists)
  • Optimizer states: 8P/N8P/N bytes (first and second moments for the local shard)
  • All-gather buffer during forward: 4P4P bytes temporarily (freed after each layer's computation)

The steady-state memory per GPU is approximately (4+4+8)P/N=16P/N(4 + 4 + 8)P/N = 16P/N bytes for the sharded quantities. The peak memory includes the temporary all-gather buffer of 4P4P bytes during computation for the current FSDP unit, but this is freed right after the unit completes. For a model split into MM FSDP units, the peak at any instant is 16P/N+4P/M16P/N + 4P/M bytes (steady-state plus one unit's full parameters).

Compare this to DDP, where each GPU holds 16P16P bytes at all times. The reduction factor from 16P16P to 16P/N16P/N for the steady state is exactly NN, the number of GPUs. This linear scaling in memory capacity with GPU count is what makes FSDP so powerful for scaling to large models.

To see this concretely: a 13B parameter model requires about 208 GB of training memory (16×1316 \times 13 GB). With FSDP on 8 GPUs of 40 GB each (320 GB total), each GPU holds only 208/8=26208/8 = 26 GB of steady-state memory, fitting comfortably within 40 GB when including activation buffers.

FSDP vs. ZeRO

FSDP and Microsoft's ZeRO optimizer (used in DeepSpeed) address the same fundamental problem and use nearly identical techniques. Understanding their relationship clarifies why FSDP exists and how to choose between them.

Historical Context: ZeRO's Origins

The ZeRO paper, "ZeRO: Memory Optimizations Toward Training Trillion Parameter Models," was published by Rajbhandari et al. at Microsoft Research in 2020. At the time, the dominant approach for large model training was model parallelism, which splits a model across GPUs by assigning different layers to different devices. Model parallelism requires careful pipeline management to keep all GPUs busy and introduces complex inter-layer communication patterns. ZeRO proposed a different insight: instead of splitting the model structurally, shard the redundant state that data parallelism was already replicating across GPUs.

The insight was that in standard DDP with NN GPUs, the optimizer states, gradients, and even parameters are all replicated NN times. Each replica is identical; the replication exists purely for convenience, not necessity. By eliminating this redundancy through sharding, you can reduce per-GPU memory by a factor of NN while preserving the data-parallel communication structure.

ZeRO introduced three progressive stages of sharding, which became the conceptual framework that FSDP adopted.

ZeRO's Three Stages

ZeRO (Zero Redundancy Optimizer) defines three progressive stages of memory optimization:

  • ZeRO Stage 1: Shards only optimizer states across GPUs. Each GPU holds 1/N1/N of the optimizer states but the full parameters and gradients. This is the least disruptive change, requiring only that the optimizer step communicate parameter shards back to all GPUs after updating.

  • ZeRO Stage 2: Shards optimizer states and gradients. Each GPU holds 1/N1/N of optimizer states and 1/N1/N of gradients but the full parameters. A reduce-scatter replaces the all-reduce for gradient synchronization.

  • ZeRO Stage 3: Shards all three components. Each GPU holds 1/N1/N of parameters, gradients, and optimizer states. All-gather operations are required before each layer's computation. This is the mode with maximum memory savings and highest communication overhead.

ZeRO Stage 3 is conceptually identical to what FSDP with FULL_SHARD implements. Both shard all three memory components and use all-gather plus reduce-scatter communication patterns. The algorithms are the same at the mathematical level.

Key Differences in Practice

The practical differences between FSDP and ZeRO Stage 3 are largely about ecosystem and integration rather than algorithmic substance.

FSDP is a native PyTorch feature. It integrates directly with the PyTorch autograd engine, requires no additional dependencies, and benefits from PyTorch's ongoing optimization work. Because it is part of the standard library, it works out of the box with tools like Hugging Face Accelerate, PyTorch Lightning, and most training frameworks that target PyTorch. The API is designed around PyTorch idioms: you wrap a model with FSDP() the same way you wrap with DDP(), and the training loop is otherwise unchanged.

DeepSpeed's ZeRO Stage 3 has a richer feature set around communication optimization, CPU offloading, and mixed precision handling. It also offers ZeRO-Infinity, which extends sharding to CPU and NVMe storage, enabling model sizes far beyond what fits in combined GPU memory. For researchers who need to push the absolute limit of model size on existing hardware, ZeRO-Infinity provides capabilities that FSDP does not yet match.

For most practical use cases, FSDP is the simpler choice if you are already in the PyTorch ecosystem. DeepSpeed is preferable when you need its extended optimization features or CPU offloading capabilities.

Comparison of FSDP and DeepSpeed ZeRO Stage 3 across key features.
FeatureFSDPZeRO Stage 3 (DeepSpeed)
IntegrationNative PyTorchRequires DeepSpeed installation
Parameter shardingYesYes
Gradient shardingYesYes
Optimizer state shardingYesYes
CPU offloadingLimitedFull ZeRO-Infinity support
Ease of useHighModerate
Communication optimizationGoodMore configurable

Sharding Strategies

PyTorch FSDP exposes several sharding strategies through the ShardingStrategy enum. Each strategy offers a different trade-off between memory savings and communication overhead. Choosing the right strategy requires understanding both the memory equations and the communication patterns, because on bandwidth-limited hardware the communication cost can exceed the memory savings benefit.

FULL_SHARD

FULL_SHARD is the most aggressive strategy. It shards parameters, gradients, and optimizer states across all GPUs. This corresponds to ZeRO Stage 3 and delivers the maximum possible memory reduction per GPU.

The communication cost is also highest with this strategy, because all-gather operations must run before every forward and backward pass through each FSDP unit. For a model split into MM FSDP units, one complete training step requires:

  • MM all-gather operations during forward pass
  • MM all-gather operations during backward pass
  • MM reduce-scatter operations during backward pass

The total communication volume per step is proportional to 3MP′3MP' where P′P' is the average parameters per FSDP unit. Since MP′=PMP' = P (total parameters), this simplifies to 3P×bytes_per_param3P \times \text{bytes\_per\_param} in aggregated data moved. Compare this to DDP, which requires one all-reduce of 2P×bytes_per_param2P \times \text{bytes\_per\_param} per step (factor of 2 comes from the ring-allreduce algorithm needing two passes). FSDP's communication volume is roughly 1.5 times larger than DDP per step, which is the price for the memory savings.

Use FULL_SHARD when:

  • The model does not fit in GPU memory without sharding
  • You have fast interconnects (NVLink, InfiniBand) that can absorb the communication overhead
  • You are training at scale and memory efficiency is the primary constraint

SHARD_GRAD_OP

SHARD_GRAD_OP shards gradients and optimizer states but keeps full parameter copies on each GPU during the forward and backward passes. This corresponds to ZeRO Stage 2.

Since parameters are not sharded, no all-gather is needed before each layer's computation. The communication happens only at the end of backward, where a reduce-scatter replaces the all-reduce of standard data parallelism. This reduces communication frequency significantly compared to FULL_SHARD: instead of 2M2M all-gathers plus MM reduce-scatters, you get just one reduce-scatter per step.

The memory savings are smaller. Parameters (4P4P bytes) are not sharded, so each GPU still holds the full parameter tensor. Only gradients (4P4P bytes) and optimizer states (8P8P bytes) are split by NN, saving:

4P+8P−(4P/N+8P/N)=12P(1−1/N)4P + 8P - (4P/N + 8P/N) = 12P(1 - 1/N)

For large NN, this approaches 12P12P bytes savings, leaving only the 4P4P parameter bytes per GPU. Compare to FULL_SHARD, which brings per-GPU memory to (16P/N+4P/M)(16P/N + 4P/M) bytes. The gap between the two strategies narrows as NN increases: on 64 GPUs, SHARD_GRAD_OP saves 12P×63/64≈11.8P12P \times 63/64 \approx 11.8P bytes, while FULL_SHARD saves nearly 15.75P15.75P bytes.

Use SHARD_GRAD_OP when:

  • The model fits in GPU memory during the forward and backward passes (just the weights)
  • You want to reduce optimizer memory usage without the communication overhead of full sharding
  • Training throughput is more important than maximum memory reduction

NO_SHARD

NO_SHARD disables all sharding. FSDP with this strategy behaves exactly like DistributedDataParallel (DDP): each GPU holds the full model, and an all-reduce synchronizes gradients after each backward pass.

This strategy exists primarily so that code written to use the FSDP API can be run without any actual sharding, for debugging or profiling, without changing the training code. It also is a useful baseline for measuring the overhead that sharding introduces, since you can compare throughput with NO_SHARD against SHARD_GRAD_OP or FULL_SHARD on the same hardware with the same code.

HYBRID_SHARD

HYBRID_SHARD is designed for multi-node training where intra-node bandwidth (via NVLink) is much higher than inter-node bandwidth (via InfiniBand or Ethernet). This asymmetry is extremely common in real clusters: eight GPUs on the same node may be connected by NVLink at 600 GB/s bidirectional, while the nodes connect to each other at 200 Gb/s InfiniBand (roughly 25 GB/s, 24 times slower).

With HYBRID_SHARD, FSDP applies full sharding within each node (using the fast intra-node interconnect) but maintains full parameter replicas across nodes. Gradient synchronization happens between nodes at the end of each backward pass, similar to DDP, but only synchronizing the gradients that each node already has aggregated.

The key benefit is that the expensive all-gather communication for parameter reconstruction stays within the node, using the fast NVLink. Cross-node communication is limited to one gradient all-reduce per step, which is the same communication pattern as standard DDP. This avoids sending all-gather communication across the slower inter-node network, which would be the bottleneck with FULL_SHARD at scale.

The cost is that memory savings are limited to within each node: if each node has 8 GPUs, parameters are sharded across those 8 GPUs but replicated across nodes. For a 70B model on a 4-node cluster of 8 GPUs each (32 GPUs total), HYBRID_SHARD gives you the memory benefit of 8-GPU sharding (saving 7/87/8 of parameter memory within each node), while FULL_SHARD would give 32-GPU sharding savings (31/3231/32).

Use HYBRID_SHARD when:

  • You are training on multiple nodes
  • Intra-node bandwidth (NVLink) significantly exceeds inter-node bandwidth
  • The model fits in memory when sharded across GPUs within a single node

_HYBRID_SHARD_ZERO2

_HYBRID_SHARD_ZERO2 combines the intra-node sharding of HYBRID_SHARD with ZeRO Stage 2 semantics: parameters are kept full within each node during computation, but gradients and optimizer states are sharded intra-node while being replicated across nodes.

The leading underscore indicates this is a less commonly used variant. It is useful when you want the communication profile of SHARD_GRAD_OP within each node combined with the cross-node replication of HYBRID_SHARD. Concretely, this means no all-gathers for parameter reconstruction during forward and backward, which is beneficial when the bottleneck is communication latency rather than memory.

Wrapping Policies

A key decision in FSDP setup is how to group model submodules into FSDP units. Each FSDP unit is the granularity at which all-gather and reduce-scatter operations occur: finer granularity means smaller all-gather buffers and lower peak memory, but more frequent communication and higher overhead from kernel launches.

Getting the wrapping policy right matters more than many practitioners expect. A policy that is too fine-grained can reduce throughput by 2-3x compared to an optimal policy, because GPU kernels have significant launch overhead that becomes dominant when individual operations are small. A policy that is too coarse can cause out-of-memory errors because the all-gather buffer for a single large unit might not fit.

Auto Wrap Policy

The transformer_auto_wrap_policy function in PyTorch's FSDP utilities wraps each transformer layer (self-attention block, FFN block, etc.) as a separate FSDP unit. This is the recommended starting point for transformer models.

The auto wrap policy accepts a set of module classes to wrap. For a standard transformer, you would typically pass the class corresponding to a single transformer block, for example GPTNeoXLayer for GPT-NeoX or LlamaDecoderLayer for Llama. The embedding and output head layers are wrapped at the top level (as part of the root FSDP module) unless you explicitly specify otherwise.

In[4]:
Code
import functools

from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy

# Example for a Llama-style model
auto_wrap_policy = functools.partial(
    transformer_auto_wrap_policy,
    transformer_layer_cls={LlamaDecoderLayer},
)

The reason this wrapping strategy works well for transformers is that each transformer layer is approximately equal in size (assuming the same hidden dimension and FFN expansion ratio across all layers), so all FSDP units have similar all-gather buffer sizes. Uniform unit sizes lead to uniform communication times, which maps well to overlapping communication with computation: while one layer runs its forward pass, FSDP can prefetch (all-gather) the next layer's parameters.

Size-Based Wrap Policy

The size_based_auto_wrap_policy wraps modules based on their parameter count. Modules with fewer than a threshold number of parameters are merged with their parent, reducing the total number of FSDP units and thus the communication frequency.

In[3]:
Code
import functools

from torch.distributed.fsdp.wrap import size_based_auto_wrap_policy

size_wrap_policy = functools.partial(
    size_based_auto_wrap_policy,
    min_num_params=1_000_000,  # wrap modules with at least 1M parameters
)

This policy is useful for models that do not follow the uniform-transformer-layer pattern, such as encoder-decoder models with asymmetric components or models with large embedding tables followed by small classification heads.

Manual Wrapping

For full control, you can manually wrap specific submodules by calling FSDP() on them before wrapping the parent model. This is useful when different parts of the model require different sharding strategies or mixed precision settings. For example, you might want the main transformer layers to use FULL_SHARD with bfloat16, but keep the embedding table in float32 with NO_SHARD because it is accessed sparsely and does not benefit from sharding.

Communication Prefetching

A subtle but important performance feature of FSDP is prefetching: the ability to issue the all-gather for the next FSDP unit while the current unit is still computing. PyTorch FSDP implements this automatically for the forward pass by default, overlapping the communication for layer k+1k+1 with the computation for layer kk. This overlap is only possible when the wrapping granularity is at the layer level (one FSDP unit per layer), not at coarser granularity where there is no "next unit" to prefetch.

The backward_prefetch parameter controls whether the same overlap happens during the backward pass. By default, PyTorch prefetches during backward as well, which can significantly improve throughput on hardware where communication and compute can run concurrently.

Mixed Precision with FSDP

FSDP has first-class support for mixed precision training through the MixedPrecision policy. This lets you specify different dtypes for parameters, gradients, and buffer tensors independently.

Why Mixed Precision Matters at Scale

On modern GPUs, bfloat16 matrix multiplications run significantly faster than float32 because the hardware has dedicated Tensor Core units optimized for reduced-precision arithmetic. On A100 GPUs, bfloat16 matrix multiplications deliver up to 312 TFLOPS, compared to 77.6 TFLOPS for float32, a 4x speedup. Beyond raw throughput, reduced-precision tensors are half the size of float32, which matters both for the all-gather buffer (which temporarily holds full-precision parameters) and for activations accumulated during the forward pass.

The challenge with reduced precision is numerical stability: float16 has a limited exponent range (maximum value ≈65,504\approx 65,504), causing overflow during training when gradients or activations become large. Bfloat16 solves this by using the same exponent range as float32 (8 bits), sacrificing mantissa precision instead (7 bits vs. 23 bits for float32). For most training scenarios, bfloat16 provides the speed benefits of fp16 without the overflow instability.

The standard mixed precision configuration for large model training stores reduced-precision parameters during computation (to speed up matrix multiplications and save memory on the all-gather buffer) while keeping full-precision optimizer states for numerical stability.

In[4]:
Code
import torch
from torch.distributed.fsdp import MixedPrecision

# BF16 mixed precision policy
bf16_policy = MixedPrecision(
    param_dtype=torch.bfloat16,  # parameters stored and communicated in BF16
    reduce_dtype=torch.float32,  # gradients reduced in FP32 for numerical stability
    buffer_dtype=torch.bfloat16,  # buffers (e.g., layer norm running stats) in BF16
)

Setting reduce_dtype=torch.float32 ensures that gradient accumulation happens in float32, preventing the loss of small gradient values that would occur if you accumulated in bfloat16 (where the minimum representable positive value is much larger than in float32). The parameters and all-gather buffers are in bfloat16, giving you the memory and compute benefits, while the gradients that drive the optimizer update retain float32 precision.

BF16 is generally preferred over FP16 for large model training because its exponent range matches FP32, eliminating the need for loss scaling that FP16 requires. Loss scaling is a workaround where gradients are multiplied by a large constant before the backward pass to prevent fp16 underflow, then divided back before the optimizer step. Managing loss scaling adds code complexity and can cause instability if the scale factor is poorly chosen.

Worked Example: Memory Savings on a 7B Model

To build intuition for the magnitudes involved, let's trace through the memory accounting for a concrete scenario: training a 7 billion parameter model on 8 A100-40GB GPUs.

Without FSDP (DDP Baseline)

With standard DDP, every GPU holds the complete model state:

  • Parameters: 4×7×109=284 \times 7 \times 10^9 = 28 GB
  • Gradients: 2828 GB
  • Optimizer states (Adam): 5656 GB
  • Total model state: 112112 GB

This already exceeds the 40 GB GPU capacity by 2.8×2.8\times. DDP cannot run this configuration at all.

With FSDP FULL_SHARD on 8 GPUs

With FULL_SHARD, each GPU holds 1/81/8 of each component:

  • Parameters (shard): 28/8=3.528 / 8 = 3.5 GB
  • Gradients (shard): 3.53.5 GB
  • Optimizer states (shard): 77 GB
  • Total steady-state: 1414 GB

During the forward and backward passes, a single FSDP unit's all-gather buffer temporarily consumes an additional 28/M28/M GB where MM is the number of FSDP units. For a 32-layer model wrapped at the layer level, each all-gather buffer is 28/32≈0.87528/32 \approx 0.875 GB, bringing peak memory to about 14.87514.875 GB per GPU, well within the 40 GB limit.

The remaining memory budget of approximately 25 GB can accommodate activations from long sequence lengths and larger batch sizes. Activations for a typical transformer layer scale as O(B×S×d)O(B \times S \times d) where BB is batch size, SS is sequence length, and dd is hidden dimension. For a 7B model (d=4096d = 4096), a batch of 4 sequences of length 2048 in bfloat16 generates roughly 4×2048×4096×2≈674 \times 2048 \times 4096 \times 2 \approx 67 MB per layer before considering the activations needed for attention score computation. With activation checkpointing (recomputing rather than storing activations), this budget extends further.

The Practical Significance

This example illustrates why FSDP enabled more open-source large model training in 2023 and 2024. Before FSDP became widely available and well-documented, training 7B models required either expensive 80 GB GPUs (A100-80GB or H100), proprietary systems like Google's TPU pods, or careful hand-engineering with DeepSpeed. With FSDP on standard 40 GB A100 GPUs, which were broadly available through cloud providers, the research community could train and fine-tune these models on accessible hardware.

Code Implementation

This section walks through a complete FSDP training setup, starting from a small model definition and building up to a full distributed training loop.

Setup and Imports

We start with the necessary imports and a simple model definition.

Defining a Simple Transformer Model

We define a small transformer model with a configurable number of layers and dimensions. This is purely for demonstration; in practice you would use a model from Hugging Face or your own architecture.

In[6]:
Code
import torch.nn as nn


class TransformerLayer(nn.Module):
    """A single transformer encoder layer."""

    def __init__(self, d_model=256, nhead=4, dim_feedforward=1024):
        super().__init__()
        self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
        self.ff = nn.Sequential(
            nn.Linear(d_model, dim_feedforward),
            nn.GELU(),
            nn.Linear(dim_feedforward, d_model),
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)

    def forward(self, x):
        attn_out, _ = self.self_attn(x, x, x)
        x = self.norm1(x + attn_out)
        x = self.norm2(x + self.ff(x))
        return x


class SimpleTransformer(nn.Module):
    """A stack of transformer layers with an embedding and output head."""

    def __init__(
        self,
        vocab_size=10000,
        d_model=256,
        nhead=4,
        dim_feedforward=1024,
        num_layers=6,
    ):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, d_model)
        self.layers = nn.ModuleList(
            [
                TransformerLayer(d_model, nhead, dim_feedforward)
                for _ in range(num_layers)
            ]
        )
        self.output_head = nn.Linear(d_model, vocab_size)

    def forward(self, input_ids):
        x = self.embedding(input_ids)
        for layer in self.layers:
            x = layer(x)
        return self.output_head(x)

Counting Parameters

Let's inspect the parameter count to understand the memory requirements.

In[7]:
Code
def count_parameters(model):
    total = sum(p.numel() for p in model.parameters())
    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
    return total, trainable


model = SimpleTransformer(
    vocab_size=10000,
    d_model=256,
    nhead=4,
    dim_feedforward=1024,
    num_layers=6,
)
total_params, trainable_params = count_parameters(model)
Out[8]:
Console
Total parameters:     9,868,560
Trainable parameters: 9,868,560
Memory (FP32):        39.5 MB
Memory (with Adam):   157.9 MB

The output shows parameter count and estimated memory footprint. The multiplier of 16 accounts for parameters (4 bytes), gradients (4 bytes), and Adam optimizer states (8 bytes: first moment plus second moment, both in float32).

FSDP Initialization

This function wraps a model with FSDP. In a real distributed job, it must be called after dist.init_process_group() so that the process group is available for communication.

In[9]:
Code
import functools

from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
)
from torch.distributed.fsdp import (
    MixedPrecision,
    ShardingStrategy,
)
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy


def wrap_model_with_fsdp(model, sharding_strategy=ShardingStrategy.FULL_SHARD):
    """
    Wrap a transformer model with FSDP.

    Parameters:
    - model: The PyTorch model to wrap
    - sharding_strategy: The FSDP sharding strategy to use

    Returns the FSDP-wrapped model.
    """
    # Define the auto-wrap policy to wrap each TransformerLayer separately
    auto_wrap_policy = functools.partial(
        transformer_auto_wrap_policy,
        transformer_layer_cls={TransformerLayer},
    )

    # Mixed precision configuration
    mp_policy = MixedPrecision(
        param_dtype=torch.float16,
        reduce_dtype=torch.float32,
        buffer_dtype=torch.float16,
    )

    # Wrap the model
    fsdp_model = FSDP(
        model,
        auto_wrap_policy=auto_wrap_policy,
        sharding_strategy=sharding_strategy,
        mixed_precision=mp_policy,
        device_id=torch.cuda.current_device(),
    )
    return fsdp_model

Simulating Memory Usage per Strategy

Since we cannot run distributed training in this notebook, we can compute the theoretical memory savings analytically for different sharding strategies.

In[10]:
Code
def memory_per_gpu(total_params, num_gpus, strategy):
    """
    Compute estimated GPU memory in bytes for each sharding strategy.

    Parameters:
    - total_params: total number of model parameters
    - num_gpus: number of GPUs in the training job
    - strategy: one of 'full_shard', 'shard_grad_op', 'no_shard'

    Returns a dict with breakdown by component.
    """
    bytes_per_param = 4  # float32

    if strategy == "no_shard":
        params = total_params * bytes_per_param
        grads = total_params * bytes_per_param
        optimizer = total_params * 2 * bytes_per_param  # Adam: 2x
    elif strategy == "shard_grad_op":
        params = total_params * bytes_per_param  # not sharded
        grads = total_params * bytes_per_param / num_gpus
        optimizer = total_params * 2 * bytes_per_param / num_gpus
    elif strategy == "full_shard":
        params = total_params * bytes_per_param / num_gpus
        grads = total_params * bytes_per_param / num_gpus
        optimizer = total_params * 2 * bytes_per_param / num_gpus
    else:
        raise ValueError(f"Unknown strategy: {strategy}")

    total = params + grads + optimizer
    return {
        "params_mb": params / 1e6,
        "grads_mb": grads / 1e6,
        "optimizer_mb": optimizer / 1e6,
        "total_mb": total / 1e6,
    }


strategies = ["no_shard", "shard_grad_op", "full_shard"]
num_gpus_options = [1, 4, 8, 16]
Out[11]:
Console
Strategy         GPUs   Params (MB)    Grads (MB)    Optim (MB)    Total (MB)  
--------------------------------------------------------------------------------
no_shard         1      39.5           39.5          78.9          157.9       
no_shard         4      39.5           39.5          78.9          157.9       
no_shard         8      39.5           39.5          78.9          157.9       
no_shard         16     39.5           39.5          78.9          157.9       

shard_grad_op    1      39.5           39.5          78.9          157.9       
shard_grad_op    4      39.5           9.9           19.7          69.1        
shard_grad_op    8      39.5           4.9           9.9           54.3        
shard_grad_op    16     39.5           2.5           4.9           46.9        

full_shard       1      39.5           39.5          78.9          157.9       
full_shard       4      9.9            9.9           19.7          39.5        
full_shard       8      4.9            4.9           9.9           19.7        
full_shard       16     2.5            2.5           4.9           9.9

The table reveals the progressive memory reduction across strategies. With full_shard on 16 GPUs, each GPU holds just 1/16 of each component, a 16x reduction in steady-state memory. Notice that shard_grad_op achieves a substantial reduction too because the optimizer states dominate: even though parameters are replicated, sharding the optimizer states (8P8P bytes) alone saves significantly more memory than sharding parameters alone (4P4P bytes).

Visualizing Memory Savings

Out[12]:
Visualization
Line chart comparing GPU memory usage vs number of GPUs for FULL_SHARD, SHARD_GRAD_OP, and NO_SHARD strategies.
Estimated GPU memory per device (MB) for three FSDP sharding strategies across different GPU counts. FULL_SHARD achieves the steepest reduction because it shards all three memory components: parameters, gradients, and optimizer states. SHARD_GRAD_OP achieves moderate savings by keeping full parameters but sharding gradients and optimizer states. NO_SHARD (equivalent to DDP) shows no reduction as GPU count increases, confirming that replication alone does not reduce per-device memory.

Visualizing Communication Volume

With FULL_SHARD, each forward pass triggers an all-gather for every FSDP unit. The total communication volume per training step is larger than with standard DDP. Let's quantify this difference.

The communication overhead is the fundamental trade-off of FSDP: you gain memory capacity in exchange for more communication. Whether this trade-off is favorable depends entirely on your hardware's interconnect speed relative to compute speed. On a cluster where GPUs are connected by slow Ethernet, the extra communication turns a compute-bound problem into a communication-bound one. On an NVLink-connected server, the all-gather operations happen at memory speed and the overhead is negligible.

In[13]:
Code
def communication_volume_gb(total_params, num_gpus, num_fsdp_units, strategy):
    """
    Estimate communication volume per training step in GB.

    For FULL_SHARD:
    - Forward: num_fsdp_units all-gathers, each moving total_params/num_fsdp_units params
    - Backward: num_fsdp_units all-gathers + num_fsdp_units reduce-scatters

    For SHARD_GRAD_OP:
    - No all-gathers during forward/backward
    - One reduce-scatter at end of backward

    For NO_SHARD (DDP):
    - One all-reduce at end of backward
    """
    bytes_per_param = 4

    if strategy == "full_shard":
        # Forward: M all-gathers, each reconstructing P/M params across N GPUs
        # Communication volume per all-gather = P * (N-1)/N * 2 (send + receive)
        # Simplified as P * 2 bytes (one full copy across ring)
        fwd_volume = (
            total_params * bytes_per_param * 2
        )  # all-gathers in forward
        # Backward: M all-gathers + M reduce-scatters
        bwd_volume = (
            total_params * bytes_per_param * 2
            + total_params * bytes_per_param * 2
        )
        total_volume = fwd_volume + bwd_volume
    elif strategy == "shard_grad_op":
        # Only one reduce-scatter at end of backward
        total_volume = total_params * bytes_per_param * 2
    elif strategy == "no_shard":
        # All-reduce = 2 * total_params * bytes (ring algorithm)
        total_volume = total_params * bytes_per_param * 2
    else:
        raise ValueError(f"Unknown strategy: {strategy}")

    return total_volume / 1e9  # GB
Out[14]:
Console
Communication volume per training step (model with 9,868,560 params):

Strategy             Volume (GB)
-----------------------------------
no_shard             0.0789 GB
shard_grad_op        0.0789 GB
full_shard           0.2368 GB

The output illustrates the communication cost trade-off: FULL_SHARD incurs more communication per step because of the all-gather operations during both forward and backward passes. For this small demonstration model, the absolute volumes are tiny. For a 70B model on 8 GPUs, FULL_SHARD would move approximately 70×109×4×6≈1.6870 \times 10^9 \times 4 \times 6 \approx 1.68 TB per step (before counting the reduce-scatters), while SHARD_GRAD_OP and NO_SHARD would each move about 560560 GB. This is why interconnect bandwidth is the deciding factor: on high-bandwidth networks, even the larger communication volumes complete in a fraction of the compute time.

Visualizing Communication vs. Memory Trade-off

Out[15]:
Visualization
Scatter plot showing communication volume vs memory per GPU for three FSDP strategies, illustrating the memory-communication trade-off.
Trade-off between communication volume per step and steady-state model-state memory per GPU for FSDP sharding strategies with 8 GPUs. FULL_SHARD minimizes memory at the cost of higher communication volume. SHARD_GRAD_OP matches DDP's communication volume while saving optimizer and gradient memory. The preferred strategy is nearest to the lower-left corner given your hardware constraints.

Sharding Strategy Selection Guide

To help you choose the right strategy for a given scenario, here is a decision framework based on model size and hardware.

In[16]:
Code
def recommend_strategy(model_params, gpu_memory_gb, num_gpus):
    """
    Recommend an FSDP sharding strategy based on model and hardware.

    Parameters:
    - model_params: number of model parameters
    - gpu_memory_gb: memory per GPU in GB
    - num_gpus: number of available GPUs

    Returns a strategy recommendation with reasoning.
    """
    # Memory requirements in GB (FP32 + Adam)
    mem_no_shard_gb = model_params * 16 / 1e9
    mem_shard_grad_op_gb = (
        model_params * (4 + 4 / num_gpus + 8 / num_gpus) / 1e9
    )
    mem_full_shard_gb = model_params * 16 / (num_gpus * 1e9)

    results = {
        "NO_SHARD fits": mem_no_shard_gb <= gpu_memory_gb,
        "SHARD_GRAD_OP fits": mem_shard_grad_op_gb <= gpu_memory_gb,
        "FULL_SHARD fits": mem_full_shard_gb <= gpu_memory_gb,
        "mem_no_shard_gb": mem_no_shard_gb,
        "mem_shard_grad_op_gb": mem_shard_grad_op_gb,
        "mem_full_shard_gb": mem_full_shard_gb,
    }

    if results["NO_SHARD fits"]:
        recommendation = "NO_SHARD (use DDP for simplicity)"
    elif results["SHARD_GRAD_OP fits"]:
        recommendation = "SHARD_GRAD_OP (moderate savings, less communication)"
    elif results["FULL_SHARD fits"]:
        recommendation = "FULL_SHARD (maximum memory reduction)"
    else:
        recommendation = "FULL_SHARD + CPU offload (model may still not fit)"

    results["recommendation"] = recommendation
    return results
Out[17]:
Console
Scenario: 7B model, 8x A100 80GB
  NO_SHARD:      112 GB/GPU  -> fits: False
  SHARD_GRAD_OP: 38.5 GB/GPU  -> fits: True
  FULL_SHARD:    14.0 GB/GPU  -> fits: True
  Recommendation: SHARD_GRAD_OP (moderate savings, less communication)

Scenario: 13B model, 8x A100 40GB
  NO_SHARD:      208 GB/GPU  -> fits: False
  SHARD_GRAD_OP: 71.5 GB/GPU  -> fits: False
  FULL_SHARD:    26.0 GB/GPU  -> fits: True
  Recommendation: FULL_SHARD (maximum memory reduction)

Scenario: 70B model, 8x A100 80GB
  NO_SHARD:      1120 GB/GPU  -> fits: False
  SHARD_GRAD_OP: 385.0 GB/GPU  -> fits: False
  FULL_SHARD:    140.0 GB/GPU  -> fits: False
  Recommendation: FULL_SHARD + CPU offload (model may still not fit)

Scenario: 70B model, 32x A100 80GB
  NO_SHARD:      1120 GB/GPU  -> fits: False
  SHARD_GRAD_OP: 306.2 GB/GPU  -> fits: False
  FULL_SHARD:    35.0 GB/GPU  -> fits: True
  Recommendation: FULL_SHARD (maximum memory reduction)

Scenario: 175B model, 64x A100 80GB
  NO_SHARD:      2800 GB/GPU  -> fits: False
  SHARD_GRAD_OP: 732.8 GB/GPU  -> fits: False
  FULL_SHARD:    43.8 GB/GPU  -> fits: True
  Recommendation: FULL_SHARD (maximum memory reduction)

The scenarios illustrate how FULL_SHARD becomes necessary as models grow. A 70B model still needs CPU offloading or more than 8 GPUs of 80 GB each because even full sharding leaves 140 GB of model state per GPU. At 32 GPUs, FULL_SHARD reduces that footprint to 35 GB per GPU. The 175B model on 64 GPUs (175×16/64=43.75175 \times 16 / 64 = 43.75 GB per GPU for the model state) fits within 80 GB with FULL_SHARD, leaving room for activation buffers.

Visualizing Strategy Recommendations

Out[18]:
Visualization
Grouped bar chart showing memory per GPU for three FSDP strategies across five model sizes with a GPU memory limit line.
FSDP strategy feasibility across model sizes (7B to 175B parameters) with 8 GPUs of 80 GB each. Bars show model-state memory per GPU, and the red dashed line marks the 80 GB limit. FULL_SHARD fits the 7B, 13B, and 30B models; none of the strategies fits 70B or 175B without CPU offloading or additional GPUs.

Complete Training Loop

Here is a self-contained training loop that demonstrates how FSDP fits into the standard PyTorch training workflow. In a real distributed setting, this function would run on each GPU process.

In[19]:
Code
def create_synthetic_batch(batch_size=4, seq_len=128, vocab_size=10000):
    """Generate a synthetic batch of token IDs and targets."""
    input_ids = torch.randint(0, vocab_size, (batch_size, seq_len))
    targets = torch.randint(0, vocab_size, (batch_size, seq_len))
    return input_ids, targets


def compute_loss(logits, targets):
    """Cross-entropy loss for language modeling."""
    # logits: (batch, seq_len, vocab_size) -> reshape for cross-entropy
    batch_size, seq_len, vocab_size = logits.shape
    return nn.functional.cross_entropy(
        logits.view(batch_size * seq_len, vocab_size),
        targets.view(batch_size * seq_len),
    )
In[20]:
Code
def simulate_training_step(model, optimizer, num_steps=20, seed=42):
    """
    Overfit one fixed synthetic batch without distributed setup.
    In practice this runs on each GPU process with an FSDP-wrapped model.
    """
    torch.manual_seed(seed)
    model.train()
    losses = []
    input_ids, targets = create_synthetic_batch()

    for step in range(num_steps):
        optimizer.zero_grad()
        logits = model(input_ids)
        loss = compute_loss(logits, targets)
        loss.backward()
        optimizer.step()
        losses.append(loss.item())

    return losses


# Run simulation on the (non-distributed) model for demonstration
optimizer = torch.optim.Adam(model.parameters(), lr=3e-4)
training_losses = simulate_training_step(model, optimizer, num_steps=30)
Out[21]:
Console
Training completed: 30 steps
Initial loss: 9.3501
Final loss:   1.9261
Loss reduction: 79.4%

The training loop is identical in structure whether you are using a plain model, DDP, or FSDP. The only differences in a real FSDP job are the process group initialization, the FSDP() wrapping call, and the checkpoint saving procedure. This is a deliberate design choice: FSDP should be a drop-in replacement for DDP in the training loop, with complexity hidden in the wrapper.

Training Loss Visualization

Out[22]:
Visualization
Line chart showing cross-entropy loss steadily decreasing while a transformer overfits one fixed synthetic batch for 30 steps.
Cross-entropy loss while deliberately overfitting one fixed synthetic batch for 30 steps with the SimpleTransformer model. The steady decline verifies that the optimizer updates the model and gradients flow through the training loop; it demonstrates mechanics, not generalization to new data.

Checkpointing with FSDP

Checkpointing is more complex with FSDP than with standard models because parameters are sharded across GPUs. In a standard PyTorch model, model.state_dict() returns a dictionary of all parameter tensors. With FSDP, each GPU only holds shards, so a naive state_dict() call would return a dictionary of shards rather than complete tensors. PyTorch provides two checkpointing modes to handle this correctly.

Full state dict mode gathers all parameter shards to rank 0 before saving, producing a standard checkpoint file identical to what you would get from a non-sharded model. The checkpoint is universally loadable: you can load it into a non-FSDP model, into an FSDP model with a different number of GPUs, or into any framework that understands standard PyTorch checkpoints. The downside is that rank 0 must temporarily hold the entire model in memory to write the checkpoint, creating a brief memory spike.

Local state dict mode saves each GPU's local parameter shard separately. This is faster and avoids memory pressure on rank 0, but the resulting checkpoint files are only loadable with the exact same FSDP configuration (same number of GPUs, same wrapping policy). This mode is useful for frequent checkpointing during training where you care about training continuity rather than portability.

In[23]:
Code
# Full state dict checkpointing (rank 0 saves, others wait)
# In a real distributed job:
#
# from torch.distributed.fsdp import FullStateDictConfig, StateDictType
#
# cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
# with FSDP.state_dict_type(fsdp_model, StateDictType.FULL_STATE_DICT, cfg):
#     state = fsdp_model.state_dict()
#     if dist.get_rank() == 0:
#         torch.save(state, "checkpoint.pt")

# For this notebook, demonstrate saving the non-distributed model
checkpoint = {
    "model_state": model.state_dict(),
    "optimizer_state": optimizer.state_dict(),
    "step": len(training_losses),
}
# torch.save(checkpoint, "model_checkpoint.pt")  # commented out to avoid file creation
Out[24]:
Console
Checkpoint contents:
  model_state: dict with 75 tensors
  optimizer_state: dict with 0 tensors
  step: 30

The offload_to_cpu=True flag in the full state dict config is important: it moves gathered tensors to CPU memory as they are assembled, preventing the GPU memory from doubling up. Without it, rank 0 would need to hold both its shard (from normal FSDP operation) and the full model (being assembled for saving) simultaneously in GPU memory.

Key Parameters

The key FSDP configuration parameters are:

  • sharding_strategy: Controls what is sharded. FULL_SHARD for maximum memory reduction, SHARD_GRAD_OP for moderate savings with less communication, NO_SHARD to disable sharding.
  • auto_wrap_policy: Determines how submodules are grouped into FSDP units. Use transformer_auto_wrap_policy for transformer architectures.
  • mixed_precision: A MixedPrecision object specifying dtypes for parameters, gradients, and buffers. BF16 is the standard choice for large model training.
  • device_id: The GPU to place the model on. Should be torch.cuda.current_device() after setting the appropriate device.
  • cpu_offload: A CPUOffload object that, when enabled, keeps parameters on CPU and only moves them to GPU when needed. Reduces GPU memory further at the cost of PCIe bandwidth.

Practical Considerations

Initializing the Process Group

In a real distributed FSDP job, you need to initialize the process group before wrapping the model. The standard pattern using torchrun (PyTorch's distributed launcher) looks like this:

In[43]:
Code
import os

import torch.distributed as dist


def setup_distributed():
    """
    Initialize the distributed process group.
    Called once at the start of each worker process.
    """
    dist.init_process_group(backend="nccl")  # NCCL for GPU-GPU communication
    local_rank = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local_rank)
    return local_rank


def cleanup_distributed():
    dist.destroy_process_group()

The torchrun launcher sets environment variables (RANK, LOCAL_RANK, WORLD_SIZE, MASTER_ADDR, MASTER_PORT) that dist.init_process_group() reads to coordinate across processes. You launch a training script on 8 GPUs with:

torchrun --nproc_per_node=8 train.py

For multi-node training across 4 nodes of 8 GPUs each (32 total):

torchrun --nproc_per_node=8 --nnodes=4 --node_rank=0 \ --master_addr=10.0.0.1 --master_port=29500 train.py

Gradient Clipping

Gradient clipping with FSDP requires a special API. The standard torch.nn.utils.clip_grad_norm_() function does not work correctly with FSDP because it does not account for the sharded nature of the parameters. FSDP provides FSDP.clip_grad_norm_() as a replacement:

In[45]:
Code
# In a real FSDP training loop, after loss.backward():
fsdp_model.clip_grad_norm_(max_norm=1.0)
optimizer.step()

The reason the standard function fails is that it computes the gradient norm locally on each GPU (seeing only shards), which underestimates the true global gradient norm. The FSDP version coordinates across GPUs to compute the correct global norm before clipping.

Activation Checkpointing

Activation checkpointing (gradient checkpointing) works with FSDP and is essential for training with long sequences or large batch sizes. The combination of FSDP and activation checkpointing is particularly powerful: FSDP reduces the persistent memory for parameters and optimizer states, while activation checkpointing reduces the peak activation memory during the backward pass.

To use activation checkpointing with FSDP, you wrap the checkpointing around the individual layers before or after the FSDP wrapping:

In[47]:
Code
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
    apply_activation_checkpointing,
    checkpoint_wrapper,
)

# Apply activation checkpointing to all TransformerLayer modules
check_fn = lambda submodule: isinstance(submodule, TransformerLayer)
apply_activation_checkpointing(
    fsdp_model,
    checkpoint_wrapper_fn=checkpoint_wrapper,
    check_fn=check_fn,
)

Activation checkpointing trades compute for memory: activations are recomputed during the backward pass rather than stored, roughly doubling the compute cost for the checkpointed layers but reducing activation memory to O(L)O(\sqrt{L}) where LL is the number of layers, instead of O(L)O(L).

Limitations and Impact

Communication Overhead on Slow Interconnects

FSDP's core trade-off is communication volume for memory capacity. With FULL_SHARD, every forward and backward pass triggers all-gather operations for each FSDP unit. On a cluster with slow interconnects, such as commodity Ethernet between nodes, this communication overhead can dominate training time and make FSDP less efficient than standard DDP even when memory fits. The technique is most effective on systems with high-bandwidth interconnects: NVLink within a node (600 GB/s bidirectional for NVLink 3.0) and InfiniBand between nodes (200 Gb/s or more per link). On such hardware, the communication overhead is a small fraction of compute time, and FSDP achieves near-linear scaling.

The HYBRID_SHARD strategy was introduced precisely to address this problem in multi-node settings. By confining the expensive all-gather communication to the fast intra-node NVLink fabric, and restricting cross-node communication to the cheaper all-reduce (gradient synchronization), HYBRID_SHARD preserves throughput on realistic clusters where inter-node bandwidth is a bottleneck.

Wrapping Granularity and Kernel Overhead

The granularity of FSDP wrapping introduces another practical consideration. Wrapping at too fine a granularity, for example wrapping every individual linear layer, causes very frequent all-gather operations with small tensors. Each all-gather has a fixed overhead for kernel launch and synchronization, making fine-grained wrapping inefficient regardless of interconnect speed. This overhead can reduce GPU utilization from the theoretical 80-90% down to 40-50% in extreme cases.

Wrapping at too coarse a granularity means larger all-gather buffers and higher peak memory, potentially defeating the purpose. The recommended practice is to wrap at the transformer layer level, where each FSDP unit contains one attention block and one FFN block together. This keeps the unit size large enough to amortize communication overhead (typically 50-200 million parameters per unit for modern models) while still enabling fine enough overlap between computation and prefetching.

Debugging and Observability

Debugging FSDP models is more involved than debugging standard PyTorch code. Because parameters are sharded at rest, you cannot inspect a model's weight tensors directly in a REPL or debugger without first gathering them. The state_dict() call with the appropriate StateDictType context is required, which adds ceremony to debugging workflows. A line of code that would normally be print(model.transformer.layers[0].self_attn.in_proj_weight.shape) requires wrapping in an FSDP state dict context to work correctly.

Memory profiling is similarly complicated. Standard PyTorch memory profiling tools report GPU memory allocations per device, but the sharded memory picture means the total model memory is distributed across all devices. Tools like PyTorch's memory profiler and the torch.cuda.memory_summary() function work correctly on individual devices, but understanding the global picture requires aggregating across all ranks.

Compatibility with Other Techniques

FSDP integrates well with several other large model training techniques but requires care with others. Activation checkpointing works cleanly with FSDP and is a standard combination. Gradient accumulation (running multiple forward-backward passes before an optimizer step) works with FSDP using the no_sync() context manager to suppress gradient synchronization during accumulation steps, similar to DDP.

Tensor parallelism (splitting individual weight matrices across GPUs, as done in Megatron-LM) requires careful integration because FSDP and tensor parallelism both modify how parameters are stored and communicated. Libraries like TorchTitan and Hugging Face Accelerate provide pre-built integrations for combining 3D parallelism (data, tensor, and pipeline) with FSDP as the data-parallel component.

Impact on the Open-Source Training Ecosystem

Despite these challenges, FSDP has become a foundational technique for frontier model training. The ability to train 70B or 175B parameter models on commodity GPU clusters, rather than requiring specialized hardware or proprietary systems, has dramatically lowered the barrier to large-scale research and production training. Models like Llama 2 and subsequent open-weight models were trained using FSDP or closely equivalent techniques, enabling the open-source community to train and reproduce frontier-scale results.

Before FSDP's broad adoption, training runs of this scale required either expensive proprietary infrastructure (Google TPU pods, AWS Trainium clusters) or deep expertise with DeepSpeed's complex configuration system. FSDP's integration with the standard PyTorch API meant that existing PyTorch users could scale to large models with modest code changes, without learning a new framework or configuration paradigm.

As models continue to grow beyond hundreds of billions of parameters, the multi-dimensional parallelism strategies that combine FSDP's data sharding with tensor and pipeline parallelism will build directly on the foundations covered here. Understanding FSDP's memory model and communication patterns gives you the conceptual tools to reason about those more complex hybrid systems.

Summary

FSDP (Fully Sharded Data Parallel) allows training of models too large to fit on a single GPU by sharding parameters, gradients, and optimizer states across all GPUs in a job.

The key ideas are:

  • Sharding all three components: Parameters, gradients, and optimizer states are each split into NN equal fragments, so each GPU's steady-state memory is 1/N1/N of the total. The full 16P16P bytes (float32 parameters, gradients, and Adam states) becomes 16P/N16P/N bytes per GPU.
  • All-gather and reduce-scatter: FSDP reconstructs full parameter tensors just in time for computation via all-gather, then discards them immediately after use. Gradients are aggregated and re-sharded via reduce-scatter instead of all-reduce, making sure no GPU ever holds a complete gradient tensor.
  • Sharding strategies: FULL_SHARD (maximum memory savings, corresponds to ZeRO Stage 3), SHARD_GRAD_OP (moderate savings, less communication, corresponds to ZeRO Stage 2), NO_SHARD (DDP equivalent for debugging), and HYBRID_SHARD (multi-node optimization that confines expensive all-gathers to the fast intra-node fabric).
  • Relationship to ZeRO: FSDP with FULL_SHARD is algorithmically equivalent to DeepSpeed ZeRO Stage 3. The practical difference is ecosystem: FSDP is native PyTorch, ZeRO offers additional features like CPU offloading via ZeRO-Infinity.
  • Wrapping policy: The granularity of FSDP units, controlled by the auto-wrap policy, determines the trade-off between all-gather frequency and buffer sizes. Wrapping at the transformer layer level is the standard recommendation.
  • Mixed precision: The MixedPrecision policy allows parameters and activations in bfloat16 while keeping gradients and optimizer states in float32, balancing compute speed with numerical stability.
  • Checkpointing: Full state dict mode produces portable checkpoints at the cost of a memory spike on rank 0. Local state dict mode is faster but tied to the specific FSDP configuration.

FSDP is the standard approach for large-scale model training in the PyTorch ecosystem and is directly supported by tools like Hugging Face Accelerate, PyTorch Lightning, and TorchTitan. Understanding its memory model and communication patterns equips you to reason about the practical constraints of large model training and to choose the right strategy for your hardware and model size.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about FSDP and fully sharded data parallel training.

FSDP Quiz

Question 1 of 80 of 8 completed
What does FSDP shard across GPUs to reduce memory per device?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026fsdpfully, author = {Michael Brenndoerfer}, title = {FSDP: Fully Sharded Data Parallel Training at Scale}, year = {2026}, url = {https://mbrenndoerfer.com/writing/fsdp-fully-sharded-data-parallel-sharding-strategies-zero}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-10-06} }
APAAcademic
Michael Brenndoerfer (2026). FSDP: Fully Sharded Data Parallel Training at Scale. Retrieved from https://mbrenndoerfer.com/writing/fsdp-fully-sharded-data-parallel-sharding-strategies-zero
MLAAcademic
Michael Brenndoerfer. "FSDP: Fully Sharded Data Parallel Training at Scale." 2026. Web. October 6, 2026. <https://mbrenndoerfer.com/writing/fsdp-fully-sharded-data-parallel-sharding-strategies-zero>.
CHICAGOAcademic
Michael Brenndoerfer. "FSDP: Fully Sharded Data Parallel Training at Scale." Accessed October 6, 2026. https://mbrenndoerfer.com/writing/fsdp-fully-sharded-data-parallel-sharding-strategies-zero.
HARVARDAcademic
Michael Brenndoerfer (2026) 'FSDP: Fully Sharded Data Parallel Training at Scale'. Available at: https://mbrenndoerfer.com/writing/fsdp-fully-sharded-data-parallel-sharding-strategies-zero (Accessed: October 6, 2026).
SimpleBasic
Michael Brenndoerfer (2026). FSDP: Fully Sharded Data Parallel Training at Scale. https://mbrenndoerfer.com/writing/fsdp-fully-sharded-data-parallel-sharding-strategies-zero

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.