观文听傑

返回

上一篇比较了 Data Parallel(数据并行)、Tensor Parallel(张量并行)与 Pipeline Parallel(流水线并行)。普通 DDP 虽然把 batch 分给多卡,却在每个 rank 保存完整参数、梯度和 optimizer state:卡越多,总副本越多,单卡容量没有下降。

本篇聚焦 Fully Sharded Data Parallel 2(全分片数据并行第二代,FSDP2)的一个核心机制:平时把模型状态分片保存,只在某个模块真正计算前临时聚齐参数,并在反向后把梯度归约回分片。

01 DDP 到底复制了多少状态?#

设参数量为 PP,数据并行 world size 为 RR。每参数字节数仍记为权重 bwb_w、梯度 bgb_g、optimizer state bob_o

普通 DDP 每卡长期占用近似

MDDP,state=P(bw+bg+bo)M_{\text{DDP,state}}=P(b_w+b_g+b_o)

FSDP 理想稳定态则近似

Mshard,stateP(bw+bg+bo)RM_{\text{shard,state}}\approx\frac{P(b_w+b_g+b_o)}{R}

但“除以 RR”不是峰值答案。计算某个参数组前,FSDP 必须 all-gather(全聚合)完整参数;通信还可能同时持有下一组的预取 buffer。更实际的粗略式是

MpeakMshard,state+Mlargest unsharded group+Mprefetch+Ma+MtempM_{\text{peak}}\approx M_{\text{shard,state}}+ M_{\text{largest unsharded group}}+M_{\text{prefetch}}+M_a+M_{\text{temp}}

因此模型总状态均分后能放下,不代表某个巨大的根分组也能安全 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 shard
mermaid

PyTorch 2.14 的 fully_shard(module) 会原地把参数转换为按第 0 维分片的 DTensor(分布式张量)。pre-forward/backward hook 暂时 all-gather 为普通完整 Tensor;post-forward/backward 再恢复分片表示。Optimizer(优化器)在 fully_shard 之后创建,因而看到并更新本地参数 shard。

03 两张卡、8 个参数的极小手算#

设一层参数向量为

w=[w0,w1,w2,w3,w4,w5,w6,w7]w=[w_0,w_1,w_2,w_3,w_4,w_5,w_6,w_7]

两路 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 分别用不同数据算出局部完整梯度 g(0),g(1)R8g^{(0)},g^{(1)}\in\mathbb R^8。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 的最小正确顺序#

示例假设各 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 GB
text

通信次数少但消息巨大;完整 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 kk 时异步 all-gather block k+1k+1,通信可以被当前 block 的矩阵乘隐藏:

时间 ->
compute k:       [==========]
gather k+1:        [------]
compute k+1:                 [==========]
text

但重叠窗口内同时存在 block kk 的完整参数、block k+1k+1 的接收 buffer 和当前激活。最大组各 3 GB 时,预取可能瞬间增加约 3 GB 峰值。出现“单步偶发 OOM”时,要把 prefetch buffer 纳入 memory snapshot,而不是只看稳定态 shard 大小。

08 与梯度累积组合时会发生什么?#

梯度累积的每个 micro-batch 都要 forward/backward。若每次 forward 都 all-gather,累积 KK 次就可能重复参数通信 KK 次。reshard_after_forward、是否在同步边界保留参数,以及 FSDP2 的梯度同步控制会影响容量与通信,不能照搬 DDP no_sync() 的直觉。

数学语义仍遵循上一篇:累积 loss sum,并用全局有效 token 数归一化。工程验证要额外统计每个 optimizer update 内的 all-gather/reduce-scatter 次数,确保为省 activation 采用的大 KK 没把网络吞吐压垮。

09 梯度为什么用 reduce-scatter 而不是 all-reduce?#

DDP all-reduce 后,每卡都得到完整平均梯度;FSDP 的 optimizer 只需要自己参数 shard 对应的梯度,因此没必要保留完整结果。

把长度 PP 的向量看成 RR 段:reduce-scatter 等价于“先跨 rank reduce,再把第 rr 段交给 rank rr”。其每卡输出只有 P/RP/R,既完成数据并行归约,也直接得到 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。生产系统应测试:

  1. 同 world size 原地恢复;
  2. 不同 world size 的 reshard 恢复;
  3. rank/节点故障后是否只承认完整提交的 checkpoint;
  4. 恢复后的下一步 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 内存会复制 RR 份。可用 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 今天真正需要记住什么?#

  1. FSDP2 平时分片参数、梯度和 optimizer state,计算某组前临时 all-gather 完整参数,反向后 reduce-scatter 梯度。
  2. 每卡稳定态约除以 world size,但真实峰值还包含最大完整参数组、预取 buffer、激活和临时 workspace。
  3. fully_shard 的模块边界就是通信分组边界:太大导致峰值高且难重叠,太碎导致大量小 collective。
  4. optimizer 创建顺序、共享参数、梯度累积和分布式 checkpoint 都必须按分片语义重新验证。

16 思考题与小练习#

  1. 模型含 12 GB 权重、12 GB 梯度和 72 GB optimizer state。用 8 路理想全分片计算每卡稳定态;若最大 unsharded group 为 3 GB、预取下一组也为 3 GB,再估算不含激活的峰值下界。
  2. 对两个 rank 的梯度 g0=[1,2,3,4]g1=[5,6,7,8],分别写出 SUM reduce-scatter 后每卡保留的 shard;若目标是全局平均,应如何缩放?
  3. 为 24 个 Transformer blocks 设计两种 fully_shard 分组方案,列出你会用哪些 profiler 指标决定选哪一种。

相关工作#

  1. Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models,系统提出对 optimizer state、梯度和参数逐级分片。
  2. Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel,总结 FSDP 的架构、通信与生产经验。
  3. Xu et al., Automatic Cross-Replica Sharding of Weight Update Computation in Data-Parallel Training,研究数据并行更新状态的自动分片。
  4. Sergeev and Del Balso, Horovod: fast and easy distributed deep learning in TensorFlow,说明基于 collective 的数据并行工程。
  5. Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training,解释 PyTorch 分布式梯度同步与通信重叠设计。

17 下一篇预告#

FSDP2 让模型长期状态不再完整复制,但某一层计算时仍会临时聚齐参数。下一篇将深入 Tensor Parallel:如何把 MLP 的上投影做列并行、下投影做行并行,并用一次必要 collective 串起正确的张量形状。

每张卡为何还要保留完整模型?FSDP2 的参数全聚合、梯度归约分片与显存峰值
https://zwjcode.cn/blog/fsdp2-parameter-gradient-optimizer-state-sharding
作者
发布于 2026年9月16日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。