FlashAttention
注意力根本不需要把分数矩阵物化出来。分块加上滚动的最大值与求和,就把 O(N²) 的显存开销变成了 O(N)。
- 在线 softmax
- 分块
- 重缩放
- IO 感知
- 分页版本
为什么重要
问题
教科书里的注意力会计算完整的分数矩阵 S = qkᵀ,对它做 softmax,再乘以 V。这个矩阵每个头是 N×N。上下文长度 8192、32 个头、fp16 时,把它物化出来要约 4 GB。而这还只是一个请求、一层的量。
比体积更糟的是流量。这个矩阵被写进 HBM,为了 softmax 又被读回来,再写一次,为乘以 value 再读第三次。注意力于是受制于内存往返,而不是算术;FLOPs 只占时间的一小部分。
障碍在 softmax。为了数值稳定它需要整行的最大值,为了归一化它需要整行的和。两者看起来都要求先看完所有分数,才能产出任何输出。
核心思路
解法
其实不用。softmax 可以增量地、精确地一趟算完,靠的是一个滚动最大值、一个滚动求和,以及每当最大值移动时施加的一个修正系数。
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
工作原理
工作原理
除了快,它还解锁了什么
一旦注意力变成了在 key tile 上的流式递推,另外几件事几乎就是白送的:
- PagedAttention(S06)—— 既然你本来就在循环遍历 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,看滚动最大值、修正系数和输出累加器如何收敛到精确答案,那个答案用绿色刻度标出。
- 尚未载入 —— 永远不会同时在内存中
- 当前 tile(位于 SRAM)
- 已处理并丢弃
任一时刻只有橙色 tile 是存在的。灰色格子从未被物化,显存的节省全部来自这里。FlashAttention 做的算术和朴素实现一模一样,所以它被称为「IO 感知」,而不是「更快的算法」。
| tile | m(滚动最大值) | 修正系数 e^(m_old−m_new) | ℓ(滚动求和) |
|---|---|---|---|
| 点击运行 | |||
当某个 tile 里出现了比之前都大的分数,滚动最大值就会移动,此前累积的所有值都要乘以 e^(m_old−m_new) 重新缩放,也就是那个小于 1 的橙色修正系数。正是这一次乘法,让流式 softmax 从近似变成了精确。
绿色刻度是真正的注意力输出。在最后一个 tile 之前,累加器都是错的:它是「到目前为止见过的 key」上的合法 softmax,而不是全部 key 上的。这没问题,但也说明 flash attention 的循环为什么永远不能提前退出。
- 内存中的分数
- 4 of 16
- 朴素实现需要保存
- 256
- 重缩放次数
- 0
- 与精确解的最大误差
- 跑完才能对比
注意累加器在每一个中间步骤都是错的:它是「到目前为止见过的 key」上的正确 softmax,而那并不是你要的东西。只有在最后一个 tile 之后它才对得上。再把 tile 大小设成 1。算法依然成立,做的算术也一样,只是几乎每一步都要做一次修正。tile 大小决定的只是「SRAM 里能装下多少」。
亲手实现
实现
你会上线的是分页版本。它就是同一个循环,只是在访问 key 之前多了一次 block table 查表。
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 / lexp(-inf − -inf) 是 NaN
-inf,所以 alpha = exp(-inf - m_new) 求值为 0,这是对的。但如果第一个 tile 被完全掩掉,m_new 也是 -inf,你就得到了 exp(nan)。所有真实 kernel 要么对第一次迭代做特判,要么把最大值初始化成一个很大的有限负数。在不规则 batch 的因果掩码下,「整块被掩掉」是家常便饭。$ python code/s12_flash_attention.py只需要 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实现 flash-decoding:把 key 切给四个「worker」,各自独立跑一遍递推,然后把它们的
(m, ℓ, o)三元组合并成一个。合并规则就是你已经写过的那次重缩放。 - 2给分块版本加上因果掩码,验证输出与朴素因果注意力一致。然后刻意让某个 tile 被完全掩掉,确认你能复现那个 NaN,再去修它。
- 3在 N = 512、2048、8192 上测量朴素与 flash 的峰值显存,实证确认「二次方 vs 常数」的缩放关系。
继续学习
接下来
kernel 已经高效了。可是在小 batch 下,GPU 大部分时间都在等 CPU 告诉它下一步做什么。S13 用算子融合和 CUDA Graph 开启解码加速这一层。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
FlashAttention 实际减少的是什么?