FlashAttention 动效 06

FA4:Blackwell 的「非对称扩展」把瓶颈挪了位置

从 Hopper 到 Blackwell,Tensor Core 吞吐翻倍,但共享内存带宽和指数单元没跟上。于是「算矩阵乘」不再是主要开销,softmax 里的指数反而成了主角。

一个 tile 的时间去向

矩阵乘
非 matmul

相对 Hopper 的吞吐变化(示意)

FA4 的三组对策

  1. 更大的 tile + 全异步 MMA:用 Blackwell 的 128×128 MMA(面积是 Hopper 64×128 的两倍)和直接写 TMEM 的异步 Tensor Core,把搬运、矩阵乘、softmax 尽量叠在一起。
  2. 把指数运算「软」下来:用 FMA 单元上的多项式近似模拟 exp,绕过吞吐最紧的多功能单元;再加条件式 softmax 重缩放,跳过不必要的缩放。
  3. 减少共享内存流量:反向把中间结果放进 256 KB 的 TMEM;用 2-CTA MMA 模式,每个 CTA 只搬一半 B 操作数,减少共享内存流量和反向的原子加。

B200 上 BF16 最高 1613 TFLOPs/s(71%),比 cuDNN 9.13 快约 1.3×、比 Triton 快约 2.7×。整个 FA4 用嵌入 Python 的 CuTe-DSL 写成,编译时间比传统 C++ 模板快 20–30×。图中的时间占比与吞吐倍数为示意,量级取自论文「非 matmul 已超过 MMA 计算 25–60%」的表述。