经典系统arXiv:2205.14135
FlashAttention:把注意力写成 IO 感知算法
FlashAttention: Fast and Memory-Efficient Exact Attention
Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré
NeurIPS · 2022
5.0k 引用 · 1.5k 阅读 · 0 收藏
摘要
按 GPU SRAM 分块重算 softmax 注意力,得到精确结果却大幅减少显存读写,从而加速 Transformer。
Recomputes softmax attention in SRAM-aware tiles to obtain exact attention with far less memory traffic, speeding up Transformers without approximation.
提要
FlashAttention 提醒我们:注意力的瓶颈经常是搬数据,不是算 FLOPs。它按 GPU SRAM 做分块并在反向重算中间量,得到与标准算法等价的精确注意力,同时大幅减少显存读写。
01背景与动机
在 Transformer 中,标准自注意力需要显式物化完整的 N×N 注意力矩阵,并在 GPU 高带宽显存与片上 SRAM 之间反复搬运中间激活与梯度。随着上下文变长,这种数据搬运的代价常常超过矩阵乘法本身的浮点运算量,使注意力成为训练与推理的主要瓶颈。
因此,FlashAttention 的动机是在不牺牲数值精确性的前提下,把标准注意力改写成贴合 GPU 存储层次的分块算法,用更少的内存流量完成等价计算,从而加速 Transformer 并支撑更长上下文。
02核心方法
该方法将查询、键、值按能放入片上 SRAM 的块划分,在 SRAM 内完成分块矩阵乘与在线 softmax 统计量更新,只把最终输出写回高带宽显存。通过维护每块的归一化信息,分块结果在数学上等价于一次全局 softmax。
反向传播时不持久化巨大的注意力矩阵,而是按需从输入重算中间量。这样既避免物化 O(N²) 中间激活,也显著降低显存读写次数,使精确注意力在 IO 意义上更高效。

03结果与证据
由于分块与重算在数学上保持与标准注意力相同的定义,FlashAttention 给出的是精确结果而非近似。其主要收益来自大幅减少的内存流量,而不是改写注意力的数值含义或引入随机稀疏等捷径。
在系统层面,更低的读写压力使注意力算子本身更快,并减轻激活显存占用,从而在相同硬件上更易训练或推理更长序列,且无需依赖低秩、稀疏等近似注意力变体。

04影响与意义
FlashAttention 已被纳入 PyTorch 及各类主流训练框架,成为注意力的默认高效实现之一。对使用者而言,往往升级库版本即可获得加速,不必改动模型结构或训练脚本中的注意力定义。
它直接抬高了可稳定训练的上下文长度,使长文档与长对话等设定在常规算力下更可行,也推动后续工作继续围绕硬件感知的算子与内核设计展开。

05局限
需要强调的是,FlashAttention 加速的是注意力算子的实现效率,并不改变其相对序列长度的二次渐近复杂度。当序列极长时,计算量本身仍会迅速增长,IO 优化无法单独消除这一标度。
因此对超长序列场景,仍需结合稀疏注意力、线性注意力或其他架构变体。硬件感知的系统优化与复杂度层面的算法创新是互补关系,而非彼此替代。

06要点
- 01算法应当对照硬件的存储层次结构来写,而不能只盯着 FLOPs。
- 02数值精确与实现高效并不必然互斥,关键在于减少不必要的数据搬运。
- 03扎实的系统与算子优化,其效果可以相当于一次架构层面的升级。
为何入典
系统论文改写模型可训练的规模。长上下文之所以可行,很大程度是 IO 被做对了。
标签
相关论文
DeepSeek-V2:MLA 与经济型 MoE
DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model
LLaMA:开放基础模型的分水岭
LLaMA: Open and Efficient Foundation Language Models
Transformer:注意力就够了
Attention Is All You Need
和这篇论文对话
开启对话后,本页的标题、摘要、导读和图注会装进上下文。模型是你自己接入的。