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,那就在它们被算出来时存下来,每步追加一列。
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,而不是之前的
工作原理
工作原理
瓶颈从 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 的主题
动手观察
动手试试
上半部分展示机制,下半部分展示账单。切换缓存开关并单步执行,看计算量从二次方翻转成线性;再用显存计算器算一算,你的 GPU 能同时容纳多少个请求。
缓存内容(单层)
- 已缓存 —— 只读
- 本步计算
每步恰好新增一列。此前每个位置的 k 和 v 都只算过一次,之后就只被读取。
每步的注意力计算量
- 有缓存
- 0
- 否则将花费
- 0
有缓存时,每步开销线性增长(对 n 个位置做注意力)。没有缓存时是二次方增长(重算全部 n 个,每个又都对 n 个做注意力)。
- 字节 / token
- 131.1kB
- 缓存总量
- 16.0 GB
- 权重
- 14.9 GB
- 最大并发
- 126
显存预算 · 80 GB
- 权重(固定)
- 激活值 + 额外开销
- KV 缓存(随负载增长)
KV 缓存占了正在使用显存的 52%。把并发拉高直到变红,再把 KV 数据类型切成 int8,看着上限翻一倍。S08 讲的就是这笔交易。
选中 GPT-3 175B,注意它的 KV 开销:96 个注意力头且完全不分组,产生的缓存比大多数模型的权重还大。再选 Llama-3-70B,它有 64 个 query 头、8 个 KV 头。同一量级的模型,缓存小 14 倍。GQA 的理由就在这里,而且是显存上的理由,与质量无关。
亲手实现
实现
朴素的分配策略,是给每个请求分配一整块连续张量,大小按最大可能长度来。这一版仍然值得先写一遍:它失败的方式,亲手体会比读到更管用。
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 就是从这个测量结果里长出来的。$ python code/s05_kv_cache.py只需要 NumPy — 查看环境准备.
生产实践
在生产环境中
- picoLM 用 fp16 存缓存以减半体积,并用软件实现 fp32↔fp16 转换,因为目标 CPU 没有硬件支持。在树莓派上,这决定了模型跑不跑得起来。
- quant.cpp 走得更远,把缓存量化到 4 bit,并为最近 128 个 token 保留一个全精度窗口。注意力集中在近期上下文,精度就该花在那里。
- vLLM 压根不用连续缓存。它改用什么,就是下一章的内容。
练习
- 1亲手测一下浪费:用符合现实分布的长度跑 200 个请求,统计预留的 KV 显存中有多大比例真正被写入过。预计会看到 60–80% 的浪费。
- 2实现一个只保留最近 W 个 token 的滑动窗口缓存。注意什么会坏掉,以及为什么模型必须为窗口注意力而训练,这个做法才站得住。
- 3
继续学习
接下来
现在 decode 快了,瓶颈换成了显存,而显存的大部分躺在根本没人用的预留里。S06 用那个让 vLLM 出名的想法收拾这些预留:把 KV 缓存当成虚拟内存,给它分页。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
加上 KV 缓存后,从 N 个 token 的 prompt 生成 T 个 token 的注意力总计算量,从 O((N+T)³) 变成了什么?