# Infra: Flash Attention

Author: Fanyi Pu

Published: 2026-09-11

Canonical: <https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention>

从 IO-aware computation、online softmax、tiling 和 recomputation 理解 FlashAttention，并梳理 FA2 到 FA4 的并行、流水线与硬件适配。

[Attention](https://pufanyi.com/blog/ml/ml-revisit/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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-dao2022flashattention)) 从这里入手：让一小块 score 在 GPU 片上产生、参与 softmax 和加权求和，随后丢弃，只留下能继续计算的状态。它保留 dense attention 的数学定义；浮点运算顺序改变后，结果允许有正常的舍入差异，并不要求逐 bit 相同。

下面沿着 UCSD CSE 291 讲义 ([Jain et al., 2026](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-jain2026kvcache)) 的 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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-liang2026kernels)) 用这个思路解释为什么算术完全相同的程序，实际运行时间可以很不一样。

但把三个算子合在一起，并不能让 $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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-milakov2018online)) 希望边读取边更新：手里暂时只有前面一段，就先用这一段的最大值作为基准；后面遇到更大的值，再调整已有的结果。

具体来说，已经处理的一段 $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$

[View diagram in the original article](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#flash-online-merge)

左右滑动查看完整示意图

这里的 $\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$

$PV$

$O$

HBM · $N\times d$

$S$ and $P$ each make a write → read round trip

Fused tiled kernel · one query tile at a time

On 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

[View diagram in the original article](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#flash-memory-flow)

左右滑动查看完整示意图

上：箭头穿过 $S,P$ 时分别发生 HBM 写入与读回。下：虚线框内只容纳当前 tile 与运行状态，$S,P$ 不形成完整的 HBM 张量。示意数据流，不按容量或耗时比例绘制。

下面采用 FlashAttention-2 ([Dao, 2023](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-dao2023flashattention2)) 的 query tile 外循环与延迟归一化形式，展示数据依赖。`rowmax`、`rowsum` 沿列归约，行向量按行 broadcast；这是片上计算的伪代码，逐行写成普通 PyTorch 操作不会自动得到一个 fused kernel。

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

Algorithm 1 **FlashAttention 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**

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

交互演示 **一条长 sequence，怎样逐块算完？**

点左侧选择 Q 块，观察它怎样扫过整条 KV 序列；也可以改变 KV 块大小。

图中 24 个 token，Q 每 4 行一块；KV 每 4 或 8 行一块，分别扫描 6 或 3 次。各段依次合入同一份运行状态，扫完后得到对应的 $O_i$。大网格只标记计算范围，不保存完整 score 矩阵；本例不加 mask。

符号与边界条件

这里的 $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

[View diagram in the original article](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#flash-causal-tiles)

左右滑动查看完整示意图

每格代表 $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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-dao2022flashattention)) 用一个两级存储模型分析 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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-dao2022flashattention))：保留 $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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-dao2023flashattention2)) 主要改进并行度与工作划分：

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)}$

$Q_i^{(0)}$ · shared $K,V$

$O_i^{(0)}$

warp 1

$Q_i\,;\ K^{(1)},V^{(1)}$

$m^{(1)},\ell^{(1)},U^{(1)}$

$Q_i^{(1)}$ · shared $K,V$

$O_i^{(1)}$

merge → $O_i$

shared memory + sync

Each warp owns its output rows

[View diagram in the original article](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#flash-warp-partition)

左右滑动查看完整示意图

以两个 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](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-dao2024flashattention3)) 进一步让不同工作重叠。Hopper 提供了异步的数据搬运与矩阵乘法指令，分别是 TMA 和 WGMMA。将搬运和计算分给不同的 warp 后，可以在计算当前 tile 时预取下一块 KV，或让一组 warp 的 softmax 与另一组的矩阵乘法同时进行。

这里的关键是找出可以重叠的工作，同时遵守数据依赖：当前 tile 的 $PV$ 必须等自己的 softmax 权重准备好。流水线也需要更多 buffer 和 registers，收益要与资源占用一起衡量。FA3 还探索 FP8；降低精度带来的误差需要单独评估，不能因为使用了 FlashAttention 就认为量化没有误差。

**FlashAttention-4** ([Zadouri et al., 2026](https://pufanyi.com/blog/ml/ml-revisit/infra/flash-attention#bib-zadouri2026flashattention4)) 的论文于 2026 年 3 月发布，重点研究 Blackwell 上的 attention。以 B200 为例，Tensor Cores 的吞吐相比 Hopper 增长更快，指数运算和 shared memory 带宽却没有同比提升。于是，softmax 和片上数据搬运可能让更快的矩阵单元等待。

作者的[实现说明](https://www.together.ai/blog/flashattention-4) 围绕这些新瓶颈展开：

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 实现](https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/cute/softmax.py) 用阈值控制这种更新：最大值相对当前基准跳得足够大时，才更新基准并重缩放状态；否则保留旧基准，缩放因子取 1。关键是 $\ell,u$ 和新 tile 始终使用同一尺度，最后再归一化。只更新基准却省掉旧状态的重缩放，会破坏这个不变量。

FA4 的[官方实现](https://github.com/Dao-AILab/flash-attention/tree/main/flash_attn/cute) 使用 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](https://pufanyi.com/blog/ml/ml-revisit/infra/paged-attention) 的关系：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](https://doi.org/10.48550/arXiv.2307.08691 "https://doi.org/10.48550/arXiv.2307.08691")

Dao, T. (2024). *FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision*. [tridao.me](https://tridao.me/blog/2024/flash3/ "https://tridao.me/blog/2024/flash3/")

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](https://arxiv.org/abs/2205.14135 "https://arxiv.org/abs/2205.14135")

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](https://haoailab.com/cse291-s26/assets/scribe_notes/may26_scribe.pdf "https://haoailab.com/cse291-s26/assets/scribe_notes/may26_scribe.pdf")

Liang, P. (2026). *CS336: Language Modeling from Scratch, Lecture 6: Kernels*. [github.com](https://github.com/stanford-cs336/lectures/blob/main/lecture_06.py "https://github.com/stanford-cs336/lectures/blob/main/lecture_06.py")

Milakov, M., & Gimelshein, N. (2018). *Online Normalizer Calculation for Softmax*. [doi.org](https://doi.org/10.48550/arXiv.1805.02867 "https://doi.org/10.48550/arXiv.1805.02867")

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](https://doi.org/10.48550/arXiv.2603.05451 "https://doi.org/10.48550/arXiv.2603.05451")
