KV Cache
KV Cache(键值缓存)
Section titled “KV Cache(键值缓存)”KV Cache 是 LLM 推理加速的核心技术。没有它,每生成一个 token 都要重新计算整个序列。
问题:自回归生成的低效
Section titled “问题:自回归生成的低效”GPT 类模型逐个生成 token:
Step 1: [A] → 预测 BStep 2: [A, B] → 预测 CStep 3: [A, B, C] → 预测 D每次都要重新计算之前所有 token 的 Attention。计算量随序列长度平方增长。
解决方案:缓存 K 和 V
Section titled “解决方案:缓存 K 和 V”观察 Attention 公式:
生成新 token 时,之前 token 的 K 和 V 不会变。把它们缓存起来,新 token 只需计算自己的 Q、K、V,然后和缓存的 K、V 拼接。
| 方式 | Step N 的计算量 | 总计算量(N 步) |
|---|---|---|
| 无缓存 | ||
| 有缓存 |
import torchimport torch.nn.functional as Fimport math
class KVCacheAttention: def __init__(self, d_k=64): self.d_k = d_k self.cache_k = None # 缓存的 K self.cache_v = None # 缓存的 V
def forward(self, q, k, v, use_cache=True): """ q, k, v: (batch, 1, d_k) -- 单步生成 """ if use_cache and self.cache_k is not None: # 拼接历史缓存 k = torch.cat([self.cache_k, k], dim=1) v = torch.cat([self.cache_v, v], dim=1)
# 更新缓存 self.cache_k = k self.cache_v = v
# 标准 Attention scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) attn = F.softmax(scores, dim=-1) output = torch.matmul(attn, v)
return output, attn
def reset(self): self.cache_k = None self.cache_v = None
# 模拟自回归生成attention = KVCacheAttention(d_k=64)batch_size = 1
for step in range(10): # 每个 step 生成一个新 token q = torch.randn(batch_size, 1, 64) k = torch.randn(batch_size, 1, 64) v = torch.randn(batch_size, 1, 64)
output, _ = attention.forward(q, k, v) seq_len = attention.cache_k.shape[1] print(f"Step {step+1}: 已缓存 {seq_len} 个 token 的 KV")KV Cache 的显存占用:
以 Llama-7B 为例:32 层 × 32 头 × 4096 tokens × 128 维 × 2 字节 ≈ 2 GB
这就是为什么长上下文需要大量显存。
| 技术 | 原理 | 效果 |
|---|---|---|
| Multi-Query Attention | 所有头共享 K、V | 显存降至 1/头数 |
| Grouped-Query Attention | 分组共享 K、V | 折中方案 |
| PagedAttention | 分页管理缓存 | 减少碎片 |
| 量化 KV Cache | INT8/INT4 存储 | 显存减半或更多 |
- Attention — KV Cache 缓存的就是 Attention 的 K 和 V
- Flash Attention — 高效 Attention 实现
- GPTQ — 权重量化,可配合 KV Cache 量化