· LLM 理论精读 · 第 1 / 3 讲

A · Flash Attention

标准 Attention 的显存墙问题,如何通过分块 (Tiling) 与在线 Softmax 将时间复杂度从 O(N^2) 降到 O(N)。

A · Flash Attention

标准 Attention 的瓶颈不在于浮点运算量,而在于 HBM(显存)的读写次数。Flash Attention 通过重新组织计算顺序,在几乎不改变数学结果的前提下,把显存访问从 O(N2)O(N^2) 降到 O(N)O(N),速度提升 2-4 倍。


一、问题的根源:内存层次结构

现代 GPU 有两级存储:

存储层容量带宽延迟
SRAM(片上缓存)~20MB(A100)~19 TB/s极低
HBM(显存)~80GB(A100)~2 TB/s

计算速度(FLOPS)的提升远快于 HBM 带宽的提升。现代 GPU 经常处于「算完了但在等数据」的状态——称为 memory-bound(访存瓶颈)。

标准 Attention 的伪代码:

# 标准实现——每一步都要读写 HBM
S = Q @ K.T / sqrt(d_k)      # 写入 HBM:N×N 矩阵
P = softmax(S)                 # 读 + 写 HBM:N×N
O = P @ V                      # 读 + 写 HBM:N×d

对于序列长度 NN,中间矩阵 SSPP 各占 O(N2)O(N^2) 显存,且每次都要在 HBM 和 SRAM 之间搬运。瓶颈不是乘法,是搬运。



## 二、Flash Attention 的核心思想:分块(Tiling)
**目标**:不把完整的 $N \times N$ 矩阵写到 HBM,在 SRAM 里分块完成 softmax 和加权聚合,只写最终结果 $O$。
### 2.1 softmax 的在线计算问题
分块计算的障碍在于 softmax——它需要看到这一行所有的值才能计算分母:
$$
P_{ij} = \frac{\exp(S_{ij})}{\sum_k \exp(S_{ik})}
$$
直接分块处理不同的 $j$ 块,分母会算错。
### 2.2 在线 softmax(Online Softmax)
关键数学技巧:softmax 可以**增量地**修正,而不需要重新过一遍所有数据。
设已处理了前 $t$ 个元素,维护两个统计量:
- $m_t = \max(x_1, ..., x_t)$:到目前为止的最大值(用于数值稳定性)
- $l_t = \sum_{i=1}^{t} \exp(x_i - m_t)$:指数和
当新来第 $t+1$ 个元素 $x_{t+1}$ 时,更新规则:
$$
m_{t+1} = \max(m_t,\ x_{t+1})
$$
$$
l_{t+1} = l_t \cdot \exp(m_t - m_{t+1}) + \exp(x_{t+1} - m_{t+1})
$$
**直觉**:新的最大值 $m_{t+1}$ 可能比旧的 $m_t$ 更大,这时旧的指数和 $l_t$ 需要乘以修正因子 $\exp(m_t - m_{t+1})$ 来统一基准。
同理,已经算好的输出 $O$ 也需要对应修正:
$$
O_{t+1} = O_t \cdot \frac{l_t \cdot \exp(m_t - m_{t+1})}{l_{t+1}} + \frac{\exp(x_{t+1} - m_{t+1})}{l_{t+1}} \cdot v_{t+1}
$$
### 2.3 分块计算流程
```javascript

将 Q 按行分成 Tr 块,每块大小 Br

将 K, V 按列分成 Tc 块,每块大小 Bc

for i = 1 to Tr:

    从 HBM 加载 Qi 到 SRAM

    初始化 Oi = 0, li = 0, mi = -∞

    for j = 1 to Tc:

        从 HBM 加载 Kj, Vj 到 SRAM

        计算局部分数 Sij = Qi @ Kj.T / sqrt(d)

        更新在线 softmax:(mi_new, li_new)

        用修正因子更新 Oi

    将最终 Oi 写回 HBM  ← 只写一次!

HBM 访问次数:每块 Q、K、V 各被读取一次,OO 被写入一次,总计 O(Nd)O(N \cdot d),远小于标准实现的 O(N2)O(N^2)

三、反向传播:重计算(Recomputation)

反向传播需要 PP(注意力权重矩阵)来计算梯度,但 Flash Attention 没有存 PP解决方法:在反向传播时,重新从 Q,K,VQ, K, V(已存在 HBM)计算 PP,而不是把 PP 从 HBM 读出来。 看起来浪费计算,但实际更快:重新计算的代价(FLOPS)远小于从 HBM 读 N2N^2 大矩阵的代价(IO)。这是 Flash Attention 反直觉但有效的地方。

四、数值稳定性

减去最大值是 softmax 的标准稳定化手段:

softmax(xi)=exp(xi)jexp(xj)=exp(xim)jexp(xjm),m=maxjxj\text{softmax}(x_i) = \frac{\exp(x_i)}{\sum_j \exp(x_j)} = \frac{\exp(x_i - m)}{\sum_j \exp(x_j - m)}, \quad m = \max_j x_j

分子分母同乘 exp(m)\exp(-m),数学上等价,但防止了 exp(xi)\exp(x_i) 上溢(当 xix_i 很大时)。在线 softmax 里,mtm_t 就扮演这个角色,且随着看到更多数据动态更新。

五、Flash Attention v2 的改进

v1 的问题:不同 warp(GPU 线程组)之间存在同步开销。 v2 的改进:


六、意义与影响

维度效果
速度比标准 Attention 快 2-4x
显存O(N2)O(N^2) 降到 O(N)O(N),使更长上下文成为可能
精度数学等价(误差在浮点精度范围内),非近似算法
影响几乎所有现代 LLM 训练和推理都使用 Flash Attention

Flash Attention 是「IO 感知算法」的典范——优化的不是数学,而是硬件访问模式。