Jellow 编辑发布,AI 辅助生成并完成技术核对。 计算示例和实测结果会在文中分别说明。
“FlashAttention 把注意力从平方复杂度降成了线性复杂度”是一个很常见的说法,也很容易让人误判优化效果。
标准的精确注意力确实包含与序列长度平方相关的工作。FlashAttention 没有删掉这些乘加运算,也不是用近似结果换速度。它改变了计算的组织方式:分块读取 Q、K、V,在更靠近计算单元的存储中完成中间步骤,避免反复把庞大的注意力矩阵写回显存。
本文先从一个放不下的中间矩阵说起,再用几十行 Python 验证分块 softmax 的核心状态。代码只用于解释算法,不是 GPU kernel,也不能用来测试 FlashAttention 的速度。
1. 真正麻烦的是中间矩阵
单个注意力头可以写成:
S = QKᵀ / √d
P = softmax(S)
O = PV
序列长度为 N 时,S 和 P 都有 N × N 个元素。假设 N = 8,192、每个元素 2 字节,单个矩阵在单个头上就有:
8192 × 8192 × 2 bytes = 128 MiB
32 个头对应 4 GiB。这只是一个中间张量的逻辑大小,没有计入 batch、反向传播或其他工作区,也不代表某个框架一定会同时保留全部张量。
矩阵大只是表面现象。GPU 计算单元附近的片上存储容量小但带宽高,显存容量大但数据搬运更贵。如果把 S 写到显存,随后为了 softmax 读回,再把 P 写出、读回,很多时间会花在搬运中间结果上。
FlashAttention 论文把注意力描述为 IO-aware 的精确算法,重点正是减少不同存储层级之间的数据读写。
2. 分块之后,softmax 怎么保持正确
不能简单地对每个块单独做 softmax,再把结果相加。因为一个位置的 softmax 分母包含这一行里的所有 key:
softmax(sᵢ) = exp(sᵢ) / Σⱼ exp(sⱼ)
稳定实现还需要先减去全局最大值。如果后面的块出现更大的分数,前面已经累加的结果必须重新缩放。
处理一段分数时,只需保存三个状态:
m:目前看到的最大分数;l:以m为基准的指数和;u:同一基准下,指数权重与 value 的加权和。
读到新分数 s 和对应的 v 后,令 m_new = max(m, s):
l_new = exp(m - m_new) × l + exp(s - m_new)
u_new = exp(m - m_new) × u + exp(s - m_new) × v
最后的输出是 u / l。最大值变化时,旧状态与新项被换算到同一基准,因此无需保存全部分数。这个思路与 online normalizer 的推导一致。
下面用标量 value 展示。真实 attention 的 value 是向量,状态 u 也相应变成向量。
from math import exp
def stable_attention(scores, values):
m = max(scores)
weights = [exp(s - m) for s in scores]
return sum(w * v for w, v in zip(weights, values)) / sum(weights)
def online_attention(scores, values):
m = float("-inf")
l = 0.0
u = 0.0
for s, v in zip(scores, values):
new_m = max(m, s)
old_scale = 0.0 if m == float("-inf") else exp(m - new_m)
new_scale = exp(s - new_m)
l = old_scale * l + new_scale
u = old_scale * u + new_scale * v
m = new_m
return u / l
scores = [1000.0, 999.0, 1002.0, 998.0]
values = [2.0, -1.0, 4.0, 3.0]
a = stable_attention(scores, values)
b = online_attention(scores, values)
print(a, b, abs(a - b))
assert abs(a - b) < 1e-12
这里故意使用接近 1,000 的分数。直接计算 exp(1000) 会溢出,减去当前最大值后则可以稳定计算。
3. 从逐元素状态到二维分块
FlashAttention 处理的不是一个分数列表,而是 Q 与 K、V 的二维块。对某一组 query,kernel 依次载入若干 K、V 块,算出当前分数块,并更新每一行的最大值、指数和与输出累加值。
假设 A、B 两个块已经分别得到状态 (mA, lA, uA) 和 (mB, lB, uB)。它们也可以合并:
m = max(mA, mB)
l = exp(mA - m) × lA + exp(mB - m) × lB
u = exp(mA - m) × uA + exp(mB - m) × uB
这说明算法可以在不物化完整 S、P 的情况下得到同一个数学结果。浮点运算顺序不同,输出不保证逐比特相同;这里的“精确”指它没有把注意力机制改成近似模型。
反向传播还可以利用保存的归一化统计量重算部分中间值,以计算换显存。在纯推理场景里,关注点主要是前向计算,但不要把训练阶段的显存数字直接套到推理服务上。
4. 它省的是 IO,不是 QKᵀ 的平方工作量
对标准全注意力,计算 QKᵀ 和 PV 的乘加量仍随 N² 增长。FlashAttention 改善的是访存复杂度和 kernel 融合方式。因此,下面两句话可以同时成立:
- 长序列时,中间矩阵的显存占用不再按 N² 物化;
- 标准全注意力的算术工作量仍然包含 N² 项。
这一区分会影响容量规划。采用 FlashAttention 后可以把更长输入放入显存,不代表序列长度翻倍后 Prefill 用时只翻倍。实际增幅还取决于硬件、维度、mask、精度和 kernel 是否覆盖当前形状。
FlashAttention-2进一步减少非矩阵乘运算,并调整 thread block 与 warp 之间的工作划分。作者的开源实现列出了支持的硬件、数据类型和 head dimension;部署时应以所用版本的支持范围为准。
5. 为什么服务总耗时不会等比例下降
假设一次请求中,attention 占原始耗时的 30%,其余 70% 来自线性层、通信、采样和调度。即使 attention 恰好加速两倍,整体加速比也只有:
1 / (0.70 + 0.30 / 2) ≈ 1.18
这是 Amdahl 定律下的计算示例,不是任何实现的实测结果。长 Prefill 中 attention 的占比可能更高,单 token Decode 中读取权重、KV Cache 和调度也可能成为主因。
因此验证 FlashAttention 时,至少分开看两层数据:
- kernel 或算子层:相同形状、精度和 mask 下的耗时与显存;
- 服务层:固定输入输出长度和请求率下的 TTFT、TPOT、吞吐与失败率。
如果算子基准明显变快,而服务指标变化很小,先看 attention 原本占总耗时多少,再排查是否走到了预期 kernel。还要把 Prefill 与 Decode 分开,因为两者交给 attention kernel 的形状并不相同。
6. Prefill 和 Decode 不能套用同一个收益预期
长 Prefill 中有很多 query 位置,完整分数矩阵会很大。分块避免物化中间矩阵,通常能同时改善峰值显存和数据搬运。因果 mask 还允许跳过整个位于对角线上方的块,但对角线下方的有效注意力工作仍随序列长度平方增长。
单条序列 Decode 时,每一步往往只有一个新 query。此时分数形状接近 1 × 历史长度,本来就没有 N × N 的新分数矩阵要保存。算子仍要读取历史 K、V,FlashAttention 类 kernel 可以改善融合和访问方式,但收益来源与长 Prefill 不同。
大 batch Decode 又会改变形状:query 长度仍短,但同时处理的序列多,历史长度还可能各不相同。因此,下面三个 benchmark 不能互相替代:
| 场景 | query 长度 | KV 长度 | 主要想回答的问题 |
|---|---|---|---|
| 长 Prefill | 8,192 | 8,192 | 是否减少中间矩阵 IO 和显存 |
| 单请求 Decode | 1 | 8,192 | 读取历史 KV 与 kernel 启动成本 |
| 批量 Decode | 1 × 多序列 | 长短混合 | 调度、padding/分页与 KV 访问效率 |
如果只用方形的 Q=K=8192 测出漂亮结果,再宣称聊天 Decode 也有相同倍数,比较对象就错了。
7. 在线状态也可以按块合并
前面的逐元素代码展示了状态更新。下面再把同一组数据切成不同块,验证块边界不会改变结果:
from math import exp
def summarize(scores, values):
m = max(scores)
weights = [exp(s - m) for s in scores]
return m, sum(weights), sum(w * v for w, v in zip(weights, values))
def merge(a, b):
ma, la, ua = a
mb, lb, ub = b
m = max(ma, mb)
return (
m,
exp(ma - m) * la + exp(mb - m) * lb,
exp(ma - m) * ua + exp(mb - m) * ub,
)
scores = [1000.0, 999.0, 1002.0, 998.0]
values = [2.0, -1.0, 4.0, 3.0]
left = summarize(scores[:2], values[:2])
right = summarize(scores[2:], values[2:])
m, denominator, numerator = merge(left, right)
blocked = numerator / denominator
whole = summarize(scores, values)
direct = whole[2] / whole[1]
print(blocked, direct)
assert abs(blocked - direct) < 1e-12
真实 kernel 会对每个 query 行保存一组状态,并让 u 成为一个 value 向量。二维 tiling、shared memory 容量、寄存器压力和并行划分决定块应该多大;上面的 Python 没有模拟这些硬件细节,但验证了分块归一化的数学接口。
8. 一套能解释结果的验证顺序
启用某个 attention backend 后,可以按以下顺序排查:
- 确认语义相同:causal mask、padding mask、dropout、精度和 head dimension 一致;
- 确认实际派发:用框架日志或 profiler 检查命中的 kernel,而不是只看配置名称;
- 固定形状测算子:预热后分别测长 Prefill、单请求 Decode 和批量 Decode;
- 记录峰值显存:区分权重、KV Cache、attention 工作区和分配器保留量;
- 回到服务指标:在相同请求到达模式下测 TTFT、TPOT、goodput 和错误率;
- 验证输出:使用允许的数值误差比较 logits 或输出,单独测试极长输入和特殊 mask。
如果没有命中预期 kernel,常见原因包括当前 GPU、dtype、head dimension 或 mask 不受该实现支持;框架可能回退到另一条路径。支持矩阵会随版本变化,部署记录里应保存框架版本、attention backend、GPU 型号和实际派发证据。
最终的决策规则很简单:长 Prefill 的显存或 attention IO 是瓶颈时,FlashAttention 值得优先验证;Decode 受权重/KV 带宽、通信或调度支配时,先处理 profile 中占比最大的部分。算法名字本身不能替代瓶颈证据。