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 利用:
- 异步数据搬运:H100 的 TMA(Tensor Memory Accelerator)能异步在 HBM 和 SRAM 间搬数据,不需要 warp 主动参与,可以和计算重叠。
- FP8 tensor core:H100 的 FP8 算力是 FP16 的 2 倍,FA2 只用 FP16/BF16。
- warp specialization:H100 支持把「搬数据」和「算」分给不同 warp,硬件级流水线化。
FA2 的 kernel 是「同步 + 单 warp group」设计,无法发挥这些特性——在 H100 上 FA2 的 MFU(model FPU utilization)远低于理论峰值。FlashAttention-3 就是为补这个缺口而生。
核心思路
FA3 的设计围绕一个核心问题:怎么让 H100 的「搬数据」和「算」真正重叠起来?
注意力计算可以粗略拆成三步循环(对每个 query block):
- 载入 K、V block 到 SRAM;
- 算 $S = Q K^\top$(GEMM);
- 算 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。
