跳转到内容

基础原理第 9 篇,共 13 篇

KV Cache:复用每一层的历史 K/V

增量解码为什么只需要处理新 token:KV Cache 的原理、显存代价、Transformers 中的实现,以及围绕它展开的系统优化

更新于 阅读约 9 分钟

自回归语言模型一次只生成一个 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 的时序。

在权重、位置编码和遮罩都固定,并且计算是确定性的前提下,因果模型里某个位置的输出只依赖它自己和它左边的位置。往序列末尾追加 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 部分的简化估计,不能直接当作整个模型的耗时比例,实际加速取决于模型结构、硬件和负载。

缓存保存的是每一层自注意力的 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 把整个 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 token1,000 token8,000 token
GPT-2 small12 层,12 个头,head_dim 64,fp327.4 MB74 MB超出上下文长度
Llama 3 8B32 层,8 个 KV 头,head_dim 128,bf1613 MB131 MB1.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 头,缓存只有不分组时的四分之一。

下面以 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 相同。

每次前向时,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 层里,缓存的使用分三步。下面是简化后的流程:

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 为 12
cache.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。

类型做法适用情况代价
DynamicCache用 torch.cat 动态增长默认选择,长度不受限每步都要重新分配内存
StaticCache预先分配最大长度的缓冲区,原地写入张量形状固定,便于配合 torch.compile始终占用最大长度的显存,收益依赖具体负载
QuantizedCache把 K/V 量化到更低的位宽显存紧张的长上下文精度略有损失
EncoderDecoderCache分别保存自注意力和交叉注意力的缓存T5、BART 这类 Encoder–Decoder 模型不适用于 decoder-only 模型

静态缓存加编译不是通用的最优解,是否更快需要在自己的负载上测量。

单条请求的缓存只是不断变长的张量,放到服务系统里就成了一个内存管理问题。这里简要列出几个方向,每个方向都值得单独展开。

内存管理。 每条请求的输出长度事先未知,缓存大小随生成过程变化。如果为每条请求预留一整块连续显存,就会产生碎片和浪费。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 完整手算。

  1. Vaswani 等. Attention Is All You Need. 2017.
  2. Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019.
  3. Leviathan 等. Fast Inference from Transformers via Speculative Decoding. 2022.
  4. Ainslie 等. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023.
  5. Zhang 等. H2O: Heavy-Hitter Oracle for Efficient Generative Inference of Large Language Models. 2023.
  6. Kwon 等. Efficient Memory Management for Large Language Model Serving with PagedAttention. 2023.
  7. Xiao 等. Efficient Streaming Language Models with Attention Sinks. 2023.
  8. Hooper 等. KVQuant: Towards 10 Million Context Length LLM Inference with KV Cache Quantization. 2024.
  9. Grattafiori 等. The Llama 3 Herd of Models. 2024.
  10. Hugging Face. Transformers 文档:Cache strategies.