基础原理第 9 篇,共 13 篇
KV Cache:复用每一层的历史 K/V
增量解码为什么只需要处理新 token:KV Cache 的原理、显存代价、Transformers 中的实现,以及围绕它展开的系统优化
自回归语言模型一次只生成一个 token。每生成一个,就要把它接回输入,再预测下一个。如果每一步都把整段前缀重新送进模型,前面已经算过的位置会被反复计算,生成越长浪费越多。KV Cache 解决的就是这个问题:它把每一层 Attention 已经算好的 key 和 value 保存下来,新 token 只需要计算自己的那一份,再去查询历史。
这个做法在 Transformer 提出时就隐含在解码器的结构里(Vaswani 等, 2017):解码器的自注意力带有因果遮罩,每个位置只能读取它左边的位置,所以历史位置的表示不会因为后面多了 token 而改变。今天几乎所有的推理框架都默认开启这项缓存,它也因此成了长上下文服务里最主要的显存开销之一。围绕它已经形成了一系列工作:减少 key/value 头数的 MQA(Shazeer, 2019)和 GQA(Ainslie 等, 2023),按页管理缓存的 PagedAttention(Kwon 等, 2023),以及对缓存做量化和淘汰的方法(Hooper 等, 2024;Zhang 等, 2023;Xiao 等, 2023)。
本文先说明缓存为什么成立、里面存的是什么,再估算它的显存代价,然后看 Hugging Face Transformers 里的具体实现,最后简要介绍缓存引出的几类系统优化。阅读前最好先了解 Attention 中的 Q、K、V 和 Prefill 与 Decode 的时序。
为什么可以缓存
Section titled “为什么可以缓存”在权重、位置编码和遮罩都固定,并且计算是确定性的前提下,因果模型里某个位置的输出只依赖它自己和它左边的位置。往序列末尾追加 token,不会改变任何历史位置在任何一层的 hidden state,也就不会改变由这些 hidden state 投影出来的 K 和 V。因此历史位置的 K/V 可以算一次、存起来、反复使用。
需要强调的是,缓存省掉的只是历史位置的重复计算。新 token 仍然要完整地经过 Embedding、每一层的 Attention 和 MLP,并不是“只查缓存”。
从复杂度看,如果每一步都对长度为 T 的整段前缀重新做 Attention,单步代价约为 O(T²),生成 N 个 token 累计为 O(N³)。使用缓存后,每一步只有一个新的 query 去读取 T 个历史 key,单步约为 O(T),累计为 O(N²)。这只是 Attention 部分的简化估计,不能直接当作整个模型的耗时比例,实际加速取决于模型结构、硬件和负载。
缓存里存的是什么
Section titled “缓存里存的是什么”缓存保存的是每一层自注意力的 K 和 V,形状为 [B, heads, seq_len, head_dim]。有三点容易混淆:
- 不保存 Q。 query 只在当前这一步使用,之后的步骤只需要新位置的 query,所以不必留存。
- 每层各存一份。 不同层的输入表示不同,K/V 的投影参数也不同,数值自然不同。并不存在一块所有层共用的 KV 张量。
- 它不是文字,也不是权重。 缓存里是中间激活值,既不是已生成 token 的列表,也不是模型参数。
用这一组文章共用的例子说明(设定见基础原理总览:B=1,2 个头,每头 4 维)。Prefill 阶段输入 [11, 23] 两个 token 后,每层的 K 和 V 各是 [1, 2, 2, 4]。模型选出下一个 token 37 时,缓存并不会立刻变长,因为 37 此时只是一个 ID。下一步把 [37] 送进模型,每层才算出它的 K 和 V,形状各为 [1, 2, 1, 4],拼接到历史后面成为 [1, 2, 3, 4]。这一步的 Q 只有新位置,形状是 [1, 2, 1, 4]。
Prefill 与 Decode
Section titled “Prefill 与 Decode”一次生成通常分成两个阶段。Prefill 把整个 prompt 一次送入,建立初始缓存;之后的每一步 Decode 只送入一个新 token。
| 阶段 | 本步输入 | Q 的形状 | 新 K/V 的形状 | 缓存长度 |
|---|---|---|---|---|
| Prefill | 整个 prompt,共 T 个 token | [B, H, T, d] | [B, H, T, d] | 从 0 增长到 T |
| Decode 第 1 步 | 1 个新 token | [B, H, 1, d] | [B, H, 1, d] | T + 1 |
| Decode 第 2 步 | 1 个新 token | [B, H, 1, d] | [B, H, 1, d] | T + 2 |
Decode 时的 Attention 是一个新 query 对全部历史 key 做点积:
Q [B, H, 1, d] × Kᵀ [B, H, d, T_total] = scores [B, H, 1, T_total]“Decode 时 Q 的长度是 1”只对普通的逐 token 解码成立。分块 prefill 一次会送入多个位置。推测解码(speculative decoding)先由小模型起草若干 token,再由大模型在一次前向中并行验证,这时一步里也有多个 query 位置(Leviathan 等, 2022)。
缓存的大小可以直接算出来:
缓存字节数 = 2 (K 和 V) × 层数 × B × KV 头数 × 序列长度 × head_dim × 每个元素的字节数以两个模型为例,batch 为 1:
| 模型 | 配置 | 100 token | 1,000 token | 8,000 token |
|---|---|---|---|---|
| GPT-2 small | 12 层,12 个头,head_dim 64,fp32 | 7.4 MB | 74 MB | 超出上下文长度 |
| Llama 3 8B | 32 层,8 个 KV 头,head_dim 128,bf16 | 13 MB | 131 MB | 1.05 GB |
GPT-2 small 的上下文长度是 1,024 个 token,Llama 3 8B 是 8,192 个 token。
缓存随 batch 大小和上下文长度线性增长。对 Llama 3 8B 来说,bf16 权重约 16 GB,单条 8,000 token 请求的缓存约 1 GB,权重仍是大头;但同时服务十几条这样的请求,缓存总量就会超过权重。缓存是否成为瓶颈,取决于模型、精度、batch 和序列长度的组合。
表中 Llama 3 8B 的 KV 头数只有 8,而它的 query 头数是 32(Grattafiori 等, 2024)。这正是为缓存做的结构设计。Shazeer (2019) 指出,增量解码的瓶颈在于反复读取体积很大的 K/V 张量,于是提出 multi-query attention(MQA):所有头共用一组 key 和 value,只保留各自的 query。Ainslie 等 (2023) 提出的 grouped-query attention(GQA)取了一个折中,把 query 头分成若干组,每组共用一套 K/V。在上面的公式里,这两种做法都直接减小了“KV 头数”这一项。Llama 3 8B 用 8 个 KV 头服务 32 个 query 头,缓存只有不分组时的四分之一。
实现:Transformers 的 DynamicCache
Section titled “实现:Transformers 的 DynamicCache”下面以 Hugging Face Transformers 的默认实现为例(官方文档)。类名和接口在不同版本间有变化,细节以所安装的版本为准。
DynamicCache 是一个容器,内部为每一层维护一个 DynamicLayer,每个 DynamicLayer 保存这一层的 keys 和 values:
DynamicCache └── layers: list[DynamicLayer] 每层一个 ├── layers[0] │ ├── keys: [B, heads, seq_len, head_dim] │ └── values: [B, heads, seq_len, head_dim] ├── layers[1] └── ...对 GPT-2 small,缓存了 10 个 token 之后,cache.layers[0].keys.shape 是 [B, 12, 10, 64],values 相同。
更新:沿序列维拼接
Section titled “更新:沿序列维拼接”每次前向时,Attention 层把新算出的 K/V 交给缓存,缓存把它们接到历史后面:
def update(self, key_states, value_states, *args, **kwargs): if not self.is_initialized: self.lazy_initialization(key_states, value_states)
# dim=-2 是序列维:把新位置接在历史位置之后 self.keys = torch.cat([self.keys, key_states], dim=-2) self.values = torch.cat([self.values, value_states], dim=-2) return self.keys, self.values缓存对象创建时还不知道张量的 dtype 和所在设备,所以采用懒初始化:第一次 update 时,根据传入的张量创建空的 keys 和 values,之后才开始拼接。
def lazy_initialization(self, key_states, value_states): self.dtype, self.device = key_states.dtype, key_states.device self.keys = torch.tensor([], dtype=self.dtype, device=self.device) self.values = torch.tensor([], dtype=self.dtype, device=self.device) self.is_initialized = True与 Attention 的配合
Section titled “与 Attention 的配合”在 Attention 层里,缓存的使用分三步。下面是简化后的流程:
def forward(self, hidden_states, past_key_values=None, ...): # 1. 只为本步输入的位置计算 Q、K、V Q, K_new, V_new = compute_qkv(hidden_states)
# 2. 把新的 K、V 追加进本层的缓存,取回“历史 + 新位置”的完整 K、V if past_key_values is not None: K, V = past_key_values.update(K_new, V_new, self.layer_idx)
# 3. 新位置的 Q 对完整的 K、V 做 Attention attn_output = attention_fn(Q, K, V, ...)layer_idx 告诉缓存当前是第几层,缓存据此操作对应的那一份 K/V。
from transformers import DynamicCache
cache = DynamicCache(config=mdl.config)
# 把缓存交给模型,模型会在前向过程中更新它with torch.no_grad(): out = mdl(input_ids=ids, past_key_values=cache, use_cache=True)
cache.get_seq_length() # 已缓存的序列长度len(cache.layers) # 层数,GPT-2 small 为 12cache.layers[0].keys.shape # [B, heads, seq, head_dim]
cache.crop(max_length=50) # 截断到前 50 个 token,用于回溯cache.batch_repeat_interleave(n) # 复制 batch 维,beam search 会用到cache.batch_select_indices(indices) # 选出指定的 batch 条目相关源码集中在 src/transformers/cache_utils.py,其中定义了 DynamicCache、DynamicLayer、StaticCache、QuantizedCache 和 EncoderDecoderCache。
几种缓存类型
Section titled “几种缓存类型”| 类型 | 做法 | 适用情况 | 代价 |
|---|---|---|---|
DynamicCache | 用 torch.cat 动态增长 | 默认选择,长度不受限 | 每步都要重新分配内存 |
StaticCache | 预先分配最大长度的缓冲区,原地写入 | 张量形状固定,便于配合 torch.compile | 始终占用最大长度的显存,收益依赖具体负载 |
QuantizedCache | 把 K/V 量化到更低的位宽 | 显存紧张的长上下文 | 精度略有损失 |
EncoderDecoderCache | 分别保存自注意力和交叉注意力的缓存 | T5、BART 这类 Encoder–Decoder 模型 | 不适用于 decoder-only 模型 |
静态缓存加编译不是通用的最优解,是否更快需要在自己的负载上测量。
缓存引出的系统问题
Section titled “缓存引出的系统问题”单条请求的缓存只是不断变长的张量,放到服务系统里就成了一个内存管理问题。这里简要列出几个方向,每个方向都值得单独展开。
内存管理。 每条请求的输出长度事先未知,缓存大小随生成过程变化。如果为每条请求预留一整块连续显存,就会产生碎片和浪费。vLLM 的 PagedAttention 借鉴操作系统的分页思路,把缓存切成固定大小的块,块之间不要求连续,按需分配(Kwon 等, 2023)。
压缩缓存。 既然缓存是激活值,就可以用更低的精度存储。KVQuant 研究了 K/V 激活的量化方法,目标是让很长的上下文也能放进显存(Hooper 等, 2024)。
淘汰缓存。 另一类方法不保留全部历史。H2O 只保留对注意力贡献大的位置和最近的位置(Zhang 等, 2023)。StreamingLLM 观察到开头几个 token 起着“注意力汇聚点”的作用,于是保留它们,再加上一个滑动窗口内的最近 token(Xiao 等, 2023)。这类方法改变了模型实际能读到的上下文,效果需要按任务评估。
从结构上减小缓存。 前面提到的 MQA 和 GQA 属于这一类,它们在模型设计阶段就减少了需要缓存的 K/V 头数。
KV Cache 的成立依赖因果遮罩:历史位置的 K/V 不随后续 token 改变,所以可以保存下来复用。Prefill 一次建好缓存,之后每步 Decode 只为新 token 计算 K/V 并追加,新位置的 query 读取全部历史。缓存按层独立保存,大小随 batch 和上下文长度线性增长,在长上下文和高并发下会成为主要的显存开销。实现上,Transformers 的 DynamicCache 就是按层保存、沿序列维拼接。减少 KV 头数、分页管理、量化和淘汰,都是针对这笔开销的优化。
想逐个数字核对缓存的增长过程,可以看 Decoder-only 完整手算。
- Vaswani 等. Attention Is All You Need. 2017.
- Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019.
- Leviathan 等. Fast Inference from Transformers via Speculative Decoding. 2022.
- Ainslie 等. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023.
- Zhang 等. H2O: Heavy-Hitter Oracle for Efficient Generative Inference of Large Language Models. 2023.
- Kwon 等. Efficient Memory Management for Large Language Model Serving with PagedAttention. 2023.
- Xiao 等. Efficient Streaming Language Models with Attention Sinks. 2023.
- Hooper 等. KVQuant: Towards 10 Million Context Length LLM Inference with KV Cache Quantization. 2024.
- Grattafiori 等. The Llama 3 Herd of Models. 2024.
- Hugging Face. Transformers 文档:Cache strategies.