为什么 FA-2 在 H100 上只用满 35% 算力,以及 FA-3 如何用异步重叠把它拉到 75%
FA-1 解决了 HBM 读写量,FA-2 解决了并行度与非 matmul 开销。到了 Hopper(H100),FlashAttention-2 仍只能达到理论峰值算力的 约 35% [PyTorch 博客]。
原因不是 HBM 流量,而是 Tensor Core 在等:
exp 由多功能单元(multi-function unit)执行,吞吐远低于 Tensor Core。FA-3 要解决的不是"算得少",而是"算的单元和搬数据的单元没有同时忙起来"。Hopper 给了三根新杠杆,FA-3 的工作是把它们用满。
| 硬件特性 | 作用 | 相比 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 都是 异步 指令,这打开了新的重叠空间。
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]。
有人会问: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。
用 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 博客]。
即使在一个 warpgroup 内,也能让 softmax 的一部分和该组自己的 GEMM 并行——因为 WGMMA 是异步的,发起后线程不必空等。这把吞吐再推到约 640–660 TFLOPS [PyTorch 博客]。
同时保留两套 GEMM 累加器 + softmax 输入输出,需要更多寄存器。FA-3 用 setmaxnreg 动态调大寄存器上限来换这个重叠——用寄存器换吞吐,划算。
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 双倍算力。
| 配置 | FA-2 | FA-3 | 提升 |
|---|---|---|---|
| FP16 前向(H100) | ~35% 峰值 | 740 TFLOPS(75–85%) | 1.5–2× |
| FP8 前向 | — | ~1.2 PFLOPS | 2.6× 更低误差 |
对比 FA-2 的 ~225 TFLOPS(A100),FA-3 在 H100 上把 attention 的 Tensor Core 利用率从 35% 拉到 75–85%。瓶颈从"Tensor Core 空闲、softmax bubble、同步等待"被进一步压缩 [FA-3 论文]。
本课推荐阅读 PyTorch 官方博客:FlashAttention-3(含 pingpong/intra-warpgroup 示意图),以及原论文 FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision 的 Section 3–4。
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)\)——任何不清楚的地方随时问我。