为什么 Blackwell 的 Tensor Core 翻倍后,exp 反而成了瓶颈——以及 FA-4 如何用软仿真和条件 rescale 把它藏起来
FA-3 在 H100 把 Tensor Core 利用率拉到 75–85%(第 4 课)。到了 Blackwell(B200/GB200),Tensor Core 吞吐再次翻倍,但论文却发现:真正卡住前向的不再是矩阵乘,而是 softmax 里的 exp [FA-4 论文 §2.2]。本课讲清这条"非对称缩放"的新瓶颈,以及 FA-4 在前向用的三招。
从 Hopper H100 到 Blackwell B200,BF16 Tensor Core 吞吐从约 1 PFLOPS 涨到约 2.25 PFLOPS。换算到每个 SM 每时钟周期:8192 ops/clock/SM,正好是 Hopper(4096)的 2 倍 [论文 §2.2] [Together AI 博客]。
但另外两个关键单元没跟着翻倍:
硬件不是等比例变快的。Tensor Core 翻倍,exp 和共享内存带宽原地踏步 → 瓶颈从 MMA 转移到 exp(前向)和共享内存流量(后向)。FA-4 的全部设计都围绕"把没翻倍的单元藏起来" [Together AI 博客]。
论文对一个 tile(M=N=d=128)做了 roofline 分析,看三类资源各要多少时钟周期 [论文 §3.1.1, Table 1]:
| 资源(M=N=d=128) | 周期数 |
|---|---|
| MMA 计算(QKᵀ + PV,两次 GEMM) | 1024 |
| exp 单元(softmax 的 exp) | 1024 |
| 共享内存读 | 768 |
MMA 翻倍后,exp 和 MMA 打成平手。换句话说:softmax 不再是"两个 matmul 之间随便塞点小活",它本身就是前向的瓶颈。这正是 FA-4 前向要重点优化的地方 [论文 §3.1.1]。
对照 FA-3(第 4 课):H100 上 exp 吞吐 3.9 TFLOPS、matmul 989 TFLOPS,差 256×,靠"重叠"把 exp 藏进 GEMM 阴影。
FA-4:Blackwell 上 MMA 又快一倍,光靠重叠已经藏不住了——必须提高 exp 本身的有效吞吐。
FA-4 前向沿用 FA-3 的 pingpong 思路(一个 CTA 算两个 Q tile,交替),但 Blackwell 改变了底层映射 [论文 §3.1.2]:
tcgen05.mma 异步写进 TMEM(每 SM 256 KB 片上内存),不再吃寄存器,于是能用更大的 128×128 tile(Hopper 是 64×128)[论文 §2.2]。流水线安排好后,剩下的瓶颈仍是 exp 本身。于是 FA-4 直接动 exp 的实现。
硬件的 MUFU.EX2 吞吐只有 16 ops/clock/SM,远低于 MMA 的 8192。FA-4 的办法是用 FMA 单元"软件仿真"一部分 exp,和硬件 MUFU 并行跑,把有效吞吐顶上去 [论文 §3.1.3] [Together AI 博客]。
核心是经典的 Cody-Waite 范围归约,把指数拆成整数和小数部分:
整个多项式用 FMA 指令评估,FMA 单元和 MUFU 是不同的硬件,所以两者真并行 → exp 的有效吞吐被抬起来。
软仿真要多占寄存器(存中间值和系数),全量仿真会 spill、反而更慢。FA-4 只对每行 10–25% 的 entry 软仿真,其余走硬件 MUFU.EX2,比例按 MMA/exp 吞吐比可调 [论文 §3.1.3]。
精度够吗?论文测了 4M 个随机输入 [论文 §3.1.3, Table 2]:degree-3 多项式的 FP32 最大相对误差 \(8.8\times10^{-5}\),比硬件差约 600×。但 round 到 BF16 后,BF16 本身的量化误差(\(\sim 3.9\times10^{-3}\))远大于多项式误差——99% 的输入和硬件结果差 ≤1 ULP。attention 的 softmax 输出本来就是 BF16 精度消费的,所以够用。
标准 FlashAttention 的 online softmax 每处理一个新 block j,都要做一次输出 rescale [论文 §3.1.4]:
其中 \(m_j\) 是 running max。这一步是个向量乘法,属于非 matmul 开销。FA-4 两个观察 [论文 §3.1.4]:
跳过时不更新 \(m\),照常用 \(m_{j-1}\)。关键在最后:用真实的最终 \(m_{\text{final}}\) 和 \(\ell_{\text{final}}\) 统一归一,输出 \(O_{\text{final}} / \ell_{\text{final}}\)。因为中间偷懒省掉的 rescale 都被这次最终归一补回来了,结果不变 [论文 §3.1.4]。
"要不要 rescale"的判断以 warp 粒度做:只要 warp 里任一线程需要 rescale,整个 warp 就 rescale。这样不会因为线程间判断不同造成分支发散 [论文 §3.1.4]。
| 配置 | FA-3(H100) | FA-4(B200) |
|---|---|---|
| BF16 前向峰值 | 740–840 TFLOPs/s(75–85%) | ~1613 TFLOPs/s(71%) |
| 对比 cuDNN 9.13 | — | 1.1–1.3× 更快 |
| 对比 Triton | — | 2.1–2.7× 更快 |
注意:FA-4 的部分技术已并入 cuDNN 9.13/9.14(作者与 cuDNN 团队合作),所以"1.3× vs cuDNN 9.13"是在 cuDNN 已经吸收了一部分 FA-4 技巧之后的比较 [Together AI 博客]。
本课推荐阅读 Together AI 官方博客:FlashAttention-4(含 feeds-and-speeds 框架和前向流水线图),再读原论文 FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling 的 §2.2(非对称缩放)和 §3.1(前向)。代码在 flash_attn/cute(CuTe-DSL 实现)。
FA-4 的 TMEM、tcgen05.mma、2-CTA、DSMEM 都是 Sm100(Blackwell)专属,不可移植到你的 PPU810E(CUTLASS 走 Sm80 路径,见 第 5 课)。本课这些 B200 数字对你都是理论值。唯一架构无关的是 LPT 调度(FA-3 已有,FA-4 强化),下节课会讲。1-4 代演进速查里有完整对比。
Cody-Waite 为什么能把 exp 拆成整数+分数、τ 为什么取 8、TMEM 256KB 怎么在多个累加器之间分——任何不清楚的地方随时问我。下一课(第 7 课)讲后向:2-CTA MMA、DSMEM、确定性模式、LPT 调度和 CuTe-DSL。