AI++
工程实践·硬核·11 分钟阅读·AI++ 编辑部

FlashAttention-3:把 H100 的异步算力榨干的注意力新基准

FlashAttention 系列靠「分块 + 重计算」避免实例化完整 N×N 注意力矩阵,把注意力从显存密集变成计算密集。 FA3 针对 H100 的异步数据搬运和 FP8 tensor core 重新设计 kernel,让 softmax 与 GEMM 重叠执行。 相比 FA2,FP16 下约 1.5~2 倍提速,FP8 下再翻倍,长上下文场景尤其受益。

一句话结论

FlashAttention-3 把注意力 kernel 重新设计成「异步数据搬运 + softmax/GEMM 流水重叠 + FP8 低精度」三件事的合奏,在 H100 上把 FA2 已经很快的注意力又提速 1.5~2 倍(FP8 下更高),让长上下文训练和推理的单位成本进一步下降。

背景问题

先回顾 FlashAttention 解决的问题。标准注意力的瓶颈是显存:要把 $Q K^\top$ 这个 $L \times L$ 矩阵完整写到 HBM(显存)再做 softmax,$O(L^2)$ 显存、$O(L^2)$ 显存读写。读写带宽是瓶颈,算力反而闲置。

FlashAttention-1/2 的解法是 tiling(分块)+ recomputation(重计算)

  • 把 Q、K、V 切成 block,每次只载入一小块到 SRAM(片上共享内存)。
  • 在 SRAM 里算 attention 的中间结果(用 online softmax 增量更新),不写回 HBM。
  • 反向传播时不存中间矩阵,重算前向(省显存换算力,但算力本来就没用满)。

效果:把 attention 从「显存密集」变成「计算密集」,FA2 在 A100 上接近理论峰值。所以 FA2 已经成了长上下文训练的标配。

到了 H100,FA2 跑不满。H100 相比 A100 有几个新特性没有被 FA2 利用:

  1. 异步数据搬运:H100 的 TMA(Tensor Memory Accelerator)能异步在 HBM 和 SRAM 间搬数据,不需要 warp 主动参与,可以和计算重叠。
  2. FP8 tensor core:H100 的 FP8 算力是 FP16 的 2 倍,FA2 只用 FP16/BF16。
  3. warp specialization:H100 支持把「搬数据」和「算」分给不同 warp,硬件级流水线化。

FA2 的 kernel 是「同步 + 单 warp group」设计,无法发挥这些特性——在 H100 上 FA2 的 MFU(model FPU utilization)远低于理论峰值。FlashAttention-3 就是为补这个缺口而生。

核心思路

FA3 的设计围绕一个核心问题:怎么让 H100 的「搬数据」和「算」真正重叠起来?

注意力计算可以粗略拆成三步循环(对每个 query block):

  1. 载入 K、V block 到 SRAM;
  2. $S = Q K^\top$(GEMM);
  3. softmax 和 $O = P \cdot V$(涉及归一化、规约)。

FA2 是串行的:载入 → 算 → 载入下一块 → 算……算的时候搬运单元闲着,搬的时候算单元闲着。

FA3 的三招:

1. 异步 TMA 载入 + warp specialization

FA3 用 TMA 把 K、V 的载入变成异步操作,并把 warp 分成两组:

  • producer warp:只负责用 TMA 发起载入,把 K、V block 搬到 SRAM,搬完发信号。
  • consumer warp:等信号,拿到 K、V 后做 GEMM 和 softmax。

这样载入下一块和算当前块真正并行——搬运单元和 tensor core 同时干活。这是 FA3 提速的主要来源。

2. softmax 与 GEMM 的流水重叠

更精细的并行:在 consumer warp 内部,把「算当前 block 的 softmax」和「算下一个 block 的 $Q K^\top$ GEMM」重叠。这要求 softmax 的归约操作不阻塞后续 GEMM,FA3 通过「两阶段 softmax」和寄存器乒乓缓冲实现。结果是 softmax 这个原本的「串行瓶颈」被藏到了 GEMM 后面。

3. FP8 低精度

H100 的 FP8 tensor core 吞吐是 FP16 的 2 倍。FA3 支持 FP8 attention,但 FP8 有精度陷阱:直接把 Q、K、V 量化到 FP8 会让 attention 分数失真(softmax 对小数值差异敏感)。

FA3 的处理:非均匀量化——对 Q、K 用 per-tensor scaling,对中间的 $S = QK^\top$ 用 per-row scaling,让关键数值落在 FP8 可表示范围内。配合「两次 softmax」(先粗算一次缩放因子,再精算)控制误差。最终 FP8 attention 在精度损失可控的前提下,吞吐比 FP16 再翻倍。

关键技术点

1. warp specialization 是 FA3 区别于 FA2 的灵魂

FA2 是「all-in-one」warp——一个 warp 既搬又算。FA3 把职责拆开,让硬件的「搬运单元 + 计算单元」并行。这是 Hopper 架构以后 GPU 编程范式的转变:从「让一个 warp 干所有事」变成「让不同 warp 干不同事,靠硬件同步原语协调」。这种范式后续被广泛用在 H100/B100 的高性能 kernel 里。

2. online softmax 仍然是基础

FA1 引入的 online softmax(增量计算 softmax 的 max 和 sum,不需要看完整行)是 FA3 仍依赖的基础——它让分块计算成为可能。FA3 在此之上加了流水,没改变「分块 + 重计算」的本质。

3. FP8 的精度管理是个独立工程

FP8 attention 不是「把数据类型改一下」那么简单。Q、K、V 的 scale、$S$ 矩阵的 scale、softmax 后的 $P$ 的 scale,每一层都要单独设计。FA3 论文给了详细的 scale 策略和误差分析。用 FP8 attention 训练大模型时,per-layer 的 scale 调优是落地门槛。

4. 长上下文是最大受益场景

FA3 的相对收益随序列长度增加而增加——长序列下 attention 占比大、tiling 重复多,流水重叠的节省更显著。在 32k~128k 上下文训练里,FA3 相对 FA2 的加速对训练成本有实质性影响。

5. 推理侧也有红利,但形态不同

推理(decode)阶段 attention 是 memory-bound(每步只算一个 query 和所有 key),FA3 的 GEMM/softmax 重叠红利变小。但 FA3 的 prefetch 设计仍能改善 decode 的访存模式,长上下文 decode 仍有可观收益。decode 的根本加速仍依赖 PagedAttention、GQA、MLA 等结构性手段。

效果与局限

效果:

  • FA3 在 H100 上:FP16 下相对 FA2 约 1.5~2 倍加速;FP8 下相对 FA2 FP16 约 3~4 倍(含精度换算)。
  • 接近 H100 FP16 tensor core 理论峰值的 75%+,FP8 下更高。
  • 数值精度:FP16 path 与 FA2 数值等价;FP8 path 在标准 benchmark 上误差可控(论文给了详细对比)。
  • 已被主流训练框架集成(PyTorch SDPA、TransformerEngine 等)。

局限:

  • 硬件绑定。FA3 的 warp specialization、TMA、FP8 都是 Hopper 及以后架构的特性,A100 及更早硬件用不上,会退化成 FA2 等价。
  • FP8 精度需要谨慎验证。在数值敏感任务(高精度数学、长程数值累积)上,FP8 attention 的误差可能放大,需要针对性测试。
  • kernel 复杂度大增。FA3 的实现远比 FA2 复杂,定制和调试成本高,新架构适配周期长。
  • 对 decode 加速有限。FA3 主要利好 prefill(训练、长 prompt 处理),decode 阶段的根本瓶颈仍是访存,需要其他手段。

谁该关心

  • 做大模型训练 infra 的人:FA3 是 H100 训练的默认 attention,理解它的 FP8 策略和 warp specialization 是榨干硬件的前提。
  • 长上下文训练团队:序列越长 FA3 红利越大,128k 训练场景几乎必用。
  • H100 集群拥有者:FA3 直接关系硬件 ROI,没接入等于浪费算力。
  • 研究 GPU kernel 的人:FA3 是 Hopper 编程范式的范本,warp specialization 的设计对写其他高性能 kernel 有普遍启发。

一个判断:FlashAttention 已经从「优化技巧」变成「基础设施」——FA1 证明了 tiling 可行,FA2 证明了能跑满 A100,FA3 证明了能在 Hopper 上把异步和低精度用到位。每一代都和硬件深度绑定,理解它要同时懂算法和 GPU 架构。对绝大多数使用者,FA3 是「装上就用」的开关;对 infra 团队,它是必须读懂的标杆 kernel。

注意力机制CUDAH100FP8