Chengshu@skadai · 2026.10.12
RSS
8,154 字 · 1,680 词 · 约 28 分钟

少搬一趟:FlashAttention 四代进化史

写给大学生的 FlashAttention 四代科普:从 HBM 搬砖讲到 FA1–FA4,在线 softmax、分块、warp 分工、乒乓与有条件重缩放。

学习笔记,据 Syfo 分享页 整理发布。

一本写给大学生的小书,讲 GPU 上最著名的一个“计算技巧”是怎么一代代长大的


写在前面

物理学家伽莫夫写《物理世界奇遇记》时,让一位叫汤普金斯的银行职员在梦里骑自行车,眼看着街道随着车速变扁,从而“亲身”体会相对论。他的诀窍是:先让读者看见现象,再告诉他背后的道理;道理一定要讲透,枝叶则大胆剪掉。

这本小书想用同样的办法讲一件近几年非常重要的事:FlashAttention。它是大语言模型(ChatGPT 这类模型)训练和推理中最核心的加速技术之一。从 2022 年到 2026 年,它出了四个版本,每一版都对应一篇论文、一代 NVIDIA GPU:

图 1 FlashAttention 四代时间线

图 1 FlashAttention 四代时间线

读完这本书,你应该能回答四个问题:

  1. 为什么注意力(attention)在 GPU 上慢?慢在哪里?
  2. 第一版用了什么数学技巧,让一个“必须看完整行才能算”的东西可以一块一块地算?
  3. 第二、三、四版又在改进什么?为什么同一个算法要改四次?
  4. 贯穿这四篇论文的那条主线是什么?

预备知识:线性代数里的矩阵乘法,微积分里的指数函数。不需要写过 GPU 程序。


第一章 慢,是因为跑腿

1.1 一个厨师的烦恼

想象一位厨师。他的灶台边有一块小砧板,切菜飞快,可砧板只放得下几样东西;所有食材都放在楼下的大仓库里,每拿一次都要跑一趟楼梯。

有一天他接到一道菜谱,步骤是这样的:

  1. 从仓库拿来两筐菜,切好,把切好的一大盆菜送回仓库;
  2. 再从仓库把那一大盆搬上来,调味,再送回仓库;
  3. 再搬上来,和第三筐菜一起下锅。

你一眼就能看出问题:他的刀工再快也没用,时间全花在楼梯上了。更聪明的做法是:一次只拿一小把菜,在砧板上切好、调味、下锅一气呵成,中间产物从不送回仓库。

这就是 FlashAttention 的全部精神。剩下的篇幅,我们要把这个比喻变成严格的数学和工程。

1.2 注意力在算什么

Transformer 模型里的注意力层,输入是三个矩阵:查询 $Q$、键 $K$、值 $V$,大小都是 $N \times d$。这里 $N$ 是序列长度(比如一篇文章有多少个词元,可以是几千到几十万),$d$ 是每个注意力头的维度(通常是 64 或 128)。

计算分三步:

$$ S = \frac{QK^\top}{\sqrt{d}}, \qquad P = \mathrm{softmax}(S), \qquad O = PV. $$

第一步算出每个词和每个词之间的“相关性打分”;第二步对每一行做 softmax,把打分变成加起来等于 1 的权重;第三步用这些权重对 $V$ 的行做加权平均。

对一行打分 $s_1, \dots, s_N$,softmax 定义为

$$ p_k = \frac{e^{s_k}}{\sum_{t=1}^{N} e^{s_t}}. $$

注意最关键的一点:$S$ 和 $P$ 都是 $N \times N$ 的大矩阵,而输入输出只有 $N \times d$。当 $N = 8192$ 时,一个注意力头的 $S$ 就有约 6700 万个数,用 16 位浮点存要 134 MB;一个模型有几十个头、几十层,还要乘上批量大小。$N$ 每翻一倍,这个矩阵就变成四倍。

1.3 GPU 的两种内存

GPU 不是一整块“算力”,它里面有一个存储金字塔:

图 2 GPU 的存储层次

图 2 GPU 的存储层次

以 FlashAttention 论文使用的 A100 为例:

  • HBM(高带宽显存):40–80 GB,带宽 1.5–2.0 TB/s。这就是“楼下的仓库”,我们平常说“显卡有 80G 显存”指的就是它。
  • SRAM(片上存储):A100 有 108 个流式多处理器(SM,可以理解为 GPU 里的 108 个“小工厂”),每个配 192 KB,合计约 20 MB,带宽估计约 19 TB/s。这就是“灶台边的砧板”,比 HBM 快一个数量级,但小了三个数量级。

所有计算都必须在片上进行。数据要先从 HBM 搬进来,算完再搬回去。

1.4 算得快不如搬得少

一块 A100 每秒能做 312 万亿次 16 位矩阵乘法运算(312 TFLOPS),而 HBM 每秒只能搬 1.5 万亿字节左右。两数相除:每从显存搬来 1 个字节,GPU 要做大约 200 次运算,才不会闲着。

这个“每字节运算次数”叫算术强度。一个操作的算术强度如果高于这个门槛(比如大矩阵乘法),它的速度由算力决定,叫计算受限;如果低于门槛(比如 softmax:每个数读进来只做一次指数、一次加法、一次除法),速度就由搬运决定,叫访存受限。

现在回头看标准注意力:

图 3 标准注意力与 FlashAttention 的数据流

图 3 标准注意力与 FlashAttention 的数据流

图 3(a)里,$S$ 被写进 HBM,又被读出来做 softmax;$P$ 被写进 HBM,又被读出来乘 $V$。两个 $N \times N$ 的大矩阵在仓库里进出了四趟,而 softmax 本身几乎不费力气。GPU 的大部分时间不是在算,而是在等数据。

这正是 FlashAttention 第一篇论文标题中那个词的意思:IO-aware,意识到输入输出(搬运)才是要优化的东西。

本章要点:注意力慢,主要不是因为运算多,而是因为两个 $N \times N$ 的中间矩阵要在慢速显存里反复进出。


第二章 FlashAttention(2022):少搬运

Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré,《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》,NeurIPS 2022

2.1 想法很简单,难点在 softmax

既然大矩阵进出 HBM 太费事,那就别让它落地:把 $Q$、$K$、$V$ 切成小块,每次只搬一小块进 SRAM,在片上把 $S$ 的一小块算出来,立刻做 softmax、立刻乘 $V$,然后丢掉。这个技巧叫分块(tiling),矩阵乘法早就这么做了。

麻烦在于 softmax。它的分母 $\sum_t e^{s_t}$ 需要整整一行的打分。如果我们一次只看到一行中的一小段,怎么知道分母是多少?

更麻烦的是,实际计算中不能直接算 $e^{s_k}$。16 位浮点数最大只能表示 65504,而 $e^{11.1}$ 就已经超过它了。所以真实的 softmax 都会先减去这一行的最大值 $m$:

$$ p_k = \frac{e^{s_k - m}}{\ell}, \qquad m = \max_t s_t, \qquad \ell = \sum_t e^{s_t - m}. $$

这样所有指数都不超过 $e^0 = 1$,不会溢出。可这又多了一个需要“看完整行”才能知道的量:最大值 $m$。

看起来我们被困住了:要分块,就看不到整行;看不到整行,就算不出 $m$ 和 $\ell$。

2.2 在线 softmax:边读边改账

解决办法是一个精巧的“改账”技巧,叫在线 softmax(online softmax)。它的思路像一个记账员:每读到一批新数据,就先按当前掌握的信息记账;一旦发现新情况,就把旧账按比例折算一下。

我们用一个小例子把它看清楚。一行打分是 $[1, 3, 2, 5]$,分两块到达:

图 4 在线 softmax 的算例

图 4 在线 softmax 的算例

读完第一块 $[1, 3]$,当前最大值是 3,于是记下 $\ell = e^{1-3} + e^{3-3} = 1.135$。

读到第二块 $[2, 5]$,发现新的最大值 5。之前那本账是“以 3 为基准”记的,现在基准变成了 5。怎么办?注意 $e^{s - 5} = e^{s - 3} \cdot e^{3 - 5}$,所以旧账整体乘上 $e^{3-5}$ 就换成了新基准,一个数都不用重算:

$$ \ell_{\text{新}} = 1.135 \times e^{3-5} + e^{2-5} + e^{5-5} = 1.204. $$

和一次性看完整行算出来的结果一模一样。

一般地,设读完前若干块后的账本是 $(m, \ell, \tilde{o})$,其中 $\tilde{o}$ 是还没除以分母的输出(加权和)。新来一块打分 $s^{(j)}$ 和对应的值 $V_j$:

$$ \begin{aligned} m_{\text{新}} &= \max\big(m,\ \max_k s^{(j)}k\big), \ \ell{\text{新}} &= e^{m - m_{\text{新}}},\ell + \sum_k e^{s^{(j)}k - m{\text{新}}}, \ \tilde{o}{\text{新}} &= e^{m - m{\text{新}}},\tilde{o} + \sum_k e^{s^{(j)}k - m{\text{新}}}, v_k . \end{aligned} $$

全部块读完后,输出就是 $o = \tilde{o} / \ell$。

请仔细体会这三行公式,它是整个 FlashAttention 家族的地基,后面三代都在它上面盖楼:

  • 每行只需要记两个数($m$ 和 $\ell$)加一个长度为 $d$ 的向量($\tilde{o}$),和 $N$ 无关;
  • 每次修正只是乘一个公共因子 $e^{m - m_{\text{新}}}$,旧数据不用重读;
  • 结果是精确的,不是近似。

最后一点值得强调。在 FlashAttention 之前,很多人为了让注意力变快,提出了各种“近似注意力”(稀疏的、低秩的),用精度换速度,但在实际墙钟时间上往往并不快。FlashAttention 算出来的结果在数学上和标准注意力完全相同,它只是换了一个计算顺序。

2.3 分块算法的全貌

有了在线 softmax,分块就顺理成章了:

图 5 分块计算示意

图 5 分块计算示意

对 $Q$ 的每一行块 $Q_i$:把它搬进 SRAM;依次搬入 $K_j, V_j$;在片上算出小块 $S_{ij} = Q_i K_j^\top$,用在线 softmax 更新这一行块的账本和 $O_i$;扫完整行后,把 $O_i$ 写回 HBM,只写这一次。

所有这些步骤被融合进一个 GPU 程序(称为一个“内核”,kernel),中间结果从不离开芯片。对照图 3(b):HBM 里只进出 $Q$、$K$、$V$、$O$,那两个 $N \times N$ 的大矩阵从未落地。

(为了好懂,我们按“先固定 $Q_i$、再扫 $K_j$”的顺序讲。第一版论文的循环顺序其实是反过来的,外层循环 $K_j, V_j$、内层循环 $Q_i$。下一章会看到,第二版正是把这个顺序调换了过来。)

2.4 反向传播:宁可重算,也不存

训练神经网络还要做反向传播,求梯度。标准做法会把前向算出的 $P$($N \times N$)存起来,留给反向用。这又是一个巨大的矩阵。

FlashAttention 的选择是:不存,到时候重算。 它只存输出 $O$,外加每行两个数:这一行最终的最大值 $m$ 和分母 $\ell$。$N \times N$ 的矩阵不存,只存 $2N$ 个数。

有了它们,任何一块 $P$ 都可以从 $Q$、$K$ 一步重算出来:

$$ P_{ij} = \frac{\exp(S_{ij} - m_i)}{\ell_i}, \qquad S_{ij} = Q_i K_j^\top / \sqrt{d}, $$

不需要再做在线修正。于是反向传播同样可以分块进行。(第二版会把 $m$ 和 $\ell$ 合并成一个数 $L = m + \log \ell$,于是 $P_{ij} = \exp(S_{ij} - L_i)$。)

重算意味着更多运算,这样划算吗?论文给了一个漂亮的实测(GPT-2 medium 的注意力,序列长度 1024,A100,前向加反向):

标准注意力 FlashAttention
运算量 66.6 GFLOPs 75.2 GFLOPs(因为重算,反而更多)
HBM 读写量 40.3 GB 4.4 GB
运行时间 41.7 ms 7.3 ms

运算多了 13%,搬运少了 9 倍,时间快了 5.7 倍。这张小表就是“IO 感知”最好的注脚:在访存受限的世界里,多算一点去换少搬一点,是好买卖。

2.5 把账算清楚:IO 复杂度

论文还从理论上算了它到底省了多少搬运。先约定记号:$N$ 是序列长度,$d$ 是每个注意力头的维度,$M$ 是 SRAM 能装下多少个数。$\Theta(\cdot)$ 只关心“搬运量随 $N$、$d$、$M$ 怎么增长”,忽略常数倍。

标准注意力:$\Theta(Nd + N^2)$。 读入 $Q$、$K$、$V$、写出 $O$,是 $Nd$ 量级;但中间的 $S$ 和 $P$ 都是 $N \times N$,要写进 HBM 再读出来,这是 $N^2$ 量级。$N$ 比 $d$ 大得多时,$N^2$ 这一项压倒一切。

FlashAttention:$\Theta(N^2 d^2 / M)$。 这个式子可以分三步拼出来:

  1. 一个 $Q$ 块能有多大? $Q$ 的每一行有 $d$ 个数,SRAM 只装得下 $M$ 个数,所以一次能搬进去的行数大约是 $B \approx M/d$。
  2. 整套 $K$、$V$ 要扫几遍? 每个 $Q$ 块都必须和全部的 $K$、$V$ 见一面,才能算出自己那几行的输出。$Q$ 共 $N$ 行,切成 $N/B = Nd/M$ 块,所以 $K$、$V$ 要从 HBM 完整读 $Nd/M$ 遍。
  3. 每遍读多少? $K$、$V$ 各是 $N \times d$,读一遍是 $Nd$ 量级。

乘起来:

$$ \underbrace{\frac{Nd}{M}}{\text{扫的遍数}} \times \underbrace{Nd}{\text{每遍读的量}} = \frac{N^2 d^2}{M}. $$

$Q$ 和 $O$ 各自只进出一次,是 $Nd$ 量级,比这一项小,可以忽略。所以分子里的 $d^2$ 来自两个 $d$:一个是“$K$、$V$ 每行有 $d$ 个数”,另一个是“每行越宽,SRAM 里放得下的 $Q$ 行越少,要扫的遍数越多”。分母的 $M$ 则说明:SRAM 越大,$Q$ 块越大,扫的遍数越少。

代入一组数。 取 $N = 4096$,$d = 64$,$M = 32768$(约 64 KB 的半精度数):

怎么算 搬运量(粗算,忽略常数)
标准注意力 $N^2$ 约 1680 万个数
FlashAttention $Q$ 块约 $M/d = 512$ 行,切成 8 块;$K$、$V$ 扫 8 遍,每遍约 $Nd \approx 26$ 万个数 约 210 万个数

两者之比正好是

$$ \frac{N^2 d^2 / M}{N^2} = \frac{d^2}{M} = \frac{4096}{32768} = \frac{1}{8}. $$

只要 $d^2$ 比 $M$ 小,FlashAttention 就搬得少。典型情况下 $d$ 为 64–128,SRAM 能放几万到十几万个数,正满足这个条件。计入常数后,论文给出的 HBM 访问量最多可减少到约九分之一。

注意一个容易误解的地方:FlashAttention 的内存占用是随 $N$ 线性增长的($N \times N$ 矩阵从不落地),但它的搬运量仍然随 $N^2$ 增长。它省下的是一个常数倍 $M/d^2$,并没有把平方变成线性。

更进一步,论文证明了一个下界:不存在这样一种精确注意力算法,它在一定范围内的所有 SRAM 大小下,HBM 访问量都渐近地少于 $N^2 d^2 / M$。也就是说,在“搬运次数”这个指标上,FlashAttention 已经是最优的。后面三代的提速,都不是从这里挤出来的。

2.6 成绩单

  • 内存:注意力占用的额外内存从随 $N^2$ 增长变成随 $N$ 线性增长,比精确注意力的基线最多省 20 倍。这让更长的上下文第一次变得可行。
  • 速度:BERT-large(序列长度 512)的训练比当时 MLPerf 1.1 的纪录快 15%;GPT-2(序列长度 1K)的训练比 HuggingFace 的实现快最多 3 倍;Long Range Arena 基准上快 2.4 倍。
  • 新能力:在 Path-X(序列长度 16K)任务上达到 61.4% 的准确率,这是 Transformer 第一次在这个任务上超过随机猜测(50%);配合块稀疏版本,在 Path-256(序列长度 64K)上达到 63.1%。

本章要点:分块 + 在线 softmax 让注意力可以一块块地精确计算,$N \times N$ 矩阵从不落地;反向传播用重算代替存储。搬运量降到了理论下界。


第三章 FlashAttention-2(2023):好分工

Tri Dao,《FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning》,2023

3.1 搬运省下来了,为什么还不够快?

FlashAttention 已经把搬运降到了理论下界,可它在 A100 上只用到了理论峰值算力的 25%–40%。作为对比,一个优化得好的矩阵乘法程序能用到 80%–90%。剩下的算力去哪了?

答案是:活儿分得不好。 要理解这一点,先认识一下 GPU 里的“劳动组织”:

  • 一个 GPU 有很多个 SM(A100 有 108 个),每个 SM 是一个独立的小工厂;
  • 程序被分成许多线程块(thread block),每个线程块分配到一个 SM 上运行;
  • 线程块内部再分成若干个 warp,一个 warp 是 32 个步调一致的线程,是 GPU 真正调度执行的单位;
  • 每个 SM 里有专门做矩阵乘法的 Tensor Core,速度远远快于做普通运算的单元。

FlashAttention-2 的算法和第一版完全一样,数学上没有任何新东西。它改的是三件“组织管理”上的事。

3.2 第一件事:能交给 Tensor Core 的,别自己算

A100 上,16 位矩阵乘法的峰值是 312 TFLOPS,而其他运算(32 位的加减乘除、指数等)只有 19.5 TFLOPS。也就是说,一次非矩阵乘运算的代价相当于 16 次矩阵乘运算。 哪怕非矩阵乘运算只占总量的一小部分,它也可能吃掉相当多的时间。

所以 FlashAttention-2 想方设法砍掉非矩阵乘运算。最典型的例子:第一版在每次读入新块时,都会把输出 $O$ 除以当前的分母 $\ell$,保持它“随时是归一化好的”。第二版意识到这没必要,只要像我们 2.2 节的公式那样,一直维护未归一化的 $\tilde{o}$,等全部读完再除一次就行了。类似地,它反向传播只需存 logsumexp $L$ 一个量,而不是 $m$ 和 $\ell$ 两个。

这些改动在数学上是平凡的,但在硬件上意义很大:每少一次除法,就省下相当于十几次矩阵乘的时间。

3.3 第二件事:让所有工厂都有活干

第一版的分工方式是:一个线程块负责一个注意力头。如果批量大小乘以头数有几百,108 个 SM 都能分到活。

可是处理长序列时,为了不爆显存,批量往往很小。比如批量为 1、16 个头,那只有 16 个线程块,108 个 SM 里有 92 个在发呆。

FlashAttention-2 的办法是沿序列长度方向也切开:同一个头里,$Q$ 的不同行块 $Q_i$ 互不相干(回看图 5 的最后一句话),完全可以交给不同的线程块。为此它把循环顺序调换为“外层 $Q_i$、内层 $K_j, V_j$”,每个线程块独立负责一行块,从头到尾自己扫完,最后写一次 $O_i$。这样即使只有一个头,也有 $N / B_r$ 个线程块($B_r$ 是行块大小)可以并行。

反向传播也类似,沿 $K_j, V_j$ 的方向切开,每个线程块负责一列块;不同线程块对 $Q$ 的梯度 $dQ$ 有贡献重叠,用原子加法汇总。

3.4 第三件事:warp 之间别开会

线程块内部,活又要分给 4 个或 8 个 warp。第一版的分法如图 6(a):

图 6 FlashAttention-2 的 warp 分工

图 6 FlashAttention-2 的 warp 分工

第一版把 $K$ 和 $V$ 切成几份分给各个 warp,大家共享 $Q$。问题是每个 warp 算出的只是 $O$ 的“部分和”,最后必须把各自的结果写进共享内存、互相同步、再加起来。这就像四个人各算了一张发票的一部分,最后还要开会对账。

FlashAttention-2 改成切 $Q$(图 6b):每个 warp 负责 $O$ 的若干整行,$K$、$V$ 大家共享读取。每个 warp 自己从头算到尾,彼此不需要交流。省下的不仅是共享内存读写,还有同步等待的时间。

3.5 成绩单

  • 比第一版快约 2 倍;
  • 前向达到 A100 理论峰值的 50%–73%,接近矩阵乘法的效率;
  • 用于端到端训练 GPT 类模型时,每块 A100 达到 225 TFLOPS,模型算力利用率 72%。

本章要点:算法不变,改的是分工:砍掉昂贵的非矩阵乘运算,沿序列方向增加并行度,让 warp 各干各的、不必互相对账。


第四章 FlashAttention-3(2024):边搬边算

Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao,《FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision》,2024

4.1 新硬件,老问题

2022 年起,NVIDIA 推出了新一代的 H100(Hopper 架构)。它的 16 位矩阵乘法峰值达到 989 TFLOPS,是 A100 的三倍多。可 FlashAttention-2 在 H100 上只能用到峰值的 35%。

原因在于 Hopper 带来了几件新“工具”,老代码根本没用上:

  • WGMMA:新的矩阵乘指令,由一个 warpgroup(4 个 warp,共 128 个线程)共同发起,而且是异步的。发出指令后线程可以去干别的,等需要结果时再回来取。
  • TMA(张量内存加速器):一个专门搬数据的硬件单元。以前搬一块数据,要很多线程各自计算地址、各自搬运;现在只要告诉 TMA“把这块搬过来”,它自己完成,线程完全空出来。
  • FP8:8 位浮点格式,矩阵乘法速度再翻一倍。

这些工具共同指向一个新的编程方式:异步。搬运、矩阵乘、其他运算可以同时进行,就看你会不会安排。

4.2 指数函数这个“慢郎中”

在安排之前,先看一个惊人的数字。H100 上,矩阵乘法每秒 989 万亿次,而指数这类特殊函数每秒只有 3.9 万亿次,相差 256 倍。

在头维度 128 的注意力里,矩阵乘法的运算次数是指数运算的 512 倍。指数运算次数少 512 倍,单价却贵 256 倍,两相抵消:256 ÷ 512 = 1/2,指数运算花掉的时间,相当于矩阵乘法的一半。

这说明 softmax 不再是可以忽略的“零头”。如果先做矩阵乘、再做 softmax、再做矩阵乘,Tensor Core 就有大约三分之一的时间在干等。理想的状态是:Tensor Core 在做矩阵乘的同时,指数单元在做 softmax,两者完全重叠。

4.3 生产者与消费者

FlashAttention-3 的第一招叫 warp 特化(warp specialization):让不同的 warp 干不同的工种。

  • 生产者 warp:只负责用 TMA 把 $K_j, V_j$ 从 HBM 搬进共享内存。它几乎不需要寄存器。
  • 消费者 warpgroup:只负责做矩阵乘和 softmax。

两者之间用一个环形缓冲区衔接:共享内存里有几个“槽位”,生产者往空槽里装数据,消费者从满槽里取数据,用完就腾出来。Hopper 还允许在 warpgroup 之间重新分配寄存器(setmaxnreg 指令),于是寄存器从闲着的生产者那里转给了消费者,消费者可以用更大的块。

这就像餐厅后厨:有人专门跑仓库,有人专门掌勺,掌勺的人永远不用自己下楼。

4.4 乒乓调度

第二招针对 4.2 节的问题:如何让 softmax 和矩阵乘同时进行?

图 7 FlashAttention-3 的乒乓调度

图 7 FlashAttention-3 的乒乓调度

让两个消费者 warpgroup 轮流使用 Tensor Core:第 1 组在做 softmax 时,第 2 组正好在做矩阵乘;然后反过来。这样 Tensor Core 几乎一刻不停,而指数运算的时间被“藏”在了另一组的矩阵乘后面。两组之间用屏障(barrier)同步,确保它们交替而不是撞车。

这个看似简单的调度,在 FP16、头维度 128、序列长度 8K 的设置下,把速度从 570 TFLOPS 提升到了 620–640 TFLOPS。

4.5 组内流水

第三招是在同一个 warpgroup 内部也做重叠:在算第 $j$ 块的 softmax 时,就提前把第 $j+1$ 块的矩阵乘 $Q K_{j+1}^\top$ 发给 Tensor Core(WGMMA 是异步的,发出去就不用等)。代价是要多占一些寄存器来存放下一块的打分。

论文的消融实验很清楚地展示了每一招的贡献:完整方法 661 TFLOPS;去掉组内流水降到 582;去掉 warp 特化降到 570。

4.6 FP8:位数减半,精度怎么保?

8 位浮点数能让矩阵乘法再快一倍,但它只有很少的有效位,精度很差。大模型的激活值里还常有离群值:少数几个数特别大。如果整个张量只用一个缩放因子,为了容纳这几个大数,其他正常的数就只能挤在很少的几个刻度里,误差很大。

FlashAttention-3 用了两个办法:

  1. 分块量化:不是整个张量一个缩放因子,而是每一块一个。FlashAttention 本来就是分块算的,这几乎是免费的。
  2. 非相干处理(incoherent processing):在量化之前,把 $Q$ 和 $K$ 都乘上同一个随机正交矩阵 $M$。因为 $MM^\top = I$,

$$ (QM)(KM)^\top = Q M M^\top K^\top = QK^\top, $$

结果完全不变。但乘上随机正交矩阵之后,原本集中在某几个分量上的大数会被“摊开”到所有分量上,离群值消失了,量化误差也就小了。这个随机矩阵由随机正负号和 Hadamard 矩阵构成,可以用类似快速傅里叶变换的方法在 $O(d \log d)$ 时间内完成,几乎不增加开销。

在一个含离群值的测试中,普通的逐张量 FP8 注意力误差(均方根误差)是 $2.4 \times 10^{-2}$,FlashAttention-3 的 FP8 版本是 $9.1 \times 10^{-3}$,降低到约 2.6 分之一。

4.7 成绩单

  • 前向比 FlashAttention-2 快 1.5–2.0 倍,反向快 1.5–1.75 倍;
  • FP16 达到 740 TFLOPS,即 H100 峰值的约 75%;
  • FP8 接近 1.2 PFLOPS(每秒 1200 万亿次)。

本章要点:Hopper 的硬件是异步的。FlashAttention-3 让专人搬运、两组轮流用 Tensor Core、组内也提前发射下一次矩阵乘,把 softmax 的时间藏在矩阵乘后面;FP8 则靠分块量化和随机旋转保住精度。


第五章 FlashAttention-4(2026):补短板

Ted Zadouri, Markus Hoehnerbach, Jay Shah, Timmy Liu, Vijay Thakkar, Tri Dao,《FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling》,arXiv 2603.05451,2026

5.1 不对称的升级

NVIDIA 的下一代 Blackwell 架构(B200)又把矩阵乘法算力翻了一番多。但这次的升级是不对称的:只有 Tensor Core 变快了,别的部件没跟上。

每个 SM、每个时钟周期 H100 B200
16 位矩阵乘运算 4096 次 8192 次
指数运算(MUFU 单元) 16 次 16 次(没变)
共享内存读带宽 128 字节 128 字节(没变)
张量内存(TMEM) 无 256 KB(新增)

用厨房的比喻:主厨的刀工又快了一倍,可帮忙调味的人还是那么几个,传菜的窗口也还是那么宽。主厨越快,就越显得别人慢。

FlashAttention-4 的标题里有个词叫 asymmetric hardware scaling(不对称的硬件扩展),说的就是这件事。

5.2 先体检:瓶颈到底在哪?

论文做的第一件事是“屋顶线”分析:处理一个块,各个部件分别要多少个时钟周期?谁最慢,谁就是瓶颈。

图 8 B200 上每个块所需的时钟周期

图 8 B200 上每个块所需的时钟周期

前向传播中,一个 $128 \times 128 \times 128$ 的块,矩阵乘要 1024 个周期,指数运算也要 1024 个周期。在 H100 上,矩阵乘要 2048 个周期,指数只要它的一半;到了 B200,矩阵乘快了一倍,指数却原地踏步,两者打了个平手。任何一点调度上的不完美,都会让指数运算成为瓶颈。

反向传播中,问题出在共享内存:读写需要 3328 个周期,比矩阵乘的 2560 个周期还多 30%。

所以 FlashAttention-4 的任务很明确:给指数运算减负,给共享内存减负。

5.3 招数一:让乘加单元帮忙算指数

GPU 里除了专门算指数的 MUFU 单元,还有大量做“乘加”运算(FMA,计算 $a \times b + c$)的普通单元,在注意力计算中它们相对空闲。能不能让它们也来算指数?

能。用的是一个古老的数值计算技巧。我们要算 $2^x$(GPU 上都用以 2 为底的指数,$e^x = 2^{x \log_2 e}$):

图 9 用乘加单元“软件模拟”指数函数

图 9 用乘加单元“软件模拟”指数函数

  1. 把 $x$ 拆成整数部分和小数部分:$x = n + f$,其中 $0 \le f < 1$。比如 $-3.3 = -4 + 0.7$。
  2. $2^n$ 完全不用算:浮点数的格式本来就是“尾数 × 2 的若干次方”,只要把 $n$ 直接写进浮点数的指数位就行了。
  3. $2^f$ 只需要在 $[0, 1)$ 这个小区间上近似,用一个三次多项式就足够好,按秦九韶算法(西方叫 Horner 法则)只要 3 次乘加。

这种“先把范围缩小、再用多项式逼近”的方法,在数值分析里叫 Cody-Waite 范围约简。

FlashAttention-4 把每行 10%–25% 的指数交给多项式去算,其余仍由 MUFU 完成。两条流水线同时开工,指数这一关的总耗时就降下来了。

精度够吗? 这里有一个很漂亮的观察。三次多项式在 32 位浮点下的最大相对误差是 $8.77 \times 10^{-5}$,比硬件指数($1.41 \times 10^{-7}$)差了约 600 倍,听上去很糟。但算出来的 $P$ 紧接着要被舍入成 BF16 格式(只有 8 位有效位)去做下一步矩阵乘,而 BF16 本身的舍入误差就有约 $3.9 \times 10^{-3}$。舍入之后,多项式版本的最大相对误差是 $3.90 \times 10^{-3}$,硬件版本是 $3.89 \times 10^{-3}$,几乎没有差别。

多花力气追求下一步马上会被舍掉的精度,是浪费。 这是工程上很值得学的一种判断力。

5.4 招数二:账本不必天天改

回忆 2.2 节的在线 softmax:每当发现更大的最大值,就要把旧账 $\ell$ 和 $\tilde{o}$ 乘上 $e^{m - m_{\text{新}}}$ 重新折算。对 $\tilde{o}$ 来说,这是一整行 $d$ 个数的乘法,而且必须在下一次累加之前做完,正好卡在关键路径上。

FlashAttention-4 问了一个问题:为什么一定要用真正的最大值做基准?

减去最大值,本来只是为了防止溢出。只要基准和真实最大值差得不太多,指数就不会大到溢出。所以它的规则是:只有当新的最大值比账本里的参考值大出超过 8(以 2 为底,即 $2^8 = 256$ 倍)时,才更新参考值并折算旧账;否则就继续沿用旧的参考值。

图 10 有条件的重缩放

图 10 有条件的重缩放

图中这一行打分的真实最大值变了 9 次,按老办法要折算 9 次;用新规则只需折算 1 次。

为什么结果仍然精确? 关键在于最后那一步 $o = \tilde{o} / \ell$。$\tilde{o}$ 和 $\ell$ 都是用同一个参考值记的账,不管这个参考值是多少,它带来的公共因子都会在相除时完全抵消。参考值只影响中间数字的大小,不影响最终结果。而中间值最多比 1 大 256 倍,对 32 位浮点数来说毫无溢出风险。

(工程细节:同一个 warp 里的线程要步调一致,所以只要其中有一个线程需要折算,整个 warp 就一起折算。)

5.5 招数三:重新设计流水线

Blackwell 还带来了新的硬件结构,FlashAttention-4 围绕它们重新设计了整条流水线:

  • 张量内存(TMEM):每个 SM 新增 256 KB 专门存放矩阵乘结果的存储。以前矩阵乘的结果放在寄存器里,寄存器很紧张;现在打分 $S$、概率 $P$、输出 $O$ 的累加器都可以放在 TMEM 中。
  • 完全异步的矩阵乘:Blackwell 的矩阵乘指令由单个线程发起,结果直接写进 TMEM,其他线程完全不必参与。
  • 更大的块:矩阵乘的基本块从 Hopper 的 64 行扩大到 128 行。

在此基础上,每个线程块同时处理两个 128 行的 $Q$ 块,像 FlashAttention-3 那样乒乓交替:一个块在做矩阵乘时,另一个块在做 softmax。线程块里的工种分得更细:

  • 一组线程专门发起矩阵乘和数据搬运;
  • 两个 softmax warpgroup,各管一个 $Q$ 块,每个线程负责一整行,求行最大值、行和都不需要和别的线程交流;
  • 一个校正 warpgroup,专门做 5.4 节里那个(已经很少发生的)旧账折算,把它从关键路径上彻底挪走。

5.6 反向传播:两个线程块合伙干

反向传播的瓶颈是共享内存。Blackwell 提供了一种 2-CTA 模式(CTA 就是线程块):相邻两个 SM 上的线程块合作完成一次更大的矩阵乘($M = 256$),每个线程块只需要准备一半的操作数。这样每个 SM 从共享内存读操作数的量就减少了。

再加上对 $dQ$ 计算顺序的重新安排,每个线程块往全局内存做原子加法的次数也减半了。共享内存总耗时从 3328 个周期降到 2688 个,只比矩阵乘多 5%。

论文还提供了确定性模式:原子加法的顺序不固定,会导致浮点舍入结果每次略有不同;确定性模式用信号量强制规定累加顺序,保证每次运行结果逐位相同,速度可达非确定性版本的 75%。这对需要可复现结果的大规模训练很重要。

5.7 顺带一提:用 Python 写内核

FlashAttention-4 全部用 CuTe-DSL 写成,这是一种嵌入在 Python 里的 GPU 编程语言,不再使用 C++ 模板。好处之一是编译快得多:一个前向内核从 55 秒降到 2.5 秒,反向从 45 秒降到 1.4 秒,快了 20–30 倍。对需要反复试验的内核开发者来说,这意味着迭代速度的质变。

5.8 成绩单

  • 在 B200 上,BF16 达到 1613 TFLOPS,即峰值的 71%;
  • 比 NVIDIA 官方库 cuDNN 9.13 快 1.1–1.3 倍,比 Triton 实现快 2.1–2.7 倍,序列越长、因果注意力下优势越明显;
  • 论文提到,NVIDIA 之后在新版 cuDNN 里也吸收了其中几项技术。

一个有趣的旁注:论文指出后续的 B300 已经把指数单元的吞吐量翻了一倍(每周期 32 次)。软件发现了硬件的短板,硬件随后就补上了。

本章要点:Blackwell 只升级了 Tensor Core,指数单元和共享内存成了新瓶颈。FlashAttention-4 用多项式让普通运算单元分担指数,用“差距大才改账”减少重缩放,用 TMEM 和 2-CTA 模式给共享内存减负。


尾声 同一个问题,问了四遍

现在我们可以把四代放在一起看了:

版本 年份 / 硬件 当时的瓶颈 核心对策 效果
FlashAttention 2022 / A100 HBM 搬运:$N \times N$ 矩阵反复进出显存 分块 + 在线 softmax,反向时重算 搬运量降到理论下界,内存从 $N^2$ 降到 $N$
FlashAttention-2 2023 / A100 分工不当:非矩阵乘运算贵、SM 闲置、warp 互相等待 延迟归一化、沿序列并行、按 Q 切分 warp 约 2 倍提速,达到峰值的 50%–73%
FlashAttention-3 2024 / H100 同步执行:搬运、矩阵乘、softmax 互相等待 warp 特化、乒乓调度、组内流水、FP8 740 TFLOPS,峰值的 75%
FlashAttention-4 2026 / B200 不对称扩展:指数单元和共享内存跟不上 软件指数、条件重缩放、TMEM、2-CTA 1613 TFLOPS,峰值的 71%

四篇论文其实都在问同一个问题:现在,到底是什么在拖后腿?

第一代发现是搬运,于是少搬。搬运降到下界后,第二代发现是分工,于是重新分工。分工理顺后,第三代发现各部件在互相等待,于是让它们同时工作。新硬件把矩阵乘做得更快后,第四代发现原来不起眼的指数运算成了短板,于是想办法补上。

瓶颈从来不会消失,只会转移。每解决一个,下一个就浮出水面。

这里还有三个值得带走的教训:

  1. 先量化,再优化。 每一代论文的起点都是一笔账:搬了多少字节,各部件要花多少周期。没有这笔账,就不知道力气该往哪使。
  2. 数学上不变,计算顺序可变。 四代 FlashAttention 算的都是精确的注意力,结果和教科书公式相同。所有加速都来自换一种顺序、换一种分工去算同一个东西。在线 softmax 这样简单的恒等变换,就是撬动整个领域的那根杠杆。
  3. 算法要贴着硬件设计。 同一个算法,在 A100、H100、B200 上的最优写法完全不同。理解硬件,才能写出好算法;好的软件也会反过来推动硬件的演进。

我们省略了什么

为了把主线讲透,下面这些内容我们有意略去了。它们在工程上都很重要,但不影响理解核心思想:

  • 因果掩码、dropout、变长序列、分页 KV 缓存、多查询/分组查询注意力(MQA/GQA)等功能的具体实现;
  • 第一版的块稀疏变体;
  • 反向传播梯度公式的完整推导;
  • FP8 中 V 矩阵在片上转置、寄存器布局重排等细节;
  • FlashAttention-4 中的负载均衡调度(LPT 调度)、TMEM 的具体布局、确定性模式的调度细节;
  • 各种具体的块大小、寄存器分配数字。

如果你想继续深入,最好的下一步是读第一篇论文的第 3 节(算法与 IO 分析),然后去看开源代码。


参考文献

  1. Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022. arXiv.14135
  2. Tri Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. 2023. arXiv.08691
  3. Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision. 2024. arXiv.08608
  4. Ted Zadouri, Markus Hoehnerbach, Jay Shah, Timmy Liu, Vijay Thakkar, Tri Dao. FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling. 2026. arXiv.05451
  5. 开源代码:github.com/Dao-AILab/flash-attention

图表数据均取自上述论文;5.2 节中 H100 的周期数,是按论文给出的每周期吞吐量推算的。

讨论

用 GitHub 账号留言;评论保存在公开仓库chengshu-blog-discussions的 Discussions 里。也可通过 RSS 订阅后续文章。