Attention 的公式很短。先只看一个 head,忽略 batch,并让 query、key、value 的维度都为 \(d\):
Softmax 沿每一行进行。\(S_{ij}\) 是第 \(i\) 个 query 对第 \(j\) 个 key 的 score,\(P_{ij}\) 是归一化后的权重;输出 \(O_i\) 是所有 value 的加权和。
直接按公式调用三个算子,会先得到完整的 \(S\),再得到完整的 \(P\),最后计算 \(O\)。两个中间矩阵都是 \(N\times N\),但调用者通常只需要 \(N\times d\) 的输出。必须计算每一对 query 和 key 的关系,是否就意味着必须同时保存所有关系?
FlashAttention (Dao et al., 2022) 从这里入手:让一小块 score 在 GPU 片上产生、参与 softmax 和加权求和,随后丢弃,只留下能继续计算的状态。它保留 dense attention 的数学定义;浮点运算顺序改变后,结果允许有正常的舍入差异,并不要求逐 bit 相同。
下面沿着 UCSD CSE 291 讲义 (Jain et al., 2026) 的 online softmax 思路展开,再连接到 GPU 的执行方式。推导默认没有 dropout,每行至少有一个有效 key。
为什么要少读写中间结果
假设 \(N=8192\)、\(d=128\),只看单个 head,张量用 BF16 存储,每个元素 2 bytes:
| 张量 | 元素数量 | 存储量 |
|---|---|---|
| 单个 \(Q\)、\(K\)、\(V\) 或 \(O\) | \(Nd\) | 2 MiB |
| 单个 \(S\) 或 \(P\) | \(N^2\) | 128 MiB |
若同时保留 \(S\) 和 \(P\),仅这两项就是 256 MiB;32 个 head 则是 8 GiB。这里没有计入其他层的 activation、梯度或临时 workspace,也没有假设真实 kernel 的所有中间量都是 BF16。
即使显存放得下,读写仍有成本。GPU 的 HBM 容量大,shared memory 和 registers 位于片上,容量小、访问快。如果中间结果能留在片上,算完后直接被下一步使用,就能少一次写回 HBM、再从 HBM 读出的往返。
在上面的三算子实现中,\(S\) 要写出、读回,\(P\) 也要写出、读回。按相同的 BF16 假设,仅这四次中间矩阵传输就有 512 MiB/head,还没算 \(Q,K,V,O\) 的流量。
这就是 kernel fusion 的出发点:把相邻计算放进同一个 GPU kernel,让它们直接传递中间结果。Stanford CS336 2026 春季的 kernels 讲义 (Liang, 2026) 用这个思路解释为什么算术完全相同的程序,实际运行时间可以很不一样。
但把三个算子合在一起,并不能让 \(N^2\) 个元素突然装进 shared memory。还需要 tiling:每次只计算一个放得下的小矩形,并尽量复用已经加载的输入。
矩阵乘法可以分块计算再累加。但 attention 中间还有 softmax:它能不能也分段计算,而且算过的 score 不必再保留?
Softmax 怎样分段计算
固定一个 query,把它与所有 key 的 score 记为 \(x_1,\ldots,x_N\)。Softmax 将它们变成权重:
假设把这组 score 分成前半段和后半段。计算前半段的任意一个权重时,分母仍是全部 score 的指数和,也包含尚未读到的后半段。因此,不能只对前半段做一次 softmax,就把结果当作最终权重。
一个自然的办法是先累加分母,等全部读完再归一化。不过,实现时还需要处理数值稳定性:\(x_j\) 较大时,直接计算 \(e^{x_j}\) 可能溢出。先看不分块的情况,通常会让所有 score 减去同一个最大值,再求指数:
分子、分母都乘了 \(e^{-m}\),所以权重不变;同时 \(x_j-m\leq 0\),指数就不会因为 score 太大而溢出。这是 stable softmax,也常称 safe softmax。
现在回到分段计算。按上面的定义直接实现,需要先读完全部 score,找到最大值,再求指数和。Online softmax (Milakov & Gimelshein, 2018) 希望边读取边更新:手里暂时只有前面一段,就先用这一段的最大值作为基准;后面遇到更大的值,再调整已有的结果。
具体来说,已经处理的一段 \(A\),记其最大值为 \(m_A\),减去 \(m_A\) 后的指数和为 \(\ell_A\);新读入的一段 \(B\) 也能求出 \(m_B,\ell_B\)。两段合起来的最大值是 \(m=\max(m_A,m_B)\)。以 \(A\) 为例,它原来的指数和是相对 \(m_A\) 算的,换成 \(m\) 后:
这个等式说明,不必重新读取 \(A\) 中的 score,只需把已有的和乘一个共同的缩放因子。 若最大值没变,因子就是 1。对 \(B\) 做同样的调整,就能把两段的和加在一起:
这样,每次只需带着 \((m,\ell)\) 继续处理下一段。不过,只得到分母还不够:如果要输出完整的 softmax 向量,最后仍需保存或再次读取各个 score,才能算出每个 \(p_j\)。
Attention 给了我们进一步简化的机会。它只需要 value 的加权和 \(O_i=\sum_jp_jV_j\),所以可以先累加尚未归一化的分子:
它与 \(\ell_A\) 使用相同的指数权重,因此合并时也乘同一个 \(\alpha\);\(B\) 的分子同理。于是
前面的问题到这里就解决了:不用提前确定每个位置的最终权重,只要把分子和分母一起累加,最后除一次。每个 query 保留两个标量 \(m,\ell\) 和一个 \(d\) 维向量 \(u\),已经算过的 score 与指数权重都可以丢弃。
仍用指数值为 \(1,2,4,8\) 的例子,并把 value 简化为标量 \(V=(1,2,3,4)\)。按完整公式计算,结果是
图中两段的基准分别是 \(m_A=\log 2\)、\(m_B=\log 8\)。合并时,把左块的 \(\ell_A,u_A\) 同时乘 \(e^{\log 2-\log 8}=1/4\),得到 \(u/\ell=6.125/1.875=49/15\),与完整计算一致。不能直接平均两段各自的 attention 输出,因为两段在完整分母中的份额不同。
把分段计算放进 GPU
上面固定一个 query,说明了如何分段处理 key。GPU 上会把 \(B_r\) 个 query 放在一起,让它们复用同一块 KV,并将计算组织成矩阵乘法。下面的 \(i,j\) 表示块编号,每块 KV 包含 \(B_c\) 个 token:
实现时,一个 thread block 可以负责一个 query tile,在 GPU 的一个 SM 计算单元上执行,用 shared memory 暂存输入,把累加状态留在 registers 中。两个矩阵乘法可以交给 Tensor Cores;最大值、分母和重缩放仍然对每个 query 独立进行。
用 \(U\) 把这些 query 的 \(u\) 排成矩阵,用 \(W\) 表示当前 score tile 减去运行中最大值后的指数权重。\(W\) 还没有除以分母,正好可以通过 \(WV_j\) 更新 \(U\)。这就把前面的分段算法变成了下图的数据流:
下面采用 FlashAttention-2 (Dao, 2023) 的 query tile 外循环与延迟归一化形式,展示数据依赖。rowmax、rowsum 沿列归约,行向量按行 broadcast;这是片上计算的伪代码,逐行写成普通 PyTorch 操作不会自动得到一个 fused kernel。
伪代码中,\(m\) 是旧的最大值,\(m'\) 是合并后的最大值。新块直接相对 \(m'\) 求指数,前面公式里的 \(\beta\) 就已经包含在 \(W\) 中,只需用 \(\alpha\) 重缩放旧状态。
Input \(Q,K,V\) in HBM,按 \(B_r,B_c\) 分块
Output \(O\),以及 backward 所需的 \(L\)
- parallel for each query tile \(Q_i\) do
- load \(Q_i\)
- \(m \leftarrow -\infty,\)\(\ell \leftarrow 0,\)\(U \leftarrow 0\)
- for each KV tile \(K_j,V_j\) do
- load \(K_j,V_j\)
- \(S \leftarrow\)\(Q_iK_j^\top / \sqrt d\)
- mask \(S\)全被 mask 的行跳过本轮更新
- \(m' \leftarrow\)\(\max\!\left(m,\operatorname{rowmax}(S)\right)\)
- \(\alpha \leftarrow\)\(\exp(m-m')\)
- \(W \leftarrow\)\(\exp(S-m')\)
- \(\ell \leftarrow\)\(\alpha\ell +\)\(\operatorname{rowsum}(W)\)
- \(U \leftarrow\)\(\alpha U + WV_j\)
- \(m \leftarrow m'\)
- end for
- \(O_i \leftarrow\)\(U / \ell\)最后才归一化
- \(L_i \leftarrow\)\(m + \log\ell\)
- store \(O_i,L_i\)
- end for
标记行:计算重缩放系数,并用它同时调整分母与加权和。
符号与边界条件
这里的 \(S,W\) 只活在当前 tile 内,累加完就能复用空间。\(L_i\) 是逐行的 log-sum-exp,训练时留给 backward 使用。
初始化时,首个有效 tile 的 \(m'\) 有限,因此 \(\alpha=0\),旧的零状态不会贡献结果。对于全被 mask 的 tile 行,需要显式跳过更新,避免计算 \(-\infty-(-\infty)\);若整个 query 行都无有效 key,则要由算子另行约定输出,不能直接套用这里的除法。
Causal attention 中,整块位于未来的 tile 可以直接跳过;与对角线相交的 tile 在片上把无效 score 设为 \(-\infty\),其指数权重就是 0。这样无需单独存一个完整的三角 mask。图中各个 query tile 可以独立计算,但每个 tile 内仍要沿着有效的 KV 块合并状态。
Tile 越大,输入复用通常越充分;但 \(Q_i\)、\(K_j\)、\(V_j\)、\(U\)、\(S_{ij}\) 和 softmax 状态都要占片上资源。Registers 或 shared memory 占得太多,会减少同一 SM 能同时驻留的工作,甚至导致 register spilling。实际选块要同时考虑数据类型、head dimension 和硬件资源。
FlashAttention 节省的是中间结果的存储和搬运,dense attention 的乘法仍要做:
| 指标(单个 head) | 三算子、显式中间矩阵 | FlashAttention |
|---|---|---|
| Dense attention 算术量 | \(\Theta(N^2d)\) | 仍为 \(\Theta(N^2d)\) |
| HBM 中保存的 \(S,P\) | \(\Theta(N^2)\) | 不保存完整矩阵 |
| 除输入输出外的归一化状态 | 取决于实现 | \(\Theta(N)\) |
| HBM 访问量 | 有 \(\Theta(N^2)\) 的中间矩阵读写 | 依赖 tile、循环顺序和片上容量 |
输入和输出本身仍占 \(\Theta(Nd)\)。所以,显存占用随序列长度线性增长,不代表计算量或总 HBM 访问量也变成了线性。
原论文怎样分析 IO 访问量
原始 FlashAttention 论文 (Dao et al., 2022) 用一个两级存储模型分析 IO。设片上容量为 \(M\) 个标量元素,取 \(d\leq M\leq Nd\)。其 KV tile 外循环令 \(B_c=\Theta(M/d)\),每个 KV 块扫描一次 \(Q\) 和输出状态;共有约 \(Nd/M\) 个 KV 块,每轮搬运 \(\Theta(Nd)\) 个元素,于是 HBM 访问量为
这是原论文循环与存储模型下的界,不是上面伪代码的逐项流量统计,也没有把不同精度、cache 命中或硬件调度算进去。它展示了片上空间如何换取数据复用:在 \(M\gg d^2\) 的适用区间,相比显式中间矩阵带来的 \(\Theta(N^2)\) 流量,收益明显。
固定 \(M,d\) 时,上面的 IO 界仍随 \(N^2\) 增长,KV 也可能被不同 query tile 多次读取。可以用一个粗略的性能下界帮助判断优化方向:
它忽略了同步、调度以及不能完全重叠的操作,不能直接当作耗时预测。但它说明,如果瓶颈是 HBM 带宽,减少中间矩阵的读写会直接降低这一项成本;优化后,还需要检查其他执行单元是否成为新的瓶颈。
Backward 为什么选择重算
训练还需要梯度。普通实现可以保留 \(P\) 给 backward 用;FlashAttention 若也这样做,前面省下的平方级存储就又回来了。
办法是 recomputation (Dao et al., 2022):保留 \(Q,K,V,O\),并保存每行一个
Backward 需要某个 tile 时,重新计算它的 score,然后用 \(P_{ij}=e^{S_{ij}-L_i}\) 恢复概率。计算完该 tile 的梯度就丢弃它。保存 \(L\) 只新增 \(N\) 个标量,但训练 activation 的总存储当然还包括 \(Q,K,V,O\)。这是将 activation checkpointing 用到 attention 内部的做法。
Recomputation 多做了 score 和指数计算,却省去了大矩阵在 HBM 中的保存与读取,因此可能同时节省显存和时间。接下来只需将梯度计算也安排成逐块执行,就能在 forward 和 backward 中都避开完整的 \(N\times N\) 中间矩阵。
Backward 的梯度如何逐块计算
记上游梯度为 \(G=\partial\mathcal{L}/\partial O\),则
Softmax 的 backward 是逐行的。令 \(H=GV^\top\),有
表面上 \(D_i\) 又要扫描一整行 \(P\)。但利用 \(O_i=\sum_jP_{ij}V_j\),可以把它改写为
于是先从已有的 \(G,O\) 算出 \(D\);之后每个 tile 只需恢复局部的 \(P\) 和 \(H\),就能得到局部的 \(\partial\mathcal{L}/\partial S\),再累加
这些是完整矩阵写法;实现时仍然逐块累加,不物化 \(P,H\) 或 score gradient。与 forward 的独立 query 行不同,多个块可能贡献给同一份梯度,backward 的工作划分还要处理这种归约。
启用 attention dropout 时,还必须重现 forward 的随机 mask,例如保存可重放的随机数状态;上面无 dropout 的梯度推导不能原样忽略这一步。
从 FA2 到 FA4:继续减少等待
到这里,FlashAttention 的基本算法已经完整了。但数据留在片上之后,GPU 还可能因为工作太少、线程间同步或某些计算太慢而等待。后续版本继续处理这些问题。
FlashAttention-2 (Dao, 2023) 主要改进并行度与工作划分:
- 少做非矩阵运算。 保留未归一化的 \(U\),最后才除以 \(\ell\),省掉反复归一化输出的工作。Tensor Cores 的矩阵吞吐很高,指数、归约、除法等操作不能按同样的 FLOPs 吞吐估算。
- 增加独立的 query 工作。 只按 batch 和 head 分配 thread block 时,小 batch、长序列可能留下许多空闲 SM。再沿 query 序列分块,就有更多互不依赖的输出行可供调度。
- 减少 warp 之间的部分和归约。 一个 warp 是一起执行的 32 个 GPU 线程。在 thread block 内,让各 warp 负责不同 query 行并共享 KV,可以让它们各自产生完整的输出行。
图中若沿 KV 拆,同一个 query 的结果散落在不同 warp,需要合并部分状态;沿 Q 拆,则每个 warp 负责自己的输出行。这减少的是输出累加所需的 warp 间通信,共享 KV 的加载仍需协调。这里的划分发生在一个 thread block 内,与把多个 query tile 分配给不同 thread block 是两个层次。
FlashAttention-3 (Dao, 2024) 进一步让不同工作重叠。Hopper 提供了异步的数据搬运与矩阵乘法指令,分别是 TMA 和 WGMMA。将搬运和计算分给不同的 warp 后,可以在计算当前 tile 时预取下一块 KV,或让一组 warp 的 softmax 与另一组的矩阵乘法同时进行。
这里的关键是找出可以重叠的工作,同时遵守数据依赖:当前 tile 的 \(PV\) 必须等自己的 softmax 权重准备好。流水线也需要更多 buffer 和 registers,收益要与资源占用一起衡量。FA3 还探索 FP8;降低精度带来的误差需要单独评估,不能因为使用了 FlashAttention 就认为量化没有误差。
FlashAttention-4 (Zadouri et al., 2026) 的论文于 2026 年 3 月发布,重点研究 Blackwell 上的 attention。以 B200 为例,Tensor Cores 的吞吐相比 Hopper 增长更快,指数运算和 shared memory 带宽却没有同比提升。于是,softmax 和片上数据搬运可能让更快的矩阵单元等待。
作者的实现说明 围绕这些新瓶颈展开:
- 用 TMEM 重排流水线。 Blackwell 的异步 MMA 将累加结果写入 tensor memory(TMEM),减轻 registers 压力。FA4 交错处理两个 query tile,让一块的矩阵乘法与另一块的 softmax 重叠,并把输出重缩放交给单独的 warpgroup。
- 让其他计算单元分担指数。 部分指数通过乘加单元上的多项式近似计算,与硬件指数单元并行。多项式也需要指令和 registers,因此要选择合适的分担比例,并评估近似误差。
- 继续减少 backward 的片上流量。 将部分中间结果保留在 TMEM,直接供矩阵单元读取;再用 2-CTA MMA 让两个 thread block 协作,配合梯度 tile 的重排,减少 shared memory 的重复读取和梯度的全局原子累加。
FA4 还通过 conditional rescaling 减少对累加向量的反复缩放:最大值只增长一点时,可以暂时沿用旧的指数基准,让分子、分母和新 tile 都在同一尺度下继续累加;增长超过阈值时再一起调整。
为什么可以暂时不更新指数基准
前面的 online softmax 每块都把 \(m\) 更新为目前的最大值。但从代数上看,共同的指数基准不必始终等于最大值。用 \(r\) 表示当前采用的基准,对已经处理的位置集合 \(A\),保留
分子、分母都有共同因子 \(e^{-r}\),所以 \(u_A(r)/\ell_A(r)\) 不随 \(r\) 改变。初始化之后,只要指数与累加值仍在安全的数值范围内,就可以暂时保留旧的 \(r\),让新 tile 也相对这个基准累加。这样,旧的向量 \(u\) 就不必每次都乘一个缩放因子。
FA4 的官方 softmax 实现 用阈值控制这种更新:最大值相对当前基准跳得足够大时,才更新基准并重缩放状态;否则保留旧基准,缩放因子取 1。关键是 \(\ell,u\) 和新 tile 始终使用同一尺度,最后再归一化。只更新基准却省掉旧状态的重缩放,会破坏这个不变量。
FA4 的官方实现 使用 CuTe DSL,以 Python 表达并编译 GPU kernel。这些实现上的变化,仍然围绕同一个目标:让分块 attention 的计算与数据流适合当前硬件。
Prefill 与 Decode 有什么不同
训练与 prefill 通常同时处理很多 query,容易形成较大的矩阵乘法,也有大量 \(N\times N\) 中间结果可以省掉。普通单 token decode 中,query 长度是 1,KV 长度是 \(N\);每个 head 当前步骤的 score 只有 \(1\times N\),计算量已经是 \(\Theta(Nd)\)。
因此 decode 仍能受益于融合与 online reduction,但常见难点转向读取长 KV cache,以及如何为少量 query 提供足够的并行工作。前面的状态合并公式不限定扫描顺序:可以将 KV 分成多段分别计算,再按树形合并结果。实数运算下结果相同,浮点下可能有舍入差异;额外的并行度也会带来部分结果存储和归约成本。不能直接把长序列训练的加速比套到 decode 上。
这也解释了它与 PagedAttention 的关系:PagedAttention 管理跨生成步骤保留的 KV 布局和共享,FlashAttention 优化一次 attention 运算内部的数据流。两者可以结合,但融合计算不会自动消除 KV cache,也不会减少 dense attention 必须关注的有效 token。
FlashAttention 值得迁移到其他算子的思考方式是:先看下游究竟需要什么,再寻找足够小的可合并状态,让中间结果在产生的位置被消费;对 backward,则重新比较计算成本与存储、搬运成本。最后还要将这个算法映射到硬件上,检查并行度、同步和各类执行单元是否成为新的限制。