大规模训练:FSDP
本教程共 60 篇 · 第 57 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:理解模型为什么装不进单张显卡,掌握 FSDP 的分片思想,学会用
fully_shard训练大规模模型。
显存都去哪了
训练一个模型,显存要装四样东西:模型参数、梯度、优化器状态、中间激活。参数和梯度好理解,优化器状态容易被忽略——Adam 要为每个参数额外存两份状态(动量),比参数本身还占地方。
算笔账。fp32 训练时,每个参数占 4 字节,梯度 4 字节,Adam 状态 8 字节,合计 16 字节。一个 70 亿参数的模型,光这三样就要 112 GB。常见的 4090 只有 24 GB 显存。
更夸张的是 700 亿参数的大模型,直接要 1.1 TB。单卡根本塞不下,这不是优化能解决的,得换思路。
Note这个估算还没算中间激活。激活在反向传播前一直占着显存,通常又是一大块。第 59 章会专门讲怎么压它。
什么信号说明该上 FSDP?训练时报 CUDA out of memory,而且调小 batch size 也救不回来;或者你按上面的公式一算,光参数和优化器就快把显存占满了。这时先别急着租更贵的卡,试试 FSDP。
DDP 的局限
第 56 章讲过 DDP:每张卡放一份完整模型副本,各算各的梯度,最后 all-reduce 同步。它的前提是模型放得进一张卡。
放不进去,DDP 就无解了。每张卡都要完整副本,模型要 112 GB,八张卡就是八个 112 GB。显存一点没省,还多了一堆通信。
换个思路:既然一张卡装不下全部,把模型拆开,每张卡只存一部分,行不行?这就是分片(sharding)思想,也是 FSDP 干的事。
FSDP:把模型拆开存
全分片数据并行(Fully Sharded Data Parallel,FSDP)把参数、梯度、优化器状态全部切碎,均匀分到每张卡上。16 张卡训练 70B 模型,每张卡只存大约 1/16。
但分片有个问题:算某一层时,这层的完整参数必须在手。FSDP 的做法是「用前收集,用后释放」:
- 前向或反向计算某层之前,通过 all-gather 把分片参数拼回完整形态;
- 这一层算完,立刻释放完整参数,恢复分片状态;
- 反向时,本地的非分片梯度通过 reduce-scatter 聚合,各卡只拿自己那份分片梯度;
- 优化器只更新本卡持有的分片参数。
本质上,FSDP 把 DDP 的 all-reduce 拆成了 reduce-scatter 加 all-gather。省显存,代价是多花通信时间。
Tip两个集合通信记不住?all-gather 是「把碎片拼成整块,发给组里每个人」;reduce-scatter 是「每人算完局部结果,汇总后各拿一份」。前者收集,后者归约。
分片后每张卡的参数侧占用只剩约 1/N(N 为卡数)。回到刚才的 70B 模型:16 张 80 GB 的卡,DDP 每张卡都要 1.1 TB,直接装不下;FSDP 每张卡约 70 GB,刚好放得下。这才是 FSDP 能训练超大模型的底气。
显存省了,代价是通信变多。每一层前向反向都要一次 all-gather 和一次 reduce-scatter,层数越多通信越频繁。所以 FSDP 的优化重点就是让通信和计算重叠,别让卡闲着等数据。
FSDP2:直接用 fully_shard
老版 FSDP(FSDP1)用 FullyShardedDataParallel 包装类,要配自动包装策略,配置繁琐,现已弃用。官方推荐 FSDP2,入口是 fully_shard 函数,不换包装类,直接在原模型上调:
from torch.distributed.fsdp import fully_shard
model = Transformer() # 你自己的模型
for layer in model.layers: # 先分片每个子层
fully_shard(layer)
fully_shard(model) # 最后分片根模型
按层分片是关键。计算某一层时只 all-gather 这一层,其余层保持分片,显存占用最小。分片后的参数会变成 DTensor(分布式张量),带着分片布局信息:
from torch.distributed.tensor import DTensor
for param in model.parameters():
assert isinstance(param, DTensor) # 分片后的参数都是 DTensor
训练循环和单卡一模一样,优化器直接拿 model.parameters() 建:
optim = torch.optim.Adam(model.parameters(), lr=1e-2)
for _ in range(epochs):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
loss = model(x).sum()
loss.backward()
optim.step()
optim.zero_grad()
启动命令也和 DDP 相同,多卡训练统一用 torchrun:
torchrun --nproc_per_node 2 train.py
配套技巧:混合精度与检查点
FSDP2 的混合精度策略很灵活。典型配置:计算时参数转成 bfloat16,梯度归约用 float32 保精度:
from torch.distributed.fsdp import MixedPrecisionPolicy
fsdp_kwargs = {
"mp_policy": MixedPrecisionPolicy(
param_dtype=torch.bfloat16, # 前向/反向计算用 bf16
reduce_dtype=torch.float32, # 梯度归约用 fp32
)
}
for layer in model.layers:
fully_shard(layer, **fsdp_kwargs)
fully_shard(model, **fsdp_kwargs)
想再省显存,还能把参数卸载到 CPU 内存,配置 offload_policy=CPUOffloadPolicy() 即可。
检查点也有讲究。分片后 state_dict() 是分布式的,直接 torch.save 会各卡存各卡,存出来的东西拼不回去。官方推荐用分布式检查点(DCP),一键拼回完整状态并落盘:
from torch.distributed.checkpoint.state_dict import (
get_model_state_dict,
StateDictOptions,
)
model_state_dict = get_model_state_dict(
model=model,
options=StateDictOptions(full_state_dict=True, cpu_offload=True),
)
torch.save(model_state_dict, "model_state_dict.pt")
WarningFSDP1 的
FullyShardedDataParallel已弃用,新项目直接用fully_shard。迁移也不难:把自动包装策略换成手动逐层调用fully_shard,use_orig_params这些参数都不用再管。
踩坑提醒
几个我踩过的坑,提前给你打预防针。
第一次训练迭代特别慢。所有层都要做首轮 all-gather,通信初始化也有开销。这是正常的,别以为卡住了。
batch size 别开太大。FSDP 省的是参数侧显存,激活一点没省。激活占大头时,配合激活检查点(第 59 章)效果才明显。
通信要和计算重叠。FSDP2 默认有隐式预取:下一层的 all-gather 会和当前层计算并行。CPU 忙不过来时,可以用 set_modules_to_forward_prefetch 手动指定预取哪些层。
多节点训练用 torchrun --nnodes=2 --nproc_per_node=8,两台机器跑同样的命令。跨节点的通信走网络,比节点内的 NVLink 慢不少,这是正常的。节点内通信量大的话,可以把 FSDP 和主机内张量并行组合(第 58 章),把慢通信挡在节点之间。
小结
FSDP 解决的核心问题:模型装不进单卡。它把参数、梯度、优化器状态分片到所有卡上,用 all-gather 和 reduce-scatter 换显存。记住三个要点:逐层 fully_shard、优化器在分片后创建、检查点用 DCP。下一章讲 GPU 更多时怎么组合多种并行策略——DeviceMesh 与张量并行。