一、500 个 token,500 次前向传播
第 21 讲结尾留了一句话:从分布里挑一个词几乎不花计算,但挑出的每一个词,都要让整个模型再完整跑一遍。
把这件事写得具体一点。用户问了一个 100 个 token 的问题,模型要生成 500 个 token 的回答:
第 1 步: 输入 100 个 token → 算出第 101 个 token 的分布 → 挑一个词
第 2 步: 输入 101 个 token → 算出第 102 个 token 的分布 → 挑一个词
第 3 步: 输入 102 个 token → ...
...
第 500 步: 输入 599 个 token → 算出第 600 个 token 的分布
每一步都把整段序列喂进模型,每一层都要对每一个位置重新算 Q、K、V、注意力、前馈层——然后只取最后一个位置的输出,其他全部扔掉。
这听起来就很浪费。这一讲讲三件让推理变便宜的事:KV Cache 省掉重复计算,量化让模型变小,批处理让一台机器同时服务很多人。讲到第二件之前,会先揭示一个反直觉的事实:逐词生成阶段,瓶颈往往不是算得慢,而是读得慢。
二、打个比方:续写一本书
你在续写一本已经写了 300 页的小说。每写一句新的,你需要回顾前文的人物、伏笔、语气。
笨办法:每写一句新的,就从第一页开始把整本书重读一遍。
好办法:第一次读的时候,给每一页做一张小卡片,记下这一页的要点。之后每写一句新的,只需要把卡片翻一遍,再给新写的这句补一张卡片。
好办法快得多,但也有代价:卡片越积越多,占满了桌面。写到第 1000 页,桌上就有 1000 张卡片。
KV Cache 就是这些卡片。下面先讲为什么"卡片"是可行的——前面的页,为什么可以只读一次。
三、KV Cache:旧位置算一次就够
为什么旧位置不用重算
关键在第 11 讲的因果掩码:位置 i 只能看到位置 1 到 i,看不到后面的位置。
那么,当第 t+1 个 token 被接到序列末尾时,前面 t 个位置会不会受影响?不会。 它们本来就看不到排在后面的位置,新来的这个 token 对它们来说根本"不存在"。所以,在每一层里,前 t 个位置算出来的所有东西——包括它们的 K 和 V——和新 token 来之前一模一样。
用第 9、11 讲的例子直接验证。x₁ = [1, 0],Wq = Wk = 单位矩阵,Wv = 2 × 单位矩阵。
第一步,序列里只有 x₁:
Q₁ = [1, 0] K₁ = [1, 0] V₁ = [2, 0]
位置 1 只能看自己 → 输出₁ = V₁ = [2, 0]
把 K₁、V₁ 存起来。
第二步,x₂ = [0, 1] 来了:
只为新位置计算: Q₂ = [0, 1] K₂ = [0, 1] V₂ = [0, 2]
从缓存里取出: K₁ = [1, 0] V₁ = [2, 0] ← 不重算
位置 2 的注意力: 用 Q₂ 和 [K₁, K₂] 打分 → 权重 [0.330, 0.670]
输出₂ = 0.330 × V₁ + 0.670 × V₂ = [0.660, 1.340]
把 K₂、V₂ 追加进缓存。
这和第 11 讲把两个位置一起算出来的结果完全一样:位置 1 是 [2, 0],位置 2 是 [0.660, 1.340]。第 11 讲里我们还特意验证过,加上掩码之后,位置 1 的输出不受 x₂ 影响——这正是它可以被缓存的原因。
为什么只存 K 和 V,不存 Q?因为一个位置的 Q,只用来计算它自己的输出。旧位置的输出已经算完了,它们的 Q 以后再也用不上;新位置需要的,只是它自己的 Q,加上所有位置的 K(用来打分)和 V(用来加权求和)。
KV Cache(键值缓存):在逐词生成时,把每一层里每个已生成位置的 Key 和 Value 保存下来。新 token 到来时,只为它计算 Q、K、V,并直接复用缓存中的历史 K、V 完成注意力计算。
省了多少
生成 n 个 token,不用缓存时,第 t 步要把 t 个位置全部处理一遍,总共处理的"位置次数"是 1 + 2 + … + n。用缓存后,每一步只处理 1 个新位置。取 n = 500:
不用缓存:1 + 2 + … + 500 = 500 × 501 / 2 = 125,250 次
用缓存: 500 次
投影矩阵和前馈层这些最重的矩阵乘法,计算量相差约 250 倍。(注意力那一部分省不了这么多:新位置的 Q 依然要和缓存里所有的 K 打分,这部分随长度增长,第 23 讲会回来算它。)
两个阶段:一次读完,然后一个一个写
有了 KV Cache,生成过程被自然地分成两个性质完全不同的阶段:
预填充(prefill):处理用户输入的那 100 个 token。它们全部已知,可以像第 13 讲的训练那样,一次前向传播并行算完,顺便把它们的 K、V 全部存进缓存。
解码(decode):之后每生成一个 token,只处理这 1 个新位置,一个接一个,严格串行。
⭐ 第 13 讲说"训练并行,推理串行"。更准确的说法是:推理的预填充阶段也是并行的,真正串行的是解码阶段。而一个回答的大部分时间,往往花在解码阶段。
卡片占了多少桌面
KV Cache 省了计算,但要占显存。一个 token 的缓存有多大?每一层、每一个头,都要存一个 K 向量和一个 V 向量:
每个 token 的缓存 = 2(K 和 V)× 层数 × 头数 × 每个头的维度 × 每个数字占的字节
用一组示意性的数字算一下:32 层,每层 32 个头,每个头 128 维(第 10 讲:总维度 32 × 128 = 4096),每个数字用 16 位(2 字节)存储:
每个 token: 2 × 32 × 32 × 128 × 2 字节 = 524,288 字节 = 0.5 MB
4096 个 token 的一段对话: 0.5 MB × 4096 = 2 GB
同时服务 16 段这样的对话: 2 GB × 16 = 32 GB
作为对照,这样规模的模型大约有几十亿个参数,16 位存储时权重本身是十几 GB。KV Cache 可以比模型本身还大。
这就是第 10 讲预告的那笔账:“头的数量和每个头的维度,直接决定了要缓存的 K、V 占多少显存”。它也是第 6 讲那个"存储换计算"权衡的又一次出现:第 6 讲是训练时存下中间结果,反向传播才不用重算;这里是推理时存下 K、V,生成下一个词才不用重算。两次都是用显存换计算。
⚠️ 关于时效性:一个被广泛采用的缓解办法是让多个查询头共用同一组 K、V。比如 32 个查询头只配 8 组 K、V,缓存就缩小到四分之一,代价是注意力的表达能力略有损失。具体共用多少,各个模型的选择不同,也在变化;但"缓存大小由 K、V 的头数决定,而不是由 Q 的头数决定"这个结论不会变。
四、真正的瓶颈:读得慢,不是算得慢
解码阶段在等什么
解码阶段每一步只处理 1 个 token。对这 1 个 token,模型要把每一层的每一个权重矩阵都用一遍——而每个权重,只和这一个 token 做一次乘法、一次加法,就再也用不上了。
问题在于,权重存放在显存里,计算发生在计算单元里,每次使用前都得先把它从显存搬过去。于是每生成一个 token,都要把全部权重搬一遍,而每搬一个数字,只做两次运算。
用示意性的数字估算:权重 14 GB,显存每秒能读出约 1 TB 的数据:
读一遍全部权重:14 GB ÷ 1000 GB/秒 ≈ 0.014 秒
每秒最多生成: 1 ÷ 0.014 ≈ 70 个 token
⭐ 不管计算单元有多快,只要每生成一个 token 都要把 14 GB 读一遍,单个请求每秒就生成不了超过约 70 个 token。计算单元大部分时间在空等数据送过来。这种状态叫受内存带宽限制(memory-bound)。
对比预填充阶段:100 个 token 一起处理,每个权重搬过来一次,就和 100 个 token 各做一次运算。同样的搬运量,做了 100 倍的活——那里瓶颈在计算,叫受计算限制(compute-bound)。
这个事实决定了接下来两个办法的方向:要么让需要搬的东西变少(量化),要么让搬一次能服务更多的 token(批处理)。
顺带算清第 12 讲的一笔账
第 12 讲说 MoE “计算量解耦了,显存没有解耦”。现在可以说得更完整:解码阶段的瓶颈是读权重,而不是算。一个请求每步只路由到少数几个专家,读的量确实少了。但一台服务器同时处理很多请求时,不同请求会被路由到不同的专家,批次一大,几乎每个专家在每一步都会被某个请求用到——所有专家的权重,每一步都得读。MoE 省下的计算,在这个瓶颈面前省不下多少时间;而全部专家都得常驻显存这一点,一个字节也省不掉。
五、量化:让要读的东西变小
用更少的位数存数字
模型的权重通常用 16 位的浮点数存储。量化(quantization) 的想法是:用更少的位数——8 位,甚至 4 位——来近似表示这些数字。
4 位只能表示 16 个不同的值。怎么用 16 个值去近似任意的小数?最基本的做法是:找一个缩放系数,把一组权重映射到一排等距的整数格子上,只存格子编号,外加这组共用的一个缩放系数。
手算一次 4 位量化
取 4 个权重:[0.12, −0.50, 0.31, 0.95]。用对称的 4 位整数,取 −7 到 7 这 15 个格子:
① 缩放系数 = 最大的绝对值 ÷ 7 = 0.95 ÷ 7 ≈ 0.1357
② 每个权重除以缩放系数,四舍五入到整数(这就是要存的 4 位数):
0.12 ÷ 0.1357 ≈ 0.88 → 1
−0.50 ÷ 0.1357 ≈ −3.68 → −4
0.31 ÷ 0.1357 ≈ 2.28 → 2
0.95 ÷ 0.1357 = 7.00 → 7
③ 使用时,乘回缩放系数:
1 × 0.1357 ≈ 0.136 误差 +0.016
−4 × 0.1357 ≈ −0.543 误差 −0.043
2 × 0.1357 ≈ 0.271 误差 −0.039
7 × 0.1357 = 0.950 误差 0
每个数从 16 位变成 4 位,要读的数据量变成四分之一。按上一节的估算,每秒能生成的 token 数上限大约提高到原来的 4 倍;原本 14 GB 的权重,现在不到 4 GB,一张小得多的显卡就能装下。代价是每个权重都有一点误差。
一个离群值能毁掉整组
现在把第 4 个权重从 0.95 换成 5.0,其他不变:
缩放系数 = 5.0 ÷ 7 ≈ 0.714
0.12 ÷ 0.714 ≈ 0.17 → 0 → 还原为 0.000 误差 −0.12
−0.50 ÷ 0.714 ≈ −0.70 → −1 → 还原为 −0.714 误差 −0.214
0.31 ÷ 0.714 ≈ 0.43 → 0 → 还原为 0.000 误差 −0.31
5.0 ÷ 0.714 = 7.00 → 7 → 还原为 5.000 误差 0
⭐ 因为一个特别大的数,缩放系数被撑大了 5 倍多,格子之间的间距也跟着变大——其余三个小权重里,有两个直接被量化成了 0。这组权重原本携带的信息,大部分丢了。
而大模型的权重和中间激活值里,恰恰存在这样的离群值。这就是为什么实际的量化方法不会让很多权重共用一个缩放系数,而是把权重切成很小的组,每组各自配一个缩放系数,并且对离群值做专门处理。
同样的思路也可以用在 KV Cache 上:把缓存的 K、V 用更少的位数存储,第三节那 32 GB 也能跟着缩小。
六、批处理:读一次,服务很多人
第四节的结论是:解码时,每一步都要把全部权重读一遍,然后只为一个 token 做很少的计算。那么,如果同时有 16 个用户在等回答,能不能读一遍权重,同时为这 16 个人各算一个 token?
可以。这就是批处理(batching):把多个请求拼成一批,一起过模型。权重搬过来一次,和 16 个 token 各做一次运算——搬运量不变,活多了 16 倍。在计算单元被喂饱之前,每一步的耗时几乎不变,但总吞吐量接近提高到 16 倍。这正是第 3 讲说的"矩阵乘法是对一批向量同时做同一种变换"在服务器上的直接应用。
批处理的上限卡在哪里?卡在第三节的 KV Cache 上。每一个请求都有自己的一整套缓存,批次每多一个请求,就多占一份。示意数字里,16 段 4096 token 的对话,光缓存就是 32 GB。显存装不下更多的缓存,就没法再往批次里加人。所以 KV Cache 的大小,直接决定了一台机器能同时服务多少人。 量化 KV Cache、共用 K/V 头,最终都是在为批处理腾地方。
⚠️ 关于时效性:实际的推理服务在批处理之上还有很多工程技巧。比如:请求的回答有长有短,不必等整批都生成完才处理下一批,而是某个请求一结束,就立刻把排队的新请求补进它的位置;KV Cache 不预先按最大长度整块分配,而是切成小块按需分配,思路和操作系统管理内存的分页一样。具体实现在快速演进,但目的都是同一个:在有限的显存里,塞进尽量多正在生成的请求。
七、代价
⚠️ KV Cache 随长度和并发数线性增长。回答越长、上下文越长、同时服务的人越多,缓存越大。它是限制上下文长度和并发数的首要因素,第 23 讲会从这里接着讲。
⚠️ 量化是有损的。手算已经显示了误差,离群值会让误差变得很严重。位数越少,损失越大;模型越小,通常越经不起量化。更麻烦的是,损失未必均匀地体现在所有能力上:常规测试的平均分可能只掉一点点,某些对精度敏感的能力——多步计算、长文本里的细节回忆——却可能明显变差,而且不容易被发现。
⚠️ 批处理用延迟换吞吐。对服务器来说,批次越大越划算;对单个用户来说,他的请求要和别人的一起排队、一起计算,每一步未必更快,有时还更慢。而且一个超长的请求会长时间占着大块缓存,挤占其他人的位置。
⚠️ 这三个办法都没有碰到注意力本身的 n²。KV Cache 让投影和前馈层不再重算,但新 token 的 Q 依然要和缓存里所有的 K 打分;上下文越长,这部分越重,缓存也越大。这个随长度增长的问题,是下一讲的主题。
八、和后面课程的关系
- 第 23 讲(上下文长度):本讲第七节的两条——缓存随长度线性增长、注意力打分依然随长度增长——加上第 8 讲的位置编码外推问题,合起来就是"长上下文为什么是工程难题"。
- 第 24 讲(测试时计算的代价):第 19 讲的推理模型要在回答前写几千个 token。按本讲的分析,这几千个 token 全部发生在串行的解码阶段、全部要占 KV Cache,这笔账下一讲之后细算。
- 第 26 讲(向量检索):大规模检索里存储海量向量,同样会用到"用更少的位数存数字"这个量化思路。
九、本讲小结
- KV Cache:因果掩码让旧位置永远看不到新 token,所以它们每一层的 K、V 算一次就不会再变,可以缓存复用。手算复用了第 9、11 讲的例子,结果与整段一起算完全一致。生成 500 个 token,投影和前馈层的计算从 125,250 次降到 500 次。
- 两个阶段:预填充一次并行处理整个输入;解码一个一个串行生成。真正串行的是解码。
- 缓存要占显存:
2 × 层数 × 头数 × 头维度 × 字节数,示意数字下每个 token 0.5 MB,16 段 4096 token 的对话要 32 GB,可以超过模型本身。这是第 6 讲"用显存换计算"的又一次出现。 - ⭐ 解码的瓶颈是读,不是算:每生成一个 token 都要把全部权重从显存读一遍,每个权重只做两次运算。示意数字下单个请求每秒最多约 70 个 token。MoE 在这个瓶颈面前省不下多少时间。
- 量化:用更少的位数存数字。手算 4 位量化,误差在 0.04 左右;一个离群值把缩放系数撑大 5 倍,两个小权重直接变成 0——所以要分小组、单独处理离群值。
- 批处理:读一次权重,同时为很多请求各算一个 token,吞吐量接近成倍提高;上限由 KV Cache 占的显存决定。
- ⚠️ 代价:缓存随长度和并发线性增长;量化有损且损失不均匀;批处理用单个用户的延迟换总吞吐;三者都没有解决注意力随长度增长的问题——下一讲。
思考题
- 用第三节的方法继续:第三个 token
x₃ = [1, 1]到来时,哪些量需要新算,哪些可以从缓存里取?写出位置 3 的注意力要用到的全部 K 和 V。 - 用第三节的缓存公式算:如果把 32 个 K/V 头改成 8 个(查询头仍是 32 个),每个 token 的缓存变成多少?16 段 4096 token 的对话一共要多少?
- 按第四节的估算,如果把权重量化到 8 位(7 GB),单个请求每秒最多生成多少 token?如果量化到 4 位呢?为什么说这个估算只是一个上限?
- 在第五节离群值的例子里,如果把 4 个权重分成两组——
[0.12, −0.50]一组、[0.31, 5.0]一组——每组各自算缩放系数,重新量化一遍。哪几个权重的误差变小了? - 一台服务器的显存,除去模型权重还剩 40 GB 给 KV Cache。按第三节的示意数字(每个 token 0.5 MB),如果每段对话平均 2000 个 token,最多能同时服务多少人?如果有一个用户的对话长达 32,000 个 token,他一个人会占掉多少个"平均用户"的位置?