第 1 课:FlashAttention 核心原理

为什么 FlashAttention 能把 attention 从 memory-bound 变成 compute-bound

1. 标准 Attention 的问题

标准 self-attention 的公式:

\[O = \text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)V\]

朴素的实现会显式构造两个中间矩阵:

当序列长度 \(N\) 增大时,这两个矩阵的显存是 \(\Theta(N^2)\),而且它们要从 HBM 读/写多次。对 transformer 来说,attention 很快变成 memory-bound:算力没用完,时间花在搬数据上。

关键观察

Attention 的 FLOPs 是 \(\Theta(N^2 d)\),但标准实现的 HBM 访问量也是 \(\Theta(N^2)\) 量级。如果能把中间矩阵 \(S, P\) 留在片上 SRAM 而不是写回 HBM,就能大幅减少显存流量。

2. GPU 内存层次:为什么 SRAM 如此重要

现代 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。

3. 障碍:Softmax 是全局归一化

Softmax 看起来无法分块:要计算某一行的 softmax,必须先知道这一行的最大值和指数和。如果按 tile 计算,后面 tile 可能出现更大的值,前面的结果就错了。

FlashAttention 的解法来自 online softmax:维护两个 running statistics——

当新 tile 的最大值 \(m_{new}\) 更大时,把旧的 \(l'\) 乘上 \(e^{m - m_{new}}\) 即可对齐到新的最大值:

\[m_{new} = \max(m, m_{tile})\]
\[l'_{new} = l' \cdot e^{m - m_{new}} + l'_{tile} \cdot e^{m_{tile} - m_{new}}\]

这样 softmax 的归一化可以在一次流式遍历中完成,不需要先看完所有 tile。

4. 核心算法:Tiling + Online Softmax + 融合

FlashAttention 把 \(Q\) 分成 query tile \(Q_i\),把 \(K, V\) 分成 KV tile \(K_j, V_j\)。外层循环遍历 query tile,内层循环遍历 KV tile。对每个 \((i, j)\) tile pair:

  1. 从 HBM 加载 \(Q_i, K_j, V_j\) 到 SRAM。
  2. 在 SRAM 内计算 \(S_{ij} = Q_i K_j^T / \sqrt{d}\)。
  3. 用 online softmax 更新 running max \(m_i\) 和 running sum \(l_i\)。
  4. 直接累加对输出的贡献:把 \(e^{S_{ij} - m_i} V_j\) 加到 \(O_i\)(并在 max 变化时 rescale 旧值)。
  5. 内层循环结束后,把 \(O_i\) 除以 \(l_i\) 得到最终输出 tile,写回 HBM。

伪代码(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)
为什么 rescale 是对的?

\(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 分母下。

5. 反向传播:Recomputation 技巧

反向传播需要 \(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 带宽便宜,所以这是划算的交易。

6. 一句话总结

FlashAttention 不是减少 FLOPs,而是减少 HBM 流量:通过 tiling 把大块拆到 SRAM,通过 online softmax 让分块归一化可行,通过 kernel fusion 避免写回中间矩阵,通过 recomputation 把 backward 内存也降到线性。

7. 快速测验

1. FlashAttention 主要减少的是哪一类开销?
2. Online softmax 解决的核心问题是什么?
3. 反向传播时 FlashAttention 为什么不保存完整的 \(P\) 矩阵?

8. 延伸阅读

本课推荐阅读 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 怎么选、伪代码哪一行不清楚,都可以继续聊。