大语言模型(LLM)推理的性能瓶颈从来不在计算,而在内存。本文深入剖析 Transformer 推理中 KV Cache 的内存管理机制,从 vLLM 的 PagedAttention 到推测解码的显存优化策略,结合生产级代码实战,揭示如何在不牺牲吞吐的前提下将显存占用压缩到极致。

1. 问题本质:为什么推理的核心瓶颈是内存?

以一个 70B 参数模型为例进行量化分析。推理过程中,每个 Token 的 KV Cache 由所有 Transformer 层的 Key 和 Value 向量组成:


单 Token KV 缓存大小 = 2 × num_layers × num_kv_heads × head_dim × dtype_size

以 Llama-2-70B 为例:80 层 × 8 个 KV Head × 128 维度 × 2 字节(FP16)= 每 Token 约 160 KB。当序列长度达到 8K 时,单个请求的 KV Cache 就高达 1.25 GB,20 个并发请求瞬间吃掉 25 GB 显存。

关键问题有三个:

  • 内存碎片:传统连续分配导致 60-80% 的显存浪费在预分配的冗余空间
  • 复用困难:Beam Search、共享前缀等场景下 KV 无法高效跨序列复用
  • 不可抢占:KV 与请求绑定,无法像 OS 那样换出到主机内存

2. PagedAttention:把操作系统虚拟内存搬进 GPU

vLLM 的 PagedAttention 是最优雅的解决方案,核心思路直接借鉴 OS 的虚拟内存 + 分页机制:

2.1 核心数据结构


# vLLM 中 BlockTable 的精简实现
class BlockTable:
    """类似于 OS 页表:虚页号 -> 物理块号"""
    
    def __init__(self, block_size: int, num_gpu_blocks: int):
        self.block_size = block_size          # 每块容纳的 token 数,通常 16
        self.block_table: List[int] = []      # 虚页 -> 物理块映射
        self.free_blocks = list(range(num_gpu_blocks))  # 空闲物理块池
        self.ref_count = [0] * num_gpu_blocks           # 引用计数(用于 CoW 和回收)
    
    def allocate(self, num_tokens: int) -> List[int]:
        """分配虚页,返回物理块号列表"""
        num_blocks = ceil(num_tokens / self.block_size)
        if len(self.free_blocks) < num_blocks:
            raise MemoryError("GPU KV Cache 耗尽,需触发抢占")
        
        physical_blocks = []
        for _ in range(num_blocks):
            block_id = self.free_blocks.pop()
            self.ref_count[block_id] = 1
            physical_blocks.append(block_id)
        
        self.block_table.extend(physical_blocks)
        return physical_blocks
    
    def append_slot(self) -> Optional[int]:
        """追加一个 token 级别的 slot,仅在当前块有余量时原地扩展"""
        last_block = self.block_table[-1]
        used_slots = self._used_slots_in_last_block()
        
        if used_slots < self.block_size:
            return last_block * self.block_size + used_slots  # 返回全局 slot 地址
        else:
            # 当前块已满,分配新块
            return self.allocate(1)[0] * self.block_size if self.free_blocks else None
    
    def free(self):
        """释放整个序列的物理块"""
        for block_id in set(self.block_table):
            self.ref_count[block_id] -= 1
            if self.ref_count[block_id] == 0:
                self.free_blocks.append(block_id)

2.2 分页对注意力的改造

标准 Self-Attention 要求 KV 连续存储。PagedAttention 需要 CUDA Kernel 支持非连续 KV 寻址:


# 概念性伪代码:PagedAttention 的 Attention 计算
def paged_attention_forward(
    query: Tensor,             # [num_heads, head_dim]
    paged_kv_cache: Tensor,    # [num_gpu_blocks, block_size, num_kv_heads, head_dim]
    block_table: Tensor,       # [num_blocks] 物理块映射
    context_len: int,          # 当前序列总 token 数
    scale: float
) -> Tensor:
    output = zeros_like(query)
    
    for block_idx in range(ceil(context_len / BLOCK_SIZE)):
        physical_block = block_table[block_idx]
        kv_block = paged_kv_cache[physical_block]  # 直接从物理块取 KV
        
        valid_slots = min(BLOCK_SIZE, context_len - block_idx * BLOCK_SIZE)
        attention_scores = query @ kv_block[:valid_slots].T * scale
        attention_probs = softmax(attention_scores)
        output += attention_probs @ kv_block[:valid_slots]
    
    return output

vLLM 实际使用 FasterTransformer 中定制的 CUDA kernel 来避免 per-block Python 开销,单 kernel 通过 block_table 指针直接寻址。

2.3 显存节省实测

通过分页机制,vLLM 消除了传统推理引擎的三大浪费源:

浪费源 传统方法 PagedAttention
预分配冗余 为 max_seq_len 预留连续空间 按需增量分配块
Beam Search 复制 逐字复制完整 KV 共享页 + Copy-on-Write
不同序列间碎片 外部碎片约 20-40% 近似零碎片

实测中,vLLM 相比 HuggingFace Transformers 显存占用减少 2-4 倍。

3. Prefix Caching:跨请求的 KV 复用

在 RAG 或多轮对话场景中,大量请求共享相同的 System Prompt 或文档上下文。vLLM 0.4+ 引入的 Automatic Prefix Caching(APC)实现了 KV Cache 的跨请求复用。

3.1 哈希索引机制


class PrefixCacheManager:
    """基于哈希的 KV Cache 复用管理器"""
    
    def __init__(self, block_size: int = 16):
        self.block_size = block_size
        self.hash_table: Dict[str, List[int]] = {}  # hash -> 物理块列表
        self.lru: OrderedDict = OrderedDict()        # LRU 淘汰
    
    def compute_block_hash(self, token_ids: Tuple[int], parent_hash: Optional[str] = None) -> str:
        """增量哈希:复用父块的哈希值"""
        if parent_hash:
            content = f"{parent_hash}:{token_ids}"
        else:
            content = str(token_ids)
        return hashlib.sha256(content.encode()).hexdigest()[:16]
    
    def lookup(self, token_ids: List[int]) -> Optional[List[int]]:
        """查找最长可复用前缀"""
        num_blocks = len(token_ids) // self.block_size
        
        for matched_len in range(num_blocks, 0, -1):
            block_tokens = tuple(token_ids[:matched_len * self.block_size])
            block_hash = self.compute_block_hash(block_tokens)
            
            if block_hash in self.hash_table:
                physical_blocks = self.hash_table[block_hash]
                # 增加引用计数防止被驱逐
                for blk_id in physical_blocks:
                    self.ref_count[blk_id] += 1
                self.lru.move_to_end(block_hash)
                return physical_blocks
        
        return None  # 未命中,需从头计算

3.2 生产环境中的收益

在实际生产场景中(System Prompt + RAG 检索结果),前缀复用率可达 70-90%:

  • 省去重复的 Prefill 计算,TTFT 减少 5-10 倍
  • 共享的物理块不重复占用显存
  • 对 10K+ token 的长 System Prompt 效果尤为显著

4. 推测解码与 KV Cache 的协同优化

推测解码(Speculative Decoding)通过小模型草稿 + 大模型验证的方式加速推理,但其对 KV Cache 的影响常被忽视。

4.1 推测解码的 KV 管理挑战

推测解码一次预测 K 个 Token,但验证可能只接受 M 个(M ≤ K)。被拒绝的 Token 对应的 KV 需要:


class SpeculativeKVManager:
    """支持推测解码的 KV 回滚"""
    
    def allocate_speculative(self, seq: Sequence, k: int) -> List[int]:
        """为 K 个推测 token 预分配 KV 块(但不提交)"""
        speculative_blocks = []
        for i in range(k):
            slot = seq.block_table.append_slot()
            if slot is None:
                return speculative_blocks
            speculative_blocks.append(slot)
        
        # 记录回滚点,验证失败后从这里截断
        seq.checkpoint_len = seq.context_len
        return speculative_blocks
    
    def accept_or_reject(self, seq: Sequence, accepted_count: int, total_speculated: int):
        """验证后处理:接受的保留,拒绝的释放"""
        rejected = total_speculated - accepted_count
        
        if rejected > 0:
            # 回滚:将拒绝的 token 对应的块释放
            seq.context_len = seq.checkpoint_len + accepted_count
            freed_blocks = seq.block_table.truncate_to_len(seq.context_len)
            
            for block_id in freed_blocks:
                self.free_blocks.append(block_id)
        else:
            # 全部接受,无需回滚
            seq.context_len += total_speculated

4.2 联合优化策略

在推测解码下,KV Cache 管理的最佳实践包括:

  1. 弹性块分配:推测阶段只分配最小必要块,验证成功后提交
  2. 投机命中率监控:当命中率低于 50% 时减小 K,减少 KV 浪费
  3. 分层的投机策略:高命中率请求使用大 K,低命中率请求使用小 K
  4. 5. 抢占与调度:当 KV Cache 不够用时

    GPU 显存终究是有限的。当多请求并发导致 KV Cache 耗尽时,需要类似 OS 页面置换的抢占机制。

    5.1 两种抢占模式

    
    class KVCacheScheduler:
        """支持抢占的 KV Cache 调度器"""
        
        def preempt_low_priority(self, num_blocks_needed: int) -> int:
            """按优先级抢占低优先级请求的 KV 缓存"""
            preemptable = sorted(
                self.running_requests,
                key=lambda r: (r.priority, -r.context_len)  # 优先级低、序列长的先抢占
            )
            
            freed_blocks = 0
            for req in preemptable:
                if freed_blocks >= num_blocks_needed:
                    break
                
                # 策略选择:Recomputation vs Swapping
                if self.swap_space_available and req.context_len > 1024:
                    # 长序列 swap out 到 CPU
                    blocks = req.block_table.swap_out_to_cpu(req.kv_blocks)
                    freed_blocks += len(blocks)
                else:
                    # 短序列直接 Recompute
                    freed_blocks += req.block_table.free_all()
                    req.state = RequestState.PREEMPTED  # 等待重新调度
                    self.swapped.add(req)
            
            return freed_blocks
    

    5.2 Swap vs Recomputation 的选择

    维度 Swap 到 CPU Recomputation
    恢复延迟 受 PCIe 带宽限制(~32 GB/s) 受 GPU 算力限制
    显存开销 需要 CPU 内存 只需存 token ids
    适用场景 长上下文(>4K tokens) 短序列
    吞吐影响 CPU-GPU 传输竞争带宽 Prefill 消耗 GPU 时间片

    经验公式:当序列长度 L × 单层 KV 大小 > PCIe 带宽 × Recompute 时间时,Swap 更优。对 70B 模型,临界点约在 2-4K tokens。

    6. 量化 KV Cache:用精度换显存

    除内存管理外,KV Cache 本身的精度也可以压缩。

    6.1 FP8 / INT8 KV Cache

    
    def quantize_kv_cache(kv_cache: torch.Tensor, bits: int = 8) -> torch.Tensor:
        """从 FP16 量化 KV 到 INT8,节省 50% 显存"""
        if bits == 8:
            # 按 head 维度分组量化
            scale = kv_cache.abs().amax(dim=-1, keepdim=True) / 127.0
            quantized = (kv_cache / scale).round().to(torch.int8)
            return quantized, scale
        
        # FP8 E4M3 (适合推理)
        return kv_cache.to(torch.float8_e4m3fn)
    

    6.2 GQA/MQA 的 Cache 共享

    • Multi-Query Attention(MQA):所有 Q Head 共享同一组 KV,KV Cache 缩减为 1/num_heads
    • Grouped-Query Attention(GQA):每 G 个 Q Head 共享一组 KV(Llama-2-70B: G=8,KV Cache 缩减为 1/8)

    当前主流模型(Llama-3、Qwen-2)都已采用 GQA,这本身就是 KV Cache 层面的优化。

    7. 工程实战:构建高吞吐推理服务

    本节给出一个基于 vLLM 的生产级配置与调优指南。

    7.1 关键配置参数

    
    # vLLM 生产环境配置示例
    model: meta-llama/Llama-3-70B-Instruct
    tensor_parallel_size: 4          # 4-way TP,4 张 A100-80GB
    gpu_memory_utilization: 0.90     # 90% 显存用于 KV Cache
    max_num_seqs: 256                # 最大并发请求数
    max_model_len: 32768             # 最大上下文长度
    block_size: 16                   # PagedAttention 块大小
    enable_prefix_caching: true      # 开启前缀缓存
    enable_chunked_prefill: true     # 分块预填充,避免大 prefill 阻塞
    max_num_batched_tokens: 2048     # 单次前向最大 token 数
    

    7.2 性能监控指标

    
    class KVCacheMonitor:
        """KV Cache 健康度监控"""
        
        def metrics(self):
            return {
                "gpu_cache_usage_percent": self.used_blocks / self.total_blocks * 100,
                "prefix_cache_hit_rate": self.prefix_hits / (self.prefix_hits + self.prefix_misses),
                "avg_blocks_per_request": mean(len(r.block_table) for r in self.active_requests),
                "preemptions_per_minute": self.preemption_counter.rate(),
                "swap_in_out_bytes": self.swap_bytes_counter.rate()
            }
    

    关键告警阈值:

    • GPU KV Cache 使用率 > 95%:需扩容或降低并发
    • Prefix Cache Hit Rate < 30%:检查请求模式,优化 System Prompt 设计
    • Preemptions > 5/min:并发过高,需限流或扩容

    7.3 瓶颈分析与优化路线

    实际部署中常见的 KV Cache 问题:

    1. 显存不足导致的吞吐塌陷:增加 tensor_parallel_size 或使用量化 KV Cache
    2. 长序列拖尾延迟:开启 Chunked Prefill,将大 prefill 拆解为 512-token 的 chunk 穿插执行
    3. Prefix Caching 失效:确保 System Prompt 无随机元素(如时间戳),将随机部分置于尾部
    4. 推测解码的回滚开销:监控接受率,动态调整 draft 长度 K
    5. 8. 前沿方向:超越传统 KV Cache

      8.1 Mamba 与 State Space Models

      Mamba 等 SSM 架构用可压缩隐状态替代了完整的 KV Cache,将复杂度从 O(n) 降低到 O(1)。在推理时,Mamba 只需维护一个固定大小的状态矩阵,与序列长度无关,从根本上解决了 KV Cache 的内存问题。代价是部分需要全局注意力的能力会在状态压缩中损失。

      8.2 Token Dropping 与动态淘汰

      H2O(Heavy-Hitter Oracle)和 Scissorhands 等研究提出动态淘汰低注意力 Token,只保留"重击者"(累计注意力分数高的 Token),在保持质量的同时将 KV Cache 压缩到原始大小的 30-50%。

      8.3 多模态 KV Cache

      多模态推理中,图像/视频 Token 的 KV 复用率远高于文本。预先计算并缓存视觉编码器输出的 KV(Vision KV Cache),可将 VL 模型的 Prefill 时间减少 80%+。

      9. 总结

      KV Cache 内存管理是 LLM 推理优化的核心战场。从 vLLM 的 PagedAttention 到 Prefix Caching,再到推测解码与抢占调度的协同,每一步都在逼近显存利用的理论极限。

      实战中的核心原则可以归纳为:

      • 分页消除碎片:借鉴 OS 虚拟内存思想,实现按需增量分配
      • 共享消除冗余:通过哈希索引和引用计数复用前缀 KV
      • 量化压缩精度:FP8/GQA 在可接受质量损失下成倍节省显存
      • 抢占保障 QoS:类似 OS 页面置换,低优先级请求的 KV 可以被换出或重算

      未来,随着 SSM 架构和 Token 淘汰技术的成熟,KV Cache 的内存效率将进一步提升。但在可预见的三年内,PagedAttention 及其派生方案仍将是生产环境中的主流选择。


      *文章标签:LLM, KV Cache, vLLM, PagedAttention, 推理优化, 显存管理*

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部