KV Cache 压缩与量化技术
KV Cache 是大语言模型推理中最大的显存开销来源,尤其在长上下文场景下。KV Cache 压缩与量化技术通过降低缓存精度和减少缓存条目来显著降低显存占用。
KV Cache 显存分析
以 Llama-2-7B 为例,KV Cache 的显存占用计算:
KV Cache 大小 = 2 × num_layers × batch_size × seq_len × num_kv_heads × head_dim × dtype_size
# FP16, seq_len=4096, batch_size=32
= 2 × 32 × 32 × 4096 × 32 × 128 × 2 bytes
≈ 16 GB显存占比
在长上下文场景(seq_len > 8K)下,KV Cache 可占 GPU 总显存的 50-80%。压缩 KV Cache 是扩展上下文长度的关键。
KV Cache 量化
FP8 量化
将 KV Cache 从 FP16 量化为 FP8,显存减半:
python
# vLLM 中启用 KV Cache FP8 量化
from vllm import LLM
llm = LLM(
model="meta-llama/Llama-2-7b-hf",
kv_cache_dtype="fp8_e5m2", # FP8 量化
gpu_memory_utilization=0.9,
)INT4/INT8 量化
更激进的 KV Cache 量化方案:
python
# KV Cache 量化示意
import torch
def quantize_kv_cache(kv_tensor, bits=4):
"""将 KV Cache 量化为低精度"""
scale = kv_tensor.abs().max(dim=-1, keepdim=True).values / (2**(bits-1) - 1)
quantized = torch.clamp(
(kv_tensor / scale).round(),
-(2**(bits-1)),
2**(bits-1) - 1
).to(torch.int8 if bits == 8 else torch.int4)
return quantized, scale
def dequantize_kv_cache(quantized, scale):
"""反量化 KV Cache"""
return quantized.float() * scaleKV Cache 驱逐策略
Sliding Window
固定窗口大小,丢弃超出窗口的旧 KV Cache:
- 优点:显存占用恒定
- 缺点:丢失长距离依赖信息
- 适用:Mistral 等原生支持 SWA 的模型
Token 蒸馏
将多个 Token 的 KV Cache 蒸馏为更少的代表向量:
python
# 简化的 Token 蒸馏
def token_distillation(kv_cache, ratio=0.5):
"""保留重要 Token 的 KV Cache"""
importance = kv_cache.norm(dim=-1) # 简单重要性度量
k = int(len(kv_cache) * ratio)
topk_indices = importance.topk(k).indices
return kv_cache[topk_indices]Heavy-Hitter Oracle (H2O)
基于注意力分数识别并保留最重要的 Token(Heavy Hitter):
- 保留最近 Token + 累计注意力最高的 Token
- 在保持模型质量的同时减少 50-80% 的 KV Cache
精度损失
KV Cache 量化会引入精度损失。FP8 量化对模型质量影响极小(<0.1%),INT4 量化可能导致 1-3% 的质量下降。建议根据应用场景选择合适的量化级别。