为什么 FlashAttention 能把 attention 从 memory-bound 变成 compute-bound
标准 self-attention 的公式:
朴素的实现会显式构造两个中间矩阵:
当序列长度 \(N\) 增大时,这两个矩阵的显存是 \(\Theta(N^2)\),而且它们要从 HBM 读/写多次。对 transformer 来说,attention 很快变成 memory-bound:算力没用完,时间花在搬数据上。
Attention 的 FLOPs 是 \(\Theta(N^2 d)\),但标准实现的 HBM 访问量也是 \(\Theta(N^2)\) 量级。如果能把中间矩阵 \(S, P\) 留在片上 SRAM 而不是写回 HBM,就能大幅减少显存流量。
现代 GPU 有两级和 attention 密切相关的内存:
| 内存层级 | 位置 | 容量 | 带宽 |
|---|---|---|---|
| HBM(High Bandwidth Memory) | GPU 全局显存 | 40–80 GB+ | ~1–2 TB/s |
| SRAM / Shared Memory | 每个 SM 片上 | ~100–200 KB/SM | ~10–20 TB/s |
SRAM 容量很小,但带宽是 HBM 的十倍左右。FlashAttention 的核心思想是:把计算拆成小块(tiling),让每次只需在 SRAM 里处理一个小 tile,从而避免把完整的 \(S, P\) 写回 HBM。
Softmax 看起来无法分块:要计算某一行的 softmax,必须先知道这一行的最大值和指数和。如果按 tile 计算,后面 tile 可能出现更大的值,前面的结果就错了。
FlashAttention 的解法来自 online softmax:维护两个 running statistics——
当新 tile 的最大值 \(m_{new}\) 更大时,把旧的 \(l'\) 乘上 \(e^{m - m_{new}}\) 即可对齐到新的最大值:
这样 softmax 的归一化可以在一次流式遍历中完成,不需要先看完所有 tile。
FlashAttention 把 \(Q\) 分成 query tile \(Q_i\),把 \(K, V\) 分成 KV tile \(K_j, V_j\)。外层循环遍历 query tile,内层循环遍历 KV tile。对每个 \((i, j)\) tile pair:
伪代码(FlashAttention-2 风格,外层 query tile):
# Inputs: Q, K, V in R^(N x d)
# Output: O in R^(N x d)
# SRAM capacity: M
# Tile sizes: B_r (query rows), B_c (key/value cols)
T_r = ceil(N / B_r)
T_c = ceil(N / B_c)
O = zeros(N, d)
m = full(N, -inf) # running row max
l = zeros(N) # running sum-exp
for i in range(T_r):
Q_i = load(Q[i*B_r:(i+1)*B_r, :])
m_i = full(B_r, -inf)
l_i = zeros(B_r)
O_i = zeros(B_r, d)
for j in range(T_c):
K_j = load(K[j*B_c:(j+1)*B_c, :])
V_j = load(V[j*B_c:(j+1)*B_c, :])
S_ij = (Q_i @ K_j.T) / sqrt(d) # B_r x B_c
m_ij = max(S_ij, axis=1)
l_ij = sum(exp(S_ij - m_ij[:, None]), axis=1)
m_new = max(m_i, m_ij)
# Rescale old and new contributions to the new max
l_i = l_i * exp(m_i - m_new) + l_ij * exp(m_ij - m_new)
O_i = O_i * exp(m_i - m_new)[:, None] \
+ exp(S_ij - m_new[:, None]) @ V_j
m_i = m_new
O_i = O_i / l_i[:, None]
store(O[i*B_r:(i+1)*B_r, :], O_i)
\(O_i\) 之前按旧 max \(m_i\) 缩放,新 max \(m_{new}\) 更大,所有旧指数都要除以 \(e^{m_{new} - m_i}\),即乘 \(e^{m_i - m_{new}}\)。新 tile 的贡献 \(e^{S_{ij} - m_i}\) 也要对齐到 \(m_{new}\),所以写成 \(e^{S_{ij} - m_{new}}\)。这等价于把所有项统一到同一个 softmax 分母下。
反向传播需要 \(P\) 和 \(dO\) 来算 \(dV, dQ, dK\)。如果保存完整 \(P\),显存又是 \(\Theta(N^2)\)。
FlashAttention 的解法:只保存每行的 running max \(m\) 和 running sum \(l\)(\(\Theta(N)\)),反向时按 tile 重新计算 \(S\) 和 \(P\)。
# Backward (per tile)
load Q_i, K_j, V_j, dO_i
m_i, l_i = saved statistics
S_ij = (Q_i @ K_j.T) / sqrt(d)
P_ij = exp(S_ij - m_i[:, None]) / l_i[:, None]
dV_j += P_ij.T @ dO_i
dP_ij = dO_i @ V_j.T
# softmax backward
dS_ij = P_ij * (dP_ij - rowsum(dP_ij * P_ij)[:, None]) / sqrt(d)
dQ_i += dS_ij @ K_j
dK_j += dS_ij.T @ Q_i
这以额外计算换取显存:FLOPs 仍是 \(\Theta(N^2 d)\),但显存降为 \(\Theta(N)\)。在 GPU 上,compute 通常比 HBM 带宽便宜,所以这是划算的交易。
FlashAttention 不是减少 FLOPs,而是减少 HBM 流量:通过 tiling 把大块拆到 SRAM,通过 online softmax 让分块归一化可行,通过 kernel fusion 避免写回中间矩阵,通过 recomputation 把 backward 内存也降到线性。
本课推荐阅读 Seth Weidman 的博客 FlashAttention: Algorithm and Pseudocode 和原论文 FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness 的 Section 3。
下节课:FlashAttention 1-4 演进速查 会先把各代优化一览讲清,然后我们再逐代深入。
我是你的老师,可以随时问我:online softmax 为什么 rescale、tile size 怎么选、伪代码哪一行不清楚,都可以继续聊。