首页 / PyTorch 入门教程 / 大规模训练:FSDP

PyTorch 入门教程

大规模训练:FSDP

本教程共 60 篇 · 第 57 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorchFSDP分布式训练显存优化DTensortorchrun

本节目标:理解模型为什么装不进单张显卡,掌握 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 的做法是「用前收集,用后释放」:

  1. 前向或反向计算某层之前,通过 all-gather 把分片参数拼回完整形态;
  2. 这一层算完,立刻释放完整参数,恢复分片状态;
  3. 反向时,本地的非分片梯度通过 reduce-scatter 聚合,各卡只拿自己那份分片梯度;
  4. 优化器只更新本卡持有的分片参数。

本质上,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")
Warning

FSDP1 的 FullyShardedDataParallel 已弃用,新项目直接用 fully_shard。迁移也不难:把自动包装策略换成手动逐层调用 fully_sharduse_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 与张量并行。