Flash Attention 是由 Tri Dao 等人于 2022 年在斯坦福大学提出的注意力计算优化算法,通过重新设计计算顺序来解决注意力机制的显存瓶颈问题。
核心问题
标准注意力计算的时间和显存复杂度是 O(n²)(n 为序列长度)。当上下文长度从 4K 扩展到 128K 甚至 1M tokens 时,显存需求爆炸。传统方案是将注意力矩阵写入 HBM(高带宽显存),读写带宽成为瓶颈。
Flash Attention 的解法
- 分块计算(Tiling):将大矩阵拆成小块,在 SRAM(芯片内高速缓存)中完成计算,减少 HBM 访问
- 重新计算替代存储:不保存中间结果,反向传播时重新计算,省显存
- IO意识:根据 GPU 存储层次(SRAM vs HBM)的速度差异设计算法
版本演进
- Flash Attention 1(2022):实现 2-4x 加速,显存降为 O(n)
- Flash Attention 2(2023):优化并行度,再提升 2x
- Flash Attention 3(2024):针对 H100 GPU 优化,利用异步执行
采用情况
PyTorch 2.0+ 内置支持,大部分现代 LLM(GPT-4、Claude、Llama 3)的训练和推理都使用某种形式的 Flash Attention。