上一章《推理与 KV Cache》说到注意力是 O(n²)——
换成真实数字就吓人了:128K 的上下文,注意力矩阵有 171 亿个元素。
这一章讲清楚——大家到底在用哪三种办法砍它。
注意力要做的事是:每个位置都要和每个位置比一次。n 个位置,就是 n×n 次比较。
互动 · 注意力矩阵到底有多大
互动 · n² 是「面积」:边长翻 32 倍,面积翻 1024 倍
所有优化都在做三件事之一。
互动 · 点开看每一类到底在改什么
互动 · 三类解法各自把哪些格子留了下来
| 类别 | 一句话 | 代表 | 质量损失 |
|---|---|---|---|
| ① 不存完整矩阵 | 算法上还是 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) 了,其实没有。 它降的是显存,不是计算量。
Softmax 有个看起来绕不过去的性质:它需要知道一整行的最大值和总和才能归一化。 你不知道整行,就不知道分母。
所以直觉上你会觉得"必须先把整行算完存下来"。FlashAttention 的突破就是用"在线 softmax" 绕过了这一点——一边读块,一边维护"到目前为止的最大值和总和",并随时修正之前的结果。
互动 · 在线 softmax 逐步演示
互动 · 分块算出来的 softmax,和完整算出来的一模一样
在线 softmax 的核心动作是:每当发现新的最大值,就把之前累加的结果整体乘上 exp(m_old − m_new)。
为什么可以这样?因为 softmax 有个不变性:给所有分数同时减掉同一个数,结果不变。
所以"先用旧的最大值算、后来发现不够大,就整体修正一下"是合法的。
这让算法只需要 O(n) 的显存,而不是 O(n²)——而结果和完整计算数学上等价。
FlashAttention 在反向传播时不保存那个 n×n 的注意力矩阵,
而是在需要时重新算一遍(这个技巧叫重计算 / activation checkpointing 的一种)。
所以它的实际计算量比标准注意力更多,只是显存少得多。
但它的净效果通常是变快——因为原来真正的瓶颈是"在显存和计算单元之间搬运那个 n×n 矩阵"
(也就是显存带宽),而不是算力本身。省掉搬运比多算一点划算得多。
上一节证明了在线 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 几乎不动。
互动 · 同一件事,三种实现的成本账单
在上面那排按钮里,在「标准注意力」和「FlashAttention」之间来回切: FLOPs 几乎不动,HBM 读写量塌下去一个量级。那个「一条数不动、另一条数塌了」的瞬间, 就是这一章的全部——它没有少算一次乘加,只是把搬家的次数砍掉了一个量级。
互动 · 少搬 9 倍,比「把带宽打满」更重要
| 存储层级(A100) | 容量 | 带宽 | 这一章里它是谁 |
|---|---|---|---|
| 寄存器 Register | 256 KB / SM | 比 SRAM 还快 | 累加器 O 就住在这里 |
| SRAM 共享内存(片上) | 192 KB / SM 全片约 20 MB |
约 19 TB/s | FlashAttention 说的"片上"就是它 |
| L2 | 40 MB | 介于两者之间 | K、V 重读时的缓冲垫 |
| HBM 显存 | 80 GB | 2.04 TB/s | 那个 n×n 矩阵的落脚处 |
SM(streaming multiprocessor)是 GPU 里一个小计算单元,自带一小块快存(片上 SRAM)。「单 SM 片上需求」量的就是一个 SM 的内存够不够装下这一趟要用的块。
SRAM 比 HBM 快约 9 倍,整片加起来却只有 20 MB——比 80 GB 的显存小约 4000 倍。
这个落差就是全部问题的根源:把数据搬到显存再搬回来,代价远高于在原地多算几次。
于是所有高效注意力的努力,最后都落在同一句话上:别让 n×n 那个矩阵落地。
| 实现 | FLOPs | HBM 读写 | 耗时 |
|---|---|---|---|
| 标准注意力 | 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 倍。
"把带宽打满"不是目标,"少搬"才是。省下的是搬运的次数,不是每次搬运的效率。
互动 · 三条复杂度曲线的真实交叉点
大 O 记号丢掉了常数因子。而常数因子在真实硬件上极其重要:
· 标准注意力是几次巨大的矩阵乘法——GPU 最擅长这个,常数因子很小
· 线性注意力 / SSM虽然复杂度低,但要逐个时间步串行递推,
每一步都有额外的开销,常数因子很大
公式里两边都是「每个 token 花的等效运算」:c₂ · n 是二次方法平均每个 token 的开销,
c_lin 是线性方法每 token 的固定开销,所以交叉点在 n = c_lin / c₂,
落在哪取决于这两个常数,而不是复杂度符号本身。
序列短到几千个 token 时,老老实实跑二次注意力往往更快。
| 方法 | 类别 | 核心机制 | 复杂度 |
|---|---|---|---|
| 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 缓存随上下文长度怎么涨
前面几节讲了三类解法,这一节只看最纯粹的那一类——FlashAttention 做的事只有一件:把那个 n×n 的矩阵拆成小块,算完一块就扔。 可 softmax 要看完一整行才知道分母——凭什么能一块一块算? 靠两个会一直滚动的数:m(最大的那个分数)和 l(指数和)。
互动 · 一块一块地读,矩阵永远不落地
于是峰值额外显存从 O(n²) 掉到 O(n)。 但「n² 那一项到底多吓人」,得放到同一根轴上看。
互动 · 峰值额外显存:一条是 n²,一条是 n×d
| 方法 | 代价 | 什么时候真的会痛 → 换什么 |
|---|---|---|
| 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×,几乎无损) |
看到一个新的"高效注意力"方法时,先问三个问题:
① 它属于哪一类?(是免费工程优化,还是有损近似)
② 它在什么序列长度上才开始赢?(交叉点在哪)
③ 它是在长上下文任务上评测的吗?
(很多方法在短序列基准上打平,一到长序列就露馅——因为长序列才是它们该被检验的地方)
三个答案连起来就是一句:先看它是有损还是免费,再看你的序列长度落在交叉点哪边,最后只信它在长序列上跑出来的数。
三类解法之所以正好是三类,是因为它们在三张不同的账单上各自动了一刀——而账单只有三张。
回到第 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 不落地,峰值额外 |
| 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²) 反而最快。下一章《投机解码与加速》:不砍矩阵,改让小模型先猜。
上面讲的都是「够用」的版本。想往下挖,这里有三个入口—— 它们不是必修内容,是给想再往前走一步的读者准备的。