推测解码工程实践
推测解码(Speculative Decoding)是一种无损推理加速技术,通过小模型(草稿模型)快速生成候选 Token,再由大模型并行验证,在保证输出质量不变的前提下显著提升推理速度。
基本原理
推测解码的工作流程:
- 草稿阶段:小模型自回归生成 K 个候选 Token
- 验证阶段:大模型一次前向传播验证所有候选 Token
- 接受/拒绝:根据大模型的概率分布决定接受哪些 Token
- 修正:拒绝位置之后重新采样
python
# 推测解码伪代码
def speculative_decode(draft_model, target_model, prompt, K=5):
tokens = prompt
while not finished:
# 1. 草稿模型生成 K 个候选
draft_tokens = []
current = tokens
for _ in range(K):
next_token = draft_model.generate_one(current)
draft_tokens.append(next_token)
current = current + [next_token]
# 2. 目标模型一次前向验证
target_probs = target_model.forward(tokens + draft_tokens)
# 3. 接受/拒绝判断
accepted = 0
for i, draft_token in enumerate(draft_tokens):
if accept_criterion(target_probs[i], draft_token):
accepted += 1
else:
# 从目标分布重新采样
resampled = sample_from(target_probs[i])
tokens = tokens + draft_tokens[:accepted] + [resampled]
break
else:
tokens = tokens + draft_tokens
return tokens加速比分析
推测解码的理论加速比取决于:
- 接受率 α:草稿 Token 被接受的概率
- 草稿步数 K:每次推测的 Token 数
- 速度比 β:大模型/小模型单步时间比
理论加速比 ≈ (1 + α·K) / (1 + K/β)
| 草稿模型 | 接受率 α | K=5 加速比 | K=10 加速比 |
|---|---|---|---|
| 7B → 70B | 0.7 | 2.3x | 2.8x |
| 1.5B → 7B | 0.6 | 1.9x | 2.2x |
| 70M → 7B | 0.4 | 1.4x | 1.5x |
草稿模型选择
草稿模型应与目标模型同系列(如 Llama-2-7B 作为 Llama-2-70B 的草稿),这样分布更接近,接受率更高。跨系列模型组合的接受率通常较低。
vLLM 中的推测解码
python
from vllm import LLM, SamplingParams
llm = LLM(
model="meta-llama/Llama-2-70b-hf",
speculative_model="meta-llama/Llama-2-7b-hf", # 草稿模型
num_speculative_tokens=5,
tensor_parallel_size=4,
)
params = SamplingParams(max_tokens=256, temperature=0.0)
outputs = llm.generate(["解释量子纠缠"], params)限制条件
推测解码在 temperature=0(贪心解码)时效果最好,因为接受率最高。高温度采样会降低接受率,加速效果减弱。此外,推测解码目前仅支持自注意力模型。