阶段 6 · 让模型跑起来

高效注意力:那个 n²
到底怎么砍

上一章《推理与 KV Cache》说到注意力是 O(n²)—— 换成真实数字就吓人了:128K 的上下文,注意力矩阵有 171 亿个元素。
这一章讲清楚——大家到底在用哪三种办法砍它。

1

先把 n² 换成真实的数字

注意力要做的事是:每个位置都要和每个位置比一次。n 个位置,就是 n×n 次比较。

互动 · 注意力矩阵到底有多大

矩阵元素数
—
FP16 显存(单层单头)
—
相对 4K 的倍数
—

互动 · n² 是「面积」:边长翻 32 倍,面积翻 1024 倍

2

所有优化只做三件事

所有优化都在做三件事之一。

互动 · 点开看每一类到底在改什么

互动 · 三类解法各自把哪些格子留了下来

类别一句话代表质量损失
① 不存完整矩阵 算法上还是 O(n²),但不把 n×n 矩阵写进显存,分块算、算完就扔 FlashAttention 零损失——数学上完全等价
② 不给每个位置都算 改变"谁看谁"的规则,让计算量真的降下来 滑窗 / 稀疏 / 线性注意力 / Mamba 有损——取决于假设是否符合真实依赖
③ 不重复存 K、V 不改注意力本身,只改缓存的组织方式 GQA / MQA / MLA(上一章) 轻微——GQA 几乎无损

GQA / MQA / MLA 都在减少 K、V 的份数:GQA 让几组头共用一份,MQA 让所有头共用一份,MLA 把 K、V 压成一个低维向量。上一章《推理与 KV Cache》讲过这三种折中。

🎯 类比

把注意力想象成一场所有人都要跟所有人打招呼的聚会。n 个人,就要发生 n² 次握手。 人数翻倍,握手次数变四倍——这就是 n² 的来源。

① 不存完整矩阵 = 照样握手,但握完就忘,不在名单上一条条记录下来。 所以"搬运名单"的开销没了——聚会本身没变,只是不用誊抄名单了。
② 不给每个位置都算 = 改规矩:只跟身边的几个人握手(滑窗), 或者只跟主持人握(全局 token)。握手次数真的少了,但你也真的漏掉了一些人。
③ 不重复存 K、V = 名单不变,只是几个人共用一本通讯录,省的是本子不是握手。

这个类比只管「谁和谁互相看」,不管「分块存取」——FlashAttention 省的是后者。

① 是工程优化:结果一模一样,只是实现方式变了。所以多数场景值得用,代价见第 3、7 节。
② 是建模假设:它在赌"远处的词不重要"或者"注意力可以用更简单的函数近似"。 赌对了省很多,赌错了质量就掉。
③ 是表示压缩:在"表达能力"和"显存"之间做交易。

最容易混淆的是 ① 和 ②——很多人以为 FlashAttention 把复杂度降到 O(n) 了,其实没有。 它降的是显存,不是计算量。

3

不存矩阵,怎么算对 softmax

Softmax 有个看起来绕不过去的性质:它需要知道一整行的最大值和总和才能归一化。 你不知道整行,就不知道分母。

所以直觉上你会觉得"必须先把整行算完存下来"。FlashAttention 的突破就是用"在线 softmax" 绕过了这一点——一边读块,一边维护"到目前为止的最大值和总和",并随时修正之前的结果。

互动 · 在线 softmax 逐步演示

互动 · 分块算出来的 softmax,和完整算出来的一模一样

💡 关键在那个"重新缩放"

在线 softmax 的核心动作是:每当发现新的最大值,就把之前累加的结果整体乘上 exp(m_old − m_new)。
为什么可以这样?因为 softmax 有个不变性:给所有分数同时减掉同一个数,结果不变。 所以"先用旧的最大值算、后来发现不够大,就整体修正一下"是合法的。
这让算法只需要 O(n) 的显存,而不是 O(n²)——而结果和完整计算数学上等价。

⚠️ 但它不是免费的:用算力换显存

FlashAttention 在反向传播时不保存那个 n×n 的注意力矩阵, 而是在需要时重新算一遍(这个技巧叫重计算 / activation checkpointing 的一种)。
所以它的实际计算量比标准注意力更多,只是显存少得多。
但它的净效果通常是变快——因为原来真正的瓶颈是"在显存和计算单元之间搬运那个 n×n 矩阵" (也就是显存带宽),而不是算力本身。省掉搬运比多算一点划算得多。

4

算的没变,搬的少了一个量级

上一节证明了在线 softmax 算得对。但“对”不是 FlashAttention 的卖点—— 如果只要算得对,标准实现早就对了。它的卖点写在论文标题里:IO-Aware—— IO 就是显存读写(论文 Dao et al. 2022)。它按“搬得少”来设计,而不是按“算得少”。

GPU 的算力比显存带宽"富余"得多:一张 A100 有 312 TFLOPs 的 FP16 算力, 却只有 2.04 TB/s 的显存带宽;而片上的 SRAM 比显存快约 9 倍、却小 4000 倍。
所以真正的瓶颈常常不是"算得多快",而是"数据搬得多快"。 下面这个互动把一次注意力的账单摊开:算多少(FLOPs)、搬多少(HBM 读写)、 片上要同时装下多少(SRAM)。
切换三种实现,盯住两条数——搬运量会塌下去,FLOPs 几乎不动。

互动 · 同一件事,三种实现的成本账单

HBM 读写量
—
FLOPs
—
峰值额外显存
—
单 SM 片上需求
—
带宽下界耗时
—

💡 这一章全部的内容都藏在这张图里

在上面那排按钮里,在「标准注意力」和「FlashAttention」之间来回切: FLOPs 几乎不动,HBM 读写量塌下去一个量级。那个「一条数不动、另一条数塌了」的瞬间, 就是这一章的全部——它没有少算一次乘加,只是把搬家的次数砍掉了一个量级。

互动 · 少搬 9 倍,比「把带宽打满」更重要

存储层级(A100)容量带宽这一章里它是谁
寄存器 Register256 KB / SM比 SRAM 还快 累加器 O 就住在这里
SRAM 共享内存(片上)192 KB / SM
全片约 20 MB
约 19 TB/sFlashAttention 说的"片上"就是它
L240 MB介于两者之间 K、V 重读时的缓冲垫
HBM 显存80 GB2.04 TB/s 那个 n×n 矩阵的落脚处

SM(streaming multiprocessor)是 GPU 里一个小计算单元,自带一小块快存(片上 SRAM)。「单 SM 片上需求」量的就是一个 SM 的内存够不够装下这一趟要用的块。

SRAM 比 HBM 快约 9 倍,整片加起来却只有 20 MB——比 80 GB 的显存小约 4000 倍。
这个落差就是全部问题的根源:把数据搬到显存再搬回来,代价远高于在原地多算几次。 于是所有高效注意力的努力,最后都落在同一句话上:别让 n×n 那个矩阵落地。

实现FLOPsHBM 读写耗时
标准注意力 66.6 GFLOPs 40.3 GB 41.7 ms
FlashAttention 75.2 GFLOPs +13% 4.4 GB −89% 7.3 ms −82%

论文实测(Dao et al. 2022 图 2):GPT-2 medium,n = 1024,d = 64,16 头 × batch 64, A100,含前向 + 反向。口径与上面单头前向的公式不同,看比例就好。

⚠️ 一个可以自己算出来的反直觉数字

把上面那张表除一下,看每毫秒搬了多少:
· 标准注意力:40.3 GB ÷ 41.7 ms = 966 GB/s,是峰值的 47%
· FlashAttention:4.4 GB ÷ 7.3 ms = 602 GB/s,只有峰值的 30%
FlashAttention 的带宽利用率其实更低——但它搬得少 9 倍,所以快 5.7 倍。
"把带宽打满"不是目标,"少搬"才是。省下的是搬运的次数,不是每次搬运的效率。

5

短序列时 O(n²) 反而最快

互动 · 三条复杂度曲线的真实交叉点

大 O 记号丢掉了常数因子。而常数因子在真实硬件上极其重要:
· 标准注意力是几次巨大的矩阵乘法——GPU 最擅长这个,常数因子很小
· 线性注意力 / SSM虽然复杂度低,但要逐个时间步串行递推, 每一步都有额外的开销,常数因子很大

公式里两边都是「每个 token 花的等效运算」:c₂ · n 是二次方法平均每个 token 的开销, c_lin 是线性方法每 token 的固定开销,所以交叉点在 n = c_lin / c₂, 落在哪取决于这两个常数,而不是复杂度符号本身。 序列短到几千个 token 时,老老实实跑二次注意力往往更快。

6

三大类里各有哪些具体方法

方法类别核心机制复杂度
FlashAttention v1/v2/v3① 不存 分块 + 在线 softmax + 重计算 O(n²) 计算,O(n) 显存
PagedAttention③ 不重复存 把 KV Cache 分页管理,消除碎片(像操作系统管内存那样按页分配) O(n²) 计算,显存利用率↑
滑动窗口
Longformer / Mistral
② 不算 每个位置只看最近 w 个位置 O(n·w)
稀疏注意力
BigBird
② 不算 滑窗 + 随机连接 + 少量全局 token(每个位置都看得到的哨兵) O(n)
线性注意力
Performer、Linformer
② 不算 用核函数把 softmax 拆开,实现结合律重排(核函数 = 一个把分数映射到非负数的简单函数) O(n)
Mamba / SSM(状态空间模型)② 不算 不用注意力,用带选择性门控的状态递推(门控按当前输入决定记多少、忘多少) O(n)
GQA / MQA / MLA③ 不重复存 多个查询头共享 K、V 计算不变,KV 显存 ↓ 视档位而定(GQA ~4~8×、MLA ~15×、MQA ~头数×)

它们不是互斥的选项。今天一个真实的大模型推理栈通常同时用着三种: GQA(少存 K、V)+ FlashAttention(不存注意力矩阵)+ PagedAttention(管好显存碎片)。 而 滑窗 / Mamba 这类是"换模型结构"级别的选择,通常由训练阶段就决定好了, 不能事后加在已训好的模型上。

互动 · ③ 不重复存:KV 缓存随上下文长度怎么涨

M

数学 · 一边读、一边滚的两个数

前面几节讲了三类解法,这一节只看最纯粹的那一类——FlashAttention 做的事只有一件:把那个 n×n 的矩阵拆成小块,算完一块就扔。 可 softmax 要看完一整行才知道分母——凭什么能一块一块算? 靠两个会一直滚动的数:m(最大的那个分数)和 l(指数和)。

互动 · 一块一块地读,矩阵永远不落地

同时留在显存里的中间结果
—
已经算过的块
—

于是峰值额外显存从 O(n²) 掉到 O(n)。 但「n² 那一项到底多吓人」,得放到同一根轴上看。

互动 · 峰值额外显存:一条是 n²,一条是 n×d

7

每类解法各自要付什么

方法代价什么时候真的会痛 → 换什么
FlashAttention 反向传播要重算注意力——多花算力换显存。实现复杂, 而且对 GPU 架构很敏感(这解释到为什么有 v1/v2/v3 这么多版本) 要自己手写 kernel、或换到非 NVIDIA 硬件 → 用 PyTorch 的 torch.nn.functional.scaled_dot_product_attention(按硬件自动选 FlashAttention 实现),把版本适配交给框架
滑动窗口 直接假设"远处不重要"。如果任务需要长距离依赖 (比如"把上面第 3 段那句话引用过来"),它就真的看不见了—— 这是信息层面的丢失,不是精度问题 要引用远处的原文(长文档问答)→ 保留几层全注意力,或给更多全局 token
线性注意力 在真实语言任务上的质量通常不如全注意力。 纯线性做主干的大模型仍少见;已有公开模型把线性 / 状态空间层和少量全注意力层交替使用,本页不比较它们的好坏。 (Mamba 用状态空间 + 选择性门控绕开了一部分问题,但在纯检索类任务上仍有短板) 要精确检索某个历史 token → 别用纯线性,至少留几层全注意力
稀疏注意力 稀疏模式是人工设计的假设("邻居重要""某些全局 token 重要"), 而真实的注意力模式是学出来的、可能和假设不符 真实依赖跟人工假设对不上 → 换成由内容决定看哪里(能被数据学出来的路由)
MQA 极度压缩 K、V,质量下降,训练也更不稳定。 所以实际多用 GQA 这个折中 质量掉得受不了 → 退到 GQA(约 4~8×,几乎无损)
⚠️ 一个判断标准

看到一个新的"高效注意力"方法时,先问三个问题:
① 它属于哪一类?(是免费工程优化,还是有损近似)
② 它在什么序列长度上才开始赢?(交叉点在哪)
③ 它是在长上下文任务上评测的吗? (很多方法在短序列基准上打平,一到长序列就露馅——因为长序列才是它们该被检验的地方)

三个答案连起来就是一句:先看它是有损还是免费,再看你的序列长度落在交叉点哪边,最后只信它在长序列上跑出来的数。

8

小结

三类解法之所以正好是三类,是因为它们在三张不同的账单上各自动了一刀——而账单只有三张。

它对应哪条线 ③ 规模会赢(七条主线里的第 ③ 条;下面暗线表用的 A、C、D 是另一套字母,不是这里的编号)——这一章是那条线最硬的一张算力账单。
注意力对序列长度收的是平方税:n 翻一倍,账单翻四倍。 「规模会赢」这句话里因此藏着一张指数级的门票, 而三类解法就是三种「怎么付得起这张票」的办法。
一句话 这一章的做法,本质上是在承认:一次注意力的成本,主要不在「算了多少次」,而在「搬了多少字节」。
于是同一个 softmax,数学一行没改,只换了个搬运顺序(分块、不落地), 就能快一个数量级。反过来,想真的把 FLOPs 从 n² 改成 n—— 就必须改数学,就必然丢东西。
这两条路的分岔,就是第 2 节那三类解法的分界线。
它牺牲了什么 三类各有各的账单,没有一类是白拿的(见上面本质卡):
② 不给每个位置都算(稀疏 / 线性):牺牲表达能力——滑动窗口看不见远处, 线性注意力把历史压成一个固定大小的状态,长序列上质量明显掉。
① 不存完整矩阵(FlashAttention):数学一点没丢,但牺牲了实现的简洁—— 要手写 kernel(GPU 上直接跑的算子)、对 GPU 架构极敏感(所以才有 v1 / v2 / v3)。
换句话说:它把「模型质量」的钱,换成了「工程复杂度」的钱。
🎬 检验一下:你现在能指着哪个互动说这句话

回到第 4 节那张「同一件事,三种实现的成本账单」。 最上面那排按钮,在「标准注意力」和「FlashAttention」之间来回切—— 切一次,盯住下面两行数字:FLOPs 和 HBM 读写量。

FLOPs 几乎一动不动,HBM 读写量塌下去一个量级。
再把上面「序列长度 n」那根滑块拖大,看两条数的差距怎么被拉开—— 搬的数量随 n 长得多快,就是三类解法各自在跟什么东西赛跑。
如果你刚才没有指着那两行数字说这句话,那这一节对你就是没用的——回去再切一次。

互动 · 交叉点公式:把两个常数拖出来(悬停公式,点亮对应滑块)

它在暗线里站在哪

暗线这一章的回答
A 信息流动 输入 [n, d] 被注意力变成 [n, n] 的分数矩阵,再乘回 V 变回 [n, d]——那个 [n, n] 就是全部的麻烦。FlashAttention 只做一件事:让 [n, n] 永远不成形,形状一路是 [n, d]。
C 参数账本 分数矩阵 = 2 · n² 字节(与头维 d 无关):n = 4096 时 33.6 MB,n = 128K 时 34.4 GB,单个头、单个层;而一个 7B 模型有 32 层 × 32 头 = 1024 个这样的矩阵。FlashAttention 不落地,峰值额外显存从 O(n²) 降到 O(n),n = 128K 时那一项从 34 GB 变成 33.5 MB(约 1000 倍)。
D 跑在什么上 这一章撞的是带宽墙:A100 的 FP16 算力 312 TFLOPs,HBM 带宽 2.04 TB/s,片上 SRAM 19 TB/s。这里的算术强度指每从显存搬 1 字节做几次运算;低于 312 / 2 = 156 的运算都撞在带宽墙上,而标准注意力的中间步(写 [n,n] 再读回来)恰好在这里——最贵的动作就是一趟 HBM 往返。完整的 Roofline 账在 《硬件与算力账本》。
一句话带走高效注意力

三类解法——① 不存完整矩阵(FlashAttention,零损失)、② 不给每个位置都算(滑窗/稀疏/线性/SSM,有损)、③ 不重复存 K、V(GQA/MLA/PagedAttention)——各自的账不同,但都由同一个 n×n 矩阵逼出来。FlashAttention 没有少算一次乘加(论文实测 FLOPs +13%),它省的是显存读写(HBM −89%、耗时 −82%),所以选型前先问自己的序列长度——短序列时 O(n²) 反而最快。下一章《投机解码与加速》:不砍矩阵,改让小模型先猜。

9

拓展阅读

上面讲的都是「够用」的版本。想往下挖,这里有三个入口—— 它们不是必修内容,是给想再往前走一步的读者准备的。

📄 这一章的说法从哪来

💻 工业界怎么写

∑ 更严格的形式