注意力瘦身术与长上下文
第 32 课的 KV 缓存救了推理速度,却留下一笔越滚越大的显存账。上下文从 8K 走向百万,这笔账怎么才能不把显卡压垮?工程师们给注意力动了四刀。
先看账单:KV 缓存到底有多大?
第 32 课我们用 KV 缓存换来了生成速度:已经算过的 K 和 V 存起来,下次直接取。这是一笔「拿空间换时间」的买卖——现在,该看看空间这一头到底有多贵了。
以一个 70B 级模型为例:80 层、64 个头、每头 128 维、FP16(2 字节)
每个 token:2 × 80 × 64 × 128 × 2 = 2,621,440 字节 ≈ 2.5 MB
上下文 128K(131,072 个 token):2.5 MB × 131,072 ≈ 320 GB
对比:模型权重本身 70B × 2 字节 = 140 GB —— 缓存比模型还大,而且每个并发用户都要一份
这就是长上下文最真实的门槛:不是算不动,而是存不下。上下文再涨 8 倍到 1M,这一项就要奔着 2.5 TB 去了。
第一刀:让几个头共用一份 K/V
标准的多头注意力(MHA)里,每个头都有自己独立的 Q、K、V(第 26 课)。可实验发现:Q 需要多样(每个头问不同的问题),但 K/V 没必要每个头都存一份。
GQA(8 组):8 个 KV 头 → 每 token 320 KB → 128K 上下文 ≈ 40 GB (小 8 倍)
MQA(1 组):1 个 KV 头 → 每 token 40 KB → 128K 上下文 ≈ 5 GB (小 64 倍)
共享得越狠,省得越多,但质量也会掉:MQA(只留 1 份)在很多任务上明显不如 MHA。GQA 是折中——分几组,就是在「省显存」和「保质量」之间拧旋钮。Llama 2 的 70B 版本起,GQA 几乎成了开放权重大模型的标配。可即便如此,8 组 GQA 的缓存仍与上下文长度成正比——想再往下压,就得换思路了。
第二刀:不存 K/V,存它们的「摘要」(MLA)
GQA 是「少存几份」。DeepSeek 在 V2 里提出的 MLA(Multi-head Latent Attention,多头潜在注意力)换了个问法:K 和 V 都是从同一个隐藏状态 h 算出来的,它们里面的信息真有那么多吗?
答案是:没有。它们大量重复,完全可以先压缩成一个很小的「潜向量」c,缓存时只存 c;真正做注意力时,再用一个矩阵把它还原成 K 和 V。这正是第 05 课线性变换的老本行——先把高维向量投到低维,需要时再升回去。
MLA:缓存 512(潜向量)+ 64(位置相关的键)= 576 个数
压缩比:32,768 ÷ 576 ≈ 57 倍
工程上还有两个精巧的细节:一是「矩阵吸收」——还原 K、V 的矩阵可以提前并进 Q 的计算里,所以推理时甚至不必真的把 K、V 还原出来;二是 RoPE 位置编码会破坏这种吸收,所以专门留出一小段 64 维不压缩的部分单独处理,这就是「512 + 64」的由来。论文报告 MLA 的效果不逊于标准 MHA,同时缓存小了一个数量级——这是近两年开放模型里最有影响力的架构创新之一。
第三刀:少看一点——学出来的稀疏注意力
前两刀都在「存得少」。第三刀问的是:每个新 token 真的需要回看全部历史吗?第 32 课的滑动窗口是个固定规则(只看最近 W 个),简单,却容易漏掉远处的关键信息。
新一代稀疏注意力让模型自己学着挑:先用一个又小又便宜的「索引器」给所有历史 token 粗打一遍分,只选出得分最高的 k 个,再对这 k 个做真正的注意力。DeepSeek 在 V3.2 里把这套做法带到了实用规模。
真正的注意力计算量:131,072 ÷ 2,048 = 缩小到 1/64
(索引器本身还要扫一遍全部历史,但它又小又快,常数很小)
挑错了就永远看不到那条被漏掉的信息——所以索引器必须和主模型对齐训练,并且要保证在各种任务里都不漏掉关键位置。另外,稀疏只砍了计算,KV 缓存本身通常还是得整个存着(除非再配合别的手段)。
第四刀:换个记法——线性注意力与状态空间
前三刀不管怎么砍,缓存都随上下文长度线性增长。第四刀更彻底:让记忆的大小根本不随长度变。
关键的数学恒等式来自矩阵乘法的结合律。标准注意力先算 (Q·Kᵀ) 得到一个 n×n 的得分表,再乘 V。如果把中间的 softmax 拿掉(近似),就可以换个顺序:先算 Kᵀ·V,得到一个固定大小的 d×d 矩阵 S,再让 Q 去乘它。而 S 可以一个 token 一个 token 地累加:
取 d = 2:token 1:k = (1, 0),v = (2, 1) → S₁ = [[2, 1], [0, 0]]
token 2:k = (0, 1),v = (1, 3) → S₂ = S₁ + [[0, 0], [1, 3]] = [[2, 1], [1, 3]]
查询 q = (1, 1):o = q · S₂ = (2 + 1, 1 + 3) = (3, 4)
不管已经读了多少个 token,S 永远只是这 4 个数——记忆大小固定。
发现了吗?这个「逐个 token 更新一个固定大小的状态」的结构,就是第 23 课的 RNN!区别在于:RNN 训练时只能串行,而线性注意力把公式展开后可以像 Transformer 一样并行训练,推理时又能像 RNN 一样每步只花常数成本。2023 年提出的 Mamba(选择性状态空间模型)走的是同一条精神路线;后来的 DeltaNet、Gated DeltaNet 等又加入了「边写边擦」的机制——写入新信息前,先擦掉状态里已有的旧关联,让有限的记忆更耐用。
固定大小的状态意味着有损压缩:读得越多,每条信息分到的「格子」就越少。它擅长抓大意,却不擅长逐字精确回忆——比如在十万字里找回「第三段提到的那个电话号码」。这个短板,在检索类任务上尤其明显。
现实的答案:混着用
既然全注意力擅长精确检索、线性层擅长便宜地读长文,何必二选一?近一年多家模型不约而同地采用了混合架构:大部分层用线性注意力或状态空间层,每隔几层夹一层全注意力,常见比例在 3:1 到 7:1 之间。
在 GQA(8 组)的基础上再乘 1/4:128K 上下文 ≈ 40 GB ÷ 4 = 10 GB
这还没算再叠加 MLA、量化——几刀叠起来,缓存能瘦下一两个数量级
把前面几刀放在一起比一比。拖动滑块改变上下文长度,再切换 GQA 的分组数,看每种方案要占多少显存:
另一头是让模型能读得那么长:位置编码(RoPE)训练时只见过有限长度,直接外推会「晕」;办法是位置插值 / YaRN 之类的缩放技巧,把更长的位置「压」回它熟悉的范围,再用少量长文本微调。
窗口大 ≠ 用得好。宣称支持百万 token 的模型,常在「大海捞针」这种简单测试上表现漂亮,但一遇到需要综合多处信息的任务就明显退化;研究还发现模型对放在中间的内容格外容易忽略(「中间迷失」)。所以有效上下文,往往比标称值短得多——这也是下一课的智能体要靠「上下文工程」而不是一味塞满的原因。
总结
KV 缓存 = 2 × 层数 × KV 头数 × 头维度 × token 数 × 字节,长上下文的门槛是存不下。四刀:GQA 少存几份 K/V,MLA 存低维摘要,稀疏注意力少看历史,线性注意力 / 状态空间换成固定大小的状态;现实里常把它们混合使用。
💡 用大白话梳理:这一课的核心直觉
- 账单:70B 级模型 128K 上下文,MHA 的 KV 缓存 ≈ 320 GB,比权重还大。
- GQA:几个 Q 头共用一份 K/V,缓存缩到 1/8 ~ 1/64,质量损失小,已是标配。
- MLA:缓存低维潜向量,用时再升维;DeepSeek-V3 每 token 每层 576 个数 vs 32,768,约 57 倍。
- 稀疏:索引器先粗筛 top-k,只对被选中的 k 个做精确注意力,计算量 ∝ n·k。
- 线性 / 状态空间:S = S + kvᵀ,固定大小的状态,回到 RNN 的老路但可并行训练;代价是有损、不擅长精确回忆。
- 混合 + 现实:少数全注意力层负责检索,其余层负责便宜地读长文;标称窗口远大于有效窗口。
学习小测验
动动脑筋:核心直觉小测验(选出你的答案后点击「提交」,即可查看生动通俗的详细解析)
38 推理提速:量化与投机解码
用更少的位数存权重,让小模型打草稿、大模型验收。