跳到正文
LLM 推理

KV 缓存

没有缓存时,生成第 N 个 token 的代价是 O(N²);有了缓存则是 O(N)。瓶颈也随之从算力转移到显存带宽。

  • 缓存布局
  • prefill/decode 分离
  • 算术强度
  • 带宽墙

为什么重要

问题

把 S01 的引擎跑在真实 prompt 上做一次性能剖析。几乎所有时间都花在重算 key 和 value 上,而这些 token 自上一轮迭代以来根本没变过。模型是确定性的,也是因果的,过去无法被改写:第 500 个 token 的 key 向量,在第 501 步与第 500 步完全相同。

朴素循环照算不误:算 500 次,每层一次,每个请求一次。在未优化的引擎里,这是最大的一处浪费,也是最容易去掉的一处。

核心思路

解法

把 key 和 value 留着。位置 t 的注意力需要所有 ≤ t 位置的 K 和 V,那就在它们被算出来时存下来,每步追加一列。

带缓存的注意力python
def attention_with_cache(x, w, cache, layer, pos):
    q = x @ w.wq                        # [1, n_heads * head_dim]  (decode: 1 token)
    k = x @ w.wk                        # [1, n_kv_heads * head_dim]
    v = x @ w.wv
    q, k = rope(q, pos), rope(k, pos)   # rotate BEFORE storing

    cache.k[layer, pos] = k             # append — this is the whole trick
    cache.v[layer, pos] = v

    # attend over everything written so far, including what we just wrote
    k_all = cache.k[layer, : pos + 1]   # [pos+1, n_kv_heads, head_dim]
    v_all = cache.v[layer, : pos + 1]
    return grouped_query_attention(q, k_all, v_all)

三个后果立刻随之而来,并且塑造了余下的每一章:

  • 前向传播的形状变了。 prefill 的序列维是 N;decode 的序列维是 1。两者差别之大,以至于引擎会为它们编译各自独立的 kernel。
  • 引擎变成有状态的。 S01 式的引擎是个纯函数。带缓存的引擎持有按请求分配的显存,需要分配、追踪和释放,S06 讲的就是怎么把这件事做好。
  • 瓶颈转移了。 你几乎消除了所有冗余算术。剩下的是从显存里读权重和缓存,而余下的算术已经不足以把这段读取掩盖过去了。

要缓存 RoPE 之后的 K,而不是之前的

旋转只依赖绝对位置,而缓存中 token 的位置永远不变。存旋转后的 K,过去那部分就再也不必碰 RoPE。存旋转前的 K,每一步都要把整个缓存重新旋转一遍;这只是用一个稍小的计算开销换掉另一个,几乎什么都没省下来。

工作原理

工作原理

示意图同一份计算,有缓存与没缓存
没有缓存 —— 第 4 步对全部 4 个位置重算 k、vq0q1q2q3k0 k1 k2 k3算了 10 个注意力单元其中 9 个和第 3 步完全一样有缓存 —— 第 4 步只为 1 个新位置算 k、v;其余直接读q3k0 k1 k2 k3算了 4 个注意力单元只存在最后一行 queryKV CACHE每层、每请求:K← 每个 token 增长 1 格V← 每个 token 增长 1 格形状:[layers, 2, n_kv_heads, seq, head_dim]决定你并发度的那个数字字节数 = 2 × layers × n_kv_heads × head_dim × dtype_bytes × tokens那个 2 是 K 和 V;其余各项都由架构固定Llama-3-8B128 KB / tokenGPT-3 175B (MHA)4.5 MB / token这就是分组查询注意力(GQA)存在的理由。把 96 个 KV head 降到 8 个,缓存缩小 12×,其他什么都不变。现在稀缺的是缓存,不再是 FLOPs,也不再是权重。S06 到 S08 讲的全是怎么把它花好。每生成一个 token,decode 都要把整个缓存读一遍,所以缓存大小同时决定了你的 token 间延迟下限。
左边算了十个格子,其中九个与上一步完全重复。右边只剩最后一行 query:一次 decode 只做一次注意力,读的是它自己没有构建过的缓存。

瓶颈从 FLOPs 转移到带宽

这是本课程最重要的一个观念,值得说清楚。decode 阶段每产生一个 token,GPU 必须读取:

  • 模型中的每一个权重:8B 模型用 fp16 就是 16 GB;
  • batch 中每个请求的 KV 缓存的每一个字节。长上下文下,这一项可能超过权重本身。

而拿着这些数据,每个权重只做大约两次浮点运算。H100 提供约 3.3 TB/s 的显存带宽和约 1000 TFLOP/s 的 fp16 算力,比值约为每字节 300 次浮点运算。一次 decode 步只有约 1:两次浮点运算,对应读取每个 fp16 权重所花的两个字节。GPU 有超过 99% 的时间闲着。

8B fp16,batch 1
~5 ms
16 GB ÷ 3.3 TB/s,由物理设定的硬下限,代码搬不动
由此得到的上限
~200 tok/s
少读字节是越过它的唯一办法
出路
×batch
一次权重读取服务 N 个请求,这正是 S09 的主题

本章之后的每一项优化,都是从这个公式的某一侧下手。少读字节: 量化(S08)、KV 压缩、GQA。把这次读取摊到更多 token 上: 批处理(S09)、投机解码(S14)、MoE(S16)。

动手观察

动手试试

上半部分展示机制,下半部分展示账单。切换缓存开关并单步执行,看计算量从二次方翻转成线性;再用显存计算器算一算,你的 GPU 能同时容纳多少个请求。

模拟器KV 缓存:机制与显存
第一部分 · 缓存做了什么
每一步只为 1 个新 token 计算 k、v,其余直接读取
第 0 / 16 步

缓存内容(单层)

012345
  • 已缓存 —— 只读
  • 本步计算

每步恰好新增一列。此前每个位置的 k 和 v 都只算过一次,之后就只被读取。

每步的注意力计算量

点「运行」开始
有缓存
0
否则将花费
0

有缓存时,每步开销线性增长(对 n 个位置做注意力)。没有缓存时是二次方增长(重算全部 n 个,每个又都对 n 个做注意力)。

第二部分 · 缓存的代价
模型
kv 数据类型
字节 / token
131.1kB
缓存总量
16.0 GB
权重
14.9 GB
最大并发
126

显存预算 · 80 GB

w
kv
  • 权重(固定)
  • 激活值 + 额外开销
  • KV 缓存(随负载增长)

KV 缓存占了正在使用显存的 52%。把并发拉高直到变红,再把 KV 数据类型切成 int8,看着上限翻一倍。S08 讲的就是这笔交易。

选中 GPT-3 175B,注意它的 KV 开销:96 个注意力头且完全不分组,产生的缓存比大多数模型的权重还大。再选 Llama-3-70B,它有 64 个 query 头、8 个 KV 头。同一量级的模型,缓存小 14 倍。GQA 的理由就在这里,而且是显存上的理由,与质量无关。

亲手实现

实现

朴素的分配策略,是给每个请求分配一整块连续张量,大小按最大可能长度来。这一版仍然值得先写一遍:它失败的方式,亲手体会比读到更管用。

code/s05_kv_cache.py(节选)python
class KVCache:
    """One contiguous buffer per request. Simple, and disastrously wasteful."""

    def __init__(self, layers, max_seq, n_kv_heads, head_dim, dtype=np.float16):
        shape = (layers, max_seq, n_kv_heads, head_dim)
        self.k = np.zeros(shape, dtype=dtype)
        self.v = np.zeros(shape, dtype=dtype)
        self.length = 0

    @property
    def wasted(self):
        """Every slot past self.length is reserved and unusable by anyone else."""
        return 1.0 - self.length / self.k.shape[1]

    def append(self, layer, k, v):
        self.k[layer, self.length] = k
        self.v[layer, self.length] = v

    def view(self, layer):
        return self.k[layer, : self.length], self.v[layer, : self.length]

max_seq 是一个你兑现不了的承诺

你无法预知这次生成会跑多久,只能按最坏情况分配。一个按 max_seq=8192 预留、却在 40 个 token 后就停下的请求,浪费掉 99.5% 的预留量,而这块显存在该请求的整个生命周期里别人都用不了。在真实负载上测量,KV 显存的 60–80% 都消失在这里。PagedAttention 就是从这个测量结果里长出来的。
在本地运行
把同一次生成跑两遍,一次关闭缓存、一次开启,并断言两次输出逐 token 完全一致。然后打印计算量剖析、加速比表格,以及若干真实模型配置的 KV 显存计算器。
$ python code/s05_kv_cache.py
预期输出: 两条路径输出完全一致(缓存是优化,不是近似)、注意力计算量大幅下降,以及一张显存表:超过约 12 万 token(8k 上下文下约 15 个并发请求)后,缓存开始超过权重。

只需要 NumPy — 查看环境准备.

生产实践

在生产环境中

  • picoLM 用 fp16 存缓存以减半体积,并用软件实现 fp32↔fp16 转换,因为目标 CPU 没有硬件支持。在树莓派上,这决定了模型跑不跑得起来。
  • quant.cpp 走得更远,把缓存量化到 4 bit,并为最近 128 个 token 保留一个全精度窗口。注意力集中在近期上下文,精度就该花在那里。
  • vLLM 压根不用连续缓存。它改用什么,就是下一章的内容。

练习

  1. 1
    亲手测一下浪费:用符合现实分布的长度跑 200 个请求,统计预留的 KV 显存中有多大比例真正被写入过。预计会看到 60–80% 的浪费。
  2. 2
    实现一个只保留最近 W 个 token 的滑动窗口缓存。注意什么会坏掉,以及为什么模型必须窗口注意力而训练,这个做法才站得住。
  3. 3
    在开关后面加上 fp16 和 int8 两种缓存数据类型,并在一段固定文本上测量困惑度变化。这是你为 S08 收集的第一个数据点。

继续学习

接下来

现在 decode 快了,瓶颈换成了显存,而显存的大部分躺在根本没人用的预留里。S06 用那个让 vLLM 出名的想法收拾这些预留:把 KV 缓存当成虚拟内存,给它分页。

自测

习题

先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。

自测 1 题 / 共 6

加上 KV 缓存后,从 N 个 token 的 prompt 生成 T 个 token 的注意力总计算量,从 O((N+T)³) 变成了什么?

得分 0/6