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 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 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 parameters stored in float32, memory consumption breaks down across four components:
- Parameters: bytes (4 bytes per float32 value)
- Gradients: bytes (same shape as parameters, one gradient per parameter)
- Optimizer states: 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 bytes. For GPT-3 with billion parameters, that amounts to:
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:
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 bytes, accounting for half of the 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 (fp16 params) (fp16 grads) (fp32 master params) (fp32 optimizer states) 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 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 of shape , containing total parameters. With GPUs, FSDP flattens into a 1D tensor of length and assigns an equal contiguous slice to each GPU:
- GPU 0 holds elements through
- GPU 1 holds elements through
- GPU holds elements through
Each GPU also holds a corresponding shard of the gradient tensor for those same parameters, and a shard of the Adam optimizer states (first moment and second moment ) 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 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 , not a shard of it. An input of shape needs to be multiplied by the full to produce output of shape . 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:
- 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.
- Each GPU uses the full parameter tensor to compute its portion of the forward pass (processing its slice of the data batch).
- 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:
- An all-gather reconstructs the full parameter tensor again (necessary for computing parameter gradients).
- The backward computation runs, producing full-parameter gradient tensors on each GPU.
- 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 receives only the portion of the gradient corresponding to its parameter shard.
- 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.
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 GPUs, the memory breakdown per GPU changes dramatically compared to DDP:
- Parameters (steady state): bytes (only the local shard persists)
- Gradients (steady state): bytes (only the local gradient shard persists)
- Optimizer states: bytes (first and second moments for the local shard)
- All-gather buffer during forward: bytes temporarily (freed after each layer's computation)
The steady-state memory per GPU is approximately bytes for the sharded quantities. The peak memory includes the temporary all-gather buffer of bytes during computation for the current FSDP unit, but this is freed right after the unit completes. For a model split into FSDP units, the peak at any instant is bytes (steady-state plus one unit's full parameters).
Compare this to DDP, where each GPU holds bytes at all times. The reduction factor from to for the steady state is exactly , 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 ( GB). With FSDP on 8 GPUs of 40 GB each (320 GB total), each GPU holds only 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 GPUs, the optimizer states, gradients, and even parameters are all replicated 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 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 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 of optimizer states and 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 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.
| Feature | FSDP | ZeRO Stage 3 (DeepSpeed) |
|---|---|---|
| Integration | Native PyTorch | Requires DeepSpeed installation |
| Parameter sharding | Yes | Yes |
| Gradient sharding | Yes | Yes |
| Optimizer state sharding | Yes | Yes |
| CPU offloading | Limited | Full ZeRO-Infinity support |
| Ease of use | High | Moderate |
| Communication optimization | Good | More 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 FSDP units, one complete training step requires:
- all-gather operations during forward pass
- all-gather operations during backward pass
- reduce-scatter operations during backward pass
The total communication volume per step is proportional to where is the average parameters per FSDP unit. Since (total parameters), this simplifies to in aggregated data moved. Compare this to DDP, which requires one all-reduce of 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 all-gathers plus reduce-scatters, you get just one reduce-scatter per step.
The memory savings are smaller. Parameters ( bytes) are not sharded, so each GPU still holds the full parameter tensor. Only gradients ( bytes) and optimizer states ( bytes) are split by , saving:
For large , this approaches bytes savings, leaving only the parameter bytes per GPU. Compare to FULL_SHARD, which brings per-GPU memory to bytes. The gap between the two strategies narrows as increases: on 64 GPUs, SHARD_GRAD_OP saves bytes, while FULL_SHARD saves nearly 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 of parameter memory within each node), while FULL_SHARD would give 32-GPU sharding savings ().
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.
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.
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 with the computation for layer . 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 ), 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.
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: GB
- Gradients: GB
- Optimizer states (Adam): GB
- Total model state: GB
This already exceeds the 40 GB GPU capacity by . DDP cannot run this configuration at all.
With FSDP FULL_SHARD on 8 GPUs
With FULL_SHARD, each GPU holds of each component:
- Parameters (shard): GB
- Gradients (shard): GB
- Optimizer states (shard): GB
- Total steady-state: GB
During the forward and backward passes, a single FSDP unit's all-gather buffer temporarily consumes an additional GB where is the number of FSDP units. For a 32-layer model wrapped at the layer level, each all-gather buffer is GB, bringing peak memory to about 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 where is batch size, is sequence length, and is hidden dimension. For a 7B model (), a batch of 4 sequences of length 2048 in bfloat16 generates roughly 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.
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.
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)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.
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_modelSimulating 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.
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]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 ( bytes) alone saves significantly more memory than sharding parameters alone ( bytes).
Visualizing Memory Savings

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.
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 # GBCommunication 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 TB per step (before counting the reduce-scatters), while SHARD_GRAD_OP and NO_SHARD would each move about 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

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.
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 resultsScenario: 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 ( GB per GPU for the model state) fits within 80 GB with FULL_SHARD, leaving room for activation buffers.
Visualizing Strategy Recommendations

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.
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),
)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)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

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.
# 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 creationCheckpoint 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_SHARDfor maximum memory reduction,SHARD_GRAD_OPfor moderate savings with less communication,NO_SHARDto disable sharding. - auto_wrap_policy: Determines how submodules are grouped into FSDP units. Use
transformer_auto_wrap_policyfor transformer architectures. - mixed_precision: A
MixedPrecisionobject 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
CPUOffloadobject 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:
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 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:
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 where is the number of layers, instead of .
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 equal fragments, so each GPU's steady-state memory is of the total. The full bytes (float32 parameters, gradients, and Adam states) becomes 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), andHYBRID_SHARD(multi-node optimization that confines expensive all-gathers to the fast intra-node fabric). - Relationship to ZeRO: FSDP with
FULL_SHARDis 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
MixedPrecisionpolicy 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
Reference
Citation details
Cite or share this article.
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 HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

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