跳到正文
LLM 推理
S12批处理与调度·183

FlashAttention

注意力根本不需要把分数矩阵物化出来。分块加上滚动的最大值与求和,就把 O(N²) 的显存开销变成了 O(N)。

  • 在线 softmax
  • 分块
  • 重缩放
  • IO 感知
  • 分页版本

为什么重要

问题

教科书里的注意力会计算完整的分数矩阵 S = qkᵀ,对它做 softmax,再乘以 V。这个矩阵每个头是 N×N。上下文长度 8192、32 个头、fp16 时,把它物化出来要约 4 GB。而这还只是一个请求、一层的量。

比体积更糟的是流量。这个矩阵被写进 HBM,为了 softmax 又被读回来,再写一次,为乘以 value 再读第三次。注意力于是受制于内存往返,而不是算术;FLOPs 只占时间的一小部分。

障碍在 softmax。为了数值稳定它需要整行的最大值,为了归一化它需要整行的和。两者看起来都要求先看完所有分数,才能产出任何输出。

核心思路

解法

其实不用。softmax 可以增量地、精确地一趟算完,靠的是一个滚动最大值、一个滚动求和,以及每当最大值移动时施加的一个修正系数。

在线 softmaxpython
def flash_attention(q, k, v, tile=64):
    d = q.shape[-1]
    m = -np.inf                        # running max
    l = 0.0                            # running sum of exp
    o = np.zeros(d)                    # running weighted sum of v

    for start in range(0, len(k), tile):
        s = (q @ k[start : start + tile].T) / np.sqrt(d)
        m_new = max(m, s.max())
        alpha = np.exp(m - m_new)      # rescale everything accumulated so far
        p = np.exp(s - m_new)
        l = alpha * l + p.sum()
        o = alpha * o + p @ v[start : start + tile]
        m = m_new

    return o / l

诀窍在 alpha。此前累积的一切都是按旧的最大值归一化的;当后面某个 tile 出现更大的分数时,一次乘法就把它们全部修正过来。任何 tile 都不会被重访,结果与完整 softmax 数值等价,而非它的某种近似。

FlashAttention 并不减少 FLOPs

它做的算术与标准注意力一模一样。它去掉的是 HBM 流量:分数矩阵在 SRAM 内部被创建、消费、丢弃,而 SRAM 在 H100 上大约比 HBM 快十倍。那 2–4 倍的加速完全来自「不搬数据」,论文标题强调 IO 感知 也正是为此。

工作原理

工作原理

示意图分块,以及让它保持精确的那条递推式
真正的瓶颈:存储层级SRAM · ~33 MB~33 TB/sHBM · 80 GB~3.35 TB/s约 10× 的带宽鸿沟。标准注意力会把N×N 的分数矩阵写进 HBM,再读回来,为了 softmax 还要读两次。N=8192 时,这个矩阵每个 head 就有 128 MB。标准注意力S = qkᵀ把 S 写进 HBMsoftmax(S)写 P,再读一遍显存:O(N²)HBM 流量占主导算术从来就不是问题所在。来回搬运才是。FLASH ATTENTION把 tile 载入 SRAM算分数 + 在线 softmax累加进 o下一个 tile显存:O(N)S 从不离开 SRAM快 2–4×,结果完全相同在线 SOFTMAX 递推 —— 为什么流式是精确的,而不是近似m_new = max(m_old, max(s_tile))α = exp(m_old − m_new)ℓ_new = α·ℓ_old + Σ exp(s_tile − m_new)o_new = α·o_old + Σ exp(s_tile − m_new)·v最终输出 = o / ℓ当后面的 tile 出现更大的分数时,此前累积的每一个值都是按错误的最大值归一化的。α 用一次乘法就把它们全部修正 ——不需要回头再访问任何一个 tile。结果与完整 softmax 在比特级别可比。FlashAttention 不是近似算法。同一个递推式让你可以把一条序列切到多张 GPU 上(ring attention),或者切成块(S06)。
左边的存储层级说明了动机。底下那条递推式,就是算法的全部。

除了快,它还解锁了什么

一旦注意力变成了在 key tile 上的流式递推,另外几件事几乎就是白送的:

  • PagedAttentionS06)—— 既然你本来就在循环遍历 tile,那这些 tile 就不必连续。每个 tile 都可以是通过 block table 取来的一块。
  • Ring attention —— 把不同 tile 放到不同 GPU 上,让滚动状态 (m, ℓ, o) 沿环传递,一条序列就能比任何单张 GPU 的显存还长。
  • Flash-decoding —— batch 为 1 时只有一个 query,通常沿 query 的并行度就消失了。改成把 key 切分到多个线程块上,最后再把各自的部分状态 (m, ℓ, o) 合并。同一条递推式,只是沿处理器而不是沿时间展开。
显存,而不是 O(N²)
O(N)
长上下文不再是显存问题
墙钟加速
2–4×
完全来自避免了 HBM 往返
不是近似
精确
与完整 softmax 的差异仅来自浮点结合律

它如何与你已经搭好的部分组合

一个生产级 decode kernel,就是 FlashAttention 的递推式 + PagedAttention 的间接寻址 + GQA 的头共享,全都在一个循环里:对 block table 中的每一块,加载这个 query 组共享的 KV 头,算分数,更新 (m, ℓ, o)。这一个循环就是 FlashInfer 和 vLLM 的 kernel 所做的大部分事情。

动手观察

动手试试

真实的数字、真实的递推。单步走过每个 tile,看滚动最大值、修正系数和输出累加器如何收敛到精确答案,那个答案用绿色刻度标出。

模拟器在线 softmax,逐块推进
第 0 / 4 步
分数 qᵀk/√d · 16 个 key,tile = 4
0.860.451.450.10
  • 尚未载入 —— 永远不会同时在内存中
  • 当前 tile(位于 SRAM)
  • 已处理并丢弃

任一时刻只有橙色 tile 是存在的。灰色格子从未被物化,显存的节省全部来自这里。FlashAttention 做的算术和朴素实现一模一样,所以它被称为「IO 感知」,而不是「更快的算法」。

在线 softmax 的递推式
tilem(滚动最大值)修正系数 e^(m_old−m_new)ℓ(滚动求和)
点击运行

当某个 tile 里出现了比之前都大的分数,滚动最大值就会移动,此前累积的所有值都要乘以 e^(m_old−m_new) 重新缩放,也就是那个小于 1 的橙色修正系数。正是这一次乘法,让流式 softmax 从近似变成了精确。

输出累加器 o / ℓ
o[0]0.952
o[1]-0.048
o[2]0.220
o[3]0.021

绿色刻度是真正的注意力输出。在最后一个 tile 之前,累加器都是错的:它是「到目前为止见过的 key」上的合法 softmax,而不是全部 key 上的。这没问题,但也说明 flash attention 的循环为什么永远不能提前退出。

内存中的分数
4 of 16
朴素实现需要保存
256
重缩放次数
0
与精确解的最大误差
跑完才能对比

注意累加器在每一个中间步骤都是错的:它是「到目前为止见过的 key」上的正确 softmax,而那并不是你要的东西。只有在最后一个 tile 之后它才对得上。再把 tile 大小设成 1。算法依然成立,做的算术也一样,只是几乎每一步都要做一次修正。tile 大小决定的只是「SRAM 里能装下多少」。

亲手实现

实现

你会上线的是分页版本。它就是同一个循环,只是在访问 key 之前多了一次 block table 查表。

code/s12_flash_attention.py(节选)python
def paged_flash_attention(q, block_table, k_cache, v_cache, seq_len, block_size):
    d = q.shape[-1]
    m, l, o = -np.inf, 0.0, np.zeros(d)

    for logical, physical in enumerate(block_table):
        n = min(block_size, seq_len - logical * block_size)
        if n <= 0:
            break
        k_blk = k_cache[physical, :n]      # S06's indirection...
        v_blk = v_cache[physical, :n]
        s = (q @ k_blk.T) / np.sqrt(d)     # ...inside S12's recurrence

        m_new = max(m, s.max())
        alpha = np.exp(m - m_new)
        p = np.exp(s - m_new)
        l, o, m = alpha * l + p.sum(), alpha * o + p @ v_blk, m_new

    return o / l

exp(-inf − -inf) 是 NaN

第一个 tile 时滚动最大值是 -inf,所以 alpha = exp(-inf - m_new) 求值为 0,这是对的。但如果第一个 tile 被完全掩掉,m_new 也是 -inf,你就得到了 exp(nan)。所有真实 kernel 要么对第一次迭代做特判,要么把最大值初始化成一个很大的有限负数。在不规则 batch 的因果掩码下,「整块被掩掉」是家常便饭。
在本地运行
实现朴素注意力、分块 flash 注意力、分页 flash 注意力和 flash-decoding,然后在随机输入、多种 tile 大小和序列长度上断言四者在浮点容差内一致,并报告各自的峰值中间显存。
$ python code/s12_flash_attention.py
预期输出: 一致性在约 1e-15 以内、朴素实现的峰值显存呈 O(N²) 而 flash 呈 O(tile),以及一段「整块被掩掉导致 NaN」的演示和防止它的那道守卫。

只需要 NumPy — 查看环境准备.

生产实践

在生产环境中

  • FlashAttention-2 / -3 —— 更好的 warp 间工作划分,v3 还加上了异步和 Hopper 上的 FP8 支持。
  • FlashInfer —— vLLM 和 SGLang 使用的 kernel 库;把分页、GQA 和 flash 合进同一个实现。
  • picoLM —— 其 C 实现把注意力融合成单次在线 softmax 遍历,因此从头到尾不写出分数数组。收益在于对 KV 只扫一遍,缓存行为也更好。这是本章的递推在真实世界里的一次现身,而且是在 CPU 上。
  • PyTorch —— 当 shape 和 dtype 允许时,scaled_dot_product_attention 会自动派发到 flash kernel。

练习

  1. 1
    实现 flash-decoding:把 key 切给四个「worker」,各自独立跑一遍递推,然后把它们的 (m, ℓ, o) 三元组合并成一个。合并规则就是你已经写过的那次重缩放。
  2. 2
    给分块版本加上因果掩码,验证输出与朴素因果注意力一致。然后刻意让某个 tile 被完全掩掉,确认你能复现那个 NaN,再去修它。
  3. 3
    在 N = 512、2048、8192 上测量朴素与 flash 的峰值显存,实证确认「二次方 vs 常数」的缩放关系。

继续学习

接下来

kernel 已经高效了。可是在小 batch 下,GPU 大部分时间都在等 CPU 告诉它下一步做什么。S13 用算子融合和 CUDA Graph 开启解码加速这一层。

自测

习题

先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。

自测 1 题 / 共 6

FlashAttention 实际减少的是什么?

得分 0/6