FlashAttention-4 是怎么把 71% 利用率做出来的?Tri Dao 的算法-内核协同设计

FlashAttention-4 是怎么把 71% 利用率做出来的?Tri Dao 的算法-内核协同设计

Tri Dao 是那种「用一个内核改动改变了整个领域」的研究者。他的 FlashAttention 系列从 2022 年第一版开始,就把注意力计算的效率问题从「算法论文」变成了「系统论文」——核心洞察是:注意力慢不是因为算得多,而是因为搬得多。此后的 FlashAttention-2、3 每一代都在把「算力利用率」往上推,而 2026 年的 FlashAttention-4 是这套思路在 Blackwell 架构上的一次集中兑现。

这一版值得单独精读,原因在于它的方法论变了。前几代的主要工作还是「把已有的算法映射到硬件上,减少内存搬运」;而 FlashAttention-4 做的事情,是**反过来根据硬件的特性去重新设计算法本身**。这种双向的重构,才是它能把峰值利用率推到 71% 的根本原因。

速览卡:这位大佬和他说了什么

Tri Dao大佬Tri Dao
身份FlashAttention 系列作者,普林斯顿大学助理教授,Together AI 首席科学家
出处论文《FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling》(arXiv 2603.05451,MLSys 2026)
时间2026 年
核心观点注意力内核的性能瓶颈已从「搬运」转向「计算单元利用率」,必须以算法与内核协同设计的方式重构:条件式 softmax 重缩放、软件模拟指数、TMEM 与 2-CTA MMA 协同,才能在 Blackwell 非对称扩展下压满算力

Tri Dao 团队说了什么

论文的核心声明集中在摘要里,先说结论部分。这一代 FlashAttention 的性能数字是这样的:

“FlashAttention-4 achieves 1.3x speedup over cuDNN 9.13 and 2.7x speedup over Triton on NVIDIA Blackwell GPUs for both forward and backward passes.”

(译文)在 NVIDIA Blackwell GPU 上,FlashAttention-4 的前向与反向传播相较 cuDNN 9.13 取得 1.3 倍加速,相较 Triton 取得 2.7 倍加速。

绝对值同样给出了:

“…reaching up to 1613 TFLOPs/s, corresponding to 71% of the peak utilization.”

(译文)……最高达到 1613 TFLOPs/s,对应峰值利用率的 71%。

这个 71% 是整篇论文里最值得琢磨的数字。在 GPU 内核优化领域,传统的注意力内核利用率通常在 30%–50% 区间,能上到 60% 已经算是很优秀的工作。71% 意味着绝大多数时间张量核都在满负荷计算,而不是在等数据、等同步、或者空转。要理解这个数字的分量,得知道它面临的阻力是什么。

论文的标题点出了那个阻力:「非对称硬件扩展」(asymmetric hardware scaling)。意思是 Blackwell 上,张量核(Tensor Core)的计算吞吐涨得很快,但内存带宽、共享内存、以及片上存储的容量并没有同比例增长。这导致一个典型的失配:计算单元能吃得进更多数据,但喂进来的速度跟不上。这不是搬多少的问题,而是单位时间内能喂多少的问题——瓶颈从「带宽」变成了「延迟」和「指令吞吐」。

紧接着的一个关键机制是编译器层面的:

“…using CuTe-DSL, which compiles 20-30x faster than CUTLASS C++ templates.”

(译文)……使用 CuTe-DSL,其编译速度比 CUTLASS 的 C++ 模板快 20–30 倍。

这条看起来是工程细节,实际上是方法论的一部分。内核开发的效率直接决定了能迭代多少轮,而 CUTLASS C++ 模板的编译时间在大型内核上动辄以分钟甚至十分钟计,严重限制了「改一行、看效果」的循环。把编译时间压掉 20–30 倍,等于把可尝试的设计空间放大了同样倍数。这一点在论文里不显眼,但它解释了为什么这版能做到这么多细粒度的硬件特化。

背景:注意力内核的瓶颈为什么会转移

要理解这一版的突破,需要先理清注意力优化这十年的主线。

**第一阶段(2022 年之前):内存墙。** 标准注意力实现会把 N×N 的注意力矩阵写回 HBM(高带宽显存)再读回来做 softmax,这个中间矩阵的大小随序列长度平方增长。序列一长,HBM 的读写就成了绝对瓶颈。FlashAttention 第一代的解法是分块(tiling)+ 在线 softmax:把 Q、K、V 分成适合片上存储的块,在 SRAM 里完成计算,只把最终结果写回 HBM。这一步把内存访问从平方级降到线性级,收益巨大。

**第二阶段(2023–2024):并行与工作划分。** FlashAttention-2 的重点转向「怎么让 GPU 的各个层级都忙起来」——重排计算顺序以减少非矩阵乘法运算,改进工作划分以提升线程块利用率,同时更好地利用 warp 级并行。这一代把利用率从第一代的水平往上推了一个台阶。

**第三阶段(2024–2025):新硬件特性的吸收。** FlashAttention-3 针对 Hopper 架构做了特化,用上了 TMA(Tensor Memory Accelerator,异步数据搬运单元)、warp specialization(不同 warp 承担不同角色)、以及 FP8 低精度计算。这一代的主题是「把新硬件给的异步能力用起来」,让搬运和计算重叠。

到 FlashAttention-4,情况变了。Blackwell 的算力增长主要落在张量核上,而 memory bandwidth、L2、SMEM(共享内存)的扩展相对滞后。这个非对称性导致:即使你把内存搬运优化到极致,计算单元依然会因为「等指令、等数据、等指数运算」而空闲。所以这一代必须解决的问题是:**在数据供给受限的前提下,如何让张量核尽量不空转**。

这就是论文标题里「co-design」的含义。它不再是「优化算法的实现」,而是「为了配合硬件特性,改变算法的形态」。

核心观点展开:四个协同设计

论文给出的技术手段有四个,每一个都是对某个具体硬件瓶颈的针对性回应。

**一、条件式 softmax 重缩放(conditional softmax rescaling)。** 传统在线 softmax 每处理一个新块,就要重新缩放之前的累积结果,以保证数值稳定。这个「每次都算一遍缩放」在算法上很自然,但在硬件上代价很高——它引入了额外的乘法和依赖链,而且在大多数情况下,缩放因子的变化其实很小,甚至不需要缩放。FlashAttention-4 的做法是加一个条件判断:只在缩放因子确实需要调整时才执行重缩放。这听起来是「作弊」,但它是有理论依据的——当处理到后期的块时,累积最大值的变化趋于收敛,无条件的重缩放大部分是浪费。这个改动的价值在于它同时减了「计算量」和「关键路径长度」,而后者在非对称扩展下更致命。

**二、软件模拟指数函数(software-emulated exponent)。** softmax 里最贵的非矩阵运算是指数函数。GPU 上通常有专用单元处理它,但当张量核在满负荷跑的时候,这些通用单元就成了瓶颈——它们成了整条流水线上最快被占满的资源。论文的做法是用软件方式模拟指数运算,把这个负载从专用单元转移到张量核或通用 ALU 上,从而让专用单元的拥塞缓解。这是一个典型的「重新分配负载」的思路:不是让某一步更快,而是让所有单元的工作量更均衡,谁都不成为短板。

**三、TMEM 的运用。** TMEM(Tensor Memory)是 Blackwell 引入的、专供张量核使用的片上存储。它的价值在于可以把累加器(accumulator)放在离张量核更近的地方,减少寄存器压力和搬运。FlashAttention-4 把中间结果的管理围绕 TMEM 重新组织,减少了寄存器溢出和额外的数据移动。这一步直接对上了「数据供给不足」的核心矛盾:让数据在离计算单元最近的地方待着。

**四、2-CTA MMA。** Blackwell 支持两个 CTA(线程块协作单元)协同完成一次 MMA(矩阵乘加)操作,相当于把两个 SM 的算力联合起来做一次大矩阵运算。这在算法层面意味着可以处理更大的块、更长的 K 维度,从而提升算术强度(每次搬运能摊到更多计算)。把它和前面的 tiling 结合,等于在「块的大小」这个关键参数上获得了新的自由度——而块越大,单位数据摊到的计算越多,越能对抗带宽受限。

把这四个手段放在一起看,会发现它们指向同一个目标:**在带宽和片上容量受限的条件下,最大化张量核的有效工作时间**。条件式重缩放减少了不必要的计算依赖,软件模拟指数平衡了单元负载,TMEM 缓解了数据供给压力,2-CTA MMA 提升了算术强度。它们不是四个独立的优化,而是同一套「协同设计」哲学的四个切面。

行业内的讨论与不同声音

FlashAttention-4 公布后,讨论主要集中在三个方向。

**第一个方向是「71% 这个数字的可比性」。** 有研究者指出,峰值利用率的定义在不同论文里并不统一——有的按理论峰值算,有的按实际可达到的峰值算;有的测试序列长度偏短,有的偏长;有的把前向和反向分开报。这些细节会显著影响数字的可比性。这个质疑是合理的,但它的结论不是「71% 注水」,而是「看这类数字时要注意测量条件」。论文在同一平台上对比 cuDNN 9.13 和 Triton 的倍数(1.3x / 2.7x),在测量口径上相对更可信,因为它是同条件对比。

**第二个方向是「收益的分配」。** 注意力内核的优化收益,最终要通过训练和推理框架传导到用户手上。而框架的集成往往滞后于论文发布——内核再好,如果框架不接入,实际收益就是零。这是这系列工作长期面临的问题:FlashAttention-2、3 都经历过「论文发布 → 框架集成 → 用户实际感受到」的长周期。所以对使用者来说,更该关注的是「什么时候能用上」,而不是「理论峰值多少」。

**第三个方向,也是最有价值的,是「协同设计这个方法本身的边界」。** 有人提出一个尖锐的问题:如果每一代硬件都需要重写一次内核,那么这种优化模式的可维护性如何?Hopper 的特化(TMA、warp specialization)、Blackwell 的特化(TMEM、2-CTA MMA)都在代码里留下架构专属的分支。当架构继续演进,这些分支会越积越多。Tri Dao 团队的应对方式是把内核描述迁移到 CuTe-DSL——用 DSL 而不是 C++ 模板来表达,让「针对不同硬件的特化」变成可组合的描述,而不是散落的 `#ifdef`。这也是为什么论文里专门提了编译速度这个看似无关的指标:它决定了这套方法能否持续。

**还有一个跨领域的共鸣值得提。** 这种「算法跟着硬件走」的思路,在近年的推理优化里越来越普遍。比如稀疏注意力、线性注意力的各种变体,本质上都是在用「算法形态的调整」换「硬件的适配性」——牺牲一点数学上的通用性,换取在真实芯片上的高利用率。但这里有一个需要警惕的张力:算法越是适配特定硬件,它的迁移成本就越高。一个在 Blackwell 上跑到 71% 的内核,换到别的架构上可能需要从头再来。这使得「内核优化」变成了一种需要持续投入的能力,而不是一次性的成果。

横向扩展:把这项工作放进更大的知识框架

要完整理解 FlashAttention-4,需要理解三个背景概念。

**概念一:算术强度与 Roofline 模型。** 算术强度(arithmetic intensity)指每搬运一字节数据能完成多少次浮点运算。Roofline 模型把芯片性能画成两条线的下包络:低于某个强度时受带宽限制,高于某个强度时受算力限制。FlashAttention 整个系列的演进,本质上是在不断提高算术强度——分块是为了让数据在 SRAM 里被复用更多次,从而「摊薄」搬运成本。FlashAttention-4 面临的问题是:在 Blackwell 上,算力线抬得比带宽线快,意味着要维持「算力受限」这个状态,需要更高的算术强度,这就是 2-CTA MMA(更大块)和 TMEM(更好复用)的意义。

**概念二:softmax 的数值稳定性与在线算法。** softmax 的标准实现需要先求全局最大值以避免 exp 溢出,但这要求看到全部输入。在线 softmax(online softmax)的核心技巧是维护一个「当前最大值」和「已累积的和」,遇到新块时按需修正之前的累积。这个「修正」就是 rescaling 的由来。它的代价是引入了跨块的依赖,破坏了纯粹的并行性。FlashAttention-4 的「条件式」重缩放,是试图减少这种依赖的触发频率——本质上是在数值稳定性的代价和并行效率之间重新找平衡点。

**概念三:硬件异步能力的演进。** 从 Ampere 到 Hopper 再到 Blackwell,一个清晰的主线是「异步化」:数据搬运(TMA)、计算(张量核)、以及各种单元之间越来越独立,可以重叠执行。这对编程模型提出了新要求——你不能再用「先搬完再算」的顺序思路写内核,而必须显式地编排「谁在什么时候做什么」。warp specialization 和 2-CTA MMA 都是这种编排手段。理解这条主线,就能理解为什么现代的 GPU 内核论文里,越来越多篇幅在讲「流水线」而不是「公式」。

总结与启示

FlashAttention-4 最值得记住的,不是 1.3x 或 1613 TFLOPs/s 这些数字,而是它示范了一种正在成为主流的工作方式:**算法设计和内核实现必须同时做,不能分成两拨人前后接力**。

对做推理和训练的工程师,第一条启示是**瓶颈会转移,优化重点要跟着换**。当序列还不长、显存还不是瓶颈时,注意力内核的优化重点是「减少搬运」;当硬件进入非对称扩展阶段,重点就变成「让计算单元不空转」。沿用过时的优化直觉(比如一味追求更小的显存占用),在新硬件上可能完全打不到点子上。

第二条启示是**关注非矩阵运算的成本**。softmax 里的指数、除法、以及对输入分布的处理,在纸面上是「次要项」,但在张量核被打满的情况下,它们反而成了限制因素。FlashAttention-4 花大力气去模拟指数函数、去减少 rescaling,正是因为这部分的边际收益变高了。做自定义内核时,先测清楚「哪一类运算占了关键路径」,比盲目优化矩阵乘法更有价值。

第三条启示是关于**工具链的乘数效应**。CuTe-DSL 编译快 20–30 倍这件事,如果只看它本身,似乎只是「开发体验更好」。但内核优化是一个大量试错的工程——可尝试的设计空间扩大 20 倍,最终得到的结果质量差距是数量级的。投资于工具链,往往比投资于某一次具体优化更有回报。这一点对所有做性能工程的团队都成立。

最后回到 Tri Dao 这系列工作的一贯特征:他不是在追求「最优雅的算法」,而是在追求「在真实硬件上跑得最快的算法」。这两者往往不一致,而选择后者意味着要接受算法层面的一些「不干净」——比如条件式的重缩放、软件模拟的指数函数,它们在数学上都显得笨拙,但它们让芯片跑满了。在算力成为稀缺资源的当下,这种「以硬件为约束的算法设计」可能比纯理论上的优雅更有价值。

参考来源

主要菜单