第 6 课:FlashAttention-4(上)非对称缩放与前向 exp 瓶颈

为什么 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 在前向用的三招。

1. 非对称缩放:MMA 翻倍,exp 和共享内存没动

从 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 博客]

2. 前向 roofline:exp 竟追平 MMA

论文对一个 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 本身的有效吞吐

3. 新流水线:双 Q tile pingpong + TMEM 累加器

FA-4 前向沿用 FA-3 的 pingpong 思路(一个 CTA 算两个 Q tile,交替),但 Blackwell 改变了底层映射 [论文 §3.1.2]

流水线安排好后,剩下的瓶颈仍是 exp 本身。于是 FA-4 直接动 exp 的实现。

4. 软仿真 exp:用 FMA 多项式和 MUFU 并行

硬件的 MUFU.EX2 吞吐只有 16 ops/clock/SM,远低于 MMA 的 8192。FA-4 的办法是用 FMA 单元"软件仿真"一部分 exp,和硬件 MUFU 并行跑,把有效吞吐顶上去 [论文 §3.1.3] [Together AI 博客]

核心是经典的 Cody-Waite 范围归约,把指数拆成整数和小数部分:

$$2^x = 2^{\lfloor x \rfloor} \cdot 2^{x - \lfloor x \rfloor}$$

整个多项式用 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 精度消费的,所以够用

5. 条件 rescale:跳过不必要的 online softmax rescale

标准 FlashAttention 的 online softmax 每处理一个新 block j,都要做一次输出 rescale [论文 §3.1.4]

$$O_j = e^{m_{j-1}-m_j}\, O_{j-1} + e^{S_j - m_j}\, V_j$$

其中 \(m_j\) 是 running max。这一步是个向量乘法,属于非 matmul 开销。FA-4 两个观察 [论文 §3.1.4]

  1. 只有当 \(m_j > m_{j-1}\)(遇到更大的值)时才真需要 rescale。
  2. 可以容忍一点"松弛":只有当跳变 \(m_j - m_{j-1} > \tau\) 才 rescale,\(\tau = \log_2 256 = 8\)(对应 rescale 因子 256×)。
$$O_j = \begin{cases} e^{m_{j-1}-m_j}\, O_{j-1} + e^{S_j-m_j}\, V_j & \text{if } m_j - m_{j-1} > \tau \\ O_{j-1} + e^{S_j - m_{j-1}}\, V_j & \text{otherwise} \end{cases}$$

跳过时不更新 \(m\),照常用 \(m_{j-1}\)。关键在最后:用真实的最终 \(m_{\text{final}}\) 和 \(\ell_{\text{final}}\) 统一归一,输出 \(O_{\text{final}} / \ell_{\text{final}}\)。因为中间偷懒省掉的 rescale 都被这次最终归一补回来了,结果不变 [论文 §3.1.4]

避免 warp divergence

"要不要 rescale"的判断以 warp 粒度做:只要 warp 里任一线程需要 rescale,整个 warp 就 rescale。这样不会因为线程间判断不同造成分支发散 [论文 §3.1.4]

6. 性能小结

配置FA-3(H100)FA-4(B200)
BF16 前向峰值740–840 TFLOPs/s(75–85%)~1613 TFLOPs/s(71%)
对比 cuDNN 9.131.1–1.3× 更快
对比 Triton2.1–2.7× 更快

注意:FA-4 的部分技术已并入 cuDNN 9.13/9.14(作者与 cuDNN 团队合作),所以"1.3× vs cuDNN 9.13"是在 cuDNN 已经吸收了一部分 FA-4 技巧之后的比较 [Together AI 博客]

7. 快速测验

1. FA-3 在 H100 已达 75–85%,FA-4 认为到了 Blackwell 前向的主要瓶颈转移到了哪里?
2. 软仿真 exp 为什么能提高有效吞吐?
3. 条件 rescale 在什么情况下跳过输出 rescale?

8. 延伸阅读

本课推荐阅读 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 实现)。

理论 vs 实测

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。