← 返回全部思考

训练 Infra 06:手把手算 MFU——8-10% 其实被低估了

用真实数字走一遍 6ND、recompute 系数 4/3、MFU 与 HFU 的区别,交叉校验到 tok/s 级别,然后指出两件被忽略的事:定长 padding 可能浪费 41%,而 6ND 在 49K 序列下漏掉的注意力项占总算力三到七成。

训练 Infra 06:吞吐与效率

本节属于大模型训练 Infra 实战课。这是全课最"算"的一节。

项目实测 MFU 是 8-10%。这个数字看起来很难看——公开分享里稠密模型在 H100 上能做到 40-50%。

这一节我们把它算清楚。结论会有点意外:8-10% 这个数其实被低估了,真实的硬件利用率大概在 10-20% 之间;同时还有一块可能高达 41% 的浪费藏在别处。

1. MFU 的定义

直觉

MFU(Model FLOPs Utilization)= 你的模型在数学上必须做的运算量,除以硬件峰值能做的运算量。

MFU = (模型必需 FLOPs / 耗时) / 硬件峰值 FLOPs每秒

关键词是"必需"。它刻意不包含:

  • recompute 重算的部分(那是为省显存自愿多付的)
  • padding 上的无效计算
  • 通信、数据加载等非计算时间(这些体现为耗时变长)

所以 MFU 衡量的是"有效产出 / 理论上限",是一个端到端的效率指标。

MFU vs HFU

另一个常见指标是 HFU(Hardware FLOPs Utilization),它把 recompute 也算进去:

MFU: 只算必需的           → 反映"有效效率"
HFU: 算硬件实际执行的      → 反映"硬件忙不忙"
HFU = MFU × recompute系数

full recompute 下 HFU = MFU × 4/3。这两个数经常被混着用,是很多"我的 MFU 是多少"讨论对不上的原因。 本节会把两个都算出来。

2. 公式:6ND 从哪来

推导

对一个参数量为 N 的模型处理 D 个 token:

前向: 每个参数对每个 token 做 1 次乘 + 1 次加 = 2 FLOPs
      → 2ND

反向: 需要算两组梯度(对输入的 + 对权重的),各约等于一次前向
      → 4ND

合计: 6ND

MoE 里 N 取哪个

取激活参数,不是总参数。 因为 MoE 每个 token 只经过被路由到的那几个专家,没被激活的专家一次乘法都没做。

Qwen3.6-35B-A3B: N = 3e9 (激活),不是 35e9

这是 MoE 算 MFU 最容易错的地方,错的代价见下面。

recompute 系数

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

系数 = 8/6 = 4/3 ≈ 1.333

3. 代入项目数字

输入

每步 token 数 D = 2.31e6   (80 样本 × 平均 28,875 token)
激活参数   N = 3e9
单步耗时   T = 120 秒
集群峰值     = 40 × 148 TFLOPs = 5.92 PFLOPs

算

必需 FLOPs = 6ND = 6 × 3e9 × 2.31e6 = 4.158e16 FLOPs

有效算力  = 4.158e16 / 120 = 3.465e14 = 346.5 TFLOPs/s
单卡      = 346.5 / 40 = 8.66 TFLOPs/s

MFU = 346.5 / 5920 = 5.85%
HFU = 5.85% × 4/3  = 7.80%

交叉校验(这一步很重要)

任何 MFU 计算都应该用另一条路验一遍。用吞吐:

tok/s = 2.31e6 / 120 = 19,250        ← 实测 19,300  ✓
单卡  = 19,250 / 40  = 481.2 tok/s   ← 实测 482     ✓
单卡  = 8.66 / 148   = 5.85%         ← 与集群口径一致 ✓

三条路互相对上,说明这套数字是自洽的。 项目报的 8-10% 对应的是 HFU 口径(7.80%),稍微乐观一点或者用了略不同的 N。

一个必须避开的错算

如果误用总参数 35B:

6 × 35e9 × 2.31e6 / 120 / 5.92e15 = 68%

68% 会让你以为已经接近硬件极限、没有优化空间了。 这个错算在 MoE 上极常见,因为 HF config 里显眼的就是总参数。

4. 被忽略的浪费一:padding

这是我看到的最大一块潜在收益。

配置的序列长度  = 49,152
实际平均 token  = 2.31e6 / 80 = 28,875

如果按定长 49,152 补齐:
  槽位总数 = 80 × 49,152 = 3,932,160
  有效占比 = 2,310,000 / 3,932,160 = 58.7%
  浪费     = 41.3%

如果确实在做定长 padding,那么四成的算力花在了 pad token 上。 完美 packing 能带来:

提速 = 3,932,160 / 2,310,000 = 1.70×
MFU: 5.85% → 9.96%
单步: 120 秒 → 70 秒

先确认这件事

我不能替你断定是不是在 padding——Megatron-SWIFT 支持 packing,也支持 varlen(用 cu_seqlens 让 FlashAttention 按真实长度做块对角注意力)。所以这一节最该带走的行动项是去量一下:

1. 看日志里每步的实际 token 数是 2.31e6 还是 3.93e6
   (2.31e6 说明已经在按有效 token 计,但计算是否也只做了有效部分要另看)
2. 确认是否开了 packing / 是否传了 cu_seqlens
3. 打印一个 batch 的 attention_mask 或 position_ids,看有没有大片 pad

三种可能的结论:

情况 说明 行动
已开 packing 41% 不存在,MFU 就是 5.85% 看下一节的注意力项
定长 padding 41% 真实存在 优先做 packing,1.7× 提速
部分 packing 介于两者之间 量化后决定

packing 的代价

packing 不是免费的:

  • 注意力必须支持块对角,否则不同样本会互相看到(严重的数据污染)。这要求 FlashAttention 的 varlen 接口 + 正确的 cu_seqlens
  • position embedding / RoPE 要按样本重置,不能连续编号
  • loss mask 要正确切分
  • 样本组合是个装箱问题,简单贪心就够用,但要注意别让某个 batch 全是长样本导致 OOM

这些 Megatron-SWIFT 基本都有支持,但每一条都值得实际验证一遍——packing 出 bug 的典型症状是 loss 看起来正常但模型学到了跨样本的假关联,很难发现。

5. 被忽略的浪费二:6ND 漏掉了注意力

问题

6ND 只算了和参数量成正比的部分(各种线性层)。但注意力里的 QK^T 和 AV 两个矩阵乘不涉及参数,它们的规模是:

每层前向 attention FLOPs ≈ 4 × seq² × hidden

注意是 seq²。短序列时这项小到可以忽略,所以 6ND 是个好近似。但 49K 序列下:

seq² = 49,152² = 2.42e9    ← 这个数很大

敏感性分析

层数和 hidden size 我不知道确切值,所以做三组假设("全注意力层"按每 4 层 1 层计):

配置 序列口径 注意力 FLOPs 占总量 真实 MFU
48层 h2048 (12层全注意力) 定长 49,152 5.70e16 58% 13.9%
48层 h2048 (12层全注意力) 实际均长 28,875 1.97e16 32% 8.6%
48层 h4096 (12层全注意力) 定长 49,152 1.14e17 73% 21.9%
48层 h4096 (12层全注意力) 实际均长 28,875 3.93e16 49% 11.4%
64层 h2560 (16层全注意力) 定长 49,152 9.50e16 70% 19.2%
64层 h2560 (16层全注意力) 实际均长 28,875 3.28e16 44% 10.5%

结论

在 49K 序列下,注意力占了总算力的三到七成。6ND 严重低估了硬件实际做的运算量。

所以那个 5.85% / 8-10% 的数字,衡量的是"按参数量折算的有效效率",而硬件真实利用率大概在 10% 到 22% 之间,取决于确切的 hidden size 和序列口径。

这不是文字游戏,它有两个实际含义:

  1. 不要拿这个 MFU 跟稠密短序列模型的 40-50% 直接比。 那是两个口径。
  2. 优化空间比 MFU 数字暗示的要小。 如果硬件真实利用率已经 20%,你能榨出的余量是有限的——H20 的 148 TFLOPs 里,注意力那部分的 GEMM 形状(长序列、小 head_dim)本身就难跑满 tensor core。

请把你们真实的 hidden_size、num_layers、全注意力层数代进去重算一遍。 这个数决定了后面所有优化的预期收益。

6. 为什么 MoE 和长序列天然拉低 MFU

MoE 的三个损耗

(1) GEMM 变瘦

MoE 把一个大 FFN 拆成很多小专家。每个专家只拿到 总token数 / 专家数 × top_k 个 token,GEMM 的 M 维大幅缩小。而 tensor core 需要足够大的矩阵才能跑满——瘦长的 GEMM 效率显著低于方阵。

稠密 FFN:  [2.31e6 token, hidden] × [hidden, 4·hidden]   ← M 巨大,跑满
MoE 专家:  [2.31e6 × top_k / 专家数, hidden] × [...]     ← M 小很多

(2) all-to-all 与 permute 开销

第 4 节算过通信只占 0.6%,但 permute/unpermute(按路由重排 token)本身是 memory-bound 的搬运操作,不产生 FLOPs 却要花时间。

(3) 负载不均

某些专家热、某些冷。热专家所在的卡成为 straggler,其他卡等它。这部分损耗随路由的不均衡程度波动。

长序列的两个损耗

(1) attention 的 GEMM 形状差

QK^T 是 [seq, head_dim] × [head_dim, seq]。head_dim 通常只有 128,是个极扁的矩阵乘,tensor core 利用率天然低于 hidden 维的方阵乘。

(2) micro batch 被压到 1

这是最隐蔽的一条。显存约束逼得 micro batch=1(第 2 节),而 batch 维是 GEMM 里最容易"喂饱"硬件的维度。

micro batch=1: 每个 GEMM 的 batch 维 = 1
micro batch=4: 同样的权重读一次,算 4 倍的数据 → 运算强度 ×4

回到第 1 节的 Roofline:H20 的 machine balance 是 37 FLOPs/byte。micro batch=1 时,很多算子的运算强度掉到这个值附近甚至以下,从 compute-bound 退化成 memory-bound——峰值算力根本用不上。

7. H20 上这是什么水平

给一个参考坐标系(都是 6ND 口径的 MFU):

场景 典型 MFU
稠密模型 + H100 + 4K序列 + 精调配置 40 - 50%
稠密模型 + 长序列(32K+) 25 - 35%
MoE + 中等序列 20 - 30%
MoE + 超长序列(49K) + micro batch 1 8 - 15%
上述 + 弱算力卡(H20 的算力/带宽比只有 H100 的 1/8) 5 - 12%

所以 5.85%(MFU)/ 7.8%(HFU)落在这个组合的正常区间下沿。 不是配置错了,是这个形状本来就难。

但"正常"不等于"没救"。下面是有序的优化路径。

8. 提升路径(按性价比排序)

① 确认并做 packing —— 潜在 1.7×

前面算过。如果确实在 padding,这是唯一一个能带来 70% 提升的单项改动。 先量再做。

② 干掉 logits 显存,换掉 full recompute —— 潜在 1.25×

第 2 节的链条:

fused/chunked cross-entropy 把 logits 峰值从 ~15GB 压到 ~2GB
  → 腾出 13GB
  → full recompute 换成 selective recompute
  → 算力系数从 4/3 (1.333) 降到约 1.1
  → 提速 1.333/1.1 = 1.21×

而且 fused CE 本身也快(省了大张量的读写)。

③ 加大 micro batch —— 潜在 1.1~1.3×

同样靠 ① 或 ② 腾出的显存。micro batch 从 1 到 2 会提高所有 GEMM 的运算强度,直接改善 Roofline 位置。

注意副作用:micro batch 变大后,每 DP rank 的 micro batch 数 m 从 2 变成 1(因为 global batch 固定 80,DP 固定 40)。这会让梯度累积消失、通信占比翻倍(但基数只有 0.8%,无所谓),并且让 PP 更加不可能(第 3 节)。

④ CP(上下文并行)—— 结构性改善

第 3 节讲过,但要注意整除约束:CP=2 会让 DP=20,EP=8 失效。这是个需要仔细权衡的改动,可能要配合 EP=5 或调整 global batch。

⑤ FlashAttention 3 / 更好的 attention kernel

考虑到注意力占总算力 32-73%,attention kernel 的效率是杠杆最大的单点。FA3 在 Hopper 上的改进(异步、warp 特化、FP8 路径)值得测。但受 cu128 约束(第 5 节),可用版本要先确认。

⑥ 通信重叠调优

第 4 节算过通信只占不到 3%,所以这里最多榨出 3%。除非前面几项都做完了,否则不值得投入。

一个重要的顺序原则

先量再改。 上面每一项的预期收益都依赖于"当前瓶颈在哪",而当前最大的不确定性是:

  1. 是否在 padding(决定 ① 值 1.7× 还是 0)
  2. logits 到底占多少(决定 ② 的可行性)
  3. hidden/layers 的真实值(决定注意力占比,从而决定 ⑤ 的杠杆)

这三个数量一天就能量出来,量完再排优先级。 否则很容易花两周优化通信,换来 2% 的提升。

9. 常见坑

  • MoE 用总参数算 MFU:35B 代进去得 68%,然后以为没优化空间了。
  • 混用 MFU 和 HFU 口径:差 33%(full recompute 下),足以让两次对比得出反向结论。
  • 拿自己的 MFU 和别人的比:模型形状、序列长度、卡型、口径全都要对齐才有意义。
  • 忽略 6ND 在长序列下的失效:49K 下注意力占大头,6ND 不再是好近似。
  • 只看 MFU 不看 padding:MFU 5.85% 和"其中 41% 是 pad"是两个独立的问题,后者更容易修。
  • 先优化通信:直觉上"分布式训练瓶颈是通信",但实测这里只占 3%。永远先量。
  • 忘了 recompute 是自愿的:如果显存腾出来了,第一个该考虑的就是关掉它。

思考题

1. 用本节的方法算一遍"如果 packing 做成了"的完整账:

  • (a) 单步耗时从 120 秒变成多少?(假设完美 packing)
  • (b) 集群吞吐 tok/s 变成多少?
  • © MFU(6ND 口径)变成多少?
  • (d) 那 2 epoch 的训练时间会从"一夜"变成多久?

2. 现在同时做两件事:packing(1.7×)+ selective recompute(1.21×)。

  • (a) 两者的收益是相乘还是相加?为什么?
  • (b) 合起来单步耗时变成多少?
  • © 到这一步之后,你觉得下一个瓶颈会是什么?(提示:回到第 1 节的 Roofline 和第 4 节的通信占比——通信占比会怎么变?)

下一节讲稳定性工程:失败模式怎么分类、监控该测什么、checkpoint 策略怎么权衡 8 分钟的写入成本,以及那次"被 kill 的进程显存没释放、新任务启动即 OOM"的脏启动竞态该怎么根治。