Skip to content
🔗 分享本题
查看我的学习进度 →
💡 答案要点

FlashAttention 的核心不是近似 Attention,而是在保持精确结果的前提下减少显存 IO。

传统 Attention 的瓶颈:

标准 Attention 步骤(每步都要读写 HBM):

Step 1: S = QK^T      → 产生 N×N 分数矩阵,写回 HBM  (O(N²) 访问)
Step 2: P = softmax(S) → 读入 S,逐行 softmax,写回 P   (O(N²) 访问)
Step 3: O = PV        → 读入 P 和 V,乘积后写回输出    (O(N²) 访问)

总 HBM 访问量: O(Nd + N²)
问题: N 越大,中间矩阵 S/P 越大,GPU 大部分时间在等内存 IO

FlashAttention 优化思路 — IO-Awareness(感知 IO 的算法设计):

FlashAttention 利用 GPU 的内存层次结构:
  HBM(高带宽慢)→ SRAM(片上快)→ Register(最快但最小)

关键观察:
  - SRAM 容量远小于完整 N×N 矩阵
  - 但单个 block 内的计算完全可以放在 SRAM 中完成
  - 不需要每次都把中间结果写回 HBM

核心优化技术:

1. Tiling(分块计算):

将 Q、K、V 矩阵切分成小 block:
  Q → [Q_0, Q_1, ..., Q_T]     每个 Q_t ∈ R^{n×d}
  K → [K_0, K_1, ..., K_T]     每个 K_t ∈ R^{m×d}
  V → [V_0, V_1, ..., V_T]     每个 V_t ∈ R^{m×d}

逐对计算 (Q_i, K_j, V_j),所有中间结果留在 SRAM
最终一次写入 HBM

2. Online Softmax(在线归一化):

问题:softmax 需要整行的最大值才能数值稳定计算

传统做法: 先求全局 max → 再算 exp(x-max) → 最后除以 sum
FlashAttention: 遍历每个 block 时动态维护
  m = 当前行最大值
  l = 归一化因子 Σexp(x-m)
  O = 未归一化的输出累积

每次遇到新 block 用以下公式增量更新:
  m_new = max(m_old, max_of_block)
  l_new = l_old * exp(m_old - m_new) + sum(exp(block - m_new))
  O_new = O_old * exp(m_old - m_new) + weighted_sum(block_V)

这等价于一次性做完整 softmax,但只用 O(1) 额外空间

3. 重计算(Recomputation)代替缓存:

反向传播时,如果保存完整的中间矩阵会消耗大量激活显存

FlashAttention: 只在正向传递时保存必要的标量统计量
               反向传播时重新计算部分分数和 softmax

成本权衡:
  - 正向多读一次 K/V 从 HBM
  - 省下的激活显存可容纳更大 batch 或更长的序列

复杂度分析:

假设序列长度 n, head 维度 d, SRAM 容量 M (M ≥ d):

标准 Attention 的 HBM 访问次数:
  Forward:  O(n·d + n²)
  Backward: O(n·d + n²)
  总计:     O(n²)

FlashAttention 的 HBM 访问次数:
  Forward:  Θ((n²/M) · n·d)  ← 受限于分块数量
  Backward: 类似

当 n >> √M 时,FlashAttention 显著优于标准实现
实际加速比: 2x~4x(取决于硬件和序列长度)

面试话术:

"FlashAttention 的本质是 IO-aware 设计——让数据尽量留在 SRAM 而不是反复往返 HBM。通过分块计算和在线 softmax,它避免了存储完整的 N×N 注意力矩阵。结果是显存占用大幅下降(可以从 O(N²) 降到接近线性),同时保持了精确注意力结果。在长序列场景下,速度提升可达 2-4 倍。这不是近似算法,而是同一个数学过程的不同组织方式。"

⭐ 面试加分项:

  • 能画出 FlashAttention 的分块流程图(外部循环跨 K/V blocks,内部循环处理 Q blocks)
  • 理解 online softmax 的增量更新公式推导
  • 知道 FlashAttention 2 引入了 persistent kernel 进一步优化
  • 了解 FlashAttention 与 KV Cache 配合使用时的效果(推理时同样受益)

📚 参考:FlashAttention: Fast and Memory-Efficient Exact Attention(原论文)