Infra: Flash Attention


Attention 的公式很短。先只看一个 head,忽略 batch,并让 query、key、value 的维度都为 \(d\)

\[ Q,K,V\in\mathbb{R}^{N\times d},\qquad S=\frac{QK^\top}{\sqrt d},\qquad P=\operatorname{softmax}(S),\qquad O=PV. \]

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 将它们变成权重:

\[ p_j=\frac{e^{x_j}}{\sum_{k=1}^{N}e^{x_k}}. \]

假设把这组 score 分成前半段和后半段。计算前半段的任意一个权重时,分母仍是全部 score 的指数和,也包含尚未读到的后半段。因此,不能只对前半段做一次 softmax,就把结果当作最终权重。

一个自然的办法是先累加分母,等全部读完再归一化。不过,实现时还需要处理数值稳定性:\(x_j\) 较大时,直接计算 \(e^{x_j}\) 可能溢出。先看不分块的情况,通常会让所有 score 减去同一个最大值,再求指数:

\[ \begin{gathered} m=\max_j x_j,\qquad \ell=\sum_j e^{x_j-m},\\ p_j=\frac{e^{x_j-m}}{\ell}. \end{gathered} \]

分子、分母都乘了 \(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\) 后:

\[ \begin{aligned} \sum_{j\in A}e^{x_j-m} &=e^{m_A-m}\sum_{j\in A}e^{x_j-m_A}\\ &=e^{m_A-m}\ell_A. \end{aligned} \]

这个等式说明,不必重新读取 \(A\) 中的 score,只需把已有的和乘一个共同的缩放因子。 若最大值没变,因子就是 1。对 \(B\) 做同样的调整,就能把两段的和加在一起:

\[ \begin{gathered} \alpha=e^{m_A-m},\qquad\beta=e^{m_B-m},\\ \ell=\alpha\ell_A+\beta\ell_B. \end{gathered} \]

这样,每次只需带着 \((m,\ell)\) 继续处理下一段。不过,只得到分母还不够:如果要输出完整的 softmax 向量,最后仍需保存或再次读取各个 score,才能算出每个 \(p_j\)

Attention 给了我们进一步简化的机会。它只需要 value 的加权和 \(O_i=\sum_jp_jV_j\),所以可以先累加尚未归一化的分子:

\[ u_A=\sum_{j\in A}e^{x_j-m_A}V_j. \]

它与 \(\ell_A\) 使用相同的指数权重,因此合并时也乘同一个 \(\alpha\)\(B\) 的分子同理。于是

\[ u=\alpha u_A+\beta u_B,\qquad O_i=\frac{u}{\ell}. \]

前面的问题到这里就解决了:不用提前确定每个位置的最终权重,只要把分子和分母一起累加,最后除一次。每个 query 保留两个标量 \(m,\ell\) 和一个 \(d\) 维向量 \(u\),已经算过的 score 与指数权重都可以丢弃。

仍用指数值为 \(1,2,4,8\) 的例子,并把 value 简化为标量 \(V=(1,2,3,4)\)。按完整公式计算,结果是

\[ O_i=\frac{1\cdot1+2\cdot2+4\cdot3+8\cdot4}{1+2+4+8} =\frac{49}{15}. \]
Online attention: merge two blocks in a common exponential scale左块的指数质量为 1、2,value 为 1、2;右块的指数质量为 4、8,value 为 3、4。左块状态按四分之一缩放,右块状态保持原样,合并得到分母 1.875、加权和 6.125,输出等于 49/15。Block A
\(\exp(x_A)=(1,2)\)
\(V_A=(1,2)\)
\(m_A=\log 2,\quad\ell_A=1.5,\quad u_A=2.5\)
\(\times \frac{1}{4}\)
Block B
\(\exp(x_B)=(4,8)\)
\(V_B=(3,4)\)
\(m_B=\log 8,\quad\ell_B=1.5,\quad u_B=5.5\)
\(\times 1\)
\(m=\max(m_A,m_B)=\log 8\)
\(\ell=1.5\times \frac{1}{4}+1.5\times 1=1.875\)
\(u=2.5\times \frac{1}{4}+5.5\times 1=6.125\)
\(O_i=\frac{6.125}{1.875}=\frac{49}{15}\approx 3.2667\)
左右滑动查看完整示意图
这里的 \(\exp(x)\) 是便于手算的未平移指数值,实际算法使用减去最大值后的指数。两个块的分母和加权和必须按同一基准合并;value 用标量展示,向量时逐分量执行相同运算。

图中两段的基准分别是 \(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:

\[ Q_i\in\mathbb{R}^{B_r\times d},\qquad K_j,V_j\in\mathbb{R}^{B_c\times d},\qquad S_{ij}=\frac{Q_iK_j^\top}{\sqrt d}\in\mathbb{R}^{B_r\times B_c}. \]

实现时,一个 thread block 可以负责一个 query tile,在 GPU 的一个 SM 计算单元上执行,用 shared memory 暂存输入,把累加状态留在 registers 中。两个矩阵乘法可以交给 Tensor Cores;最大值、分母和重缩放仍然对每个 query 独立进行。

\(U\) 把这些 query 的 \(u\) 排成矩阵,用 \(W\) 表示当前 score tile 减去运行中最大值后的指数权重。\(W\) 还没有除以分母,正好可以通过 \(WV_j\) 更新 \(U\)。这就把前面的分段算法变成了下图的数据流:

Attention intermediates: HBM tensors versus on-chip tiles上方三个独立算子将完整 score S 和概率 P 分别写入并读出 HBM。下方将一个 score tile 的矩阵乘法、online softmax 和 value 加权求和融合在片上,保留运行状态并复用临时空间,最终输出 O。
Separate kernels · \(Q,K,V\) start in HBM
\(\frac{QK^\top}{\sqrt d}\)
matmul
\(S\)
HBM · \(N\times N\)
\(\operatorname{softmax}\)
row reduction
\(P\)
HBM · \(N\times N\)
\(PV\)
matmul
\(O\)
HBM · \(N\times d\)
\(S\) and \(P\) each make a write → read round trip
Fused tiled kernel · one query tile at a timeOn chip · shared memory + registers
\(\frac{Q_iK_j^\top}{\sqrt d}\)
score tile
online softmax
update \(m,\ell\)
\(WV_j\)
rescale + update \(U\)
\(O_i=\frac{U}{\ell}\)
write to HBM
Repeat with the next \(K_j,V_j\) · reuse tile storage
左右滑动查看完整示意图
上:箭头穿过 \(S,P\) 时分别发生 HBM 写入与读回。下:虚线框内只容纳当前 tile 与运行状态,\(S,P\) 不形成完整的 HBM 张量。示意数据流,不按容量或耗时比例绘制。

下面采用 FlashAttention-2 (Dao, 2023) 的 query tile 外循环与延迟归一化形式,展示数据依赖。rowmaxrowsum 沿列归约,行向量按行 broadcast;这是片上计算的伪代码,逐行写成普通 PyTorch 操作不会自动得到一个 fused kernel。

伪代码中,\(m\) 是旧的最大值,\(m'\) 是合并后的最大值。新块直接相对 \(m'\) 求指数,前面公式里的 \(\beta\) 就已经包含在 \(W\) 中,只需用 \(\alpha\) 重缩放旧状态。

Algorithm 1FlashAttention forward

Input \(Q,K,V\) in HBM,按 \(B_r,B_c\) 分块

Output \(O\),以及 backward 所需的 \(L\)

  1. parallel for each query tile \(Q_i\) do
  2. load \(Q_i\)
  3. \(m \leftarrow -\infty,\)\(\ell \leftarrow 0,\)\(U \leftarrow 0\)
  4. for each KV tile \(K_j,V_j\) do
  5. load \(K_j,V_j\)
  6. \(S \leftarrow\)\(Q_iK_j^\top / \sqrt d\)
  7. mask \(S\)全被 mask 的行跳过本轮更新
  8. \(m' \leftarrow\)\(\max\!\left(m,\operatorname{rowmax}(S)\right)\)
  9. \(\alpha \leftarrow\)\(\exp(m-m')\)
  10. \(W \leftarrow\)\(\exp(S-m')\)
  11. \(\ell \leftarrow\)\(\alpha\ell +\)\(\operatorname{rowsum}(W)\)
  12. \(U \leftarrow\)\(\alpha U + WV_j\)
  13. \(m \leftarrow m'\)
  14. end for
  15. \(O_i \leftarrow\)\(U / \ell\)最后才归一化
  16. \(L_i \leftarrow\)\(m + \log\ell\)
  17. store \(O_i,L_i\)
  18. end for

标记行:计算重缩放系数,并用它同时调整分母与加权和。

符号与边界条件

这里的 \(S,W\) 只活在当前 tile 内,累加完就能复用空间。\(L_i\) 是逐行的 log-sum-exp,训练时留给 backward 使用。

初始化时,首个有效 tile 的 \(m'\) 有限,因此 \(\alpha=0\),旧的零状态不会贡献结果。对于全被 mask 的 tile 行,需要显式跳过更新,避免计算 \(-\infty-(-\infty)\);若整个 query 行都无有效 key,则要由算子另行约定输出,不能直接套用这里的除法。

Causal attention: traverse KV tiles for one query tile序列长 12,每块 2 个 query 和 2 个 key。当前 query 块为第 3 块,已处理前 2 个 KV 块,正在计算第 2 块,然后处理对角块。未来的整块跳过,对角块只计算有效的下三角。
\(K,V\) tiles →
\(0\)
\(Q_0\)
\(1\)
\(Q_1\)
\(2\)
\(Q_2\)
\(3\)
\(Q_3\)
\(4\)
\(Q_4\)
\(5\)
\(Q_5\)
\(Q_i\) stays on chip
\(B_r\times d\)
\(m,\ell,U\)
\(B_r,\ B_r,\ B_r\times d\) · running state
\(O_i=\frac{U}{\ell}\)
after all valid KV tiles
Conceptual score grid · never stored in HBM
左右滑动查看完整示意图
每格代表 \(2\times 2\) 个 score,块编号从 0 开始。粗框表示当前 query tile 的有效范围,✓ 表示已累加,● 表示当前 tile,斜线右上方是被 mask 的位置。右侧状态每处理一块就更新一次。

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 访问量为

\[ \Theta\!\left(\frac{N^2d^2}{M}\right). \]

这是原论文循环与存储模型下的界,不是上面伪代码的逐项流量统计,也没有把不同精度、cache 命中或硬件调度算进去。它展示了片上空间如何换取数据复用:在 \(M\gg d^2\) 的适用区间,相比显式中间矩阵带来的 \(\Theta(N^2)\) 流量,收益明显。

固定 \(M,d\) 时,上面的 IO 界仍随 \(N^2\) 增长,KV 也可能被不同 query tile 多次读取。可以用一个粗略的性能下界帮助判断优化方向:

\[ t\gtrsim\max\left( \frac{\text{FLOPs}}{\text{compute throughput}},\; \frac{\text{HBM bytes}}{\text{HBM bandwidth}} \right). \]

它忽略了同步、调度以及不能完全重叠的操作,不能直接当作耗时预测。但它说明,如果瓶颈是 HBM 带宽,减少中间矩阵的读写会直接降低这一项成本;优化后,还需要检查其他执行单元是否成为新的瓶颈。

Backward 为什么选择重算

训练还需要梯度。普通实现可以保留 \(P\) 给 backward 用;FlashAttention 若也这样做,前面省下的平方级存储就又回来了。

办法是 recomputation (Dao et al., 2022):保留 \(Q,K,V,O\),并保存每行一个

\[ L_i=m_i+\log\ell_i=\log\sum_j e^{S_{ij}}. \]

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\),则

\[ \frac{\partial\mathcal{L}}{\partial V}=P^\top G,\qquad \frac{\partial\mathcal{L}}{\partial P}=GV^\top. \]

Softmax 的 backward 是逐行的。令 \(H=GV^\top\),有

\[ \frac{\partial\mathcal{L}}{\partial S_{ij}} =P_{ij}\left(H_{ij}-D_i\right),\qquad D_i=\sum_j P_{ij}H_{ij}. \]

表面上 \(D_i\) 又要扫描一整行 \(P\)。但利用 \(O_i=\sum_jP_{ij}V_j\),可以把它改写为

\[ D_i=\sum_jP_{ij}(G_i\cdot V_j) =G_i\cdot\left(\sum_jP_{ij}V_j\right) =G_i\cdot O_i. \]

于是先从已有的 \(G,O\) 算出 \(D\);之后每个 tile 只需恢复局部的 \(P\)\(H\),就能得到局部的 \(\partial\mathcal{L}/\partial S\),再累加

\[ \frac{\partial\mathcal{L}}{\partial Q} =\frac{1}{\sqrt d}\frac{\partial\mathcal{L}}{\partial S}K, \qquad \frac{\partial\mathcal{L}}{\partial K} =\frac{1}{\sqrt d}\left(\frac{\partial\mathcal{L}}{\partial S}\right)^\top Q. \]

这些是完整矩阵写法;实现时仍然逐块累加,不物化 \(P,H\) 或 score gradient。与 forward 的独立 query 行不同,多个块可能贡献给同一份梯度,backward 的工作划分还要处理这种归约。

启用 attention dropout 时,还必须重现 forward 的随机 mask,例如保存可重放的随机数状态;上面无 dropout 的梯度推导不能原样忽略这一步。

从 FA2 到 FA4:继续减少等待

到这里,FlashAttention 的基本算法已经完整了。但数据留在片上之后,GPU 还可能因为工作太少、线程间同步或某些计算太慢而等待。后续版本继续处理这些问题。

FlashAttention-2 (Dao, 2023) 主要改进并行度与工作划分:

  1. 少做非矩阵运算。 保留未归一化的 \(U\),最后才除以 \(\ell\),省掉反复归一化输出的工作。Tensor Cores 的矩阵吞吐很高,指数、归约、除法等操作不能按同样的 FLOPs 吞吐估算。
  2. 增加独立的 query 工作。 只按 batch 和 head 分配 thread block 时,小 batch、长序列可能留下许多空闲 SM。再沿 query 序列分块,就有更多互不依赖的输出行可供调度。
  3. 减少 warp 之间的部分和归约。 一个 warp 是一起执行的 32 个 GPU 线程。在 thread block 内,让各 warp 负责不同 query 行并共享 KV,可以让它们各自产生完整的输出行。
Within a thread block: split KV versus split Q左侧两个 warp 处理相同的 query 和不同的 KV 范围,输出是同一行的部分结果,需要归约。右侧两个 warp 处理不同的 query 行和共享的 KV,各自得到完整的输出行,无需在 warp 之间归约这些输出。
Split \(K,V\) · partial outputs
Split \(Q\) · independent output rows
warp 0
\(Q_i\,;\ K^{(0)},V^{(0)}\)
\(m^{(0)},\ell^{(0)},U^{(0)}\)
warp 0
\(Q_i^{(0)}\) · shared \(K,V\)
\(O_i^{(0)}\)
warp 1
\(Q_i\,;\ K^{(1)},V^{(1)}\)
\(m^{(1)},\ell^{(1)},U^{(1)}\)
warp 1
\(Q_i^{(1)}\) · shared \(K,V\)
\(O_i^{(1)}\)
merge → \(O_i\)
shared memory + sync
Each warp owns its output rows
左右滑动查看完整示意图
以两个 warp 示意 forward 的工作划分,依据 FlashAttention-2 的 split-K 与 split-Q 思路重绘。上标 \((w)\) 标记 warp \(w\) 对应的分片。右侧省掉输出部分和的跨 warp 合并,加载共享 \(K,V\) 时仍需必要的同步。

图中若沿 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 和片上数据搬运可能让更快的矩阵单元等待。

作者的实现说明 围绕这些新瓶颈展开:

  1. 用 TMEM 重排流水线。 Blackwell 的异步 MMA 将累加结果写入 tensor memory(TMEM),减轻 registers 压力。FA4 交错处理两个 query tile,让一块的矩阵乘法与另一块的 softmax 重叠,并把输出重缩放交给单独的 warpgroup。
  2. 让其他计算单元分担指数。 部分指数通过乘加单元上的多项式近似计算,与硬件指数单元并行。多项式也需要指令和 registers,因此要选择合适的分担比例,并评估近似误差。
  3. 继续减少 backward 的片上流量。 将部分中间结果保留在 TMEM,直接供矩阵单元读取;再用 2-CTA MMA 让两个 thread block 协作,配合梯度 tile 的重排,减少 shared memory 的重复读取和梯度的全局原子累加。

FA4 还通过 conditional rescaling 减少对累加向量的反复缩放:最大值只增长一点时,可以暂时沿用旧的指数基准,让分子、分母和新 tile 都在同一尺度下继续累加;增长超过阈值时再一起调整。

为什么可以暂时不更新指数基准

前面的 online softmax 每块都把 \(m\) 更新为目前的最大值。但从代数上看,共同的指数基准不必始终等于最大值。用 \(r\) 表示当前采用的基准,对已经处理的位置集合 \(A\),保留

\[ \begin{aligned} \ell_A(r)&=\sum_{j\in A}e^{x_j-r},\\ u_A(r)&=\sum_{j\in A}e^{x_j-r}V_j. \end{aligned} \]

分子、分母都有共同因子 \(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,则重新比较计算成本与存储、搬运成本。最后还要将这个算法映射到硬件上,检查并行度、同步和各类执行单元是否成为新的限制。

References

Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. doi.org
Dao, T. (2024). FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision. tridao.me
Dao, T., Fu, D. Y., Ermon, S., Rudra, A., & Ré, C. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. Advances in Neural Information Processing Systems, 35, 16344–16359. arxiv.org
Jain, M., Kandar, T., Srivastava, S., Zhu, J., Chaudhary, J., Srivastava, D., & Kong, J. (2026). CSE 291A/DSC 291, Lecture 15: KV Cache Management and Flash Attention. haoailab.com
Liang, P. (2026). CS336: Language Modeling from Scratch, Lecture 6: Kernels. github.com
Milakov, M., & Gimelshein, N. (2018). Online Normalizer Calculation for Softmax. doi.org
Zadouri, T., Hoehnerbach, M., Shah, J., Liu, T., Thakkar, V., & Dao, T. (2026). FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling. doi.org

Cite this post

@misc{pu2026mlmlrevisitinfraflashattention,
  author = {Pu, Fanyi},
  title  = {Infra: Flash Attention},
  year   = {2026},
  month  = {9},
  url    = {https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention}
}