标准注意力要先算出完整的 注意力矩阵再做 softmax—— 显存占用 。序列一长,先炸显存,再炸速度。 FlashAttention 的结论出人意料:同样的数学,换个执行顺序, 显存降到 ,速度提升数倍。
两个核心观察
- 瓶颈是显存带宽: 和 softmax 都是内存密集操作, GPU 的算力根本没吃饱。
- 重计算比存中间结果便宜:注意力矩阵算完就丢,反传时重算一遍, 省下的显存换成更大的分块与更少的 HBM 读写。
分块 softmax 的 trick
标准 softmax 需要全行的最大值做归一化。分块计算时,维护 运行中的最大值与修正因子,逐块更新:
合并两个块的统计量即可得到与全局 softmax 相同的结果——这是 FlashAttention 数学上成立的关键,也是后续 FlashAttention-2/3 做并行分块的基础。
为什么值得关注
- 长上下文(32k+)推理的标配组件;
- 几乎所有推理框架(vLLM、SGLang)与训练框架(Megatron、HuggingFace)内置;
- 思想通用:「以计算换显存」在算子层面比换算法更普适。
实现细节可看原论文 FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., 2022)。