← 返回全部思考

训练 Infra 02:显存账——把 85.6GB 峰值逐项拆开算

从一个参数在训练时占多少字节开始,反推出 35B MoE 的专家/非专家参数拆分,算清 bf16 权重、梯度 buffer、distributed optimizer 分片、激活与 logits 各吃多少,并复盘 FusedAdam 惰性初始化导致 39 卡静默死锁 5 小时的事故。

训练 Infra 02:显存账

本节属于大模型训练 Infra 实战课。上一节建立了带宽地图,这一节算显存。

上一节我们说 H20 的 96GB 显存"买通"了 TP=1 的设计,实测峰值 85.6GB。这一节把这 85.6GB 逐项拆开——你会发现它几乎每一个 GB 都能被解释,而解释不了的那部分,正好是优化空间。

1. 一个参数的"随从队伍"

直觉

新手最容易犯的错,是把"模型有多大"等同于"要多少显存"。

35B 参数、bf16、每参数 2 字节 = 70GB——听起来两张 H20 就够了。但训练时一个参数不是孤身上路,它带着一支随从队伍:

随从 干什么 精度 字节/参数
权重 前向计算用 bf16 2
梯度 反向累积 bf16 或 fp32 2 或 4
fp32 主参数 优化器更新的"真身" fp32 4
Adam 一阶动量 m 梯度的滑动平均 fp32 4
Adam 二阶动量 v 梯度平方的滑动平均 fp32 4

为什么要有 fp32 主参数

这是混合精度训练里最反直觉的一件事:明明用 bf16 训练,为什么还要留一份 fp32 权重?

因为更新量太小了。bf16 只有 8 位尾数,能表示的相对精度约 1/256。当 lr × grad 相对于权重本身小于这个精度时,w = w + Δ 这个加法结果等于 w 本身——更新被直接吞掉。训练后期学习率衰减,这种情况会大面积发生,表现为 loss 停滞。

所以标准做法是:前向/反向用 bf16 算(省显存、跑 tensor core),更新在 fp32 主参数上做(保住小更新),更新完再 cast 回 bf16 权重。代价就是那多出来的 4 字节/参数。

公式

混合精度 + Adam 的显存开销(每参数):
  bf16 权重      2 B
  bf16 梯度      2 B      (若用 fp32 梯度则 4 B)
  fp32 主参数    4 B
  Adam m         4 B
  Adam v         4 B
  ------------------------
  合计          16 B/参数  (fp32 梯度时 18 B/参数)

代入 35B,不做任何分片:

35e9 × 16 B = 560 GB

单卡 96GB,差 5.8 倍。所以分片不是优化,是能不能跑起来的前提。 这就是下一节讲并行的动机。

2. 反推:这 35B 里专家占多少

项目实测:EP=8 分完专家后,每卡 6.9B 参数,bf16 权重 13.8GB。

这两个数字里藏着一个可以反推出来的信息。EP 只切专家参数——非专家部分(attention、embedding、norm、router)每张卡都要有一份完整副本。所以:

设 专家参数 = E,  非专家参数 = X

X + E     = 35B      (总参数)
X + E/8   = 6.9B     (EP8 后每卡参数)

两式相减:  E × (1 - 1/8) = 35 - 6.9 = 28.1
        →  E ≈ 32.1B ,  X ≈ 2.9B

校验:X + E/8 = 2.9 + 4.0 = 6.9B,× 2 字节 = 13.8GB —— 和实测完全对上。

这个反推有实用价值:它告诉你这个模型 92% 的参数都在专家里。所以 EP 是这个模型唯一真正有效的参数切分手段,而 TP 切那 2.9B 的非专家部分,省不下多少显存却要付每层两次 all-reduce——这是第 3 节"为什么 TP=1"的又一条依据。

3. 梯度 buffer:那次减半的 13.8GB

Megatron 会为梯度开一块连续的大 buffer(为了 reduce-scatter 能一次搞定,且能和反向计算重叠)。

fp32 梯度 buffer: 6.9e9 × 4 B = 27.6 GB   ← 项目实测 27.7GB
bf16 梯度 buffer: 6.9e9 × 2 B = 13.8 GB   ← 改完之后

项目里把它从 fp32 改成 bf16,账面省 13.8GB,实测峰值从 95GB 降到 85.6GB(降 9.4GB)。

差额约 4.4GB 值得解释一下,因为这是个常见困惑:账面省的和实测省的往往不等。梯度 buffer 降精度后,reduce-scatter 的输出、以及给优化器喂的那一份,通常仍需要 fp32 的中间落点;此外峰值时刻可能不在同一处。所以显存账要"算 + 量"结合——算出来的是上界,量出来的才是真相。

常见坑:bf16 梯度会不会掉精度

会,但通常可接受,因为:

  1. 梯度 reduce 时的累加在 NCCL 内部/或用 fp32 累加器完成
  2. 动量 m、v 仍是 fp32,起到了平滑作用
  3. 真正危险的是梯度累积步数很大时,bf16 累加会有明显舍入误差

项目里梯度累积只有 2 步,风险很低。如果哪天把 global batch 拉到 640(累积 16 步),就要重新评估——那时候更该考虑 fp32 累加 + bf16 存储的方案。

4. distributed optimizer:专家 DP 组只有 5 的代价

原理

Megatron 的 distributed optimizer 做的事:把优化器状态(fp32 主参数 + m + v,共 12 字节/参数)沿 DP 维切开,每个 rank 只持有 1/DP 份,更新完再 all-gather 回权重。这本质上就是 ZeRO-1(第 3 节细讲两者差异)。

分片效果直接由 DP 组的大小决定。而这里有个 MoE 特有的陷阱:

非专家参数: 每张卡都有完整副本 → DP 组 = 40
专家参数  : 已被 EP8 切成 8 份 → 剩下的 DP 组 = 40 / 8 = 5

专家参数的 DP 组只有 5,不是 40。 分片效果差 8 倍。

代入真实数字

非专家: X = 2.9B , DP组 40
  2.9e9 × 12 B / 40 = 0.87 GB

专家  : 每卡 E/8 = 4.0B , DP组 5
  4.0e9 × 12 B / 5  = 9.63 GB

优化器状态合计 ≈ 10.5 GB

对照一下"如果专家 DP 组也是 40"的假想情况:

4.0e9 × 12 B / 40 = 1.20 GB
→ 现状多吃 9.63 - 1.20 = 8.4 GB

8.4GB。 这是 EP=8 这个选择的显存代价,而 EP=8 换来的是 all-to-all 全部走 NVLink。这笔交易在 H20 上是划算的(96GB 付得起 8.4GB),但在 80GB 的卡上就可能是压垮骆驼的那一根——又一次印证了上一节的结论。

常驻显存小结

bf16 权重          13.8 GB
bf16 梯度 buffer   13.8 GB
优化器状态         10.5 GB
--------------------------------
常驻小计           38.1 GB

实测峰值           85.6 GB
→ 激活 + 临时缓冲 ≈ 47.5 GB

47.5GB 花在激活和临时缓冲上,比常驻的模型状态还多。 这就是长序列训练的真面目,也是下一小节的主题。

5. 激活显存:和序列长度的关系

直觉

激活(activation)是前向过程中为了反向能算梯度而必须留下的中间结果。它和参数量无关,和 batch × 序列长度 × 模型宽度 × 层数 有关。

三种量级

不开任何优化时,每层要存的中间量很多(QKV、attention 输出、FFN 中间层、各种 norm 的输入),系数能到十几倍 hidden。开了 recompute 之后,只存"每层的输入"这一个张量:

full recompute 下的激活显存 ≈ layers × seq × hidden × 2 B × micro_batch

代入 micro_batch=1, seq=49152, 假设 48 层 / hidden 2048:
  48 × 49152 × 2048 × 2 B ≈ 9.7 GB

(层数和 hidden 请代入你们的实际配置,这里只是量级示意。)

那剩下的约 38GB 去哪了?重点怀疑对象是 logits。

隐形巨兽:logits

输出层要把每个位置映射到词表:

logits bf16: 49152 × 151936 × 2 B = 14.9 GB
logits fp32: 49152 × 151936 × 4 B = 29.9 GB

一个张量就 15 到 30GB。 长序列训练里,logits 经常比整个模型的激活加起来还大,而且它常常被忽略,因为在短序列时它小得看不见。

对策是成熟的:

  • fused cross-entropy / chunked loss:把序列切块,算一块 loss 释放一块 logits,峰值降到 1/chunk
  • vocab 并行(Megatron 的 --tensor-model-parallel-size 会顺带切 vocab,但 TP=1 时没有这个红利)
  • 避免把 logits 升到 fp32 后整块留存

这是我建议你们优先量一下的地方。 如果 logits 确实占了 15GB 以上,换 fused CE 是纯赚——不损精度、不加通信,直接腾出显存给更大的 micro batch,而更大的 micro batch 直接提 MFU(第 6 节)。

attention 与序列长度:为什么不是平方爆炸

朴素 attention 的显存是 O(seq²)——49K 序列下 attention 矩阵单头就是 49152² × 2B = 4.8GB,完全不可行。

救命的有两件事:

  1. FlashAttention:不落地 attention 矩阵,分块在 SRAM 里算,显存降到 O(seq)
  2. GDN 线性注意力:本身就是 O(seq),用状态递推替代全局 attention

Qwen3.6-35B-A3B 是"每 4 层 1 层全注意力 + 其余 GDN"的混合结构。这意味着只有 1/4 的层需要付全注意力的代价——这是它能在 49K 序列上跑起来的结构性原因。

6. recompute:省显存,费算力

原理

前向时不保存中间激活,反向需要时重新前向算一遍。

不 recompute : 前向 2ND + 反向 4ND = 6ND
full recompute: 前向 2ND + 重算 2ND + 反向 4ND = 8ND

计算量放大系数 = 8/6 = 4/3 ≈ 1.33

这个 4/3 会在第 6 节算 MFU 时直接用到——它是 MFU 和 HFU 两个指标的差别来源。

取舍

策略 显存 算力 什么时候用
不 recompute 最高 1.0× 显存富裕、追极限吞吐
selective recompute 中 ~1.1× 只重算便宜的部分(如 attention),跳过贵的
full recompute 最低 1.33× 显存是硬约束时

项目用的是 full recompute,因为 49K 序列下没有选择。但一旦 logits 那 15GB 被 fused CE 释放出来,selective recompute 就重新变成可选项——那可能是 20% 左右的吞吐提升。这是一条清晰的优化链:省 logits 显存 → 换掉 full recompute → 提 MFU。

7. 精度感知优化器:bf16 动量

再往下省,可以动优化器状态本身:

标准:  fp32 主参数 4 + fp32 m 4 + fp32 v 4 = 12 B/参数
bf16 动量: fp32 主参数 4 + bf16 m 2 + bf16 v 2 = 8 B/参数

代入专家侧:4.0e9 × 8 / 5 = 6.4GB,比现在的 9.63GB 再省 3.2GB。

风险在 **v(二阶动量)**上:v 是梯度平方的滑动平均,动态范围极大(跨很多个数量级),bf16 的 8 位尾数容易在 sqrt(v) + eps 这一步引入可观误差,表现为训练后期不稳。所以工程上通常的顺序是:先降 m,观察若干百步的 loss 曲线,再考虑降 v。不要两个一起改——出问题时你分不清是谁的锅。

8. 事故复盘 ①:优化器惰性初始化,39 卡静默死锁 5 小时

这是本课程最值得反复讲的一次事故。

现象

任务启动,前向正常、反向正常,第一个 step 走完之后,个别 rank 报 OOM。剩下 39 个 rank 没有任何报错,GPU 利用率显示 100%。整个任务就这么"跑"了 5 小时,什么都没产出。

根因

FusedAdam(以及大多数融合优化器)的状态是惰性分配的。

exp_avg(m)和 exp_avg_sq(v)不是在优化器构造时 alloc 的,而是在第一次 optimizer.step() 真正碰到某个参数时才 alloc。

于是显存曲线长这样:

时刻            显存占用
---------------------------------------------
权重加载后      13.8 GB
前向峰值        13.8 + 激活
反向峰值        13.8 + 13.8(梯度) + 激活   ← 很多人以为这就是峰值
第一次 step     + 优化器状态 10.5 GB       ← 真正的峰值在这里

你在反向结束时看到的显存曲线是骗人的。 真正的峰值在第一次 step,而那时前面的激活可能还没完全释放,两个峰叠在一起。

更阴的是它的非均匀性:不同 rank 分到的参数分片大小不同(尤其 MoE 的专家分配不均、以及 padding 差异),所以只有最"倒霉"的那一两个 rank 会 OOM。

为什么会静默死锁

这是分布式训练里最重要的一个心智模型:

集体通信是一场约会。all-reduce、all-gather、all-to-all 都要求组内每一个 rank 都到场。少一个人,其余所有人无限期等待。

OOM 的 rank 抛异常、退出了。剩下 39 个 rank 进入下一个 all-gather,然后永远等那个已经死掉的第 40 个。

而 NCCL 的等待是 busy-wait(自旋)——它会拿一个 CUDA kernel 死循环轮询。所以:

nvidia-smi 显示 GPU 利用率 100%
实际上一条有用指令都没执行

这就是"100% 空转"。第 4 节会把这个"GPU 利用率形态学"讲透:100% 空转 = 卡在 NCCL 等人;0% = 卡在 CPU 侧(数据加载、Python 死锁、文件 IO)。两种形态指向完全不同的根因。

怎么防

按性价比排序:

  1. 启动即做一次 dry-run step:用真实 shape 但立刻丢弃结果,把优化器状态在训练开始前就 alloc 出来。OOM 会在第 10 秒暴露,而不是第 1 个 step 之后。
  2. 把峰值显存算出来再启动(就是本节做的事),而不是"跑起来看看"。
  3. 调短超时:ddp_timeout 默认经常是 30 分钟到数小时,把它降到 10 分钟量级。死锁不可避免,但没必要让它烧 5 小时(第 4 节详解)。
  4. 每个 rank 打显存日志:记录 max_memory_allocated,而且要记 step 1 之后的值。只看 step 0 的显存报告等于没看。
  5. 给自己留余量:85.6 / 96 = 89% 已经相当满了。多模态数据、变长序列、偶发的长样本都可能让某一步多吃几个 GB。

常见坑清单

  • 只看 rank0 的日志:rank0 往往正是没 OOM 的那个,日志一片祥和
  • 看 nvidia-smi 的利用率判断健康:100% 可能是死锁
  • 用 step 0 的显存做容量规划:漏掉优化器状态
  • 假设各 rank 显存对称:MoE 下几乎从不对称

9. 本节账本汇总

项目 数值 由什么决定
bf16 权重 13.8 GB 参数量 ÷ EP
bf16 梯度 buffer 13.8 GB 同上,精度可选
优化器状态 10.5 GB 12 B/参数 ÷ DP组(专家仅 5)
常驻小计 38.1 GB 模型状态
激活(full recompute) 约 10 GB layers × seq × hidden
logits 等临时 约 37 GB(待量) seq × vocab ← 优化重点
实测峰值 85.6 GB
单卡容量 96 GB 占用率 89%

一句话总结:模型状态只占 38GB,一半以上的显存被序列长度带来的激活和 logits 吃掉了。 所以这个任务真正的显存优化方向不在模型侧,在序列侧。

思考题

1. 假设你把 logits 的峰值从 15GB 压到 2GB(fused chunked CE),腾出 13GB。你有三种花法,各会带来什么?

  • (a) micro batch 从 1 提到 2
  • (b) full recompute 换成 selective recompute
  • © 梯度 buffer 从 bf16 退回 fp32,换取更稳的大批量累积

2. 现在把 EP 从 8 改成 4(专家切 4 份,all-to-all 仍在单机内,只用 4 卡一组):

  • (a) 每卡参数量变成多少?bf16 权重多少 GB?
  • (b) 专家的 DP 组变成多少?优化器状态是变大还是变小?
  • © 综合起来,峰值显存会上升还是下降?

下一节讲并行策略全景:DP/ZeRO/TP/PP/EP/CP 逐个拆开,然后用本项目的约束条件,把"为什么是 TP1/PP1/EP8/DP40"完整推演一遍。