跳到主要内容

FlashAttention:把 O(n²) 显存压成 O(n)

不换算法,只靠分块与重计算,让标准注意力在长序列上快一个量级。

标准注意力要先算出完整的 n×nn \times n 注意力矩阵再做 softmax—— 显存占用 O(n2)\mathcal{O}(n^2)。序列一长,先炸显存,再炸速度。 FlashAttention 的结论出人意料:同样的数学,换个执行顺序, 显存降到 O(n)\mathcal{O}(n),速度提升数倍。

两个核心观察

  1. 瓶颈是显存带宽QKQK^\top 和 softmax 都是内存密集操作, GPU 的算力根本没吃饱。
  2. 重计算比存中间结果便宜:注意力矩阵算完就丢,反传时重算一遍, 省下的显存换成更大的分块与更少的 HBM 读写。

分块 softmax 的 trick

标准 softmax 需要全行的最大值做归一化。分块计算时,维护 运行中的最大值与修正因子,逐块更新:

mnew=max(m,mi),lnew=emmnewl+emimnewlim_{new} = \max(m, m_i),\qquad l_{new} = e^{m - m_{new}}\,l + e^{m_i - m_{new}}\,l_i

合并两个块的统计量即可得到与全局 softmax 相同的结果——这是 FlashAttention 数学上成立的关键,也是后续 FlashAttention-2/3 做并行分块的基础。

为什么值得关注

  • 长上下文(32k+)推理的标配组件;
  • 几乎所有推理框架(vLLM、SGLang)与训练框架(Megatron、HuggingFace)内置;
  • 思想通用:「以计算换显存」在算子层面比换算法更普适。

实现细节可看原论文 FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., 2022)。

评论