返回博客

FlashAttention 分块之后,如何保留全局 softmax?

从块内归一化的反例推导最大值、分母与未归一化加权和的合并,检验掩码和空行,并区分反向重算、HBM 访问、prefill 与 decode 的真实边界。

把注意力矩阵切成小块,似乎只是一个内存优化。但 softmax 的分母要看整行:当前块还不知道后面会出现多大的分数,怎么能够先算?如果每块各自归一化再拼接,结果又为什么不对?这才是理解 FlashAttention 的入口。

本文深入的是 FlashAttention(2022) 与 FlashAttention-2(2023) 的基础机制,不把它们写成新发布事件。我们从一个查询行出发,推导可合并状态,再落到掩码、训练重算和成本。本文实际执行了可复跑的纯 Python CPU 前向数值检验;没有运行 FlashAttention GPU 内核、模型训练或延迟基准。

1. 先固定算子:省内存之前,不能偷换注意力

先看一个 batch 的一个注意力头。查询、键和值分别是下面三个二维张量。\(d_k\) 是 query/key 维度,\(d_v\) 是 value 维度;两者不必相等。\(J_i\) 是第 \(i\) 行允许访问的 key 集合,\(b_{ij}\) 是预先固定、与本文求导变量无关的可选偏置。合法位置的 logits 均假设为有限实数,先关闭 dropout。

\[\begin{aligned}Q&\in\mathbb R^{N_q\times d_k},\quad K\in\mathbb R^{N_k\times d_k},\quad V\in\mathbb R^{N_k\times d_v},\\s_{ij}&=\frac{q_i^\top k_j}{\sqrt{d_k}}+b_{ij},\\p_{ij}&=\frac{\exp(s_{ij})}{\sum_{t\in J_i}\exp(s_{it})}\quad(j\in J_i),\qquad o_i=\sum_{j\in J_i}p_{ij}v_j.\end{aligned}\]

所有非空行的权重和为 1,屏蔽位置的权重为 0。公式中的比例 \(1/\sqrt{d_k}\) 不随 tile 大小改变;同样,换实现时也要保持 Q/K/V、位置处理和偏置完全一致。本文是在同一个算子上换执行顺序,而不是重新训练一种注意力结构。

最直接的实现先把 \(S\) 写入显存,逐行 softmax 得到 \(P\),再算 \(PV\)。当 \(N_q=N_k=N\) 时,两个中间矩阵都有 \(N^2\) 个元素。显存不仅要容纳它们,还要承受跨算子读取与写入。FlashAttention 的问题是:能否只在小块上生成分数,消耗它们之后就丢弃,同时保留整行归一化的结果?

2. 为什么“块内 softmax 后相加”会改变答案?

取一个查询、两个 key,并让每个块只有一个 key。两个 logits 为 \((0,\log 3)\),对应标量 value 为 \((0,1)\)。全局 softmax 权重是 \((1/4,3/4)\),输出为 \(3/4\)。每个单元素块的局部权重却都为 1,块输出分别为 0、1;相加得到 1,平均得到 \(1/2\)。两种做法都丢失了不同块在全局分母中的质量。

这不是一个“小误差”。把两个块的 logit 差拉大,正确输出会进一步靠近某一个 value,而局部输出依旧不知道另一个块有多强。即使改成按照块长度平均也无济于事:长度相同并不意味着指数质量相同。

所以 tile 不能只交出一个已归一化输出。它必须同时交出足够的信息,让下一块能够恢复“这个输出在全局分母中应该占多少”。数值稳定性又要求这些信息不直接使用可能溢出的 \(\exp(s_j)\)。

3. 一个块保留什么:最大值、分母、未归一化加权和

固定一个查询行,暂时省略行下标。对非空的合法 key 子集 \(A\),定义三个状态:标量最大值 \(m_A\)、标量归一化量 \(\ell_A\)、长度为 \(d_v\) 的向量 \(u_A\)。

\[\begin{aligned}m_A&=\max_{j\in A}s_j,\\\ell_A&=\sum_{j\in A}\exp(s_j-m_A),\\u_A&=\sum_{j\in A}\exp(s_j-m_A)v_j,\qquad o_A=\frac{u_A}{\ell_A}.\end{aligned}\]

这里的 \(u_A\) 不是最终 attention 输出:它仍带着本块的指数质量。最大值是数值坐标的原点,分母说明这个坐标下有多少权重,加权和说明这些权重携带什么 value。\(m_A\) 一旦改变,后两项都必须同步变换。

对 \(n_A\) 个有限有效 logits,\(1\le\ell_A\le n_A\),因为至少一个指数为 1,其余位于 0 到 1 之间。这能防止指数正向溢出,却不是对全部浮点错误的承诺:很小的项仍可能下溢,带正负 value 的加权和仍可能相消,QK 点积本身也可能溢出。CPU 检验使用有限输入,真实低精度内核还要另查累加精度。

Online normalizer calculation for softmax(2018) 已给出最大值与归一化量的在线更新。注意力多了一项向量加权和,但合并时使用同一个坐标变换。下面的图展示这个关系,完整公式保留在正文中。

原创前向机制示意,非性能测量。对一个查询行,将合法键分成两个不相交块;每块保留最大分数 m、相对该最大值计算的指数权重和 l,以及相同未归一化权重下的值向量加权和 u。块 A 和 B 各有自己的三项状态。合并时选择两个块最大值中较大者作为共享的新最大值,再由各块旧最大值与新最大值之差计算指数缩放系数。同一块的 l 和 u 使用相同系数缩放,分别相加得到合并状态,最后将合并 u 除以合并 l。不能直接相加两块已归一化的输出。无合法键的空块跳过;若整行都没有合法键,输出需由实现单独约定。完整注意力仍计算所有合法 query-key 配对,分块省去完整二次规模分数或概率中间矩阵的 HBM 读写,不减少 dense attention 的合法配对数量。
原创机制图:单行合法键分为两块,各自保留最大分数、相对指数权重和及未归一化的值向量加权和。合并时先统一最大值,再对各块分母与分子使用相同的指数缩放,分别求和后归一化。不能直接相加块内归一化输出。空块跳过,全空行另定策略。该重排避免完整注意力中间矩阵落入 HBM,不减少完整注意力的合法配对数量;图不是性能测量。

4. 合并不是拼输出,而是统一指数坐标

假设 \(A\)、\(B\) 不相交、都非空,并且合起来正好覆盖我们当前要处理的合法 key。将两块的最大值统一为 \(m\),定义两个缩放因子,再分别合并分母与向量。

\[\begin{aligned}m&=\max(m_A,m_B),\\\alpha&=\exp(m_A-m),\qquad\beta=\exp(m_B-m),\\\ell&=\alpha\ell_A+\beta\ell_B,\\u&=\alpha u_A+\beta u_B,\qquad o=\frac{u}{\ell}.\end{aligned}\]

为什么正确?只要展开旧块的加权和,两个指数中的 \(m_A\) 便消去;分母同理。

\[\begin{aligned}\alpha u_A&=\sum_{j\in A}\exp(m_A-m)\exp(s_j-m_A)v_j\\&=\sum_{j\in A}\exp(s_j-m)v_j,\\\alpha\ell_A&=\sum_{j\in A}\exp(s_j-m).\end{aligned}\]

新块也得到关于同一个 \(m\) 的表达。因此 \(u\) 和 \(\ell\) 就是并集上指数加权和与分母,最后一次相除得到全局输出。按这个不变量继续处理任意多块,结果对应全部合法 key,而不是任意某个块的局部 softmax。

在精确实数算术中,这个合并对不相交子集具有结合性,也允许交换处理顺序;因为任意合并树都代表同一集合。它不能为重复访问 key 开脱:一个 key 被纳入两次,相当于在分母和加权和中重复计数。有限精度中的求和顺序仍会改变舍入误差,所以“精确注意力”表示没有稀疏化或低秩近似,不表示所有后端逐比特相同。

还有一个更隐蔽的错误:更新分母时重标定,更新 \(u\) 时却忘了。令 logits 仍为 \((0,\log 3)\),value 改为 \((1,0)\)。正确输出为 \(1/4\)。新最大值产生 \(\alpha=1/3\),旧分母和旧加权和都应乘它;如果只重标定分母,输出会变成 \(3/4\)。分母数值看起来正常,并不代表答案正确。

5. 从一行递推到张量 tile:最后再除一次

实际计算通常一次处理 \(B_q\) 个查询与 \(B_k\) 个 key。当前分数 tile 形状为 \(B_q\times B_k\),每个查询维护独立的 \(m,\ell,u\)。矩阵乘法负责生成分数和 value 加权和,行归约负责更新统计;旧状态的缩放沿着行广播,不能误用整块共享的一个最大值。

将当前块的指数直接写在新最大值坐标下,就得到便于实现的顺序更新。此处 \(B\) 只包含当前行在这个 tile 内的合法 key。

\[\begin{aligned}m'&=\max\!\left(m,\max_{j\in B}s_j\right),\qquad a=\exp(m-m'),\\\ell'&=a\ell+\sum_{j\in B}\exp(s_j-m'),\\u'&=a u+\sum_{j\in B}\exp(s_j-m')v_j.\end{aligned}\]

非空行首次更新时,旧状态视为零质量,不需要真正计算空状态的指数。之后每一步保持同一个不变量。循环结束才做 \(o=u/\ell\),避免每个 tile 都对整个输出向量做一次除法再在下一步撤销归一化。FlashAttention-2 第 3.1 节 正是将未归一化累积量留到循环末尾;它还重新安排查询块并行与 warp 内工作分配,减少非矩阵乘运算和共享内存通信。递推成立与 GPU 调度高效,是两个分别需要验证的层面。

for each query tile:
    initialize per-row m, l, u as empty states
    for each key/value tile:
        compute scaled scores and the global-position mask
        for each row with valid keys in this tile:
            choose the new maximum
            rescale BOTH the old l and old u
            add this tile's mass and weighted values
    normalize each nonempty row once; handle empty rows explicitly

这是算法骨架,不是 CUDA 内核。把相同循环写成若干 Python tensor 操作,仍可能把中间 tile 写回显存、反复启动内核,或让自动微分保存所有 tile。数学上在线,不代表系统已经达到 FlashAttention 的内存访问行为。块大小也不能无限增大:寄存器、共享内存和同时驻留的工作块都会限制可用并行度。

6. 掩码要跟全局位置走,空块要单独处理

方形自注意力中,因果条件是 key 的全局位置不晚于 query 的全局位置。一个查询 tile 与一个 key tile 的局部行列号不能直接比较:它们可能来自完全不同的位置。为了效率跳过完全屏蔽的 tile,可以;把每个 tile 都当作从位置零开始,则改变了可见历史。

\[j_{\mathrm{global}}\le i_{\mathrm{global}},\qquad j\le i+N_k-N_q\quad\text{(bottom-right alignment)}.\]

矩形右下对齐的第二个条件是特定接口约定,不是所有框架统一规则。这里 \(i,j\) 是从零开始的矩阵下标,它把两个序列的末端位置对齐;查询序列不长于键序列时,可将 query 视为 key 序列最后的 \(N_q\) 个位置。官方仓库固定提交的 README 记录:FlashAttention 2.1 将长度不等时的 causal 掩码改成右下对齐。对 \(N_q=2,N_k=5\),两行分别允许前 4 与前 5 个 key;左上对齐只允许前 1 与前 2 个,输出当然会不同。

如果 \(N_q>N_k\),右下对齐会产生全空的前部行。更一般地,padding 也可能屏蔽整行。softmax 在空集合上没有上述分母定义;把全部位置写成负无穷,再机械地减去负无穷,会产生 NaN。本文代码明确采用“全空行输出零”的约定,官方上述接口也记录了零输出,但不能外推为所有框架、所有损失或所有后端的规则。

一个块对某行没有合法 key 时,直接跳过这行的更新;若先前已有有效状态,就完整保留它。空状态作为合并的单位元必须由分支处理,不能求 \(\exp(-\infty-(-\infty))\)。本文关闭 dropout;若训练中启用它,前向与反向重算必须复现同一随机掩码及缩放。位置旋转、滑窗、偏置和 padding 都需要同样的对齐审计,不能在 tile 合并后再补一个大概相同的掩码。

7. 训练为什么可以不保存整个概率矩阵?

前向不保存 \(P\),训练时是否就没有梯度了?关键是保存每个非空行的 log-sum-exp \(L_i\) 与输出。给定原来的 Q/K 和固定掩码,反向按 tile 重建分数,再重建概率。对损失 \(\mathcal L\),令上游梯度的第 \(i\) 行为 \(g_i\),组成矩阵 \(G\);定义 \(D\) 为对 scaled logits 的梯度。

\[\begin{aligned}L_i&=m_i+\log\ell_i,\qquad p_{ij}=\exp(s_{ij}-L_i),\\g_i&=\frac{\partial\mathcal L}{\partial o_i},\qquad c_i=g_i^\top o_i,\\D_{ij}&=p_{ij}(g_i^\top v_j-c_i),\\\nabla_Q\mathcal L&=\frac{DK}{\sqrt{d_k}},\qquad\nabla_K\mathcal L=\frac{D^\top Q}{\sqrt{d_k}},\qquad\nabla_V\mathcal L=P^\top G.\end{aligned}\]

中间标量 \(c_i\) 的来源是 softmax Jacobian。它本来写作 \(\sum_j p_{ij}(g_i^\top v_j)\),交换求和与点积就得到 \(g_i^\top o_i\)。于是不用保存一整行概率,也能得到 Jacobian 所需的行归约。每个 tile 重建 \(P,D\) 后累加 Q/K/V 梯度,计算结束便可以丢弃 tile。注意 Q/K 梯度仍要保留原来的缩放因子。

上述推导只覆盖非空行、关闭 dropout、固定掩码与固定偏置;若偏置可训练,它还需要自己的梯度,若 Q/K 来自 RoPE 或投影层,还需继续应用相应链式法则。本文没有执行自动微分或低精度反向检验,也不把自定义空行零输出自动解释为标准 softmax 在空集合上的导数。FA1 附录 B 给出了含缩放、掩码和 dropout 的完整训练算法。

重算增加了一些算术,却减少保存和搬运大矩阵。这解释了为什么计算量增加与墙钟时间下降可以同时出现;是否真的下降,依赖设备与工作负载。若一个实现仍在反向中保存每个概率 tile,或与前向使用不同的掩码/随机状态,它就没有实现这里承诺的路径。

8. 省的是哪些字节,仍要算哪些配对?

对单头 dense attention,主导算术量仍是 \(O(N_qN_k(d_k+d_v))\)。在方形、同维度时为 \(O(N^2d)\);因果结构减少约一半合法配对,却不改变这个阶。在线 softmax 没有让所有 key 对 query 的影响变成线性计算。它主要免去 \(N_qN_k\) 分数/概率中间量在大显存中的完整物化。

一种直接 tile 工作区的元素数量可写为下式:分数 tile、查询、键值和输出累积量。常数还会受双缓冲、寄存器、数据类型与具体内核影响,不能据此宣称整张卡的峰值显存。

\[O\!\left(B_qB_k+B_qd_k+B_k(d_k+d_v)+B_qd_v\right)\]

此外要计算完整输入与输出存储,以及覆盖所有查询行的总计 \(O(N_q)\) 个标量统计量;训练还有梯度、优化器、其他层激活。说“线性额外状态”时,必须指出相对于输入/输出之外,避免把它写成完整模型只需要几个标量。

FA1 定理 2 在理想化存储层级中比较 HBM 访问。它针对方形同维度问题,\(M\) 是快速存储可容纳的数值元素数,不是未经换算的字节数。在给定范围内:

\[\begin{aligned}d\le M\le Nd:\qquad\mathrm{IO}_{\mathrm{materialized}}&=\Theta(Nd+N^2),\\\mathrm{IO}_{\mathrm{FA1}}&=\Theta\!\left(\frac{N^2d^2}{M}\right).\end{aligned}\]

固定 \(M,d\) 时,这个 IO 表达仍对 \(N\) 二次增长;它描述减少访问次数的系数与存储容量的关系。若把 \(M\) 当作随 \(N\) 增长的量,才会得到其他缩放情景。这个定理也不等于任意现代内核、任意显卡上的实测带宽模型,更不保证任意序列长度都有同一速度收益。

prefill 时通常 \(N_q\approx N_k=N\),适合通过 tile 复用 Q/K/V,避免大中间矩阵。单 token decode 则是 \(N_q=1,N_k=L\):本来只有一行概率,主导工作常是读入长 KV cache。该步配对量为 \(O(L(d_k+d_v))\),保留的 cache 仍为 \(O(L(d_k+d_v))\)。官方 README 的 2.2 说明了面向短 query 的 KV 加载拆分。可合并状态使拆分后的归约有数学依据,额外内核与同步开销仍需测量。

这也解释了它与 MLA 缓存压缩 的关系:在线归约改变 attention 中间量的执行路径,不自动压缩历史 KV 表示。结合两种机制时,要同时保留各自的张量与缩放约定。FA4 的 2026-03-05 原始报告 进一步讨论硬件算力增长不对称带来的指数计算、共享内存与管线瓶颈;它提醒我们,消除一种瓶颈之后,下一种瓶颈不会自动消失。

9. 可检验的边界:本文实际测了什么?

下载完整 CPU 检验代码。脚本只用 Python 标准库,固定随机种子;同一 Q/K/V、缩放与显式布尔掩码(关闭额外偏置)分别交给 dense 和 online 前向路径。dense 路径保存分数以作参考,online 路径按块计算并合并;输入、参考矩阵和测试报告本身的内存不计作生产 GPU 内核的峰值。

实际运行环境为 CPython 3.12.14 / binary64 (53 位有效精度)。逐元素判定采用 \(|x-y|\le 2\times10^{-12}+2\times10^{-12}\max(|x|,|y|)\);这是本文小规模检验的容差,不是所有设备和 dtype 的误差保证。28 个案例共 59 次对照全部通过。下面形状列依次为 query 数、key 数、key 维度、value 维度,误差是该检验组内与同输入 dense 参考的最大绝对差;没有给出耗时或 GPU 显存数字。

检验组Q/K/V 形状参数最大绝对误差
非因果、不规则分块5 / 9 / 4 / 32.22045e-16
方形全局因果掩码7 / 7 / 5 / 22.22045e-16
矩形右下对齐与空前缀行2 / 6 / 1 / 1
5 / 2 / 1 / 1
0
padding、空块与全空行2 / 4 / 3 / 2
4 / 7 / 3 / 4
2.22045e-16
空 key 轴2 / 0 / 3 / 20
公共 logit 平移与大跨度1 / 5 / 1 / 1
1 / 5 / 1 / 2
4.44089e-16
全部六种三块顺序3 / 6 / 4 / 32.22045e-16
16 组固定种子形状16 组形状见代码4.44089e-16

检验还验证 key 分区完整且不重复,允许空块、不规则块和改变遍历顺序;针对 padding、全空行以及中间无有效 key 的块检查零输出或旧状态保留。极端有限 logit 案例检验减去最大值后的归一化,不能证明对任意溢出的 QK 点积都安全。两个错误实现也实际产生了第 2、4 节的反例。这些是算法前向的教育验证,不是 FlashAttention 内核测试、梯度检验、训练质量评估或服务加速结果。

真正部署前,应分三层验收。第一层固定权重与输入,检查输出、掩码/位置、dtype 和梯度,覆盖不等长、全空行、padding、dropout 及极端 logits。第二层在同一 GPU、后端版本、batch/head、长度、dtype 与 causal 设置下,预热并同步,分别测 attention 内核、prefill 与 decode、峰值显存与失败率。第三层才测试端到端延迟、吞吐和训练/任务效果,明确计算预算与基线。一个默认已选择融合后端的 API,不应被错误地标成“未经优化的 dense 基线”。这些 GPU 与模型层实验是建议方案,本文没有执行。

如果 tile 大小一改,输出就显著变化,应先追查缩放、行广播、掩码和状态重标定;如果数值通过但没有加速,应检查后端是否真正融合、tile 是否落回 HBM、并行度、cache 布局与同步开销。前者否定实现正确性,后者否定该工作负载下的加速判断。只有将两者分开,才能判断下一步应该修数学还是修系统。

这套方法的可迁移之处是:先找可以合并的状态,再把每一步的输入、输出、不变量与失效条件写清楚。FlashAttention 保留的是全局注意力的数学定义,改变的是中间量的存放与流动方式。省去大矩阵物化,不意味着省去全部配对,也不意味着缓存、梯度和交付成本已经得到验证。

参考文献元数据使用 citation-management 工具核查;软件署名为 Scientific Agent Skills,它不是本文注意力机制或性能判断的证据。

参考资料

  1. Dao, Fu, Ermon, Rudra and Ré — FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (NeurIPS 2022; v2 June 23, 2022) · 2022-05-27 · 查阅 2026-10-08
  2. Tri Dao — FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning (first preprint July 17, 2023; author PDF July 18; ICLR 2024) · 2023-07-17 · 查阅 2026-10-08
  3. Milakov and Gimelshein — Online normalizer calculation for softmax (first preprint May 8, 2018; v1) · 2018-05-08 · 查阅 2026-10-08
  4. Dao-AILab — FlashAttention README, pinned commit 94e22c906678e5483fa0e9e24d8e787bc2c0ed4c; causal alignment and decoding interfaces · 查阅 2026-10-08
  5. Zadouri et al. — FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling (v1) · 2026-03-05 · 查阅 2026-10-08
  6. Kassis, Agarwal, He, Patel and Brueckner — Scientific Agent Skills: A Library of Procedural Knowledge for Research Agents (reference-metadata tooling) · 2026-08-30 · 查阅 2026-10-08
利友诚

关于作者

利友诚 · Youcheng Li

北京大学智能学院人工智能专业博士研究生,导师为王立威教授;Isoplex Intelligence(壹索智能)联合创始人兼 CTO。

研究关注医疗人工智能、生成式基础模型、诊断推理与科学智能体。以第一作者或共同第一作者身份在 Nature Biomedical Engineering、Scientific Data、KDD 和 PLOS Computational Biology 发表研究。