跳到正文
Together AI Blog·· 2026-03-05精选AI 评分68

FlashAttention-4:面向非对称硬件扩展的算法与内核流水线协同设计

FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling

AI 导读

FlashAttention-4 发布,针对 Blackwell 的非对称硬件扩展做算法与内核协同设计,在 B200 BF16 上达到最高 1605 TFLOPs/s(71% 利用率),比 cuDNN 9.13 快最多 1.3 倍、比 Triton 快 2.7 倍。

推荐理由

原文给出 B200 上注意力核的瓶颈拆解与算法内核协同设计,可看硬件非对称扩展下的优化思路。

正文 · AI 翻译

像 Blackwell GPU 这样的现代加速器延续了非对称硬件扩展的趋势,其中张量核心吞吐量的增长速度远超其他资源,例如共享内存带宽、用于指数等超越运算的特殊功能单元(SFU),以及通用整数和浮点 ALU。例如,从 Hopper H100 到 Blackwell B200,BF16 张量核心吞吐量从 1 PFLOPs 提升到 2.25 PFLOPs,而 SFU 数量和共享内存带宽均保持不变。

这种扩展的不对称性对为 Blackwell 架构优化注意力等复杂内核具有深远影响。注意力的核心由两个 GEMM 组成 $(S=Q \cdot K^T$ 和 $O=P \cdot V)$,中间夹着 softmax;在实践中,它还涉及大量辅助和簿记工作:数据移动、同步、布局变换、逐元素操作、调度、掩码等。

对注意力的一种朴素看法可能是,GEMM 的速度完全决定了内核性能,人们可以有效地忽略这些其他注意力组件,至少在一阶近似下如此。然而,对 B200 进行“供给与速度”分析实际上显示了相反的结果:主要性能瓶颈不在于张量核心执行 MMA 的速度,而在于 (a) 前向计算期间用于 softmax 指数的 SFU 单元,以及 (b) 反向计算期间的共享内存流量。

在这篇博客文章中,我们介绍 FlashAttention-4,这是一种算法与内核协同设计,可最大化矩阵乘法与这些其他资源瓶颈之间的重叠。在 B200 上使用 BF16 时,它达到高达 1605 TFLOPs/s(71% 利用率),比 cuDNN 版本 9.13 快高达 1.3×,比 Triton 快 2.7×。

我们的主要算法和内核协同设计思路如下:

  1. 新的流水线以实现最大重叠:新的前向和反向软件流水线,利用 Blackwell 完全异步 MMA 和更大的瓦片尺寸,重叠张量核心、softmax 指数和内存操作。
  2. 前向(FWD)传递:通过 FMA 单元上的多项式近似实现指数函数的软件仿真,以缓解指数瓶颈,外加条件在线 softmax 重缩放。‍
  3. 反向(BWD)传递:将中间结果存储在张量内存中以减轻共享内存流量,结合 Blackwell 新的 2-CTA MMA 模式进一步减少共享内存流量,并将原子归约减半,以及对确定性执行模式的额外支持以实现可重现训练。‍
  4. 调度:新的瓦片调度器以缓解因果掩码和可变序列长度带来的负载不平衡。

Blackwell 上的新硬件特性

  1. 张量内存(TMEM):在 B200 上,148 个 SM 中的每一个都有 256 KB 的 TMEM,这是一种连接到张量核心的片上暂存器,用于 warp 同步中间存储。
  2. 完全异步的第 5 代张量核心:tcgen05.mma 是异步的,并在 TMEM 中累加。对于 BF16 和 FP16,最大的单个 CTA UMMA 块是 128×256×16,大约比最大的 Hopper WGMMA 原子大 2 倍。UMMA 由单个线程启动,缓解了寄存器压力,使得更大的块和更深的流水线变得可行,而没有 Hopper warpgroup MMA 的溢出痛点。这也使 warp 专门化更加可行,一些 warp 移动块,而其他 warp 发出 MMA,以将矩阵乘法累加与 softmax 和内存流量重叠。tcgen05.mma 也可以从 TMEM 获取操作数 A。
  3. 2-CTA MMA:Blackwell 可以在同一集群中的一对 CTA 上执行一个 UMMA,跨越两个对等 CTA 的 TMEM。领导者 CTA 中的一个线程启动 MMA,但两个 CTA 在它执行期间必须保持活跃。这通过在对之间拆分 M 和 N,将 MMA 块维度扩展到 256×256×16,减少了冗余流量并降低了每个 CTA 的占用。CTA 组大小(1 或 2)在内核中的 TMEM 和张量核心操作中必须保持不变。

馈送与速度

对于 M=N=D=128

B200 上的馈送(每个 SM):

  1. 张量核心(BF16): $\frac{8192 \text{ ops}}{cycle}$
  2. 指数单元:$\frac{16 \text{ ops}}{cycle}$
  3. 共享内存流量:$\frac{128 \text{ bytes}}{cycle}$

速度(每个块的时钟周期):

  1. Forward (2 MMAs + MN exp)
    1. 张量核心:$1024$
    2. Exp:$1024$
    3. SMEM:$768$
  2. Backward (5 MMAs + MN exp): 1-CTA
    1. 张量核心:$2560$
    2. Exp:$1024$
    3. SMEM:$3328$

要点:前向受计算和指数瓶颈限制,后向受共享内存带宽瓶颈限制。因此,我们在前向传递中将 softmax 与 MMA 重叠,并在后向传递中减少共享内存流量。

前向传递:带有条件重缩放的新 softmax 流水线

前向传递有两个矩阵乘法,QK^T 和 PV。在 Blackwell 上,张量核心变得快得多,但指数单元(MUFU.EX2)没有。因此 softmax 不再是“两个矩阵乘法之间的东西”,它是一个必须仔细流水线化的瓶颈。

简而言之,前向传递:

  • 乒乓调度 $2x$ Q 和 $2x$ O 块每个 CTA:最大化 MMA 和 Softmax 之间的重叠
  • 2x softmax warpgroups: per tile softmax with synchronization to not overlap when computing exponential
    • $2^x$ 的软件模拟:将 exp 计算分布在硬件的 MUFU 和软件模拟的 FMA 上
    • 分阶段将 P 存储在 TMEM 中:缓解寄存器压力
  • Correction warpgroup: designated "correction" warpgroup to perform rescaling to remove from critical path
    • 在线 softmax(条件)重缩放:减少重缩放频率以最小化非矩阵乘法操作

流水线:乒乓 Q 块加上一个专用的校正阶段

FlashAttention-4 每个 CTA 计算两个查询块——$Q^H$ 和 $Q^L$——每个覆盖 128 个查询令牌,并以乒乓调度交替它们。

Blackwell 改变了 softmax 映射。S = QK^T 的累加器块是 128×128,位于张量内存中;然而,在读取到寄存器时,根据硬件的分区,我们有每行一个线程。我们使用两个 128 线程的 warpgroup,每个 Q 块一个,每个 softmax warpgroup 执行以下操作序列:

  1. 每个线程从张量内存中加载 S 的一行 128 个元素到寄存器
  2. 归约 rowmax 和 rowsum
  3. 使用一个可调参数,决定 128 个元素中哪部分使用硬件的 MUFU 与软件模拟的 $e^x$
  4. 计算 P = softmax(S) 并转换为 BF16 精度
  5. 分阶段将 P 存回张量内存以缓解寄存器压力(而不是同时保存 128 个 S 元素和 64 个(BF16)P 元素)
  6. 一旦存储了 P 的 $\frac{3}{4}$ 块,就触发相应的 PV 矩阵乘法

关键细节在于 exp 是瓶颈部分。我们显式同步两个 softmax warpgroup,使它们不同时计算 exp,从而减少 MUFU 争用。

为了将重缩放保持在关键路径之外,内核将其分配给专用的 warpgroup。校正 warpgroup 计算:

  1. 仅当最大跳跃较大时才重缩放:
  2. $O_j =\begin{cases}\exp(m_{j-1}-m_j)\,O_{j-1} + \exp(S_j-m_j)\,V_j, & \text{if } m_j - m_{j-1} > \tau,\\O_{j-1} + \exp(S_j-m_{j-1})\,V_j, & \text{otherwise.}\end{cases}$
  3. 在迭代结束时应用最终归一化 $O_{final} = \frac{O}{l_{final}}$
  4. 可选地计算并存储 LSE

最后我们仍然使用真实的最终统计量进行归一化,因此跳过小的重缩放步骤可以保持最终输出不变,同时从关键路径中删除许多向量计算。我们在 warp 粒度上做出决定以避免分歧。

更快的指数:将 $2^x$ 分布到 MUFU.EX2 和 FMA(软件仿真)

Softmax 需要大量指数运算,而 MUFU 吞吐量远低于张量核心吞吐量。FlashAttention-4 通过将 exp2 的软件仿真与硬件 MUFU.EX2 路径并行运行,利用原本未充分利用的 FMA 单元,提高了有效 exp 吞吐量。

范围缩减(Cody-Waite): 我们使用经典的 Cody-Waite 范围缩减技术,将指数计算分解为整数部分和小数部分:$2^x = 2^{n} \cdot 2^{f}$。在 IEEE 754 float32 中,乘以 $2^{n}$ 只是指数更新。

$2^{x_{frac}}$ 的多项式近似(Horner 方法): 为了 ****近似 $2^{f}$,我们将其重写为 Horner 形式以便高效求值。

$$2^{x_{\mathrm{frac}}} \approx p_0 + p_1 x_{\mathrm{frac}} + p_2 x_{\mathrm{frac}}^{2} + p_3 x_{\mathrm{frac}}^{3}$$

系数 p0 = 1.0、p1 ≈ 0.6951、p2 ≈ 0.2276、p3 ≈ 0.0771 使用 Sollya 软件包选择,以最小化 $[0, 1)$ 上的相对近似误差。

指数位移和加法: 最后一步是将整数部分 n 和小数近似 2^{f} 组合成 2^{x}  \approx 2^{n}\cdot 2^{f}。由于 2^f \in[1,2) 的 float32 指数为 127,乘以 2^{n} 只是将整数 n 移位到指数字段,然后加上 2^{f} 的尾数位。

反向传播:共享内存流量占主导地位

优化 FlashAttention 反向传播就像把一条过大的地毯塞进房间:压平一个角,另一个角又弹出来。反向传播的张量核心工作量约为前向传播的 2.5 倍,链接五个 MMA 操作以重新计算 S,并为 dQ、dK、dP 和 dV 运行 QK 和 PV 梯度 MMA,再加上 P 和 dS 的逐元素工作。在 Blackwell 上,FLOPs 不是反向传播的限制因素;共享内存带宽才是

流水线:将 MMA 与 softmax 重叠

Hopper 时代的 FlashAttention-3 将 MMA 累加器保留在寄存器中,因此寄存器压力常常迫使调度更加串行。在 Blackwell 上,累加器位于 TMEM 中,这使得在 CUDA 核心处理 P 和 dS 的逐元素工作时,可以实际保持多个 MMA 在飞行中。由于在我们的 roofline 中,指数吞吐量与两个 MMA 相当,因此隐藏它是值得的。

关键重叠很简单:当我们为 tile j 计算 softmax 时,我们已经为 tile j−1 发出了 dK 和 dQ MMA

为了减少共享内存流量,反向传播在前向传播的转置 tile 中重新计算 S 和 P,因此中间结果已经是 $S^T$ 和 $P^T$。然后我们可以将 $P^T$(以及稍后的 $dS^T$)直接存储在 TMEM 中,采用 dV 和 dK MMA 分别消耗的操作数 A 布局。

TMEM 无法同时容纳五个完整的累加器和中间结果,因此 FA4 跨阶段重用 TMEM 列:S 和 P 共享一组列,dP、dS 和 dQ 共享另一组。

2-CTA 反向传播:减少共享内存流量和全局原子加

共享内存流量。 即使有了改进的流水线,并且十个 GEMM 操作数中的两个保留在张量内存中,反向传播仍然受限于共享内存带宽。我们通过 Blackwell 2-CTA MMA 模式来缓解这个问题,该模式将输出累加器划分到 CTA 对中。当 M=256 且 N=K=128 时,两个 CTA 作为一个 tile 协作:每个 CTA 暂存一半的操作数 B,并仅保留自己的累加器切片。这大约将操作数 B 的共享内存流量减半。

归约轴冲突。 我们在五个反向 GEMM 中使用 M=256 和 N=K=128 的 MMA tile 来减少 B 流量,但 dQ MMA 的性质引入了不匹配。在 FlashAttention 反向传播中,每个 CTA 拥有固定的 KV tile(外层循环在 N 个 CTA 上并行化),并在内层循环中迭代 M 个 tile。dQ 更新在外层循环中对 KV 序列进行归约。2-CTA MMA 分割的是输出 tile,而不是归约,并且 dQ 归约维度是 N,它已经在 CTA 对之间分割。每个 CTA 仍然需要对其拥有的行进行完整归约。

解决方案:DSMEM 交换。 我们通过使用集群内的分布式共享内存在两个 CTA 之间交换一半的 dS 来解决这个问题。这重新打包了 dS,使其沿非归约轴进行分区:每个 CTA 拥有 M/2 行,同时持有完整的 2N 归约。每个 CTA 的 dQ MMA 变为 (M/2, 2N)(2N, d),在张量内存中累加一个 (M/2, d) tile。在 2-CTA 模式下,S、dP、dV 和 dK MMA 保持 M=256,而 dQ 使用 M=128,归约加倍为 2N=256。然后我们重新排序流水线以隐藏 DSMEM 延迟:在计算前一个 tile 的 dQ 之前,先计算当前 tile 的 dP。由于 dQ tile 与 P 一起适合 TMEM,它可以重用用于 S 的 TMEM 区域,因此 dP 和 dQ 不再像 1-CTA 模式那样共享区域。通过这种排序,当前 tile 的逐元素 dS 与前一次迭代的 dQ MMA 重叠。

dQ 原子加。 作为附带好处,dQ 分解将全局原子归约的数量减半。原子操作是非确定性的且开销大,并且它们出现在每个内层循环迭代中。因此,在 2CTA 反向传播中,每个 CTA 只写入一半的 dQ tile,并执行与 1-CTA 对应物相比一半的全局原子归约。

确定性模式:可复现的 dQ 而不牺牲吞吐量

非确定性的来源是 dQ 的全局原子累加。FA4 提供了一种确定性模式,通过信号量风格的锁和内存栅栏来串行化全局归约,以强制固定的累加顺序。然而,确定性并不意味着“一切都停下来”。FA4 通过 CTA swizzling 减少锁竞争,并使用最短处理时间优先(SPT)排序进行因果掩码以减少停顿。在实践中,在我们的基准测试中,确定性反向传播可达到非确定性吞吐量的约 85-90%。

调度

因果掩码和可变序列长度使注意力负载不均衡,因为不同的工作块具有不同的主循环长度,因此 FA4 改进了网格线性化,并应用最长处理时间优先(LPT)调度来减少尾部。事实上,这些想法并非 Blackwell 或任何特定 GPU 架构所特有,我们也在 FA3 中使用了它们。

对于因果掩码,标准的 (mblocks, heads, batches) 网格顺序会以从最短到最长的次优方式处理块,因此 FA4 将 batch-heads 混洗为 L2 大小的区段,并按 batch-head 区段遍历网格,以逆序迭代 mblocks,然后遍历每个区段内的 batch-heads。

对于可变序列长度,由于不同批次涉及的工作量不同,从 LPT 调度启发式的角度来看,给定的批处理顺序通常是次优的。为了纠正这一点,我们可以启动一个预处理内核,按每个工作块的最大执行时间对批次进行排序,并写入一个虚拟到实际批次索引的映射,注意力内核使用该映射按排序顺序遍历批次;此外,元数据可以被缓存,因此排序不会带来性能损失。在撰写本文时,我们已经验证了这一想法并在 FA3 中实现了它,我们预计在不久的将来会更普遍地将排序和其他元数据准备纳入 F4。

语言和框架:CuTe-DSL

FA4 完全用 CuTe-DSL 实现,这是 CUTLASS 的 Python 内核 DSL。内核用 Python 编写;DSL 降低到 PTX,然后 CUDA 工具包编译为 GPU 机器码。编程模型镜像了 CuTe/CUTLASS 抽象,并带有 PTX 逃生舱,同时与 C++ 模板相比,编译时间缩短了约 20–30 倍。

注意力基准测试

我们展示了 FlashAttention-4 在 B200(BF16)上的结果,并将其与 FlashAttention-2 以及 Triton、Gluon 和 cuDNN 中的实现进行比较。对于 cuDNN,我们与 cuDNN 9.13 和最新版本 9.19.1.2 进行比较。从 9.13 和 9.14 版本开始,我们与 cuDNN 团队合作,将 FlashAttention-4 的一些技术融入 cuDNN,以便我们的工作能够惠及尽可能多的从业者。对于反向传播,FlashAttention-4 在长序列长度下始终优于其他基线。在前向传播中,FlashAttention-4 比 cuDNN 9.13 快 1.1-1.3 倍,比 Triton 快 2.1-2.7 倍。

Bar chart comparing cuDNN versions and FA4 on Forward TFLOPS across sequence lengths from 1K to 32K.

致谢

我们感谢 Together AI、Meta、xAI 和 Princeton Language and Intelligence (PLI) 提供的计算支持。我们还要进一步感谢 Nvidia 的以下团队:CuDNN、TensorRT-LLM 和 CUTLASS 团队,感谢他们不断的讨论、想法和反馈。

来源:Together AI Blog · together.ai