探索/系统

经典系统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 意义上更高效。

A100上不同序列长度相对标准PyTorch注意力的加速比。
1. A100上不同序列长度相对标准PyTorch注意力的加速比。

03结果与证据

由于分块与重算在数学上保持与标准注意力相同的定义,FlashAttention 给出的是精确结果而非近似。其主要收益来自大幅减少的内存流量,而不是改写注意力的数值含义或引入随机稀疏等捷径。

在系统层面,更低的读写压力使注意力算子本身更快,并减轻激活显存占用,从而在相同硬件上更易训练或推理更长序列,且无需依赖低秩、稀疏等近似注意力变体。

A100上头维度128时,不同序列长度相对标准PyTorch注意力的加速比。
2. A100上头维度128时,不同序列长度相对标准PyTorch注意力的加速比。

04影响与意义

FlashAttention 已被纳入 PyTorch 及各类主流训练框架,成为注意力的默认高效实现之一。对使用者而言,往往升级库版本即可获得加速,不必改动模型结构或训练脚本中的注意力定义。

它直接抬高了可稳定训练的上下文长度,使长文档与长对话等设定在常规算力下更可行,也推动后续工作继续围绕硬件感知的算子与内核设计展开。

RTX 3090上不同序列长度相对标准PyTorch注意力的加速比。
3. RTX 3090上不同序列长度相对标准PyTorch注意力的加速比。

05局限

需要强调的是,FlashAttention 加速的是注意力算子的实现效率,并不改变其相对序列长度的二次渐近复杂度。当序列极长时,计算量本身仍会迅速增长,IO 优化无法单独消除这一标度。

因此对超长序列场景,仍需结合稀疏注意力、线性注意力或其他架构变体。硬件感知的系统优化与复杂度层面的算法创新是互补关系,而非彼此替代。

T4上不同序列长度相对标准PyTorch注意力的加速比。上:前向+反向;下:仅前向。
4. T4上不同序列长度相对标准PyTorch注意力的加速比。上:前向+反向;下:仅前向。

06要点

  1. 01算法应当对照硬件的存储层次结构来写,而不能只盯着 FLOPs。
  2. 02数值精确与实现高效并不必然互斥,关键在于减少不必要的数据搬运。
  3. 03扎实的系统与算子优化,其效果可以相当于一次架构层面的升级。

为何入典

系统论文改写模型可训练的规模。长上下文之所以可行,很大程度是 IO 被做对了。

标签

attentionGPUsystems

相关论文

和这篇论文对话

开启对话后,本页的标题、摘要、导读和图注会装进上下文。模型是你自己接入的。

登录后可以接入自己的模型,向这篇论文提问。