0%

FlashAttention 原理

FlashAttention 原理

一句话定位:Attention 慢的真正原因不是算力不够,而是HBM 读写太多(IO-bound)。FlashAttention 通过分块 + 片上融合计算,避免把 N×N 的注意力矩阵写进显存,从而同时省下时间和显存——而且是精确计算,不是近似

1. 关键洞察:瓶颈是 HBM 读写而非算力

标准 Attention 的执行过程会**物化(materialize)**中间矩阵:

  1. S = QKᵀ(N×N)→ 写入 HBM;
  2. 读回 S,算 P = softmax(S)(N×N)→ 再写入 HBM;
  3. 读回 P,算 O = PV → 写出。

问题:N×N 矩阵随序列长度平方级增长,反复在 HBM 与计算单元之间搬运。由于 HBM 带宽远低于 GPU 算力,实际耗时由访存量决定——即 Attention 是 memory-bound / IO-bound,GPU 算力大量空等。

2. 做法:Tiling 分块 + 片上融合 softmax

  • Tiling(分块):把 Q、K、V 切成能放进**片上 SRAM(共享内存)**的小块,按块循环计算。
  • 算子融合:在 SRAM 内对一个块连续完成 QKᵀ → softmax → 乘 V 的全部步骤,中间结果不落 HBM
  • Online Softmax:softmax 需要整行的最大值与求和做归一化,分块后无法一次看全行。解决办法是采用增量式(online)softmax——边遍历块边维护running max 与 running sum,并对已累积的输出做相应的重新缩放(rescale),最终得到与全局 softmax 完全一致的结果。

于是 N×N 的注意力矩阵从未被完整写入显存

3. 收益

  • 减少 HBM 读写量 → 墙钟时间显著下降(访存是瓶颈,所以省访存就是省时间);
  • 显存占用从 O(N²) 降到 O(N) → 可支持更长上下文;
  • 精确而非近似:区别于稀疏注意力、低秩近似等方法,FlashAttention 的输出与标准 Attention 数学上等价,这是它能被广泛默认启用的根本原因。
  • FlashAttention-2 进一步优化并行划分与工作分配(减少非矩阵乘操作、改进 warp 间划分),提升 GPU 占用率。

面试要求:不需要推公式,但必须能说清优化动机——“Attention 是 IO-bound,所以要减少 HBM 往返,用分块把计算搬到片上并融合,且靠 online softmax 保证结果精确”。

4. 可迁移话术

这与大数据里”减少 Shuffle 落盘与网络往返”是同一套思维:瓶颈在数据搬运而非计算,就把计算推到数据所在的高速层级并做算子融合(对应 1.1 中 Shuffle 磁盘/网络开销 ≈ HBM 读写瓶颈)。

参考:论文《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》(Dao et al., 2022) 及《FlashAttention-2》