跳到正文
Modal Blog· Charles Frye·· 2025-09-26精选AI 评分60

Modal 逆向解析 Flash Attention 4 内核实现

We reverse-engineered Flash Attention 4

AI 导读

Modal 团队阅读已开源的 Flash Attention 4 源码,逆向解析了这款针对 Nvidia Blackwell 架构优化的 CUDA 内核。

推荐理由

通过逐层拆解 FA4 内核的异步流水线与 warp 分工,读者可理解 Blackwell 上注意力计算的工程实现思路。

正文 · AI 翻译

这篇博文登上了 HackerNews 首页!讨论见此处。

这篇博文已在 GPU MODE Discord 上展示!观看录像请点击此处。

一个月前,在 Hot Chips 大会上,Tri Dao 展示了 Flash Attention 4 的初步成果,这是 Flash Attention 系列 CUDA 内核的最新成员。这些内核用于 Transformer 神经网络的注意力层。与更标准的矩阵乘法一样,这些计算是当代生成式 AI 工作负载的主要瓶颈。数十亿美元和吉瓦级的电力正被投入到 GPU 上,以更快地运行更多此类计算。而 Flash Attention 4 正是尽可能快地运行大量此类计算的方法。这篇博文将解释其工作原理。

新的 FA4 内核针对 Nvidia 新的 Blackwell 流式多处理器架构进行了优化,据报道,相比之前的最先进技术——Nvidia cudnn 库中的注意力内核,实现了约 20% 的加速。

A chart depicting the ~20% performance improvement of Flash Attention 4 over cudnn attention kernels.

cudnn 内核是闭源的,所以 Jensen 只知道里面发生了什么。

目前也还没有关于 FA4 工作原理的官方技术报告。但 Flash Attention 4 的源代码已经在网上发布,见此处。我们最近一直在为开源 LLM 推理引擎做贡献,因此我们阅读了代码并逆向工程了该内核的工作原理,包括两个数学技巧(更快的近似指数和更高效的在线 softmax),这些都是 Dao 的经典手法。本文包含了我们的发现。

也许令人惊讶的是,FA4 的架构对于普通软件工程受众来说很容易理解。

这是因为 FA4 最大的变化不是(非常酷的)数学——而是其异步操作“流水线”复杂度的巨大增加。这种异步编程在 CUDA 世界中相当新,但管道在 Unix 中已经存在了大约 40 年。有并行和并发程序经验的程序员,比如高性能数据库和 Web 服务器,会感到非常熟悉(除了某些新奇的 GPU 技术词汇)。

因此,我们将文章分为两部分。

第一部分是“快速导览”,通过追踪一个输入块如何转化为输出块的过程,介绍 FA4 的架构。其编写目的是让没有 CUDA 编程经验的执业软件工程师也能理解。我们简要解释了 CUDA 概念和硬件,如 warp 和 warp 调度器,但详细解释请参阅我们的 GPU 术语表(全文链接)。

第二部分是“深入探讨”,依次介绍每个子组件,解释其功能,并为特别勇敢的探索者提供源代码链接。

Flash Attention 4 快速导览:“一个 Tile 的生命周期”

我们从全局内存(即 GPU RAM)中的 bf16 查询、键和值张量开始。我们的目标是生成一个同样位于全局内存中的 bf16 输出张量。输出是查询与键的相似度加权后的值。计算这种加权需要矩阵乘法、指数运算和归一化。

作为优秀的工程师,我们通过将这个大问题分解成更小的部分来解决它。在这种情况下,这相当字面化:我们取非常大的输入张量,将其分割成相邻行和列的“瓦片”,每个瓦片都为计算一个输出瓦片做出贡献。

具体来说,我们的内核程序的一个运行实例(即一个线程的 “协作线程数组”)通过读取查询张量的两个瓦片来生成输出张量的两个瓦片。在此过程中,它为每个查询瓦片流式传输所有的键和值。键和值也以瓦片形式读取。如果你是数据库爱好者,可以将其视为针对键值存储的一批聚合查询的向量化顺序扫描。

A diagram showing a tile of queries combining with a stream of key and value tiles to produce an output tile.

通过并发地多次运行这个瓦片级程序(通常是 massively 并行),我们生成整个输出张量。这是一个 “单程序,多数据”执行模型,其中每个数据是一对瓦片。这种跨程序实例的并发是 CUDA 编程模型 的日常操作,并由 CUDA 运行时 为程序员透明处理。

但在最快的当代内核中,如 Flash Attention 3 和 4 以及所有最先进的矩阵乘法,我们的程序内部也存在并发。每个程序实例设置一个异步操作流水线,这些操作共同实现上述瓦片级计算。我们编写内核,使得在处理瓦片时,所有流水线步骤都能尽可能并发运行。在 Flash Attention 4 中,我们通过将流水线块映射到称为 warp 的 32 线程组来实现这一点(这种技术称为 warp 特化)。

然后我们依赖 warp 调度器 在每个时钟周期内在程序实例中的流水线步骤之间切换,当某个步骤停滞时换出,当某个步骤的下一个输入就绪时换回。可以想象成 CPU 的 同时多线程/“超线程”,但更强大。下面的图表来自我们的 GPU 性能术语表,描绘了四个周期跨四个并行槽,总共十六个 执行槽,其中十五个由于这种快速的 warp 切换而填充了正在执行指令的 warp。详情请参阅 相关文章。

A diagram depicting sixteen execution slots. Fifteen of them are colored in, indicating that they are filled with an active warp.

这种执行模型在以下意义上与 CPU 异步程序的工作方式 是“对偶”的。在异步 CPU 程序中,单个线程通过状态机(例如读取、解析、写入)实现单个数据(例如请求)的整个旅程,在数据就绪时在转换之间切换。在像 FA4 这样的异步 GPU 程序中,单个 warp 在类似的状态机中实现单个转换(例如从查询和值到注意力分数)。

cpu-async-vs-gpu-async.drawio.png

流水线采用生产者/消费者模型组织,并使用屏障进行同步。

与跨程序实例的并发不同,内部流水线并发都是手动实现的。这导致代码相当复杂——尽管控制流对任何 编写过自己的事件循环 的人来说都会很熟悉。

因此,像大多数异步代码一样,理解 FA4 内核的最佳方式是追踪单个 tile 的路径:即“一个 tile 的生命周期”,类似于浏览器渲染管线中的“一个像素的生命周期”。特别地,让我们跟随该 tile 在GPU 的内存层次结构中的路径,看它如何从初始的查询 tile 转变为最终的输出 tile。

在高层面上,忽略一些关于多重缓冲以增加并发性和并行性的细节,大致如下:

fa4-life-of-tile.drawio.png

这隐约类似于微服务架构图。如上所述,下亦如此!

具体来说,就是:

  • 一个查询 tile 由 Load warp 从全局内存(mQ)加载到共享内存(sQ)。共享内存是由程序员管理的“暂存器”L1 缓存。
  • 键(mK)和值(mV)的 tile 同样由 Load warp 流式加载到共享内存(sK、sV)。注意,如果工作集大小允许,这些 tile 为其他查询 tile 的未来加载将由硬件管理的 L2 缓存(未画出)提供服务。
  • 当每个键 tile 就绪时,MMA warp 使用 Tensor Core 将其与我们的查询 tile 相乘,在 Tensor Memory 中生成一个未归一化的注意力分数 tile(tS)。Tensor Core是用于运行矩阵乘法的专用硬件。Tensor Memory是另一个由程序员管理的 L1 缓存,旨在在一系列 Tensor Core 操作期间保存和累积中间结果。
  • When each tile of unnormalized attention scores is ready, a Softmax warp produces normalized attention scores for that tile in Tensor Memory (tP) without using the Tensor Core and updates a scaling factor used for numerical stability (in shared memory, not pictured).
  • When each tile of normalized attention scores is ready, a Correction warp checks if the normalization scaling factor has changed and, if necessary, rescales the final output tile in Tensor Memory (tO).
    • ⚡️ Flash Attention 4 新特性:何时重新缩放的决策变得更加智能,据报道将输出重新缩放操作减少了 10 倍。大致来说:缩放因子过去是一个简单的运行最大值。现在仅当最大值变化足够大以至于影响数值稳定性时才应用更新。这似乎是一个很好且非常可移植的想法。
  • 当每次重新缩放更新完成时,MMA warp 通过将值 tile(sV)按注意力分数 tile(tP)缩放后与 Tensor Memory 中的输出 tile(tO)累加,来更新输出 tile。
  • 当每个最终输出值 tile 就绪时,Correction warp 将其存储到共享内存(sO),然后 Epilogue warp 将其存储到全局内存(mO),该 tile 的处理就完成了。

我们这种高层面、以 tile 为中心的视角省略了许多细节,例如分配给每个流水线步骤的 warp 数量以及用于存储不同 tile 的缓冲区使用。它还遗漏了屏障同步的所有细节,这在每个生产者/消费者关系的两侧都是必需的(即图中箭头尖端与箭头尾部相接之处)。这些对性能至关重要。

下面我们将以“以 warp 为中心”的内核视角来详细讨论这些细节,重点关注每个 warp 中的操作,而不是 tile 的移动,并包含指向源代码的链接。这必然更加技术化,并且会以更快的速度介绍一些 GPU 特有的特性,因此不太适合一般的软件工程受众。

但在此之前,对于只关心高层次的读者,还有最后一个要点。

GPU 编程从这里走向何方?

当 Ian Buck 等人设计 CUDA C 时,他们被一个北极星所驱动:能否用它来编写一个单精度向量加法(saxpy),使其性能可观,并且是一行简洁的代码,让 C 程序员容易理解?当时奠定的 CUDA 编程模型的核心,并在 2008 年 Lindholm 等人的论文中描述,至今仍然存在。

过去几年(在 Hopper 和 Blackwell 架构中)的新变化是越来越依赖程序员管理的异步性,比如 FA4 的多阶段、多缓冲流水线。这相对于 FA3 更简单的“乒乓”流水线(为了利用 Hopper GPU 的异步能力而添加)来说,复杂性有了重大飞跃。

而且就像在其他设计良好的语言中一样,CUDA C/C++ 在适应异步性的引入方面一直很挣扎。这是一个普遍公认的事实,异步编程烂透了。当你需要管理自己的事件循环时尤其如此,就像我们在这里实际上所做的那样。而 CUDA 编程模型和 PTX 机器模型的以线程为中心和 warp 统一性使其变得更难,而不是更容易。

难怪Triton 团队放弃了编写 Blackwell 注意力机制,并在更低层次上添加了新的 Gluon 前端!

尽管 Triton 遇到了麻烦,但这个内核是向基于 tile、warp 专用编程转变的一个明显例子。Nvidia 正大力押注于许多新语言和库来试图让这变得更容易,从本内核中使用的 CuTe DSL 和 CUTLASS C++,到即将推出的 CuTile。不管你对聊天机器人炒作浪潮有什么看法,对于高性能数值计算来说,这是激动人心的时代!

面向 GPU 爱好者的深入探讨:在 Flash Attention 4 中,每个 warp 做什么?

在 Flash Attention 4 内核中,warp 有五种不同的专用化。下面列出了它们,并附有指向其源代码的链接。

  1. 一个加载 warp,用于将查询、键和值 tile 从全局内存加载到共享内存
  2. 一个MMA warp,用于从查询和键 tile 计算未归一化的注意力分数,并将分数加权的值 tile 累加到输出 tile 中
  3. 八个Softmax warp,用于计算归一化的注意力分数并跟踪运行统计量(最大值、总和)
  4. 四个校正 warp,用于监视归一化比例的更新并重新归一化输出 tile
  5. 一个或两个收尾 warp,用于将完成的输出 tile 从共享内存存储到全局内存

在上述讨论中,我们暗示每个 CTA 仅处理两个查询瓦片并仅产生两个输出瓦片。这在某些设置下是正确的,但瓦片与 CTA 之间的映射在技术上是由 TileScheduler 抽象出来的。为了获得最佳性能,你需要使用 StaticPersistentTileScheduler,它每个流式多处理器最多启动一个 CTA,然后将瓦片调度到这些 SM 上。这减少了 CTA 启动开销,并允许更细粒度的并发(例如,将一个瓦片的 Epilogue warp 与下一个瓦片的 Load 和 MMA warp 重叠)。

内核的核心工作是一样的——只是工作到线程构造的映射并不清晰,这使得解释工作更加困难。从这里开始,我们将回到像每个 CTA 只处理两个瓦片那样谈论代码(如果你使用 SingleTileScheduler,这实际上是正确的)。

此外,从这里开始我们将使用一些简写,与代码和惯例保持一致:Q 表示查询,K 表示键,V 表示值,O 表示输出,S 表示未归一化的注意力分数,P 表示归一化的注意力分数/“概率”。

Load warp 加载两个 Q 瓦片并流式传输所有 K 和 V 瓦片。

Load warp 操作全局内存中 Q、K 和 V 张量的指针,并写入共享内存中的 Q、K 和 V 张量。它通过可选的“页表”张量支持分页的键和值(如 Paged Attention 中那样,而不是操作系统页面)(再次强调,不是由操作系统、CPU 和 MMU 共同管理的页表)。

它使用 张量内存加速器 (TMA) 来减少多维数组访问带来的寄存器压力,并异步发起复制。这也避免了加载时非常长的 warp 停顿,否则将需要更多的 warp 专门化来隐藏延迟。

Load warp 加载两个 Q 瓦片。它在循环中加载所有 K 和 V 块。它是这些瓦片的“生产者”(在生产者/消费者设置中)。它可以并发加载最多三个 K 和 V 块。

当完成这些加载时,Load warp 通过共享内存中的屏障数组向 MMA warp 发出完成信号。所有屏障(不仅仅是用于 Load/MMA 同步的屏障)都通过它们在此数组中的偏移量来引用,以支持不同配置设置下的可变屏障数量。

MMA warp 计算未归一化的注意力分数和输出值。

MMA warp 操作共享内存中 Q、K 和 V 张量的指针。对于每个 K/V 瓦片,它运行两个 matmul 来创建 S 瓦片和两个 matmul 来生成 O(Q/K 用于 S 瓦片,P/V 用于 O 瓦片)。这些 matmul 是以内联 PTX 汇编形式发出的,这是 CUDA C/C++ 程序使用 Hopper 和 Blackwell 中 Tensor Core 所必需的。此内核中绝大多数 FLOPS 由这些行驱动;其他大部分都是内存管理。

使用的具体 PTX 指令是 tcgen05.mma.cta_group::1。mma 是矩阵乘累加。tcgen05 表示第 5 代 tensor core,即 Blackwell,对应 sm100/计算能力 10.0。cta_group::1 表示我们仅使用单个 CTA 来运行矩阵乘法,避免了基于 TPC 的 2SM/2CTA 矩阵乘法在 Blackwell 中可用的复杂性。这可能会引入轻微的内存吞吐量损失,但简化了 CTA/瓦片调度。有趣的是,ThunderKittens Blackwell 注意力内核做出了不同的选择。

同样在调度/简化方面:只有单个 leader_thread 发出指令。而且我们仅从单个 warp 工作。这与高性能的 Hopper MMA 有重要区别,后者是在整个 warpgroup 中协调的。

在获取 Q 瓦片和我们的第一个 K 瓦片后,我们运行第一次矩阵乘法以生成 S 的第一个结果。然后我们循环处理剩余的 K 和 V 瓦片并更新 S 和 O。这些 S 和 O 张量位于 Tensor Memory 中。这是Tensor Memory 的“预期”用途,作为 Tensor Cores 读取和写入的累加器存储。

由于 K 和 V 瓦片是缓冲的,我们需要在每次使用完它们后向 Load warp 发送信号(例如这里,表示一旦用于构建第二个 O 瓦片后,包含 V 的内存就可以重用)。这里还有一些额外的协调(围绕 S、P 和 O),我们将在其他 warp 中讨论时再谈。

八个 Softmax warp 生成归一化的注意力分数。

Softmax warp 生成归一化的注意力分数(P,即“概率”),供 MMA warp 使用。忽略这个名字,不要试图将注意力分数解释为随机变量的概率分布;这会让你头疼,并给关于 Transformer 的直觉带来误导。它们更好地被视为来自 V 的向量线性组合的权重。

核心的softmax 操作由两个 warpgroup 实现,即八个 warp。这两个 warpgroup 映射到两个查询/输出瓦片工作流。Warpgroup 由四个相邻的 warp 组成,warp 索引对齐为四。使用它们对于 Hopper GPU 中的快速 warpgroup MMA 至关重要,如 Flash Attention 3 中所示,但我们在该内核中没有看到任何明确使用它们的地方。Warpgroup 对齐可能导致工作更均匀地分布在 SM 的 warp 调度器/子单元上,正如在 Hopper 中那样,Hopper 每个 SM 有四个 warp 调度器。据我们和维基百科所知,关于 SM100 Blackwell GPU(如 B200)的这种细节水平尚未在任何地方发布(但对于 SM120 RTX Blackwell GPU 确实如此)。

我们也不确定为什么某些流水线阶段被分配了比其他阶段更多的 warp,以及为什么是这种特定比例。据推测,这有助于确保不同阶段之间的吞吐量平衡,但我们关于矩阵乘法和注意力操作之间相对操作负载、带宽和延迟的粗略计算并没有找到确凿证据。我们推测这是通过基准测试确定的。

每个 warp 一次运行在线 softmax 计算的单个步骤,同时循环处理由 MMA warp 生成的 S 瓦片。

深入单个softmax 步骤:未归一化的注意力分数存储在 Tensor Memory 中,只有 Tensor Cores 才能直接对其进行操作。但 Tensor Cores 只能做矩阵乘法。因此 Softmax warps 必须将分数复制到寄存器中以执行指数运算,然后再将结果复制回去。

指数运算的实现方式与之前版本的 Flash Attention 不同。FA3 及更早版本使用 GPU 的 Special Function Units 来执行硬件加速的指数运算。具体来说,它们使用 exp2 CUDA PTX 内建函数,该函数通常由(闭源的)ptxas 编译器映射到 MUFU.EX2 SASS 指令。

FA4 内核也这样做,但对于较小的注意力头尺寸,它还会以可调的频率在某些迭代中混入一种不同的指数运算算法。该实现使用这段内联 PTX 来计算 2 ** x。该算法将指数运算拆分为两部分:简单的整数部分(2 ** floor(x))和困难的有理部分(2 ** (x - floor(x)))。它使用三次多项式在单位区间上近似 2 ** x(可在 Wolfram Alpha 上此处查看该近似)。该近似在 bf16 取值范围内与 SFU 的输出相匹配。

三次多项式的计算按照 Horner 方法进行线性时间多项式求值,使用三个融合乘加(fma):

注意 f32x2 意味着我们操作的是一个包含两个 32 位值的向量(如向量通道)。你可以在 Stack Overflow 上此处阅读关于 CPU 向量指令的类似实现。

除了仅在某些迭代中应用此方法外,它还会在最后可配置数量的 S 块上停止应用该方法。综合来看,这表明应用该方法的目的是避免 SFU 出现瓶颈(由于波量化效应,这在最后的块中不太相关)。

Softmax warps 还会跟踪用于重新缩放和归一化注意力分数的运行时统计信息,供下文讨论的 Correction warps 使用。

这里还有另一个重要变化。所有 softmax 算法都需要处理大数值指数运算引起的数值不稳定问题。在 Flash Attention 之前,通常的做法是找到每一行中的最大值,并在指数运算前将其从该值中减去。所有 Flash Attention 内核都使用流式或在线 softmax 算法,而最大值事先并不知道——遍历分数来找到它会破坏使用流式算法的意义!相反,它们使用运行最大值来保证数值稳定性,并在遇到新最大值时更新缩放因子。这确保了持续的数值稳定性并避免了额外的扫描,但每当观察到新最大值时,都需要对先前的值进行代价高昂的修正(由 Correction warps 处理)。

这样做效率低下。我们只需要当新最大值变化足够大以致威胁数值稳定性时才更新缩放因子,而不是每次出现新最大值时都更新。该逻辑在此处实现。在 Hot Chips 演讲中,Dao 表示这将修正次数减少了 10 倍。

还额外支持 attention sinks 以及存储反向传播中使用的 log-sum-exp 张量。在撰写本文时(2025 年 9 月下旬),该 kernel 的反向版本尚不可用,但预计很快就会推出。

四个 Correction warp 会在归一化发生变化时对先前的输出进行重新缩放。

Correction warp 会随着数值稳定性缩放因子的变化,更新来自 MMA warp 的过往输出结果。Correction warp 需要与 MMA warp 协调它们对 Tensor Memory 中 O 值的访问(例如此处,表明这些值已被消费,内存可以被回收)。

与 Softmax warp 一样,四个 Correction warp 组成一个 warpgroup。同样与 Softmax warp 一样,它们需要从 Tensor Memory 加载到寄存器,以应用其非 matmul 重缩放操作。

Correction warp 还负责将输出从 Tensor Memory 写入共享内存,并应用按行求和的最终缩放。这被称为 correction_epilogue。这里的“Epilogue”与“Epilogue” warp 名称中的含义相同——即在一系列对存储于某一内存中的值进行的操作结束时、在结果被写入另一内存之前发生的操作。但在本例中,它指的是在数据被存储到共享内存之前对 Tensor Memory 中的数据进行的操作,而 Epilogue warp 则是从共享内存中取出数据并存储到全局内存。

这一点尤其令人困惑,因为这个 epilogue 的完成是 Epilogue warp 开始其工作的信号。

Correction warp 的参数中包含全局内存输出张量,但只在被注释掉的代码中使用它。

Epilogue Warp 将完整的输出 tile 存回全局内存。

根据是否启用 TMA,会有一个或两个 Epilogue warp。

在Epilogue warp 可以使用 TMA 的情况下,只有一个,而且它的工作很简单。它等待某个输出 tile 的 correction 循环完成,然后执行一次 TMA 拷贝,接着发出信号表示它已读完共享内存中的 O 张量,该缓冲区可以被复用。

如果它们无法使用 TMA,它们的工作会更复杂——它们需要处理切片和打包,这相当困难。这还会消耗相当多的额外寄存器。

如果你读到了这里,你可能会喜欢在 Modal 工作。

在 Modal,我们正在构建像巨型 Transformer 这样的计算密集型工作负载所需的云基础设施。我们的平台被 Suno、Lovable、Ramp 和 Substack 等公司使用。我们正在招聘。

作者感谢 vLLM 的 Simon Mo、RedHat AI 的 Michael Goin 以及 SemiAnalysis 的 Kimbo Chen 对本文草稿的评论。我们还要感谢 Tri Dao 写出了又一个超赞的 kernel。

来源:Modal Blog · modal.com