一、引言:当模型参数放不进单卡显存时
随着大语言模型规模的不断膨胀,从 BERT-base 的 1.1 亿参数到 Llama 3 的 4050 亿参数,单张 GPU 的显存早已无法承载完整的模型训练。即便是拥有 80GB HBM3 显存的 H100,在面对超过 30B 参数的模型时也显得捉襟见肘。分布式训练不再是可选方案,而是必修课。
在分布式训练策略中,Data Parallelism(数据并行)是最直观的思路:每个 GPU 持有完整的模型副本,各自处理不同的数据批次。但当模型规模超过单卡显存时,数据并行就失效了。ZeRO(Zero Redundancy Optimizer)系列方法由 Microsoft 在 2019 年提出,通过消除冗余的显存占用来突破单卡容量天花板。PyTorch 的 FSDP(Fully Sharded Data Parallel)正是 ZeRO-3 策略的工程化实现,它已经成为 PyTorch 生态中训练大模型的事实标准。
本文将从源码层面深入剖析 FSDP 的核心机制,覆盖分片策略、通信调度、内存管理与生产环境调优等关键维度。阅读本文之前,建议对 PyTorch 的 DistributedDataParallel(DDP)有基本了解,因为 FSDP 可以理解为”在 DDP 之上叠加了模型状态的显存分片”。

二、从 DDP 到 FSDP:显存冗余的本质
2.1 DDP 的显存分布与问题
在标准的 DDP 训练中,每张 GPU 上存储了完整的模型参数(Parameters)、梯度(Gradients)和优化器状态(Optimizer States,如 Adam 中的 momentum 和 variance)。对于一个参数量为 Ψ 的模型:
- 模型参数:占用 4Ψ 字节(FP32)或 2Ψ 字节(BF16/FP16)
- 梯度:与参数大小相同,占用 4Ψ 或 2Ψ 字节
- 优化器状态(以 Adam 为例):两个状态各 4Ψ 字节,共 8Ψ 字节(FP32 下)
- 中间激活值(Activations):取决于 batch size 和序列长度,通常与参数显存相当甚至更大
以 70B 参数模型为例,仅参数 + 梯度 + Adam 状态在 FP32 下就需要 4×70×4 = 1120 GB 显存——远超过任何单张 GPU 的容量。即便使用混合精度训练(BF16 参数 + FP32 优化器状态),也需要约 2Ψ + 2Ψ + 8Ψ = 12Ψ = 840 GB。这显然行不通。
2.2 ZeRO 的三阶段优化
ZeRO 的核心洞察是:数据并行中,每个 GPU 都存储了完整的优化器状态和参数,但实际上它们都是同步的——我们只需要在每个 GPU 上保留自己负责的那一份,其他 GPU 的数据可以通过通信来获取。ZeRO 将优化分为三个阶段:
| 阶段 | 分片内容 | 每 GPU 显存(70B 模型, BF16+FP32 Adam) | 通信开销 |
|---|---|---|---|
| ZeRO-1 | 优化器状态(Optimizer States) | ~120 GB | 与 DDP 相同 |
| ZeRO-2 | 优化器状态 + 梯度(Gradients) | ~60 GB | 与 DDP 相同 |
| ZeRO-3 | 优化器状态 + 梯度 + 参数(Parameters) | ~30 GB(64 GPU 时) | 增加 ~50% |
FSDP 实现了 ZeRO-3 的全部功能。在 N 张 GPU 上,每张 GPU 只持有约 1/N 的完整模型状态,从而实现了接近线性的显存缩放。
三、FSDP 核心原理:前向与反向传播中的动态参数收集
3.1 分片的粒度:Flattening 与 FlatParameter
FSDP 并没有对每个独立参数进行分片——那样会产生大量的细粒度通信操作(All-Gather),通信效率极低。相反,FSDP 将整个模型的参数进行展平(Flatten),形成一个或几个连续的 FlatParameter 张量,然后在这些大块上执行分片。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20 # 简化源码示意:FlatParameter 的构造
# 来源:torch/distributed/fsdp/_flat_param.py
class FlatParameter(nn.Parameter):
"""将多个 nn.Parameter 展平为一个连续张量。"""
def __init__(self, params: List[nn.Parameter]):
# 获取所有参数的数据指针
flat_data = torch.cat([p.data.view(-1) for p in params])
super().__init__(flat_data, requires_grad=True)
self._param_infos = [] # 记录每个子参数在 flat 中的偏移
offset = 0
for p in params:
numel = p.numel()
self._param_infos.append({
"start": offset,
"end": offset + numel,
"shape": p.shape,
"dtype": p.dtype,
})
offset += numel
这样做的好处有两个:一是减少了通信元的数量(All-Gather 一个大张量比对数千个小张量分别 All-Gather 高效得多);二是方便将 FlatParameter 均匀切分到各个 rank。
3.2 前向传播:All-Gather + 计算 + Discard
在前向传播过程中,FSDP 将整个计算图划分为多个”单元”,每个单元对应模型中的一个子模块(通常是 Transformer Block)。对每个单元:
- All-Gather:从所有 rank 上收集完整的参数,形成完整的 FlatParameter
- 前向计算:用收集到的完整参数执行该子模块的前向传播
- Discard(丢弃):前向计算完成后,立即释放那些不属于本 rank 的参数片段,只保留本 rank 拥有的原始分片
这意味着在前向传播过程中,同一时刻只有一个”单元”的参数是完整的。这大幅降低了峰值显存占用。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18 # 前向传播中的 All-Gather 流程(简化伪代码)
def forward(self, input):
if self._is_root or self._has_params:
# Step 1: All-Gather 收集完整参数
with torch.cuda.stream(self._streams["all_gather"]):
self._all_gather_params()
# Step 2: 同步流,等收集完成
torch.cuda.current_stream().wait_stream(self._streams["all_gather"])
# Step 3: 执行前向计算
output = self._module(input)
if not self._is_root:
# Step 4: 释放非本 rank 的参数
self._free_unsharded_params()
return output
3.3 反向传播:重新收集参数 + 计算梯度 + Reduce-Scatter
反向传播的逻辑更加精巧。由于前向传播结束后参数已被丢弃,反向传播时需要重新通过 All-Gather 收集当前单元的参数来计算梯度。这也是 ZeRO-3 “以通信换显存”的核心代价:前向和反向各需要一次 All-Gather。
- 重新 All-Gather:再次从所有 rank 收集完整的参数
- 计算梯度:使用收集到的参数计算梯度
- Reduce-Scatter:对所有 rank 上的梯度执行 Reduce-Scatter 操作——先 All-Reduce 求和,再将结果按分片分散到各 rank,使得每个 rank 只保留自己分片的梯度
- 释放完整参数:与参数类似,释放非本 rank 的参数
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23 # 反向传播中的 Reduce-Scatter(简化伪代码)
# 来源:torch/distributed/fsdp/_flat_param.py::_reduce_scatter
def _reduce_scatter(self):
fsdp_group = self.process_group
world_size = fsdp_group.size()
shard_size = self.full_numel // world_size
# 对梯度做 Reduce-Scatter
# 每个 rank 得到 1/world_size 的梯度碎片
output = torch.empty(shard_size, dtype=self.dtype, device=self.device)
# 使用 NCCL 的 reduce_scatter
dist.reduce_scatter_tensor(
output=output,
input=self._gradient,
group=fsdp_group,
op=dist.ReduceOp.AVG # 或 SUM,取决于配置
)
# 释放完整梯度
self._gradient = None
self._sharded_gradient = output
四、分片策略(Sharding Strategy):如何选择最优粒度
FSDP 提供了三种分片策略,通过
1 | ShardingStrategy |
枚举来配置:
| 策略 | 枚举值 | 描述 | 适用场景 |
|---|---|---|---|
| FULL_SHARD | 1 | 模型参数、梯度、优化器状态全部分片(ZeRO-3) | 超大模型训练,显存极度紧张 |
| SHARD_GRAD_OP | 2 | 仅分片梯度和优化器状态,参数副本全部保留(ZeRO-2) | 显存够用但想降低优化器状态开销 |
| NO_SHARD | 3 | 不分片,相当于标准 DDP | 模型参数可放入单卡,仅需数据并行 |
| HYBRID_SHARD | 4 | 节点内 FULL_SHARD,节点间 NO_SHARD(即跨节点复制完整模型) | 多节点训练,节点内用 NVLink,节点间用 IB/RoCE |
HYBRID_SHARD 是一个特别实用的策略。在多节点训练场景中,同一节点内的 GPU 通过 NVLink 互联(带宽可达 600 GB/s),但跨节点只有 InfiniBand(通常 200-400 Gbit/s ≈ 25-50 GB/s)。HYBRID_SHARD 在节点内部做 FULL_SHARD(利用高带宽进行频繁通信),而在节点之间做 NO_SHARD(只在节点间同步梯度,通信量与 DDP 相同)。这很好地在通信效率和显存效率之间取得了平衡。
五、CPU Offload:将显存压力转移到内存
当显存仍然不够时,FSDP 支持将优化器状态甚至参数卸载(Offload)到 CPU 内存中。这是突破 GPU 显存物理极限的最后一招。
1
2
3
4
5
6
7
8
9
10 from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import CPUOffload
model = FSDP(
model,
cpu_offload=CPUOffload(offload_params=True),
# 参数卸载到 CPU
# forward/backward 时 All-Gather 从 CPU 取参数
# 计算完成后参数又释放回 CPU
)
CPU Offload 的工作原理:
- 参数卸载:分片后的参数存储在 CPU 固定的(pinned)内存上,通过
1cudaMemcpyAsync
异步拷贝到 GPU 进行计算
- 优化器状态卸载:优化器中的 momentum 和 variance 存储在 CPU 内存中,优化器更新步骤在 CPU 上执行
- 梯度卸载:Reduce-Scatter 后的梯度碎片也存放在 CPU 上
代价是什么?PCIe 带宽。一张 GPU 到 CPU 的 PCIe 4.0 x16 带宽约为 32 GB/s,而 NVLink 带宽是 600 GB/s——相差近 20 倍。启用 CPU Offload 后,All-Gather 的延迟会急剧增加,训练吞吐可能下降 50%-70%。
因此,CPU Offload 应作为”最后的逃生出口”:只有当模型尺寸实在无法靠 GPU 显存容纳时,才考虑启用。一般来说,优先增加 GPU 数量或使用梯度检查点(Gradient Checkpointing)来降低激活值显存,而不是直接启用 CPU Offload。
六、通信调度与重叠(Communication Overlap)
每次 All-Gather 和 Reduce-Scatter 都是同步操作,如果顺序执行会严重拖慢训练速度。FSDP 通过以下机制来减少通信对计算的影响:
6.1 预取(Prefetch)机制
FSDP 在前向传播时会预取下一个单元的完整参数,与当前单元的计算重叠执行。这通过
1 | _prefetch_forward_params |
和
1 | _prefetch_backward_params |
实现:
1
2
3
4
5
6
7
8
9 # 前向预取逻辑(简化)
def _prefetch_forward_params(self, next_unit):
if next_unit is not None and next_unit._has_params:
# 使用单独的计算流执行 All-Gather
torch.cuda.current_stream().record_event(self._prefetch_event)
next_unit._all_gather_params(
stream=self._prefetch_stream,
wait_event=self._prefetch_event
)
通过这种方式,当 FSDP 在处理第 N 个 Transformer Block 的前向计算时,第 N+1 个 Block 的参数 All-Gather 已经在后台完成了。理想情况下,通信开销可以被完全隐藏。
6.2 梯度同步的时机:NO_WAIT / WAIT / POST_ORDER
在反向传播中,Reduce-Scatter 的调度时机对性能影响很大。FSDP 的
1 | BackwardPrefetch |
参数控制:
- BACKWARD_PRE(默认):在当前单元的梯度计算完成后立即启动 Reduce-Scatter,同时预取下一个单元的参数。这是推荐配置,通信和计算重叠效果最好。
- BACKWARD_POST:在当前单元的计算和梯度处理完全结束后才启动通信,重叠效果较差,但显存占用稍低。
1
2
3
4
5
6 from torch.distributed.fsdp import BackwardPrefetch
FSDP(
model,
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
)
6.3 Limit All-Gathers(限制并行通信数)
虽然预取能有效隐藏通信延迟,但过多的并行 All-Gather 操作也会占用显存(因为被预取到 GPU 上的完整参数需要空间暂存)。FSDP 通过
1 | limit_all_gathers=True |
来限制同时进行的 All-Gather 数量,防止显存溢出:
1
2
3
4
5
6 FSDP(
model,
limit_all_gathers=True, # 默认 False
# 启用后,最多同时执行 2 个 All-Gather
# 防止预取过多导致 OOM
)
七、生产环境调优实战指南
7.1 Auto Wrap Policy:自动决定 FSDP 单元边界
FSDP 需要知道在哪里进行参数 All-Gather 的边界。手动指定每个子模块是否 wrap 是繁琐且容易出错的。PyTorch 提供了自动封装策略(Auto Wrap Policy):
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28 from torch.distributed.fsdp.wrap import (
size_based_auto_wrap_policy,
transformer_auto_wrap_policy,
lambda_auto_wrap_policy,
)
# 策略一:基于参数大小自动 wrap
# 当一个模块的参数超过 min_num_params 时自动封装
auto_wrap_policy = partial(
size_based_auto_wrap_policy,
min_num_params=1e8, # 1 亿参数为阈值
)
# 策略二:基于模块类型
# 只对 Transformer Block 类型进行封装
auto_wrap_policy = partial(
transformer_auto_wrap_policy,
transformer_layer_cls={
LlamaDecoderLayer,
GPT2Block,
T5Block,
},
)
FSDP(
model,
auto_wrap_policy=auto_wrap_policy,
)
最佳实践:对于 Transformer 架构,推荐使用
1 | transformer_auto_wrap_policy |
按层(Layer)划分 FSDP 单元。每个 Transformer Block 作为一个 FSDP 单元,这样通信粒度和计算粒度匹配良好。
7.2 Mixed Precision 配置:FP32 Master Weights 与 BF16 计算
FSDP 内建了混合精度支持,无需外挂 AMP(Automatic Mixed Precision):
1
2
3
4
5
6
7
8
9
10
11
12
13
14 from torch.distributed.fsdp import MixedPrecision
mp_config = MixedPrecision(
param_dtype=torch.bfloat16, # 前向传播时 All-Gather 到 GPU 的参数 dtype
reduce_dtype=torch.bfloat16, # Reduce-Scatter 时梯度的 dtype
buffer_dtype=torch.bfloat16, # buffer 的 dtype
)
FSDP(
model,
mixed_precision=mp_config,
# 优化器仍使用 FP32 精度的主权重
# FSDP 内部自动维护 FP32 的 shadow weights 在 optimizer states 中
)
需要注意的是:即使前向传播使用 BF16 参数计算,FSDP 内部仍然为优化器保存了 FP32 版本的参数分片。这意味着优化器步骤仍然在 FP32 精度下更新,保证了训练精度不下降——这正是 ZeRO 的理念:用通信保精度。
7.3 执行单元测试:诊断通信瓶颈
下面是一个完整的 FSDP 训练脚本框架,包含关键的 profiler 工具:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58 import torch
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.profiler import profile, record_function, ProfilerActivity
from functools import partial
def train_fsdp():
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
model = LlamaForCausalLM.from_pretrained("meta-llama/Llama-2-7b")
auto_wrap_policy = partial(
transformer_auto_wrap_policy,
transformer_layer_cls={LlamaDecoderLayer},
)
fsdp_model = FSDP(
model,
auto_wrap_policy=auto_wrap_policy,
sharding_strategy=ShardingStrategy.FULL_SHARD,
mixed_precision=MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
),
device_id=local_rank,
)
optim = torch.optim.AdamW(fsdp_model.parameters(), lr=3e-5)
# Profiling:观察 All-Gather / Reduce-Scatter 耗时
with profile(
activities=[ProfilerActivity.CUDA, ProfilerActivity.CPU],
profile_memory=True,
record_shapes=True,
) as prof:
for step, batch in enumerate(dataloader):
with record_function("forward"):
loss = fsdp_model(**batch).loss
with record_function("backward"):
loss.backward()
with record_function("optimizer"):
optim.step()
optim.zero_grad()
if step >= 10: # warm up后采样
break
# 打印通信相关 op 的耗时统计
print(prof.key_averages().table(
sort_by="cuda_time_total",
row_limit=20
))
# 观察 AllGather 和 ReduceScatter 是否被计算隐藏
# 理想状况下它们在 time_line 上应与 compute 重叠
7.4 常见陷阱与解决策略
| 问题 | 现象 | 根因与解决 | ||||
|---|---|---|---|---|---|---|
| OOM(显存溢出) | 启动后不久即报 CUDA OOM | 启用
;减小 batch size;启用 Gradient Checkpointing;确认 forward_prefetch 未过度预取 |
||||
| 训练速度慢 | 吞吐远低于预期 MFU | 检查通信是否与计算重叠(Profiler);确认 NCCL 环境变量已优化(NCCL_IB_HCA, NCCL_SOCKET_IFNAME);尝试 HYBRID_SHARD 减少跨节点通信 | ||||
| Loss 不收敛 | Loss 震荡或发散 | 检查 MixedPrecision 配置是否丢失精度;确认 reduce_dtype 与 param_dtype 匹配;验证学习率是否需要缩放 | ||||
| All-Gather 死锁 | 训练卡住,GPU 利用率 0% | 常见于嵌套 FSDP + activation checkpointing 场景。尝试设置
或
|
||||
| 显存碎片 | 总显存有剩余但分配大张量仍 OOM | 启用
(PyTorch 2.0+) |
八、FSDP 与 DeepSpeed ZeRO-3 的对比
DeepSpeed ZeRO-3 是 FSDP 的主要竞争对手。两者在算法层面等价,但工程实现上有显著差异:
| 维度 | PyTorch FSDP | DeepSpeed ZeRO-3 |
|---|---|---|
| 集成方式 | 原生 PyTorch 组件,pip install 即用 | 需要安装 deepspeed 库,模型需封装为 DeepSpeedEngine |
| 分片粒度 | FlatParameter(整体展平后分片) | 参数组(Param Group)级别分片 |
| 通信重叠 | 通过预取流,支持前向/反向预取 | 更激进的梯度 Bucketing 策略 |
| CPU Offload | 支持(CPUOffload) | 支持(ZeRO-Infinity/NVMe Offload) |
| 训练优化器 | 需用户自行创建,FSDP 自动 params() 接口 | DeepSpeed 自动管理优化器状态 |
| Hugging Face 兼容性 | 优秀(被 Trainer 原生支持) | 良好(需 DeepSpeed 配置 JSON 文件) |
| 梯度检查点 | 由 PyTorch 原生的 activation checkpointing 提供 | 内建 deepspeed.checkpointing 模块 |
| 显存效率 | 相近(实现细节差异引入 ±5%) | 相近 |
| 易用性 | ⭐⭐⭐⭐(API 简洁,文档完善) | ⭐⭐⭐(配置略复杂) |
选择建议:
- 如果你正在使用 Hugging Face Transformers Trainer,FSDP 是最自然的选择——几行配置即可启用
- 如果你需要更极致的优化(如 ZeRO-Infinity + NVMe Offload),或者正在进行非常规的分布式训练拓扑,DeepSpeed 提供了更多底层控制
- 对于大多数场景,FSDP 已经足够成熟,而且减少了对外部库的依赖
九、未来展望:FSDP 的演进方向
PyTorch 社区没有停下 FSDP 的迭代脚步。在即将发布的 PyTorch 2.x 系列中,以下几个方向值得关注:
- FSDPv2(torch.distributed.fsdp_v2):更高效的通信调度,原生支持异步 Reduce-Scatter 和更智能的预取策略,减少调度器层面的开销
- HSDP(Hybrid Sharded Data Parallel)正式化:HYBRID_SHARD 策略将在 v2 中获得更完善的支持,包括自动检测节点拓扑
- DTensor 集成:将 FSDP 的底层分片逻辑与 DTensor(Distributed Tensor)对齐,实现统一的分布式张量抽象
- 自动调优:自动探测最优的分片策略(FULL_SHARD vs HYBRID_SHARD)、batch size 和通信预取参数
十、总结
PyTorch FSDP 是当前训练大语言模型的基础性工具。它将 ZeRO-3 的理论优势转化为可用的工程实现,让拥有数十甚至数百张 GPU 的研究团队能够训练出千亿参数级别的模型。
理解 FSDP 的核心原理——FlatParameter 分片、前向/反向的 All-Gather 与 Reduce-Scatter 动态调度、通信与计算的重叠——是从事 AI Infra 工作的必备知识。在排查分布式训练的性能问题时,熟练使用 PyTorch Profiler 来观察通信操作的时间线、判断通信是否被计算掩盖,是 AI Infra 工程师的核心技能。
最后,回到本文开篇的问题:当模型参数放不进单卡显存时?答案是:分而治之(Shard & Conquer)。FSDP 正是这一思想的工程典范——它让大规模分布式训练不再是少数几家巨头公司的专利,而是每一个 AI 工程师都能触达的技术能力。
希望本文能帮助各位读者更深入地理解 FSDP 的内在机制,在实际工作中更高效地训练你的大模型。
汤不热吧