KV Cache 到底是什么?面试被问住之后我翻了源码
Transformer 推理优化的基础,但很多人只背了八股文。从源码讲清楚 KV Cache 的原理、显存开销、PagedAttention,以及什么时候不用 KV Cache。
# KV Cache 到底是什么?面试被问住之后我翻了源码
昨天面试一个后辈,问他 Transformer 推理的时候有没有什么优化手段。对方说了个"KV Cache",然后就没下文了。
我问:"能说说它为什么能加速吗?" 他沉默了三秒。
我自己也想过,很多人知道 KV Cache 这三个字,但真要讲清楚,十个人里八个说不明白。今天来把这个东西掰开揉碎讲一讲。
先搞清楚:没有 KV Cache 的时候发生了什么
Transformer 推理是逐个 token 生成的。每生成一个新 token,模型都要重新跑一遍整个序列的 Attention 计算。
假设有 1000 个 token 的输入,已经生成了 50 个输出 token。生成第 51 个 token 时,如果没有 KV Cache,模型会重新计算全部 1050 个 token 的 QKV。生成第 100 个 token 时,要重新算 1100 个。
越往后越贵。这个复杂度是 O(n²) 的——你每多生成一个 token,Attention 的计算量就线性增长。
KV Cache 的本质:别算重复的
KV Cache 的核心思想很简单:**Key 和 Value 矩阵是历史状态的快照,不需要重算。**
每个 Transformer 层维护两组缓存:
生成新 token 时,只需要:
2. 把新的 K、V 追加到缓存末尾
3. 用 Q 和完整的 K_cache、V_cache 做 Attention
这样每次推理的计算量就从 O(n²) 降到 O(n)——因为当前 token 的 Q 只需要和已有缓存做乘法,不需要重新处理整个历史序列。
代码说话
用一个最小化的 PyTorch 实现来看:
import torch
import torch.nn as nn
import torch.nn.functional as F
class SimpleKVCache:
def __init__(self, num_heads, head_dim, max_len):
self.k_cache = torch.zeros(1, num_heads, max_len, head_dim)
self.v_cache = torch.zeros(1, num_heads, max_len, head_dim)
self.used_len = 0
def update(self, k, v):
"""k, v shape: (batch, heads, seq_len, head_dim)"""
seq_len = k.shape[2]
start = self.used_len
end = start + seq_len
self.k_cache[:, :, start:end] = k
self.v_cache[:, :, start:end] = v
self.used_len += seq_len
return self.k_cache[:, :, :self.used_len], self.v_cache[:, :, :self.used_len]
def attention_with_kv_cache(q, k, v, kv_cache, num_heads, head_dim):
k_cached, v_cached = kv_cache.update(k, v)
# Q @ K^T / sqrt(d)
scores = torch.matmul(q, k_cached.transpose(-2, -1)) / (head_dim ** 0.5)
attn = F.softmax(scores, dim=-1)
out = torch.matmul(attn, v_cached)
return out
注意 update 方法——这里用的是原地赋值,不是拼接。这是 vLLM 和 llama.cpp 都用的做法:预分配好 buffer,写的时候直接覆盖。拼接的开销比你想的大,尤其是 CUDA 显存操作。
坑:KV Cache 不是免费的
很多人以为用了 KV Cache 就万事大吉,实际上有几个要命的细节:
1. 显存占用
KV Cache 的显存大小 = num_layers × seq_len × num_heads × head_dim × 2 × 2(float16 占 2 字节,K 和 V 各一份)。
以 LLaMA-7B 为例:
seq_len=4096 时,光是 KV Cache 就要占 **64MB**。seq_len=32768 时,**512MB**。模型推理占几 GB,KV Cache 也能轻松吃掉几 GB。
2. PagedAttention:vLLM 的真正杀手锏
传统 KV Cache 分配方式是连续内存,容易导致显存碎片。vLLM 引入 PagedAttention——把 KV Cache 按 page 分块管理,类似操作系统的虚拟内存分页。
# 伪代码:vLLM 的 KV Cache 管理
class PagedAttentionCache:
def __init__(self, num_layers, num_blocks, block_size, head_num, head_dim):
self.pages = {} # block_id -> [layer_id] -> page_buffer
self.free_blocks = list(range(num_blocks))
self.allocations = {} # req_id -> [block_ids]
这个设计让 vLLM 的吞吐比 naive 实现高 2-4 倍。不是理论上的,是实测的。
3. 量化 KV Cache
INT8 量化 KV Cache 在 llama.cpp 里是标配。从 float16 到 int8,显存直接减半。质量损失呢?实测下来对大多数场景几乎没有感知影响——Attention 的 softmax 本身有容错性,KV 的精度要求没那么高。
// llama.cpp 量化 KV Cache 的核心思路
void kv_cache_quantize(struct ggml_context *ctx, struct ggml_tensor *k, struct ggml_tensor *v) {
// float16 -> int8, 记录 scale 值
// 计算时反量化回 float16 再做 attention
}
什么时候不用 KV Cache?
说实话,这个场景很少,但存在:
写在最后
KV Cache 是 LLM 推理优化的基础,但不是什么银弹。真正的高性能推理服务(vLLM、TGI、llama.cpp)是 KV Cache + PagedAttention + 量化 + CUDA Graph 的一堆组合拳。
面试的时候如果对方只说了"KV Cache"三个字,可以追问:它的显存开销是多少?PagedAttention 解决了什么?INT8 量化后的精度下降大概多少?
这三个问题,能筛掉大部分只背过八股的人。
*我是做 AI 基础设施的,踩过不少坑。如果觉得有用,可以关注后续的推理优化系列。*