FlashAttention
FlashAttention
1. 使用背景
Transformer 很强,但 attention 一直有一个很现实的问题:又慢又吃显存。
如果序列长度是NNN,标准自注意力里最核心的中间量是:
S=QK⊤ S = QK^\topS=QK⊤
它的形状是N×NN\times NN×N。
这意味着,随着序列变长,attention 的中间结果会迅速膨胀。早期很多人一看到 attention 是二次复杂度,就会自然想到:是不是只能去做近似、稀疏化、低秩化。
但 FlashAttention 提出了一个非常关键的观点:很多时候 attention 慢,不只是算得多,更是因为显存读写太重。
也就是说,问题不只是 FLOPs,不只是数学复杂度,还包括一个很重要的系统因素:IO(输入输出 / 内存访问)成本。
FlashAttention 就是在这个背景下提出来的。它最核心的思想可以概括成一句话:
不改变标准 attention 的数学结果,而是通过分块、融合和在线 softmax,把大规模中间矩阵的显存读写降下来。
2. 理论基础
(1)标准 attention 到底在算什么
给定输入隐藏状态,先通过线性变换得到:
Q=XWQ,K=XWK,V=XWV Q=XW_Q,\qquad K=XW_K,\qquad V=XW_VQ=XWQ,K=XWK,V=XWV标准 scaled dot-product attention 写成:
Attention(Q,K,V)=softmax(QK⊤d)V \text{Attention}(Q,K,V)=\text{softmax}\left(\frac{QK^\top}{\sqrt{d}}\right)VAttention(Q,K,V)=softmax(dQK⊤)V如果把它拆开看,其实主要有三步:
- 先算分数矩阵:
S=QK⊤d S=\frac{QK^\top}{\sqrt{d}}S=dQK⊤ - 再做 softmax,得到注意力权重:
P=softmax(S) P=\text{softmax}(S)P=softmax(S) - 最后和VVV相乘:
O=PV O=PVO=PV
- 这个流程数学上很清楚,但在实现上会有一个大问题:
中间矩阵SSS和PPP都是N×NN\times NN×N。
(2)为什么这会成为瓶颈
当序列长度NNN很大时,SSS和PPP的存储、读写都非常重。
即使 GPU 算力很强,如果每一步都要频繁把这些大矩阵在 HBM(高带宽显存)和片上 SRAM / cache 之间来回搬运,速度也会被拖慢。
所以 attention 的瓶颈并不只是“乘法很多”,而是:
中间结果太大,内存访问代价太高。
(3)为什么普通实现会浪费
- 在常规实现里,通常会:
- 先把S=QK⊤S=QK^\topS=QK⊤整个算出来并写到显存;
- 再从显存里读出来做 softmax,得到PPP;
- 再把PPP写回显存;
- 最后再读出PPP和VVV去算输出。
这个过程最大的问题在于:
很多中间结果只是临时用一下,却被完整写回了大显存。
这就是 FlashAttention 最想解决的地方。
3. FlashAttention 的核心思想
(1)一句话理解
FlashAttention 最核心的想法就是:
不要把整个 attention matrix 显式落到显存里,而是在片上做小块计算,边算边归约,直接得到输出。
也就是说,它不是先完整构造:
S=QK⊤,P=softmax(S) S=QK^\top,\qquad P=\text{softmax}(S)S=QK⊤,P=softmax(S)
再去算输出,而是把这几个步骤融合在一起,以分块方式完成。
(2)为什么叫 IO-aware
FlashAttention 论文里一个非常关键的词就是IO-awareness。
传统算法通常更关注算术复杂度,比如时间复杂度是不是O(N2)O(N^2)O(N2);
FlashAttention 则进一步强调:
在现代 GPU 上,显存和片上缓存之间的数据搬运,本身就是主要开销。
所以 FlashAttention 并不是去近似 attention,而是去优化 attention 的数据流。
(3)最重要的一点:它是 exact attention
这一点特别重要。
FlashAttention 不是 Performer、Linformer 那种近似 attention 方法。
它并没有改 attention 的数学定义,而只是改了实现方式。
所以它的核心价值不是“牺牲精度换速度”,而是:
在保持精确等价的前提下,把显存访问和运行时间都降下来。
4. FlashAttention 是怎么做的
(1)分块(tiling)
假设把Q,K,VQ,K,VQ,K,V都按块切开。
比如把QQQ按行分成若干 query block,把K,VK,VK,V按列分成若干 key/value block。
那么 attention 就不再是一次性拿整块大矩阵去算,而是按 block 逐块处理。
直观上理解,就是:
整张大表一次性装不下,那就分小块搬进片上缓存里,一块一块算。
(2)问题:softmax 不是一个纯局部操作
如果只是矩阵乘法,分块很好理解。
但 attention 里还有 softmax,而 softmax 涉及整行归一化:
softmax(si)j=esij∑kesik \text{softmax}(s_i)_j=\frac{e^{s_{ij}}}{\sum_k e^{s_{ik}}}softmax(si)j=∑kesikesij这里的分母要看整行所有元素,所以它不像普通矩阵乘法那样天然可以局部分开。
这也是 FlashAttention 真正巧的地方:
它用在线(online)方式维护 softmax 所需的统计量。
(3)在线 softmax
对于某一行 attention score,常规 softmax 会先拿到整行所有值,再求最大值、再求指数和。
FlashAttention 则是在分块扫描时,逐步维护:
- 当前见过的最大值;
- 当前对应的归一化系数;
- 当前输出累积值。
这样一来,即使一整行分散在多个 block 中,也不需要把整行完整写回显存后再统一 softmax。
所以 FlashAttention 的关键点不只是“分块”,而是:
分块 + 在线 softmax + 输出累积。
5. 在线 softmax 的直观理解
(1)为什么 softmax 通常要先减最大值
普通 softmax 为了数值稳定,常写成:
softmax(si)j=esij−mi∑kesik−mi \text{softmax}(s_i)_j=\frac{e^{s_{ij}-m_i}}{\sum_k e^{s_{ik}-m_i}}softmax(si)j=∑kesik−miesij−mi
其中
mi=maxjsij m_i=\max_j s_{ij}mi=jmaxsij这样做是为了防止指数爆炸。
所以如果想分块做,就必须一边扫描 block,一边正确维护这个最大值和归一化项。
(2)FlashAttention 的核心统计量
对于某一行,FlashAttention 会在扫描 block 的过程中不断更新:
- 行最大值mmm;
- 行归一化和lll;
- 输出向量累计值。
每看到一个新的 key block,就先算当前 block 对这一行的局部分数,再和旧的m,lm,lm,l合并更新。
这样到最后,虽然没有显式存整个SSS和PPP,但最终得到的输出和标准 attention 完全一致。
(3)本质上在做什么
如果用一句话概括这个过程,我觉得最合适的是:
把原来“先得到完整注意力矩阵,再做 softmax”这件事,改写成“扫描中间块时同步完成归一化和加权求和”。
所以它不是把 attention 简化了,而是把计算顺序重新排了一遍。
6. FlashAttention 到底省了什么
(1)不是把二次复杂度变成线性
这是最容易误解的一点。
FlashAttention 并没有把 attention 的理论计算复杂度从O(N2)O(N^2)O(N2)变成O(N)O(N)O(N)。
从数学上讲,所有 query-key 对之间的交互还是要算。
它真正优化的是:
中间矩阵的显式物化和反复显存访问。
(2)它主要省的是 HBM 访问
在 GPU 上,HBM 访问代价远高于片上 SRAM / register 的访问。
FlashAttention 通过块级计算,让很多中间值留在片上,而不是反复写回大显存。
所以它真正省的,不是“注意力不用算了”,而是:
很多原本没必要写出的大矩阵,不再落到 HBM。
(3)为什么这会带来显著提速
因为在长序列下,attention 往往不是单纯 compute-bound,而是明显受内存访问限制。
一旦把显存读写压下去,GPU 的有效利用率就能明显提高。
所以 FlashAttention 的加速来源,本质上是:
更合理的数据流,而不是近似。
7. FlashAttention 和普通 fused kernel 有什么区别
(1)不是简单 kernel fusion
看到这里,很多人会觉得:
“这不就是把几个 kernel 融合一下吗?”其实比普通 fusion 更进一步。
普通 kernel fusion 更多是把几个连续操作合并,以减少 launch overhead 和部分中间读写;
FlashAttention 则是从 attention 的算法结构本身出发,重新设计了 block-wise 的 IO 流程。
所以它不只是“把 kernel 拼起来”,而是:
把 attention 的执行顺序按 GPU 内存层级重新组织了一遍。
(2)为什么这很重要
因为 attention 里的问题不只是 kernel 多,而是中间矩阵太大。
单纯 fusion 并不能自动解决“整块 attention matrix 落显存”的问题。
FlashAttention 真正解决的是这个更深层的问题。
8. FlashAttention 的训练和推理意义
(1)训练里最直接的收益
- 在训练阶段,FlashAttention 可以明显降低 attention 中间激活带来的显存占用。
- 这意味着:
- 能撑更长序列;
- 能用更大 batch;
- 或者在同样显存下训练更大的模型。
- 所以 FlashAttention 一开始最重要的影响,其实是在训练侧。
(2)推理里也很重要
在推理阶段,尤其是 prefill 阶段,长上下文 attention 本身依然很重。
FlashAttention 仍然能带来明显收益。
不过在 decode 阶段,很多时候瓶颈又会更多落在 KV Cache 访问上。
所以在推理里,FlashAttention 和 KV Cache 往往是互补关系,而不是替代关系。
(3)和 KV Cache 怎么配合理解
我觉得最简单的理解方式是:
FlashAttention 主要优化“单次 attention 怎么算”;
KV Cache 主要优化“历史 K/V 怎么复用”。
二者关注点不同,但在大模型推理里通常会一起出现。
9. FlashAttention-2 在改什么
(1)为什么还要有 FlashAttention-2
初代 FlashAttention 已经把 IO 问题解决得很好了,但它并不意味着 GPU 利用率就已经到头。
FlashAttention-2 继续往前推进,关注的重点更偏工程实现和并行划分:
怎样让 GPU 上的线程块、warp 和 matmul 工作分配得更合理。
(2)核心改进方向
- FlashAttention-2 的主线大致可以概括成三件事:
- 减少不必要的非 matmul FLOPs;
- 提高不同线程块之间的并行度;
- 优化单个线程块内部不同 warp 的工作划分,减少共享内存通信。
所以它相对初代的重点已经有点变化:
初代更像“把 attention 变成 IO-aware”;
二代更像“在这个框架下,把 GPU 并行度和吞吐继续榨出来”。
(3)为什么它还是很重要
因为很多时候,真正影响实际速度的不是“有没有一个好算法”,而是“这个算法在硬件上吃得够不够满”。
FlashAttention-2 做的就是把这件事继续推进。
10. FlashAttention-3 又在改什么
(1)为什么还会有 3
到 FlashAttention-3 时,重点已经更偏向新硬件特性,尤其是 Hopper 架构 GPU。
它关注的是:
怎样利用新一代 GPU 的异步执行能力、Tensor Memory Accelerator(TMA)、以及低精度能力,把 attention 再往前推。
(2)主线变化
如果很粗地说:
- FlashAttention 1:解决 IO-aware exact attention;
- FlashAttention 2:优化并行划分和吞吐;
- FlashAttention 3:进一步拥抱新硬件特性与低精度执行。
所以这一条线越往后,越明显能看出来:
它已经不只是一个“attention 小技巧”,而是在逐渐变成 attention 的高性能实现主线。
11. FlashAttention 为什么这么重要
(1)它改变了大家看 attention 的方式
早期一提 attention,很多人第一反应就是二次复杂度,然后想到各种近似。
FlashAttention 提醒大家:
在真实硬件上,算法快不快,不只取决于 FLOPs,也取决于 IO。这一点其实很重要,因为它把问题从“数学复杂度”推进到了“系统复杂度”。
(2)它证明了 exact 也可以很快
以前很多人会下意识觉得:
想让 attention 快,就必须近似。FlashAttention 的价值就在于,它说明了一件事:
即使不近似,只要实现方式对,exact attention 也能快很多。
这个结论对后面的很多工作都很有启发。
(3)它已经变成基础设施级组件
现在很多训练框架、推理引擎、长上下文模型,一提高性能 attention,基本都会碰到 FlashAttention。
这说明它已经不只是某篇论文里的方法,而是逐渐变成了 Transformer 生态中的基础设施级组件。
12. FlashAttention 和近似注意力怎么理解
(1)它们解决的是不同问题
近似注意力方法的核心通常是:
从数学形式上减少 attention 的计算或存储开销。
比如稀疏化、低秩化、核方法等,都在试图减少理论复杂度。
FlashAttention 则不同,它基本保留标准 attention 数学形式,重点放在:
如何把标准 attention 在真实 GPU 上算得更高效。
(2)所以它不是“替代一切近似方法”
FlashAttention 很强,但它并不意味着以后所有长上下文问题都只靠 FlashAttention 就够了。
当序列真的极长时,二次计算本身仍然是问题。
所以在更极端的长度下,稀疏化、压缩、窗口化等方法仍然有意义。
更准确地说:
FlashAttention解决的是“标准 attention 该怎么高效实现”;
近似方法解决的是“当标准 attention 本身太贵时,要不要改数学形式”。
13. 从更高一层看,FlashAttention 到底是什么
(1)它不是新的注意力定义
我觉得理解 FlashAttention 最关键的一点是:
它不是新的 attention 公式,而是新的 attention 执行方式。
这个区别很重要。
因为它并没有改模型定义,而是改了模型在硬件上的落地方式。
(2)它本质上是“attention 的 IO 重排”
如果只用一句话概括 FlashAttention,我觉得最准确的是:
FlashAttention 是一种基于分块和在线 softmax 的 IO-aware exact attention 实现。
这个定义比单纯说“更快 attention”更准确。
(3)它代表了一种非常重要的研究思路
我觉得 FlashAttention 这条线真正厉害的地方,不只是把 attention 做快了,而是它代表了一种思路:
算法设计不能只盯数学公式,还得盯硬件的数据流。
这个思路在大模型时代会越来越重要。
14. 一点理解
(1)FlashAttention 最漂亮的地方
我觉得 FlashAttention 最漂亮的地方在于,它不是靠“偷懒少算一点”来变快,而是靠“别把没必要的东西搬来搬去”来变快。
这个想法其实非常朴素,但也非常有效。
(2)它为什么看起来像工程优化,却又很有方法味
因为它既有工程味,也有算法味。
如果只看表面,你会觉得它像 kernel 优化;
但再往里看,它其实重新设计了 attention 的执行顺序和 softmax 归约方式。
所以它不是纯工程 patch,而是一个很完整的方法。
(3)怎么记 FlashAttention
- 如果只是为了学习,我觉得可以把 FlashAttention 记成四句话:
- 标准 attention 慢,不只是因为二次复杂度,还因为大矩阵显存读写很重;
- FlashAttention 不显式存整个 attention matrix,而是按块计算并在线完成 softmax;
- 它保持和标准 attention 数学等价,因此是 exact attention;
- 所以它本质上是在做 attention 的 IO 优化,而不是近似替代。
15. 参考鸣谢
FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
https://arxiv.org/abs/2205.14135FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
https://arxiv.org/abs/2307.08691FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision
https://arxiv.org/abs/2407.08608
16. 注
- 这篇主要是个人学习整理,重点放在主线理解;
- 文中没有展开很多实现细节,比如 backward pass 的重计算策略、causal mask 的具体 block 处理方式、以及不同 GPU 架构下 kernel 设计的差异;
- 才疏学浅,欢迎批评、指导和交流;
- 有错误望大家及时指正!
