每张卡为何还要保留完整模型?FSDP2 的参数全聚合、梯度归约分片与显存峰值
从普通 DDP 的复制成本出发,手算 FSDP2 如何分片参数、梯度和优化器状态,拆解 all-gather、reduce-scatter、reshard、分组边界与分布式 checkpoint。
上一篇比较了 Data Parallel(数据并行)、Tensor Parallel(张量并行)与 Pipeline Parallel(流水线并行)。普通 DDP 虽然把 batch 分给多卡,却在每个 rank 保存完整参数、梯度和 optimizer state:卡越多,总副本越多,单卡容量没有下降。
本篇聚焦 Fully Sharded Data Parallel 2(全分片数据并行第二代,FSDP2)的一个核心机制:平时把模型状态分片保存,只在某个模块真正计算前临时聚齐参数,并在反向后把梯度归约回分片。
01 DDP 到底复制了多少状态?#
设参数量为 ,数据并行 world size 为 。每参数字节数仍记为权重 、梯度 、optimizer state 。
普通 DDP 每卡长期占用近似
FSDP 理想稳定态则近似
但“除以 ”不是峰值答案。计算某个参数组前,FSDP 必须 all-gather(全聚合)完整参数;通信还可能同时持有下一组的预取 buffer。更实际的粗略式是
因此模型总状态均分后能放下,不代表某个巨大的根分组也能安全 all-gather。
02 FSDP2 的一层在时间上怎样变化?#
sequenceDiagram
participant S as 每卡参数 shard (DTensor)
participant C as Collective
participant M as 本地模块计算
S->>C: pre-forward all-gather
C->>M: 临时完整参数 Tensor
M->>M: forward
M->>S: post-forward reshard / 释放完整参数
S->>C: pre-backward 再 all-gather
C->>M: 用完整参数计算 backward
M->>C: 完整局部梯度 reduce-scatter
C->>S: 每卡保留梯度 shard
S->>S: optimizer 更新本地 state shardmermaidPyTorch 2.14 的 fully_shard(module) 会原地把参数转换为按第 0 维分片的 DTensor(分布式张量)。pre-forward/backward hook 暂时 all-gather 为普通完整 Tensor;post-forward/backward 再恢复分片表示。Optimizer(优化器)在 fully_shard 之后创建,因而看到并更新本地参数 shard。
03 两张卡、8 个参数的极小手算#
设一层参数向量为
两路 FSDP2 沿第 0 维切分:
rank 0 常驻: [w0,w1,w2,w3]
rank 1 常驻: [w4,w5,w6,w7]
forward 前 all-gather:
rank 0 临时得到 [w0,...,w7]
rank 1 临时得到 [w0,...,w7]text两 rank 分别用不同数据算出局部完整梯度 。reduce-scatter(归约分片)先按元素求和或平均,再让每卡只保留一半:
rank 0 留: reduce(g0[0:4], g1[0:4])
rank 1 留: reduce(g0[4:8], g1[4:8])text若每卡 AdamW 为自己的 4 个参数保存一阶矩、二阶矩和 master weight,optimizer state 也自然只占一半。下一次 forward 再从两个新 shard 聚齐更新后的完整参数。
04 fully_shard 的最小正确顺序#
import os
import torch
from torch import nn
from torch.distributed.fsdp import fully_shard
torch.distributed.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
model = Transformer().cuda(local_rank)
# 自底向上:先给每个 block 建通信组,再处理根模块剩余参数
for block in model.blocks:
fully_shard(block, reshard_after_forward=True)
fully_shard(model, reshard_after_forward=True)
# 必须在 fully_shard 之后,让 optimizer 看到 DTensor parameter shards
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
for input_ids, labels in loader:
optimizer.zero_grad(set_to_none=True)
logits = model(input_ids.cuda(local_rank)) # [B_r,L,V]
loss = nn.functional.cross_entropy(
logits.transpose(1, 2), labels.cuda(local_rank)
)
loss.backward() # grad reduce-scatter
optimizer.step() # 更新本地 shardspython示例假设各 rank 的 batch shape 与有效元素数相同,并且 loader 已用 DistributedSampler(分布式采样器)无重复地切分数据;可变 token 数仍要使用上一篇的全局精确分母。多机环境必须使用 LOCAL_RANK 选择本机设备,不能把全局 rank 直接当 GPU 编号。
fully_shard 应自底向上调用:子模块先成为各自通信组,根模块最后接管尚未分组的参数。若只对根模型调用一次,所有参数会成为一个巨大组,forward 开始前一次性 all-gather,全程几乎无法与计算重叠,峰值也接近重新放入完整模型。
05 分片组边界为何同时控制内存与吞吐?#
假设四个 block 参数量分别为 [2,2,2,2] GB。
只 shard 根模块#
all-gather 8 GB -> 依次算 block 1..4 -> reduce-scatter 8 GBtext通信次数少但消息巨大;完整 8 GB 参数长时间驻留,计算与通信难重叠。
每个 block 单独 shard#
gather block1 2 GB -> compute1 -> free
gather block2 2 GB -> compute2 -> free
...text完整参数峰值下降,并可预取下一组;但组太碎会产生很多小 collective,延迟和 hook 开销上升。正确边界通常靠 Transformer block 或若干相邻 block,而不是每个小 Linear 都单独分组。
选择边界时测四个量:最大 unsharded group、预取时双 buffer 峰值、collective 次数与大小、通信被计算隐藏的比例。
06 reshard_after_forward 在换什么?#
reshard_after_forward=True 会在 forward 后释放完整参数,backward 前再 all-gather 一次:
- 优点:forward 与 backward 之间只保留 shards,激活高峰期显存更低;
- 代价:每轮该组通常多一次参数 all-gather。
设 False 则可能让完整参数跨过 forward 保留到 backward,减少一次通信但提高峰值。对最外层根模块、共享参数、梯度累积和不同 mesh,行为还需结合当前 API 契约验证,不能仅凭布尔值猜内存。
一个稳妥实验是对同一 batch 分别记录:max_memory_allocated()、每组 all-gather 次数和总字节、forward/backward 时间。若显存仍有余量而网络成为瓶颈,可以探索保留;若 activation 峰值已贴近上限,优先 reshard。
07 前向预取为何可能更快也可能 OOM?#
若计算 block 时异步 all-gather block ,通信可以被当前 block 的矩阵乘隐藏:
时间 ->
compute k: [==========]
gather k+1: [------]
compute k+1: [==========]text但重叠窗口内同时存在 block 的完整参数、block 的接收 buffer 和当前激活。最大组各 3 GB 时,预取可能瞬间增加约 3 GB 峰值。出现“单步偶发 OOM”时,要把 prefetch buffer 纳入 memory snapshot,而不是只看稳定态 shard 大小。
08 与梯度累积组合时会发生什么?#
梯度累积的每个 micro-batch 都要 forward/backward。若每次 forward 都 all-gather,累积 次就可能重复参数通信 次。reshard_after_forward、是否在同步边界保留参数,以及 FSDP2 的梯度同步控制会影响容量与通信,不能照搬 DDP no_sync() 的直觉。
数学语义仍遵循上一篇:累积 loss sum,并用全局有效 token 数归一化。工程验证要额外统计每个 optimizer update 内的 all-gather/reduce-scatter 次数,确保为省 activation 采用的大 没把网络吞吐压垮。
09 梯度为什么用 reduce-scatter 而不是 all-reduce?#
DDP all-reduce 后,每卡都得到完整平均梯度;FSDP 的 optimizer 只需要自己参数 shard 对应的梯度,因此没必要保留完整结果。
把长度 的向量看成 段:reduce-scatter 等价于“先跨 rank reduce,再把第 段交给 rank ”。其每卡输出只有 ,既完成数据并行归约,也直接得到 optimizer 所需的本地梯度 shard。
调试归一化时,不要假设库内部是 SUM 还是 AVG。构造两 rank 单参数例子,给 rank 0/1 不同输入,手算全局 batch 梯度,再与更新后的 full parameter 比较,比只观察 loss 更可靠。
10 分布式 checkpoint 为什么不能每卡随便 torch.save?#
FSDP2 参数是 DTensor shards。每个 rank 单独保存 model.state_dict() 的结果并不自动构成一个可迁移的完整 checkpoint;还需要保存 shard 元数据、world-size/mesh 信息、optimizer shards、随机状态、数据游标和成功更新计数。
PyTorch 官方建议通过 Distributed Checkpoint(分布式检查点)相关接口处理 sharded state dict,或显式用 DTensor API materialize full tensor。生产系统应测试:
- 同 world size 原地恢复;
- 不同 world size 的 reshard 恢复;
- rank/节点故障后是否只承认完整提交的 checkpoint;
- 恢复后的下一步 loss、学习率、scaler 和数据位置是否连续。
不要在每卡同时 gather 完整 state dict 后保存:它可能把 CPU/GPU 内存峰值重新推回未分片规模,并造成多份重复文件。
11 Shared Parameter(共享参数)与初始化陷阱#
语言模型常让 token embedding 与 LM head 权重绑定。同一参数若被两个不一致的 shard group 重复管理,可能破坏共享关系或触发难懂错误。应用 fully_shard 前后应检查 data_ptr/参数对象关系以及 state dict key,确保 tied weight 仍是同一个逻辑参数。
超大模型也不应先在每个 rank 的 CPU 上完整随机初始化,再搬到 GPU:CPU 内存会复制 份。可用 meta device 建结构、先分片再 materialize,并明确参数初始化如何在 ranks 间保持一致。不同 rank 各自用不同随机数填 shard 并不天然等价于同一个全局初始化。
12 如何证明分片训练仍等价?#
一个最小测试流程:
# 1. 固定种子,构造极小模型与完整全局 batch
# 2. 路径 A:单卡 FP32 得到 loss、full grads、更新后参数
# 3. 路径 B:两 rank FSDP2,各取 batch 一半,执行一步
# 4. 将参数 shards 聚成 full tensor 后比较
torch.testing.assert_close(full_after_fsdp, after_single,
rtol=1e-5, atol=1e-7)python若失败,依次排查:初始 full parameter 是否一致、sampler 是否无重复无遗漏、loss 分母、collective group、共享权重、dropout/RNG、混合精度,最后才归因于浮点归约顺序。
13 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| shard 后仍在 forward 开始 OOM | 根分组一次 all-gather 全模型 | 列出每个 FSDP group 参数字节数 |
| 显存偶发尖峰 | prefetch 与当前完整参数双重驻留 | 对照 collective 时间线与 memory snapshot |
| optimizer state 仍像完整模型 | optimizer 在 fully_shard 前创建 | 打印参数类型、本地 shape 与 state 元素数 |
| 卡数增加反而更慢 | 模块分组太碎或网络延迟高 | 汇总 collective 次数、大小、未重叠时间 |
| checkpoint 单卡可读但整体不能恢复 | 只保存了局部 shard,缺少元数据 | 做全流程故障恢复演练 |
| 两 rank 更新后参数不等 | sampler、归一化或 process group 错 | gather full gradient 对单卡基准 |
| tied embedding 不再共享 | 分片边界错误处理共享参数 | 对比分片前后参数身份与 state keys |
14 FSDP2 不解决什么?#
FSDP2 不能让单层计算天然跨卡:层执行时仍需要完整参数;巨型 embedding、超宽 Linear 或 attention 临时张量若单卡放不下,需 Tensor/Sequence Parallel。它也不会消除激活,长上下文仍可能需要 checkpointing、Flash Attention 或序列切分。
当模型很小、网络较慢或 batch 已足够大时,重复 all-gather 的成本可能高于内存收益。普通 DDP 代码更简单、通信模式更成熟,能放下时通常应先作为基线。
15 今天真正需要记住什么?#
- FSDP2 平时分片参数、梯度和 optimizer state,计算某组前临时 all-gather 完整参数,反向后 reduce-scatter 梯度。
- 每卡稳定态约除以 world size,但真实峰值还包含最大完整参数组、预取 buffer、激活和临时 workspace。
fully_shard的模块边界就是通信分组边界:太大导致峰值高且难重叠,太碎导致大量小 collective。- optimizer 创建顺序、共享参数、梯度累积和分布式 checkpoint 都必须按分片语义重新验证。
16 思考题与小练习#
- 模型含 12 GB 权重、12 GB 梯度和 72 GB optimizer state。用 8 路理想全分片计算每卡稳定态;若最大 unsharded group 为 3 GB、预取下一组也为 3 GB,再估算不含激活的峰值下界。
- 对两个 rank 的梯度
g0=[1,2,3,4]、g1=[5,6,7,8],分别写出 SUM reduce-scatter 后每卡保留的 shard;若目标是全局平均,应如何缩放? - 为 24 个 Transformer blocks 设计两种
fully_shard分组方案,列出你会用哪些 profiler 指标决定选哪一种。
相关工作#
- Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models ↗,系统提出对 optimizer state、梯度和参数逐级分片。
- Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel ↗,总结 FSDP 的架构、通信与生产经验。
- Xu et al., Automatic Cross-Replica Sharding of Weight Update Computation in Data-Parallel Training ↗,研究数据并行更新状态的自动分片。
- Sergeev and Del Balso, Horovod: fast and easy distributed deep learning in TensorFlow ↗,说明基于 collective 的数据并行工程。
- Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training ↗,解释 PyTorch 分布式梯度同步与通信重叠设计。
17 下一篇预告#
FSDP2 让模型长期状态不再完整复制,但某一层计算时仍会临时聚齐参数。下一篇将深入 Tensor Parallel:如何把 MLP 的上投影做列并行、下投影做行并行,并用一次必要 collective 串起正确的张量形状。