跳转至

投机解码与 KV Cache 优化

更新日期:2026-04-15


一、投机解码 (Speculative Decoding)

1.1 核心思想

用一个快但不准的小模型做草稿,用慢但准的大模型做验证。在一次大模型前向传播中验证多个 draft token,加速 2-3x。参考 Speculative Decoding (Leviathan et al., 2022)

flowchart LR
    draft["Draft Model<br/>(小, 快)"]
    propose["提议 K 个 token"]
    target["Target Model<br/>(大, 慢)"]
    verify["一次前向<br/>验证 K 个"]
    accept["逐个接受/<br/>拒绝重采样"]

    draft --> propose --> target --> verify --> accept

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    class draft,propose,target,verify,accept stage

数学保证:rejection sampling 让最终输出分布严格等于纯 target 模型采样分布(不是近似)。

1.2 完整算法

def speculative_decoding(target_model, draft_model, prompt, K=5):
    output = [prompt]

    while not done:
        # Step 1: Draft model 快速生成 K 个 token
        draft_tokens = []
        draft_probs = []
        for _ in range(K):
            logits = draft_model.forward_one(context=output + draft_tokens)
            p = softmax(logits)
            t = sample(p)
            draft_tokens.append(t)
            draft_probs.append(p)

        # Step 2: Target model 一次前向, 验证 K 个 token
        target_logits = target_model.forward_batch(
            context=output, new_tokens=draft_tokens)
        # target_logits: [K+1] 个位置的分布

        # Step 3: 逐个接受/拒绝
        accepted = []
        for i in range(K):
            p_draft = draft_probs[i][draft_tokens[i]]
            p_target = softmax(target_logits[i])[draft_tokens[i]]

            # 以 min(1, p_target/p_draft) 的概率接受
            if random() < min(1.0, p_target / p_draft):
                accepted.append(draft_tokens[i])
            else:
                # 拒绝: 从调整分布重采样
                adjusted = max(0, softmax(target_logits[i]) - draft_probs[i])
                adjusted /= adjusted.sum()
                new_token = sample(adjusted)
                accepted.append(new_token)
                break  # 拒绝后, 后续 draft 全部作废

        # Bonus: 如果全部接受, 还能多一个 token
        if len(accepted) == K:
            bonus = sample(softmax(target_logits[K]))
            accepted.append(bonus)

        output.extend(accepted)

    return output
    # 数学保证: 输出分布和直接用 target 完全一致!

1.3 加速效果


二、投机解码变体

2.1 Medusa 结构

class MedusaModel:
    def __init__(self, base_model, n_heads=5):
        self.base = base_model
        # 添加多个预测头, 每个预测不同未来位置
        self.medusa_heads = [MedusaHead() for _ in range(n_heads)]

    def forward(self, x):
        hidden = self.base(x)                    # 标准前向
        base_logits = self.base.lm_head(hidden)  # 位置 +0

        # 多个头并行预测 +1, +2, +3, +4
        future_logits = [head(hidden) for head in self.medusa_heads]

        # 一次前向得到 5 个未来位置的预测
        return base_logits, future_logits

# 优势: 不需要独立的 draft 模型
# 劣势: 需要额外训练预测头

三、KV Cache 优化深入

3.1 KV Cache 为什么是核心

decode 阶段的瓶颈不是计算,而是内存带宽。每生成一个 token 都要读取整个 KV Cache。KV Cache 越大,decode 越慢。

KV Cache 大小 = 2 × n_layers × n_kv_heads × d_head × seq_len × batch × dtype_bytes

LLaMA-70B 示例:
= 2 × 80 × 8 (GQA) × 128 × seq_len × batch × 2 bytes
= 327,680 × seq_len × batch bytes

8K 上下文, batch=32: 
= 327680 × 8192 × 32 = 86 GB  (!!!)

128K 上下文, batch=1:
= 327680 × 128000 × 1 = 42 GB

3.2 KV Cache 优化方法

flowchart LR
    arch["架构层<br/>MQA/GQA/MLA"]
    quant["量化<br/>FP8/INT8/INT4"]
    evict["淘汰<br/>H2O/StreamingLLM"]
    offload["卸载<br/>CPU/NVMe"]

    arch --> quant --> evict --> offload

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    class arch,quant,evict,offload stage
方法 节省比例 代价 典型应用
MQA / GQA 4-8x 训练时引入,需要重训 LLaMA-⅔ (GQA), PaLM (MQA)
MLA ~10x 改架构,需重训 DeepSeek-V2/V3
FP8 / INT8 量化 2x 精度略降,长上下文质量受影响 vLLM kv_cache_dtype=fp8
H2O / 重要性淘汰 2-5x 早期 token 丢失 长对话场景
StreamingLLM (sink) 任意长度 中段内容遗忘 永久流式聊天
CPU/NVMe Offload 内存换显存 带宽下降,TTFT 变慢 离线长文档
PagedAttention 碎片 60% → <5% 实现复杂 vLLM 默认
RadixAttention 共享前缀复用 树管理开销 RAG / 多轮

3.3 PagedAttention (vLLM 核心)

类比操作系统虚拟内存管理:KV Cache 分成固定大小的 block(如 16 tokens/block),按需分配,不连续存储。

class PagedKVCache:
    def __init__(self, block_size=16, num_physical_blocks=1000):
        self.block_size = block_size
        # 物理 block 池
        self.physical_blocks = [KVBlock() for _ in range(num_physical_blocks)]
        self.free_blocks = list(range(num_physical_blocks))
        # 逻辑 → 物理映射 (per request)
        self.page_table = {}

    def allocate(self, request_id, n_tokens):
        n_blocks = ceil(n_tokens / self.block_size)
        physical_ids = [self.free_blocks.pop() for _ in range(n_blocks)]
        self.page_table[request_id] = physical_ids

    def append(self, request_id, token_kv):
        # 追加 token, 可能需要新 block
        current_block = self.page_table[request_id][-1]
        if self.physical_blocks[current_block].is_full():
            new_block = self.free_blocks.pop()
            self.page_table[request_id].append(new_block)
        self.physical_blocks[current_block].append(token_kv)

    def fork(self, src_request, new_request):
        # Copy-on-Write: 共享 block, 修改时才复制
        self.page_table[new_request] = list(self.page_table[src_request])
        self.increment_ref_count(self.page_table[new_request])

效果: 内存碎片从 60-80% 降到 <5%,支持 2-4x 更多并发请求。

3.4 RadixAttention (SGLang)

RadixAttention 用 Radix Tree 管理 KV Cache,自动发现共享前缀,对 RAG/chatbot 等有共享前缀的场景极为有效。

场景: system_prompt (2K tokens) + 不同 queries - 不用 RadixAttention:每个请求独立 prefill,每次重算 system_prompt - 用 RadixAttention:system_prompt 只算一次,所有请求共享

前缀共享场景下,比 vLLM 快 6.4x。

3.5 KV Cache 量化

KV Cache 占用大,量化可以大幅节省内存。 vLLM 中的实现:

vllm_args = {
    'kv_cache_dtype': 'fp8',  # 或 'int8'; KV Cache 用 FP8, 计算时仍用 BF16
}

3.6 StreamingLLM: 注意力 Sink

观察: 序列前几个 token 的 attention 权重特别大 (注意力 sink)。即使丢弃中间的 token, 只要保留前几个 + 最近的, 就能继续生成。参考 StreamingLLM (Xiao et al., 2023)

def streaming_kv_cache(kv_cache, max_size=4000, sink_size=4):
    """
    保留: 前 sink_size 个 token + 最近 (max_size - sink_size) 个 token
    丢弃: 中间的 token
    """
    if len(kv_cache) <= max_size:
        return kv_cache

    sink = kv_cache[:sink_size]
    recent = kv_cache[-(max_size - sink_size):]
    return sink + recent
    # 代价: 中间内容"忘记"
    # 收益: 可以无限长流式生成

3.7 KV Cache 压缩方法对比


四、Prefill 优化

4.1 Chunked Prefill

# 长 prompt (100K+ tokens) 的 prefill 可能需要数秒
# Chunked Prefill: 分块处理, 每次一小块

def chunked_prefill(prompt, chunk_size=2048):
    n_chunks = ceil(len(prompt) / chunk_size)
    for i in range(n_chunks):
        chunk = prompt[ichunk_size : (i+1)chunk_size]
        kv_cache = extend_kv(kv_cache, chunk)
        # 这里可以插入其他请求的 decode step!
        yield_to_other_requests()

    return kv_cache

# 好处: prefill 不再阻塞 decode, 延迟更稳定

4.2 Prompt Caching

缓存已计算的 KV Cache,后续请求直接复用。Anthropic/OpenAI API 均已支持,缓存命中时成本减少约 90%。

典型使用场景:

  • 大 system prompt(所有请求共享)

  • RAG(多次查询同一个文档)

  • Agent(多步调用,上下文累积)


五、综合优化效果

场景: LLaMA-70B 服务 100 并发请求, 8K 上下文

基础: BF16, 无优化 - 需 8×H100, ~500 tok/s 总吞吐

优化 1: PagedAttention (vLLM) - 内存利用率 2x → 可以 16 并发 → ~1000 tok/s

优化 2: FP8 量化 - 模型内存减半 → 可以 32 并发 → ~2000 tok/s

优化 3: 投机解码 (draft=8B) - 2x 加速 → ~4000 tok/s

优化 4: Chunked Prefill + Prompt Caching (RAG 场景) - 共享 system prompt → ~6000 tok/s

综合: 12x 吞吐提升


参考文献

  • [1] Leviathan et al. Speculative Decoding. 2022. 论文

  • [2] Cai et al. Medusa. 2024. 论文

  • [3] Li et al. EAGLE. 2024. 论文

  • [4] Kwon et al. PagedAttention / vLLM. SOSP 2023. 论文

  • [5] Zheng et al. SGLang / RadixAttention. 2024. 论文

  • [6] Xiao et al. StreamingLLM. ICLR 2024. 论文

  • [7] Zhang et al. H2O (Token Eviction). NeurIPS 2023. 论文

  • [8] Speculative Decoding in 2026 博客


上级 · G. 推理与部署