FlashAttention-4:面向 Blackwell 的算法-流水线协同设计

6632 字
33 分钟
FlashAttention-4:面向 Blackwell 的算法-流水线协同设计

AI 生成内容声明

背景:注意力计算的演进困境#

FlashAttention 系列是过去四年里对 Transformer 推理和训练影响最大的系统优化之一。回顾其演进脉络,可以清晰地看到一个模式:每一代 FlashAttention 都是为了解决上一代在新硬件上暴露出的瓶颈。

FlashAttention-1(2022,Dao et al.)的核心贡献是提出了 IO-aware 的注意力算法:利用 online softmax 将注意力计算分解为分块操作,避免将完整的 N×NN \times N 注意力矩阵写入 HBM。在 A100 上,这带来了 2-4 倍的加速,并使得标准注意力首次能在长序列上以合理速度运行。

FlashAttention-2(2023)在同样的 Ampere 架构上进一步优化了并行策略:将序列长度维度的并行从 batch/head 维度移到外层循环,减少非矩阵乘操作,并将 warp 调度从 4-way 分块改为 8-way 分块以提升 occupancy。这些改动在 A100 上额外带来了约 2 倍的加速。

FlashAttention-3(2024,Shah et al.)是面向 Hopper 架构的全面重写。Hopper 引入了 WGMMA(warp group matrix multiply-accumulate)指令和 TMA(Tensor Memory Accelerator),使得异步数据搬运成为可能。FA3 利用这些特性实现了 warp group 级别的生产者-消费者流水线:一个 warp group 通过 TMA 异步加载下一块数据,另一个 warp group 同时执行矩阵乘。但 FA3 的流水线仍然受限于 Hopper 的硬件约束——每个 SM 最多 8-12 个 warp,流水线深度有限,且在 softmax 指数计算上存在串行瓶颈。

到了 2025 年,NVIDIA 发布了 Blackwell 架构(B200/GB200)。Blackwell 带来了超大规模的张量核心吞吐提升——BF16 下的理论峰值从 H100 的 1 PFLOPS 跃升至 2.25 PFLOPS,但其他关键硬件资源几乎没有增长。这种非对称缩放使得上一代的设计策略完全失效,催生了 FlashAttention-4(FA4)。

FA4 于 2026 年 3 月发布,由 Tri Dao、Ted Zadouri 等人联合 Princeton、Together AI、Meta、NVIDIA 和 Colfax Research 共同开发,发表于 MLSys 2026(Oral)。它不是对 FA3 的增量改进,而是从算法到流水线的联合重设计,使得注意力计算在 Blackwell 上首次达到矩阵乘法级别的速度。

核心问题:非对称硬件缩放#

理解 FA4 的设计动机,必须先理解 Blackwell 架构的硬件特性及其对注意力计算的深远影响。

Blackwell 微架构的关键变化#

Blackwell B200 相对于 Hopper H100 的主要变化可以概括为:张量核心吞吐翻倍有余,其余几乎原地踏步

硬件资源H100 (Hopper)B200 (Blackwell)增幅
BF16 Tensor Core 吞吐~1.0 PFLOPS~2.25 PFLOPS2.25×
SM 数量1321601.21×
每 SM 张量核心吞吐512 TFLOPS1024 TFLOPS2.0×
共享内存带宽128 bytes/cycle/SM~128 bytes/cycle/SM~1×
SFU/MUFU 吞吐16 ops/cycle/SM16 ops/cycle/SM
TMEM(张量内存)256 KB/SM全新
每 SM 最大 warp 数8-12161.33-2.0×

这个表格揭示了核心矛盾:张量核心的吞吐翻了 2 倍以上,而执行 softmax 指数运算的 SFU(Special Function Unit)/MUFU(Multi-Function Unit)吞吐完全没有增加。共享内存带宽也几乎不变。

这意味着什么?在 H100 上,注意力计算的主要瓶颈是矩阵乘——张量核心一直在满负荷工作,softmax 和内存搬运可以隐藏在矩阵乘的延迟之后。而在 B200 上,矩阵乘不再是瓶颈,softmax 的指数计算共享内存带宽分别成为了前向和反向传播的关键路径。

前向传播的瓶颈转移#

考虑一个典型的前向注意力计算(以 head dim = 128,序列长度 = 8K 为例):

每个注意力 tile 需要完成的操作为:

  • 两个矩阵乘:S=QKTS = QK^TO=PVO = PV
  • 一次 softmax:对 SS 的每一行计算指数、求和、归一化

在 H100 上,矩阵乘大约需要 512 个时钟周期,softmax 的指数运算大约需要 512 个周期(因为 SFU 一次只能处理 16 个元素,而一个 tile 有 128 行需要逐行计算指数)。但由于矩阵乘和 softmax 不能完全并行(softmax 必须等 SS 算完才能开始),且 H100 的张量核心速度刚好能让 softmax 几乎不成为瓶颈。

但在 B200 上,矩阵乘只需约 256 个周期(张量核心快了一倍),而 softmax 的指数运算仍然需要约 512 个周期(SFU 没变快)。结果是指数计算需要的时间是矩阵乘的 2 倍,成为前向传播的绝对瓶颈。

反向传播的瓶颈转移#

反向传播的计算图更加复杂——它涉及 5 次矩阵乘(计算 dQdQdKdKdVdV 需要 dS=dOVTdS = dO \cdot V^TdP=dOVdP = dO \cdot V 的转置等操作)和一次对 SS 的指数运算(用于重算 softmax 梯度)。在 H100 上,这 5 次矩阵乘大约需要 2560 个周期,远大于共享内存的读写时间。

但在 B200 上,5 次矩阵乘只需约 1280 个周期(张量核心快了一倍),而中间结果的共享内存读写仍然需要约 3328 个周期(共享内存带宽没变)。结果是共享内存带宽成为了反向传播的主导瓶颈——与普遍认为”注意力计算受限于矩阵乘”的直觉完全相反。

FA4 的应对策略#

面对这个局面,FA4 的设计哲学可以概括为:既然矩阵乘不再是瓶颈,就把之前隐藏在其他计算后面的非矩阵乘操作重新设计,使其不再拖累整体吞吐。

具体来说:

  • 前向传播:不再靠硬件 SFU 算指数,改用 FMA(Fused Multiply-Add)单元的软件近似,将有效指数吞吐提升数倍。
  • 反向传播:利用 Blackwell 新引入的 TMEM(Tensor Memory)存放中间结果,避免经过共享内存,绕开共享内存带宽瓶颈。
  • 整体流水线:利用 Blackwell 每 SM 支持 16 个 warp 的能力,构建更深的多级异步流水线,最大化各硬件单元的并发利用率。

前向传播:Warp 特化流水线与混合指数计算#

FA4 前向传播的设计包含三个核心创新:warp 特化的多级异步流水线、基于 FMA 的混合指数计算、以及条件 softmax 重缩放。

Warp 特化架构#

FA4 的前向 kernel 使用了 16 个 warp(每 warp 32 线程,共 512 线程),划分为 6 种功能角色:

Warp ID角色职责
0-3Softmax 阶段 0计算 QKTQK^T 和 softmax,生成概率矩阵 PP 的第一部分
4-7Softmax 阶段 1计算 QKTQK^T 和 softmax,生成概率矩阵 PP 的第二部分
8-11校正 warp执行条件重缩放,修正 OO 累加器
12MMA warp执行 P×VP \times V 矩阵乘法
13收尾 warp通过 TMA 将最终 OO 写回全局内存
14加载 warp通过 TMA 从全局内存异步加载 KKVV 分块
15空闲 warp预留,按 kernel 配置动态复用

这种精细的角色划分在之前的 GPU 架构上是不可行的。H100 每 SM 只有 8-12 个 warp,无法支撑 6 种角色的流水线深度。Blackwell 将每 SM 的 warp 容量提升到 16 个,为这种细粒度特化提供了硬件基础。

每个 warp 角色的设计都围绕一个原则:让不同硬件单元同时处于忙碌状态。当加载 warp 正通过 TMA 从 HBM 搬运下一批 KKVV 数据时,softmax warp 正利用张量核心计算 QKTQK^T,同时 MMA warp 正将上一个 tile 的 PPVV 相乘。三种操作——内存搬运、矩阵乘、指数计算——在不同的硬件单元上并发执行。

FlashAttention-4 前向传播的 warp 特化流水线(来源:FA4 论文图 1,arXiv:2603.05451)
FlashAttention-4 前向传播的 warp 特化流水线(来源:FA4 论文图 1,arXiv:2603.05451)

流水线同步机制#

多级流水线的正确性依赖于精确的同步。FA4 使用 Blackwell 的命名屏障(named barriers)来协调各 warp 之间的数据依赖:

加载阶段 ──→ Softmax 阶段 ──→ MMA 阶段 ──→ 校正阶段 ──→ 收尾阶段
(TMA load) (QK^T + softmax) (P×V) (rescale O) (TMA store)
mbar_load_KV → mbar_P_full → mbar_O_full → mbar_O_epi → TMA store
mbar_S_full_P_full_O_rescaled

每一步数据传输都由显式的屏障保护。当加载 warp 完成数据搬运后,它通过 mbar_load_KV 通知 softmax warp 可以开始计算。Softmax warp 完成概率矩阵 PP 的计算后,通过 mbar_P_full 告知 MMA warp 可以开始 P×VP \times V。MMA warp 完成后通过 mbar_O_full 触发校正 warp。校正 warp 完成重缩放后通过 mbar_O_epi 让收尾 warp 执行 TMA store。

关键细节:PP 矩阵被分为两个阶段写入 TMEM(阶段 0 占用偏移 64,阶段 1 占用偏移 192),这使得 MMA warp 可以在 Softmax 阶段 0 完成后立即开始消费 PP 矩阵的前半部分,无需等待 Softmax 阶段 1 完成。这是减少串行等待时间的关键设计。

混合指数计算:绕过 SFU 瓶颈#

FA4 最具突破性的创新之一是混合指数计算。传统的注意力实现依赖 GPU 的硬件指数单元(MUFU.EX2)来计算 softmax 中的 exp\exp 操作。但在 Blackwell 上,MUFU.EX2 的吞吐(16 ops/cycle/SM)仅为张量核心吞吐(8192 ops/cycle/SM)的约 1/500,成为前向传播的绝对瓶颈。

FA4 的解决方案是:用 FMA 单元上的三次多项式近似来分担指数计算压力。具体来说,将一个 tile 中 128 个元素的指数计算按可调比例分配给硬件 MUFU 和软件 FMA 两个路径,两条路径并行执行。

软件路径的核心算法如下:

第一步:Cody-Waite 范围归约

将指数函数重新表述为整数部分和分数部分的乘积:

2x=2x2xx2^x = 2^{\lfloor x \rfloor} \cdot 2^{x - \lfloor x \rfloor}

其中 x\lfloor x \rfloorxx 的整数部分(向负无穷取整),xx[0,1)x - \lfloor x \rfloor \in [0, 1) 是分数部分。整数部分的 2x2^{\lfloor x \rfloor} 可以通过 IEEE 754 浮点数的位操作高效实现——直接将 x\lfloor x \rfloor 加上 bias 后移入指数位即可。这完全不需要任何函数计算。

第二步:三次多项式逼近

对于分数部分 f=xx[0,1)f = x - \lfloor x \rfloor \in [0, 1),使用三次多项式逼近 2f2^f

2fp0+p1f+p2f2+p3f32^{f} \approx p_0 + p_1 f + p_2 f^2 + p_3 f^3

其中系数 p0=1.0p_0 = 1.0p10.6951p_1 \approx 0.6951p20.2276p_2 \approx 0.2276p30.0771p_3 \approx 0.0771,由数值软件包 Sollya 在 [0,1)[0, 1) 区间上最小化相对逼近误差得到。

多项式使用 Horner 方法 求值以减少乘法次数:

p0+f(p1+f(p2+fp3))p_0 + f \cdot (p_1 + f \cdot (p_2 + f \cdot p_3))

只需要 3 次 FMA 指令即可完成。FMA 指令的吞吐远高于 MUFU.EX2,且 FMA 单元在前向传播的 softmax 阶段本来就有空闲(因为矩阵乘已经由张量核心异步执行了),所以这个软件路径几乎不增加额外延迟。

第三步:组合结果

将整数部分和分数多项式的结果组合:将 x\lfloor x \rfloor 移入浮点数的指数位,乘以多项式逼近得到的 2f2^f 尾数值,得到最终的 2x2^x

精度论证:三次多项式的 FP32 相对误差约为 10510^{-5} 量级,高于硬件 MUFU 的精度。但注意力计算通常使用 BF16 精度,其有效尾数只有 7 位(约 10210^{-2} 量级的分辨率)。一旦结果被舍入到 BF16,软件逼近和硬件指数的舍入行为几乎无法区分。换句话说,在 BF16 精度下,三次多项式逼近等价于硬件指数

算法还需要处理边界情况。x<127x < -127 时,2x2^x 在 FP32 中下溢为零,因此将所有输入 clamp 到 x127x \ge -127xx 过大时直接返回 ++\infty(这在 softmax 中通常不会发生,因为 QKTQK^T 的值域有限)。

两个 softmax warp 组被显式同步,以避免同时使用 MUFU.EX2 造成争用——当一组使用硬件 MUFU 时,另一组使用 FMA 路径,交替进行。

条件 Softmax 重缩放#

Online softmax 算法在遍历 KVKV 序列块时维护两个运行统计量:

  • mjm_j:当前已看到的最大 logit 值
  • j\ell_j:当前已累加的指数和

当遇到更大的 mjm_j 时,需要将之前累加的部分输出 Oj1O_{j-1} 按比例 exp(mj1mj)\exp(m_{j-1} - m_j) 重缩放。在标准的 online softmax 实现中(包括 FA3),每次 mjm_j 变化都要重缩放。

FA4 的关键洞察:不需要每次 mjm_j 变化都重缩放。如果 mjmj1m_j - m_{j-1} 很小,跳过重缩放带来的数值误差可以忽略。FA4 引入了一个阈值 τ\tau

Oj={exp(mj1mj)Oj1+exp(Sjmj)Vj,如果 mjmj1>τOj1+exp(Sjmj1)Vj,否则O_j = \begin{cases} \exp(m_{j-1} - m_j) \cdot O_{j-1} + \exp(S_j - m_j) \cdot V_j, & \text{如果 } m_j - m_{j-1} > \tau \\[6pt] O_{j-1} + \exp(S_j - m_{j-1}) \cdot V_j, & \text{否则} \end{cases}

mjmj1τm_j - m_{j-1} \le \tau 时,直接跳过重缩放,用旧的 mj1m_{j-1} 代替 mjm_j 进行局部的指数计算。在所有 KVKV 块遍历完成后,最终做一次全局归一化:Ofinal=O/finalO_{\text{final}} = O / \ell_{\text{final}},保证数学上的正确性。

这个简单但高效的技巧消除了约 90% 的重缩放操作。在 FA3 中,重缩放位于关键路径上——校正 warp 必须等 MMA warp 完成上一轮、且 softmax warp 完成新 tile 的 softmax 之后才能执行,造成流水线停顿。FA4 的重缩放操作大幅减少后,校正 warp 很少成为瓶颈,流水线得以保持满载。

反向传播:TMEM 与 2-CTA MMA 模式#

反向传播的计算量远大于前向——需要计算 dQdQdKdKdVdV 三个梯度,涉及 5 次矩阵乘法和一次 softmax 梯度的重算。如前所述,在 Blackwell 上,反向传播的瓶颈不在矩阵乘法(张量核心足够快),而在共享内存带宽——中间结果(SSPPdSdS 等)在张量核心和共享内存之间反复搬运,占用了大量带宽。

FA4 在反向传播中引入了两项创新来缓解这个瓶颈。

FlashAttention-4 反向传播计算图与 1-CTA MMA 软件流水线(来源:FA4 论文图 2,arXiv:2603.05451)
FlashAttention-4 反向传播计算图与 1-CTA MMA 软件流水线(来源:FA4 论文图 2,arXiv:2603.05451)

TMEM:直接连接张量核心的本地存储#

Blackwell 架构引入了一种新的片上存储——TMEM(Tensor Memory),每 SM 256 KB。TMEM 的特殊之处在于它直接连接到张量核心,是张量核心的专用工作内存,不像共享内存需要通过交叉开关(crossbar)路由。

FA4 在反向传播中利用 TMEM 存放中间结果,从而避免了这些结果经过共享内存。具体来说,S=QKTS = QK^T 的计算结果直接写入 TMEM,随后被 softmax 梯度重算和 dSdS 的计算直接消费——整个过程不经过共享内存。这大幅减少了共享内存的读写压力,使得共享内存带宽不再成为瓶颈。

在前向传播中,TMEM 也被用于存放 SSPP 矩阵,使 softmax warp 和 MMA warp 之间的数据交换在 TMEM 内部完成:

TMEM 缓冲区TMEM 偏移生产者消费者
SS tile 00Softmax warp 0-3MMA warp
SS tile 1128Softmax warp 4-7MMA warp
PP tile 064MMA warp(覆盖 SSMMA warp
PP tile 1192MMA warp(覆盖 SSMMA warp

要注意 TMEM 只在数据中心级 Blackwell GPU(B200、B300)上可用。消费级 Blackwell GPU(RTX 5090、RTX PRO 6000)没有 TMEM,因此 FA4 在这些 GPU 上不可用。

2-CTA MMA 模式#

2-CTA MMA 是 Blackwell 张量核心的一项新功能:两个 CTA(Cooperative Thread Array,即线程块)协作完成一次矩阵乘法。两个 CTA 共享各自的 TMEM,操作数在两者之间分配。

FA4 在反向传播的 dQdQ 计算步骤中利用了 2-CTA MMA 模式。传统的做法是每个 CTA 独立加载完整的 KK(用于计算 dQ=dSKdQ = dS \cdot K 时所需的操作数),导致每个 CTA 都要从共享内存读取一遍 KK,造成重复的共享内存访问。在 2-CTA MMA 模式下,两个配对的 CTA 各自只加载 KK 的一半,然后共享操作数,将共享内存访问量减半

此外,传统方法中每个 CTA 需要对自己的部分 dQdQ 结果执行原子加法(atomic add)将其累加到全局内存中的 dQdQ 梯度张量。2-CTA MMA 模式将两个 CTA 的输出在 TMEM 中合并后再执行一次原子加法,将原子操作数减半

反向传播 dQ 步骤的 2-CTA MMA 模式:CTA 对通过 DSMEM 交换半个 dS tile(来源:FA4 论文图 3,arXiv:2603.05451)
反向传播 dQ 步骤的 2-CTA MMA 模式:CTA 对通过 DSMEM 交换半个 dS tile(来源:FA4 论文图 3,arXiv:2603.05451)

这个优化虽然看起来简单,但对反向传播的整体吞吐有显著影响——共享内存带宽是反向传播的瓶颈,任何减少共享内存流量的措施都能直接转化为加速。论文在 B200 上测得的反向 TFLOPS 印证了这一点:借助 TMEM 与 2-CTA MMA 这两项针对性优化,FA4 的反向 kernel 在长序列下同样达到了远超 Triton、领先 cuDNN 的吞吐水平。

FlashAttention-4 反向 TFLOPS(B200,BF16,head dim 128,因果注意力):TMEM 与 2-CTA MMA 共同缓解共享内存带宽瓶颈,反向吞吐同样显著领先(来源:FA4 论文图 6 右,arXiv:2603.05451)
FlashAttention-4 反向 TFLOPS(B200,BF16,head dim 128,因果注意力):TMEM 与 2-CTA MMA 共同缓解共享内存带宽瓶颈,反向吞吐同样显著领先(来源:FA4 论文图 6 右,arXiv:2603.05451)

确定性模式#

FA4 的反向传播还支持确定性执行模式。在默认的非确定性模式下,由于浮点加法的非结合性,原子加法的执行顺序会影响最终 dQdQ 的值,导致每次运行产生微小差异。确定性模式通过固定归约顺序来保证每次运行得到完全一致的结果,代价是约 10-15% 的吞吐损失。

在诸如强化学习(RL)等对可复现性要求严格的场景中,确定性模式非常关键——训练过程中任何非确定性都可能导致策略梯度估计的偏差累积。

论文对确定性反向传播做了细致的消融:在因果注意力下,SPT(短路径优先)、LPT 及其反向 mblock 顺序等调度策略的相对收益被逐一拆解;在非因果注意力下,则对比了带 batch/head swizzle 与朴素调度。这些消融说明确定性模式带来的 10-15% 吞吐损失并非固定代价,调度策略的选择可以在保证可复现性的同时尽量缩小这一差距。

确定性反向传播的因果注意力消融:SPT、反向 mblock 顺序的 LPT、标准 LPT 与无 swizzle 的朴素调度对比(来源:FA4 论文图 7,arXiv:2603.05451)
确定性反向传播的因果注意力消融:SPT、反向 mblock 顺序的 LPT、标准 LPT 与无 swizzle 的朴素调度对比(来源:FA4 论文图 7,arXiv:2603.05451)

确定性反向传播的非因果注意力消融:batch/head swizzle 相对朴素调度的收益(来源:FA4 论文图 8,arXiv:2603.05451)
确定性反向传播的非因果注意力消融:batch/head swizzle 相对朴素调度的收益(来源:FA4 论文图 8,arXiv:2603.05451)

高阶优化:LPT 调度与变长序列#

除了核心的算法和流水线设计外,FA4 还引入了多项针对实际部署场景的优化。

LPT 调度#

在实际推理中,不同请求的 prompt 长度可能差异很大——从几个 token 的短对话到数十万 token 的长文档。此外,因果掩码(causal mask)使得不同 query 位置对应的 KVKV 长度不同(位置 ii 只需关注前 ii 个 token)。这两种情况都会导致各 query 的计算量不均衡,而 GPU 的并行执行模型(SIMT)要求同一 warp 内的线程同时完成才能进入下一阶段,负载不均意味着快的线程在等慢的线程。

FA4 采用 LPT(Longest Processing Time First)调度策略处理变长序列和因果掩码:将计算量最大的 query 块优先分配给处理单元,减少尾部等待时间。在多 query 并行处理时,LPT 调度遵循”先处理最重的任务”原则,使得整体完成时间逼近理论下界。

论文报告的实验数据显示,LPT 调度为多头注意力(MHA)带来 4-8% 的额外 FLOPs 增益,对于 MQA/GQA(Multi-Query / Grouped-Query Attention)则可达到 14%。这是因为 MQA/GQA 中 KVKV 头数远少于 QQ 头数,QKTQK^TKK 维度上的负载不均更加显著。

FlexAttention 集成#

PyTorch 的 FlexAttention API 允许用户在 Python 层面通过 score_mod 函数定义自定义的注意力变体(ALiBi、滑动窗口、文档掩码、soft-capping 等),而不需要编写 CUDA 代码。FA4 实现了作为 FlexAttention 后端的集成,在 Hopper 和 Blackwell GPU 上可用。

当用户定义一个 score_mod 函数后,FlexAttention 将其 JIT 编译为 FA4 kernel 的一个变体。相比 Triton 后端,FA4 后端的加速在 1.2× 到 3.2× 之间,具体取决于注意力模式的复杂度。这意味着研究者可以快速实验新的注意力变体,同时享受到手写 CUDA kernel 级别的性能——在此之前,这两者是不可兼得的。

性能评估#

Kernel 级别性能#

在 B200 GPU 上的 BF16 精度下,FA4 的前向 kernel 达到了最高 1613 TFLOPs/s,约为 B200 理论峰值(约 2250 TFLOPS BF16)的 71%。对于注意力计算而言,这是一个里程碑式的数字——以往的注意力 kernel(包括 FA3 在 H100 上)很难突破 50-60% 的硬件利用率,因为非矩阵乘操作(softmax、内存搬运)长期占据大量时间。

对比基线:

  • 1.3× 于 cuDNN 9.13(NVIDIA 的官方注意力库,同样针对 Blackwell 优化)
  • 2.7× 于 Triton 实现的注意力 kernel
  • FA4 在序列长度 4K 以上时优势最为显著,因为此时 softmax 和共享内存带宽的瓶颈效应最强

FlashAttention-4 前向 TFLOPS(B200,BF16,head dim 128,因果注意力):FA4 在 4K+ 序列长度上对 cuDNN 9.13.0 达 1.1–1.3×、对 Triton 达 2.1–2.7× 加速,长序列下逼近 B200 理论峰值(来源:FA4 论文图 4 右,arXiv:2603.05451)
FlashAttention-4 前向 TFLOPS(B200,BF16,head dim 128,因果注意力):FA4 在 4K+ 序列长度上对 cuDNN 9.13.0 达 1.1–1.3×、对 Triton 达 2.1–2.7× 加速,长序列下逼近 B200 理论峰值(来源:FA4 论文图 4 右,arXiv:2603.05451)

FA4 的性能优势并不局限于标准的 head dim = 128 配置。在 DeepSeek V3 架构常用的 head dim (192, 128) 配置下,FA4 同样保持了对 cuDNN 的稳定领先——这说明它对非标准头维度的兼容并非”只对默认配置调优”,而是算法层面的通用设计。

DeepSeek V3 架构 head dim (192, 128) 下的前向 TFLOPS 对比,FA4 在非标准头维度配置下依然领先 cuDNN(来源:FA4 论文图 5,arXiv:2603.05451)
DeepSeek V3 架构 head dim (192, 128) 下的前向 TFLOPS 对比,FA4 在非标准头维度配置下依然领先 cuDNN(来源:FA4 论文图 5,arXiv:2603.05451)

端到端推理性能#

在实际 LLM 推理中(BF16,batch size = 32,8K 上下文):

模型H100 + FA3B200 + FA4吞吐提升
Llama 3.1 8B~75,000 tok/s~140,000 tok/s~1.9×
Llama 3.1 70B~18,000 tok/s~35,000 tok/s~1.9×
Llama 3.1 405B (8 GPU)~3,200 tok/s~6,400 tok/s~2.0×

在更长序列(32K+)上,FA4 的优势进一步扩大,因为注意力计算在总延迟中的占比随序列长度平方增长。对于需要处理超长上下文的 RAG、代码库问答和 Agent 工作负载,FA4 在 B200 上的 2× 级别加速直接转化为用户体验的显著改善。

成本效率#

虽然 B200 的按需价格高于 H100,但由于单卡吞吐约为 H100 的两倍(尤其是在长序列下),部署相同吞吐量所需的 GPU 数量减少,从而降低了跨 GPU 通信开销。在 spot 实例定价下,B200 + FA4 的每百万 token 成本与 H100 + FA3 几乎持平,但延迟显著更低。

局限与未解决的问题#

FA4 的局限主要来自两个方面:硬件依赖和算法边界。

硬件依赖:FA4 仅支持 Blackwell 和 Hopper 架构。在消费级 Blackwell GPU(RTX 5090 等)上,由于缺少 TMEM,FA4 无法运行。在 Ampere/Ada 架构(A100、RTX 4090 等)上,用户仍需使用 FA2。在 Hopper(H100、H200)上可以使用 FA4,但由于 Hopper 没有 TMEM 且每 SM 的 warp 容量较小,其性能优势相对 FA3 不大。

序列长度的隐性限制:虽然 FA4 支持任意长度的序列,但混合指数计算中三次多项式的精度论证依赖于 BF16 的舍入误差掩盖逼近误差。如果用户使用 FP32 或 FP16 精度(而非 BF16),软件指数路径可能产生可察觉的精度损失。论文推荐在非 BF16 场景下降低软件路径的元素比例,以精度换取速度。

对自定义注意力的限制:FlexAttention 集成虽然灵活,但 score_mod 函数在编译为 FA4 kernel 时受到一定限制——某些复杂的逐元素操作可能无法被高效映射到 warp 特化的流水线中,导致回退到 Triton 后端。

TMEM 容量:256 KB 的 TMEM 对于标准注意力配置(head dim = 128)足够,但对于某些使用更大 head dim(256 或 512)的模型(如某些视觉 Transformer 或 State Space Model 的变体),TMEM 可能无法同时容纳完整的 SSPP tile,需要额外的分块策略,降低流水线效率。

非 GPU 平台:FA4 的设计深度绑定 NVIDIA GPU 的硬件特性(TMEM、TMA、warp 特化、张量核心指令)。在 AMD、Intel GPU 或专用 AI 加速器上,FA4 的思想可以参考,但实现需要从零开始。

小结#

FlashAttention-4 的意义远超一个 kernel 库的版本更新。它揭示了一个更深刻的趋势:随着 GPU 张量核心吞吐持续高速增长,而其他硬件资源(SFU、共享内存、寄存器文件)的增长相对滞后,注意力计算——乃至更广泛意义上的非矩阵乘操作——将越来越多地从”可隐藏的次要开销”变为”关键路径瓶颈”。

FA4 给出的答案不是等待硬件厂商解决这个问题,而是通过算法-系统的协同设计主动适应硬件特性:用 FMA 单元的闲置能力弥补 SFU 的不足,用 TMEM 的直连通路绕开共享内存的带宽瓶颈,用更深的 warp 特化流水线榨取有限的并发资源。

从工程角度看,FA4 证明了一个已经被多次验证但常被忽视的原则:在新硬件上,“把旧代码移植过去”远不如”为硬件特性重新设计算法”有效。在 Blackwell 上,FA4 相对于 cuDNN 的 1.3 倍加速看起来不多,但考虑到 cuDNN 本身也是 NVIDIA 工程师针对自家硬件高度优化的结果,这 30% 的差距恰恰反映了算法-硬件协同设计的价值。

参考资料#

  1. FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling — FA4 论文,MLSys 2026 Oral
  2. FlashAttention-4: Algorithm and Kernel Co-Design for Blackwell GPUs — 上海纽约大学的技术概述
  3. FlashAttention-4 gives the NVIDIA Blackwell platform its most optimized attention kernel yet — Lambda Labs 的详尽评测
  4. FlashAttention-4 on GPU Cloud: Blackwell Inference Guide (2026) — Spheron 的推理部署指南
  5. Reverse Engineering FlashAttention-4: Why It Matters for AI Engineering Teams — Propel Code 对 FA4 的反向工程分析
  6. Making FlashAttention-4 faster for inference — Modal 的推理优化实践
  7. Dao-AILab/flash-attention GitHub Repository — 官方开源实现
  8. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision — FA3 论文,理解 FA4 设计动机的前置阅读
  9. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — FA1 原始论文
  10. NVIDIA Blackwell Architecture Technical Overview — Blackwell 架构官方文档

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

FlashAttention-4:面向 Blackwell 的算法-流水线协同设计
https://pinghaoyang.com.cn/aigc/posts/flashattention-4/
作者
平昊阳
发布于
2026-08-10
许可协议
CC BY-NC-SA 4.0

评论区

Profile Image of the Author
平昊阳
乘长风,破巨浪, 展鸿图于未央!
--
总访问量
--
访客数
公告
欢迎来到我的个人博客!欢迎关注交流吖!
更多相关公告,见
社交-留言」。
音乐
封面

音乐

暂未播放

0:000:00
暂无歌词
站点统计
文章
66
分类
16
标签
93
总字数
477,284
运行时长
0
最后活动
0 天前

文章目录