A · Flash Attention
标准 Attention 的显存墙问题,如何通过分块 (Tiling) 与在线 Softmax 将时间复杂度从 O(N^2) 降到 O(N)。
A · Flash Attention
标准 Attention 的瓶颈不在于浮点运算量,而在于 HBM(显存)的读写次数。Flash Attention 通过重新组织计算顺序,在几乎不改变数学结果的前提下,把显存访问从 降到 ,速度提升 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
对于序列长度 ,中间矩阵 和 各占 显存,且每次都要在 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 各被读取一次, 被写入一次,总计 ,远小于标准实现的 。
三、反向传播:重计算(Recomputation)
反向传播需要 (注意力权重矩阵)来计算梯度,但 Flash Attention 没有存 。 解决方法:在反向传播时,重新从 (已存在 HBM)计算 ,而不是把 从 HBM 读出来。 看起来浪费计算,但实际更快:重新计算的代价(FLOPS)远小于从 HBM 读 大矩阵的代价(IO)。这是 Flash Attention 反直觉但有效的地方。
四、数值稳定性
减去最大值是 softmax 的标准稳定化手段:
分子分母同乘 ,数学上等价,但防止了 上溢(当 很大时)。在线 softmax 里, 就扮演这个角色,且随着看到更多数据动态更新。
五、Flash Attention v2 的改进
v1 的问题:不同 warp(GPU 线程组)之间存在同步开销。 v2 的改进:
- 减少非矩阵乘法运算的占比(softmax 修正步骤)
- 重新设计并行策略:在序列长度维度上并行(而非 batch 维度),更好利用 GPU 并行度
- 前向传播速度在 A100 上达到理论峰值的 72%(v1 约 35%)
六、意义与影响
| 维度 | 效果 |
|---|---|
| 速度 | 比标准 Attention 快 2-4x |
| 显存 | 从 降到 ,使更长上下文成为可能 |
| 精度 | 数学等价(误差在浮点精度范围内),非近似算法 |
| 影响 | 几乎所有现代 LLM 训练和推理都使用 Flash Attention |
Flash Attention 是「IO 感知算法」的典范——优化的不是数学,而是硬件访问模式。