LESSON 37 · 卷IV 大语言模型

注意力瘦身术与长上下文

第 32 课的 KV 缓存救了推理速度,却留下一笔越滚越大的显存账。上下文从 8K 走向百万,这笔账怎么才能不把显卡压垮?工程师们给注意力动了四刀。

第 1 站

先看账单:KV 缓存到底有多大?

第 32 课我们用 KV 缓存换来了生成速度:已经算过的 K 和 V 存起来,下次直接取。这是一笔「拿空间换时间」的买卖——现在,该看看空间这一头到底有多贵了。

KV 缓存 = 2(K 和 V)× 层数 × KV 头数 × 每头维度 × token 数 × 每个数占的字节
以一个 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 去了。

上面这条乘法里有 6 个因子,哪些能想办法砍掉?
第 2 站

第一刀:让几个头共用一份 K/V

标准的多头注意力(MHA)里,每个头都有自己独立的 Q、K、V(第 26 课)。可实验发现:Q 需要多样(每个头问不同的问题),但 K/V 没必要每个头都存一份。

MHA 多头注意力Q 头(8 个)K/V 头(8 份要缓存)缓存 100%GQA 分组(8 头 → 2 组)Q 头(8 个)K/V 头(2 份要缓存)缓存 25%MQA 全体共享Q 头(8 个)K/V 头(1 份要缓存)缓存 12.5%
图 37-1三种共享方式。MHA:8 个 Q 头对应 8 份 K/V;GQA:分成 2 组,每组 4 个 Q 头共用 1 份 K/V;MQA:所有 Q 头共用 1 份。缓存量按 K/V 头数等比例缩小。
MHA:64 个 KV 头 → 每 token 2.5 MB → 128K 上下文 ≈ 320 GB
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 的缓存仍与上下文长度成正比——想再往下压,就得换思路了。

第 3 站

第二刀:不存 K/V,存它们的「摘要」(MLA)

GQA 是「少存几份」。DeepSeek 在 V2 里提出的 MLA(Multi-head Latent Attention,多头潜在注意力)换了个问法:K 和 V 都是从同一个隐藏状态 h 算出来的,它们里面的信息真有那么多吗?

答案是:没有。它们大量重复,完全可以先压缩成一个很小的「潜向量」c,缓存时只存 c;真正做注意力时,再用一个矩阵把它还原成 K 和 V。这正是第 05 课线性变换的老本行——先把高维向量投到低维,需要时再升回去。

隐藏状态 h7168 维降维潜向量 c512 维 · 只缓存它RoPE 键 · 64 维(也缓存)用时升维K:128 头 × 128 维(不落盘,用完即弃)V:128 头 × 128 维(不落盘,用完即弃)每个 token 每层:MHA 缓存 32,768 个数 → MLA 只缓存 512 + 64 = 576 个数(约 57 倍压缩)思路 = 存「摘要」,用的时候再展开——第 05 课的线性变换:先投到低维,再还原
图 37-2MLA 的思路:缓存的是压缩后的潜向量(加一小段位置信息),而不是庞大的 K 和 V。DeepSeek-V3 里是 128 个头 × 128 维,压缩后每个 token 每层只需缓存 576 个数。
MHA:每 token 每层缓存 2 × 128 头 × 128 维 = 32,768 个数
MLA:缓存 512(潜向量)+ 64(位置相关的键)= 576 个数
压缩比:32,768 ÷ 576 ≈ 57 倍

工程上还有两个精巧的细节:一是「矩阵吸收」——还原 K、V 的矩阵可以提前并进 Q 的计算里,所以推理时甚至不必真的把 K、V 还原出来;二是 RoPE 位置编码会破坏这种吸收,所以专门留出一小段 64 维不压缩的部分单独处理,这就是「512 + 64」的由来。论文报告 MLA 的效果不逊于标准 MHA,同时缓存小了一个数量级——这是近两年开放模型里最有影响力的架构创新之一。

第 4 站

第三刀:少看一点——学出来的稀疏注意力

前两刀都在「存得少」。第三刀问的是:每个新 token 真的需要回看全部历史吗?第 32 课的滑动窗口是个固定规则(只看最近 W 个),简单,却容易漏掉远处的关键信息。

新一代稀疏注意力让模型自己学着挑:先用一个又小又便宜的「索引器」给所有历史 token 粗打一遍分,只选出得分最高的 k 个,再对这 k 个做真正的注意力。DeepSeek 在 V3.2 里把这套做法带到了实用规模。

上下文 n = 131,072,每个新 token 只精挑 k = 2,048 个历史位置
真正的注意力计算量:131,072 ÷ 2,048 = 缩小到 1/64
(索引器本身还要扫一遍全部历史,但它又小又快,常数很小)
⚠️ 遇到的拦路虎 · 方案局限

挑错了就永远看不到那条被漏掉的信息——所以索引器必须和主模型对齐训练,并且要保证在各种任务里都不漏掉关键位置。另外,稀疏只砍了计算,KV 缓存本身通常还是得整个存着(除非再配合别的手段)。

第 5 站

第四刀:换个记法——线性注意力与状态空间

前三刀不管怎么砍,缓存都随上下文长度线性增长。第四刀更彻底:让记忆的大小根本不随长度变。

关键的数学恒等式来自矩阵乘法的结合律。标准注意力先算 (Q·Kᵀ) 得到一个 n×n 的得分表,再乘 V。如果把中间的 softmax 拿掉(近似),就可以换个顺序:先算 Kᵀ·V,得到一个固定大小的 d×d 矩阵 S,再让 Q 去乘它。而 S 可以一个 token 一个 token 地累加:

St = St−1 + kt · vtᵀ  ot = qt · St
取 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 等又加入了「边写边擦」的机制——写入新信息前,先擦掉状态里已有的旧关联,让有限的记忆更耐用。

⚠️ 遇到的拦路虎 · 方案局限

固定大小的状态意味着有损压缩:读得越多,每条信息分到的「格子」就越少。它擅长抓大意,却不擅长逐字精确回忆——比如在十万字里找回「第三段提到的那个电话号码」。这个短板,在检索类任务上尤其明显。

第 6 站

现实的答案:混着用

既然全注意力擅长精确检索、线性层擅长便宜地读长文,何必二选一?近一年多家模型不约而同地采用了混合架构:大部分层用线性注意力或状态空间层,每隔几层夹一层全注意力,常见比例在 3:1 到 7:1 之间。

线1线2线3全4线5线6线7全8线9线10线11全12全注意力层(精确回看,需要 KV 缓存)线性注意力 / 状态空间层(固定大小状态)混合架构示例:每 4 层里只有 1 层是「全注意力」(3 : 1)
图 37-3混合架构:12 层里只有 3 层是全注意力,其余是线性层。需要 KV 缓存的只有那 3 层,其余层的状态大小固定。
80 层里只有 1/4 是全注意力层(3 : 1)→ 需要缓存 KV 的只有 20 层
在 GQA(8 组)的基础上再乘 1/4:128K 上下文 ≈ 40 GB ÷ 4 = 10 GB
这还没算再叠加 MLA、量化——几刀叠起来,缓存能瘦下一两个数量级

把前面几刀放在一起比一比。拖动滑块改变上下文长度,再切换 GQA 的分组数,看每种方案要占多少显存:

LAB · 31-A
KV 缓存计算器:一个 70B 级模型,上下文越长,各种方案要占多少显存?
上下文长度–
GQA 的 K/V 组数–
模型设定:80 层、64 个 Q 头、每头 128 维、FP16(2 字节)。条形按对数刻度绘制(每格 10 倍),单位 1 GB = 1024³ 字节;橙色 = 超过一张 80GB 显卡。MLA 按 DeepSeek-V3 的 512 + 64 维潜向量估算;线性注意力按每头 128×128 的固定状态估算。注意:这里的模型只有 64 个头,所以 MLA 约压缩 28 倍;DeepSeek-V3 有 128 个头,正文里的「约 57 倍」是按它的配置算的。

另一头是让模型能读得那么长:位置编码(RoPE)训练时只见过有限长度,直接外推会「晕」;办法是位置插值 / YaRN 之类的缩放技巧,把更长的位置「压」回它熟悉的范围,再用少量长文本微调。

⚠️ 遇到的拦路虎 · 方案局限

窗口大 ≠ 用得好。宣称支持百万 token 的模型,常在「大海捞针」这种简单测试上表现漂亮,但一遇到需要综合多处信息的任务就明显退化;研究还发现模型对放在中间的内容格外容易忽略(「中间迷失」)。所以有效上下文,往往比标称值短得多——这也是下一课的智能体要靠「上下文工程」而不是一味塞满的原因。

第 7 站

总结

💡 章节速记 · 本课核心

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 的老路但可并行训练;代价是有损、不擅长精确回忆。
  • 混合 + 现实:少数全注意力层负责检索,其余层负责便宜地读长文;标称窗口远大于有效窗口。
小测验

学习小测验

动动脑筋:核心直觉小测验(选出你的答案后点击「提交」,即可查看生动通俗的详细解析)

Q1一个 70B 级模型做 128K 上下文推理,KV 缓存往往比模型权重还大。下面哪一项不是缩小 KV 缓存的办法?
Q2线性注意力为什么能让「记忆大小不随上下文长度增长」?它主要的代价是什么?