投机解码与 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 中的实现:
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. 推理与部署