大模型推理加速(三):FlashAttention 为什么更快,又没有省掉哪些计算

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 融合方式。因此,下面两句话可以同时成立:

  1. 长序列时,中间矩阵的显存占用不再按 N² 物化;
  2. 标准全注意力的算术工作量仍然包含 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 后,可以按以下顺序排查:

  1. 确认语义相同:causal mask、padding mask、dropout、精度和 head dimension 一致;
  2. 确认实际派发:用框架日志或 profiler 检查命中的 kernel,而不是只看配置名称;
  3. 固定形状测算子:预热后分别测长 Prefill、单请求 Decode 和批量 Decode;
  4. 记录峰值显存:区分权重、KV Cache、attention 工作区和分配器保留量;
  5. 回到服务指标:在相同请求到达模式下测 TTFT、TPOT、goodput 和错误率;
  6. 验证输出:使用允许的数值误差比较 logits 或输出,单独测试极长输入和特殊 mask。

如果没有命中预期 kernel,常见原因包括当前 GPU、dtype、head dimension 或 mask 不受该实现支持;框架可能回退到另一条路径。支持矩阵会随版本变化,部署记录里应保存框架版本、attention backend、GPU 型号和实际派发证据。

最终的决策规则很简单:长 Prefill 的显存或 attention IO 是瓶颈时,FlashAttention 值得优先验证;Decode 受权重/KV 带宽、通信或调度支配时,先处理 profile 中占比最大的部分。算法名字本身不能替代瓶颈证据。

继续阅读


本站的内容生成、审核与纠错说明

发表评论