KV Cache:推理加速核心

自回归生成时,每个新 token 都要对前面所有 token 做注意力。若每次都重算历史 key/value,开销随长度平方增长。KV Cache 把已算出的 K、V 缓存复用,使每步只需计算新 token,是推理加速的核心。

原理

Transformer 注意力的 K、V 只依赖各自输入,与后续 token 无关。生成第 t 步时,前 t-1 步的 K、V 已固定,直接拼接新 token 的 K、V 即可参与注意力,省去全序列重算。

# 伪代码:增量解码时复用缓存
past_kv = None
for step in range(max_len):
    logits, past_kv = model(input_ids, past_key_values=past_kv)
    next_id = logits[:, -1].argmax(-1, keepdim=True)
    input_ids = next_id

代价与优化

KV Cache 随层数、头数、序列长度线性增长显存占用,长上下文会成为瓶颈。vLLM 的 PagedAttention 把 KV 分页管理,减少碎片;MQA/GQA 则压缩 KV 头数降低显存。

小结

KV Cache 缓存历史键值以增量方式解码,把每步复杂度从序列平方降到线性,是推理提速基础,但显存随上下文增长,需用分页或共享注意力头优化。

参考与延伸阅读

本文累计阅读