第 4 课:FlashAttention-3 Hopper 异步执行与 FP8

为什么 FA-2 在 H100 上只用满 35% 算力,以及 FA-3 如何用异步重叠把它拉到 75%

1. 瓶颈转移:从 HBM 流量到 Tensor Core 空闲

FA-1 解决了 HBM 读写量,FA-2 解决了并行度与非 matmul 开销。到了 Hopper(H100),FlashAttention-2 仍只能达到理论峰值算力的 约 35% [PyTorch 博客]

原因不是 HBM 流量,而是 Tensor Core 在等

关键直觉

FA-3 要解决的不是"算得少",而是"算的单元和搬数据的单元没有同时忙起来"。Hopper 给了三根新杠杆,FA-3 的工作是把它们用满。

2. Hopper 的三根新杠杆

硬件特性作用相比 Ampere
WGMMA
Warpgroup MMA
异步发起矩阵乘加,可从 shared memory 直接读操作数。 老的 mma.sync 只能到 Hopper Tensor Core 峰值的约 ⅔ [arXiv:2402.13499]
TMA
Tensor Memory Accelerator
专用硬件异步搬数据(global↔shared),自动算索引、自动处理越界。 释放寄存器(不再用线程算地址),允许更大 tile。
FP8 低精度 Tensor Core,吞吐翻倍。 H100:FP16 约 989 TFLOPS → FP8 约 1978 TFLOPS [PyTorch 博客]

光是把 FA-2 改写成用 WGMMA + TMA,前向就从 ~350 TFLOPS 提升到约 540–570 TFLOPS [PyTorch 博客]。但 WGMMA 和 TMA 都是 异步 指令,这打开了新的重叠空间。

3. Warp Specialization:生产者 / 消费者分工

FA-2 里所有 warp 干一样的活:既 load 又算。FA-3 把 warp 分成两类 [PyTorch 博客]

# FA-3 warp specialization(概念伪代码)
# 一个 warpgroup = 4 个 warp

# producer warpgroup
for j in range(T_c):
    cp_async_bulk(K_j, smem_K[next])   # TMA 异步搬 K_j 到 shared
    cp_async_bulk(V_j, smem_V[next])
    arrive_barrier()                    # 通知 consumer:数据就绪

# consumer warpgroup
for j in range(T_c):
    wait_barrier()                      # 等 producer 搬完
    S_ij = wgmma(Q_i, smem_K[j])        # 异步 MMA,不阻塞 softmax 单元
    # softmax 更新可与下一轮 wgmma 重叠……

这样 搬数据 同时进行,内存延迟被计算量藏了起来。这是通用的 GEMM 技巧,FA-3 把它搬进了 attention [CUTLASS: Hopper Warp Specialization]

4. 为什么必须重叠 GEMM 和 Softmax

有人会问:FLOPs 大头不都在 matmul 里吗?只要 matmul 够快不就行了?问题在于 非 matmul 单元慢得多 [PyTorch 博客]

H100 SXM5:
FP16 matmul ≈ 989 TFLOPS
特殊函数(exp 等)≈ 3.9 TFLOPS(差 256×

对 head_dim=128,matmul FLOPs 是 exp 的 512 倍。算下来 exp 能花掉 matmul 一半的时间。FP8 更糟:matmul 快一倍,exp 速度不变。理想情况是 Tensor Core 忙 matmul 时,多功能单元同时算 exp

5. 两层重叠:Pingpong + Intra-warpgroup

5.1 Pingpong(warpgroup 之间)

用 2 个 warpgroup(记为 1、2),靠 bar.sync 屏障交替:
warpgroup 1 算当前 tile 的 MMA + 下一 tile 的 MMA 时,warpgroup 2 做上一个 tile 的 softmax,反之亦然。

时间线 →
WG1: [GEMM₁][GEMM₀']      [GEMM₂][GEMM₁']      ...
WG2:       [softmax₁][GEMM₂']   [softmax₂][GEMM₃'] ...
        ↑ softmax 藏在另一组的 GEMM 阴影里

同色代表同一轮迭代。这让 softmax 跑在"另一组 GEMM 的影子"里,FP16 前向从 ~570 提升到约 620 TFLOPS [PyTorch 博客]

5.2 Intra-warpgroup(warpgroup 内部)

即使在一个 warpgroup 内,也能让 softmax 的一部分和该组自己的 GEMM 并行——因为 WGMMA 是异步的,发起后线程不必空等。这把吞吐再推到约 640–660 TFLOPS [PyTorch 博客]

代价:寄存器压力

同时保留两套 GEMM 累加器 + softmax 输入输出,需要更多寄存器。FA-3 用 setmaxnreg 动态调大寄存器上限来换这个重叠——用寄存器换吞吐,划算。

6. FP8 + Incoherent Processing:低精度不丢精度

FP8 算力翻倍,但 LLM 激活有 outlier(少数值远大于其他),量化误差被放大 [outlier 研究]。FA-3 借用量化文献的 incoherent processing [QuIP]

实验(0.1% outlier 的正态分布):incoherent processing 把 FP8 量化误差降低 2.6×
FP8 前向接近 1.2 PFLOPS [PyTorch 博客]

配合 block quantization(按 tile 给缩放因子),FA-3 在保持数值稳定的前提下吃下了 Hopper 的 FP8 双倍算力。

7. 性能小结

配置FA-2FA-3提升
FP16 前向(H100)~35% 峰值740 TFLOPS(75–85%)1.5–2×
FP8 前向~1.2 PFLOPS2.6× 更低误差

对比 FA-2 的 ~225 TFLOPS(A100),FA-3 在 H100 上把 attention 的 Tensor Core 利用率从 35% 拉到 75–85%。瓶颈从"Tensor Core 空闲、softmax bubble、同步等待"被进一步压缩 [FA-3 论文]

8. 快速测验

1. FA-2 在 H100 上只用了约 35% 峰值,FA-3 认为主要瓶颈转移到了哪里?
2. Pingpong 调度之所以能把 softmax 藏起来,靠的是什么?
3. FA-3 用 Hadamard 变换做 incoherent processing 的目的是?

9. 延伸阅读

本课推荐阅读 PyTorch 官方博客:FlashAttention-3(含 pingpong/intra-warpgroup 示意图),以及原论文 FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision 的 Section 3–4。

理论 vs 实测

FA-3 的 Hopper 异步(WGMMA/TMA/warp-specialization)是 Sm90 专属。下一课我们会看一份真实模型的 kernel-trace 实测——在非 Hopper 平台上 FA-3 走的是另一条 CUTLASS 通用路径,收益和陷阱完全不同。第 5 课:IAW 模型 FA2 vs FA3 实测

有疑问?

pingpong 的时间线图、setmaxnreg 的寄存器权衡、Hadamard 为什么是 \(O(d \log d)\)——任何不清楚的地方随时问我。