KV Cache 优化与内存管理
KV Cache 是自回归推理的核心瓶颈——随着序列增长,Key 和 Value 的缓存占用线性增加,成为长上下文推理的主要内存开销。
优化策略
python
# GQA 减少 KV 头数
class GroupedQueryAttention(nn.Module):
def __init__(self, n_heads, n_kv_heads):
# n_heads 个 Q 头共享 n_kv_heads 个 KV 头
self.n_rep = n_heads // n_kv_heads
# KV Cache 量化
def quantize_kv_cache(cache, bits=4):
scale = cache.abs().amax(dim=-1, keepdim=True) / (2**(bits-1) - 1)
return (cache / scale).round().clamp(-2**(bits-1), 2**(bits-1)-1), scale
# PagedAttention
class PagedKVCache:
def __init__(self, block_size=16):
self.blocks = {} # 物理块池
self.tables = {} # 虚拟到物理的映射PagedAttention
vLLM 的 PagedAttention 借鉴了操作系统的虚拟内存分页机制,将 KV Cache 分配为固定大小的块,按需分配,避免了预分配导致的内存浪费。