报错非常突然:CUDA out of memory。我盯着终端看了半天,想不明白——模型只有十几 GB,显卡是 32GB 的 V100,怎么算都不该爆。直到我把序列长度从 32k 换回 4k,一切恢复正常。那一刻我才意识到,问题根本不在模型大小,而在 attention 自己。

你也许知道 Transformer 的注意力复杂度是 O(N²)。但 O(N²) 意味着什么,很多人(包括从前的我)没真正想过:它不只是算得慢,它还要在显存里实实在在放下一张 N×N 的矩阵。序列一长,光这张矩阵就能把显存撑爆。
FlashAttention 这篇论文 2022 年由斯坦福的 Tri Dao 等人发表,核心承诺是:把注意力中间结果的显存占用从 O(N²) 降到 O(N)。所以标题说“砍掉一半”其实太保守了——N 每增大 10 倍,标准实现要多吃 100 倍显存,Flash Attention 只多吃 10 倍。
一张 N×N 的大表,就是显存爆掉的元凶
注意力计算的标准实现长这样:假设序列有 N 个 token,每个 token 的表示维度是 d。Query、Key、Value 三个矩阵都是 [N, d]。第一步算 S = QKT(通常还会除以 √d 做缩放,这里不影响主线),得到 [N, N] 的分数矩阵;第二步对每行做 softmax,得到 P = softmax(S);第三步用 P 去加权 V,得到输出 O。
问题出在 S 和 P 这两个中间矩阵上。它们会完整地写进显存(HBM)。PyTorch 里,你在内存里就能看到它们的形状。N = 32k,fp16 精度,一个矩阵就是 32K × 32K × 2 字节 ≈ 2GB。S 和 P 两份,就是 4GB。而实际模型还有多头——每个 head 都是一张独立的 N×N 矩阵,显存占用直接再乘上一个 head 数。序列长度一旦上万,显存爆掉完全不奇怪。
这还只是前向传播。训练的时候,反向传播需要 S 和 P 的梯度,标准实现会把中间激活值原样存起来。所以训练时的峰值显存比推理更狠。
真正的瓶颈不是算力,是从仓库到灶台的搬运
GPU 里有两种存储,它们的差距经常被忽略:
| 对比维度 | SRAM | HBM(显存) |
|---|---|---|
| 位置 | 芯片内部缓存 | 独立显存颗粒 |
| 容量(A100) | 约 20MB | 40GB 或 80GB |
| 读写带宽 | 约 19 TB/s | 约 1.5-2 TB/s |
| 读一次 N=32k 的 S 矩阵 | 放不下 | 约 4GB,约 2.7ms |
一句话:SRAM 快得吓人但小得可怜,HBM 大得慷慨但慢得让人焦虑。为什么要关心带宽?因为 attention 是典型的 memory-bound 算子。A100 的 FP16 算力约 312 TFLOPs,HBM 带宽却只有 1.5TB/s 左右。GPU 从 HBM 读一个数要等很久,算完一个数却只要一瞬间。再加上标准实现里 S 和 P 在 HBM 中的往返读写,你会发现大部分时间根本不是在算,而是在等数据。
打个比方:一位厨师站在灶台前,面前是一块 20MB 的案板(SRAM),食材存放在仓库(HBM)。标准实现的做法是:切一批菜,放回仓库,再搬一批,再放回去。厨师的手艺再好,时间也全耗在路上了。
Flash Attention 的全部秘密:分块、在线归一化、重计算
Flash Attention 的工程实现思路一句话:别把中间结果搬回仓库,在案板上就地处理完。具体是三个动作:
- 分块(Tiling):把 Q、K、V 切成小块,让一块 Q 和一块 K 的乘积以及后续的中间结果,整个能塞进 SRAM。
- 在线归一化(Online Softmax):每一块先算自己的 softmax 分子分母,同时维护一个全局修正系数,保证分块结果和一次算完整个矩阵的结果完全一致。
- 重计算(Recomputation):反向传播时不保存 S 和 P,只保存每行的最大值和总和,反向时重新算一遍。
三个动作加在一起,HBM 的访问量从 O(N²) 降到了 O(N)。那张 N×N 的矩阵,从头到尾没离开过芯片。
分块算出来的 softmax,凭什么和整行算的分毫不差
我第一次看 online softmax 的时候,第一反应是:等等,softmax 要看到整行才能算,你分块怎么知道前面有没有一个更大的数?
答案是:不知道,但可以修正。数值稳定的 softmax 写法通常是先找这行的最大值 m,然后算 exp(x − m) 的和。假设你第一块拿到 [90, 80, 70],以为 m = 90;结果第二块蹦出来一个 95。没关系,你只要把之前所有按 90 平移的结果,整体乘上一个 exp(90 − 95),就等于它们一开始就是按 95 平移的。
Flash Attention 就是不断做这件事:每来一块新分数,就更新当前最大值 m 和累加总和 l,如果新最大值更大,就把之前的输出整体乘一个修正因子。全部块走完,结果和一次算完 softmax 分毫不差。这个“在线归一化”技巧不是 Flash Attention 首创——2018 年的论文就提出了,但把它嵌入分块 attention 并做到工程级优化,是这篇论文的贡献。
反向传播也玩了小心思:不存矩阵,重新算
PyTorch 训练时会把中间激活值存下来,这个显存需求同样是 N×N。Flash Attention 的做法是:反向传播时用保存的 m 和 l 把 S 重新算一遍,再求梯度。
代价是多花约三成的前向计算量。但这笔交易非常划算——因为前面说了,GPU 的算力本来就闲着,真正紧张的是显存带宽。用富余的资源换短缺的资源,这是系统设计里最经典的一课。
它不是一个近似算法——这是最容易误解的地方
我最早读 Flash Attention 时,默认它是“稀疏注意力”那一类的东西:通过某种近似跳过计算,换速度,牺牲精度。直到我读完论文,才发现自己错得离谱。
FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
论文标题里的关键词是 Exact。它和 Longformer、Linformer、Performer 那些高效注意力方法有本质区别:那些方法改变了注意力矩阵的结构(稀疏化、低秩化),而 Flash Attention 没有。它输出的数学结果和标准 softmax attention 完全一致,唯一的差别只是浮点数舍入顺序。
为什么这点很重要?因为“省显存”往往意味着“失真”。Flash Attention 的省法,跟精度没有半点关系。你可以无痛替换现有模型里的注意力实现,不需要重训,效果完全不变。这正是它被整个行业快速接受的原因。
为什么这么顺理成章的想法,到 2022 年才出现?
分块、在线归一化、重计算,单拆开来看,都是计算机系统里的老技术。为什么 2022 年才有人把它们用在 attention 上?难在交叉。做机器学习的人操作的是 PyTorch 的抽象,很少会去关心矩阵乘法在 GPU 内部怎么流转;做系统的人又未必熟悉 Transformer 的数学细节。Flash Attention 的作者 Tri Dao,当时在 Christopher Ré 的组里,做的就是打通系统和 ML 边界的事。这个生态位,之前恰好没有人站。
2024 年,这篇论文拿到了 NeurIPS 的 Test of Time Award。这不是偶然——它打开的是一条“IO-aware”的优化路线。后来的 FlashAttention-2、FlashAttention-3 都是沿着这条路继续走,vLLM 里的 PagedAttention 也是对数据物理路径的极致抠搜。
Flash Attention 的边界,和它没解决的那个问题
说几句它的局限。Flash Attention 不是一个通用优化:它深度依赖 NVIDIA GPU 的 SRAM 层级和 CUDA 生态,换到其他硬件效果完全两说。如果你用的是 CPU 推理,或者模型很小、序列很短,收益基本为零,反而引入额外复杂度。
更大的局限在这里:它把注意力中间的显存成本从 O(N²) 降到了 O(N),但 attention 的计算量 O(N²) 依然在。FlashAttention-2/3 只是把常数因子压得更低。序列长度到几十万、上百万的时候,即使显存不再爆,计算时间照样爆炸。长上下文模型的未来,还需要比“IO-aware attention”更根本的突破。
补充:Flash Attention 前向计算的完整流程
- 把 Q 切成 N/M 块,K 和 V 切成 N/M 块(M 由 SRAM 容量决定)。
- 对每个 Q 块,初始化行最大值 m = −∞、归一化总和 l = 0、输出块 O = 0。
- 遍历所有 K/V 块:在 SRAM 里算 S_block;更新 m 为当前见过的最大值;用修正系数把 l 和 O 重新缩放;再把 P_block × V_block 累加进 O。
- 所有块遍历完,O 写回 HBM。
原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/651.html