深度长文:LLM 推测解码(Speculative Decoding)工程化实战——从原理到 3 倍加速的完整实现
一、背景:自回归解码的「慢」是一个系统性问题
2026 年的今天,大语言模型早已不再停留在聊天玩具的阶段。ChatGPT Work、Claude Cowork、GitHub Copilot Agent——AI 正在真实地嵌入开发者的日常流水线。但无论模型多么聪明,推理速度始终是一道绕不过去的坎。
你在终端里敲下一条指令,光标闪烁 2 秒才开始吐字,然后一个字一个字往外蹦——这种体验用久了会让人产生一种错觉:「AI 就是慢的,习惯了就好。」
但这个「慢」真的无可避免吗?
1.1 量化那个「慢」
我们来算一笔账。以 Llama-3-70B 为例,模型参数约 140GB(FP16),而一块 NVIDIA A100-80G 的显存带宽大约是 2TB/s。这意味着每生成一个 token,你至少需要把全部参数从显存搬运到计算单元一次——耗时约 70ms。这还没算上 KV Cache 的读写和 Attention 计算的开销。
单次 token 生成 = 参数加载 ~70ms
+ KV Cache 读取 ~15ms
+ Attention 计算 ~10ms
+ FFN 前向 ~20ms
+ KV Cache 写入 ~5ms
= 约 120ms
每秒 ≈ 8 tokens
如果你的产品经理告诉你:「用户需要首 token 延迟 < 500ms,生成速度 > 30 tokens/s」,靠传统自回归解码,即使把模型量化到 INT4,70B 模型也顶多拉到 20 token/s——差距仍然巨大。
但这还不是最要命的。更大的痛点是:内存带宽是铁桶里的水龙头,瓶颈根本不在算力(FLOPs),而在数据搬运(Memory Bound)。
这里有个非常反直觉的真相:运行一个 70B 模型的 GPU,其计算单元利用率可能还不到 20%。剩下的 80% 时间在干什么?在等数据。每次生成一个 token,都要把 140GB 的参数从 HBM(高带宽显存)搬到 SM(流式多处理器),这个过程就是所谓的「内存墙」(Memory Wall)。
1.2 为什么 KV Cache 帮不了这个忙
有人在想:既然 Attention 的 KV Cache 能复用,那能不能扩大它的范围,减少参数加载次数?
答案是不能。KV Cache 只缓存了 Attention 层的 Key 和 Value,占用的显存虽然大(4K 上下文中约 2-4GB),但占总推理时间的比例相对较小。参数量才是带宽瓶颈的主宰。不管上下文多长,每生成一个 token,你都需要把整个模型的所有参数加载一遍——因为每个 token 的计算路径都可能不同,FFN 层的权重无法缓存。
Transformer Decoder 单步计算流:
1. Input Embedding + Position Encoding
2. Self-Attention (利用 KV Cache,只算当前 token 的 Q)
3. Residual + LayerNorm
4. FFN (两个全连接层,权重要全部加载!)
5. Residual + LayerNorm
6. LM Head (大词汇表投影)
其中 Step 4 的参数量占整个模型的 2/3
1.3 现有方案的局限
工业界解决推理性能有几种常用手段,但各有代价:
- 模型量化(INT4/INT8):参数体积直接减半或减到四分之一,但精度有损,对数学推理和代码生成任务的影响不可忽视
- KV Cache 量化:只解决 Attention 部分的显存和带宽,对 FFN 部分无影响,加速幅度有限
- 批量推理(Batching):通过增大 batch size 摊薄每 token 的固定开销,但单用户场景下 batch 无法做大
- 结构剪枝/MoE 激活:直接减少参数量,但模型结构发生了变化,需要重新训练或微调
有没有一种方案,不修改模型、不损失精度、只靠优化计算流程就能拿到 2-3 倍加速?
这就是 Speculative Decoding 的核心价值主张。
二、核心原理:猜-验范式
2.1 直觉:为什么「猜」比「算」快?
想象两个人在写代码:
- 主模型:是一位资深架构师,代码写得稳,但每次写一行都要思考半天,读几十页参考文档。
- 草稿模型:是一位刚入职的实习生,手速飞快,但代码质量一般。
如果没有 Speculative Decoding,架构师自己一行一行写,虽然每行都对,但效率极低。
有了 Speculative Decoding,流程变成:
- 实习生先唰唰唰写出 5 行代码(草稿)
- 架构师扫一眼,把每行标记为「✓」或「✗」
- 从第一个错的开始,架构师自己写一行
- 重复
关键点在于:架构师「检查」5 行代码的时间,和「自己写」1 行的时间几乎一样。因为 GPU 计算是高度并行的——一次 forward pass 处理 5 个 token 的验证,和一次 forward pass 生成 1 个 token,时间成本差不多。
这就是 Speculative Decoding 的底层逻辑:利用 GPU 的并行计算能力,用一次推理的代价验证多个候选 token。
2.2 数学本质:拒绝采样的并行化
Speculative Decoding 的核心算法是 拒绝采样(Rejection Sampling)的并行化版本。
设目标模型为 $p(x)$(大模型,精度高但慢),草稿模型为 $q(x)$(小模型,快但精度低)。
每一步,草稿模型 $q$ 自回归地生成 $\gamma$ 个候选 token:
$$
\hat{x}_1, \hat{x}2, ..., \hat{x}\gamma \sim q(\cdot | \text{context})
$$
然后目标模型 $p$ 用一次 forward pass 计算这 $\gamma$ 个 token 的概率分布:
$$
p(\cdot | \text{context}), p(\cdot | \text{context}, \hat{x}_1), ..., p(\cdot | \text{context}, \hat{x}1, ..., \hat{x}{\gamma-1})
$$
对于每个候选 token $\hat{x}_t$,以概率 $\min(1, \frac{p(\hat{x}_t | \cdot)}{q(\hat{x}_t | \cdot)})$ 接受它。如果被拒绝,则从调整后的分布 $p'(\cdot) = \text{norm}(\max(0, p(\cdot) - q(\cdot)))$ 中采样一个 token 填补,并停止本轮验证。
这个过程的数学保证是:最终输出分布与纯目标模型 $p(x)$ 完全一致。
所以它不是「近似加速」,而是「无损加速」——这点非常关键。如果你的 CTO 怀疑加速会影响质量,Speculative Decoding 的答案很明确:数学上保证一致。
2.3 为什么要并行验证
传统自回归解码为什么慢?因为它每一串行步骤都要重新加载模型权重:
Step 1: Load weights (70ms) → Compute (20ms) → Emit token
Step 2: Load weights (70ms) → Compute (20ms) → Emit token
...
Step N: Load weights (70ms) → Compute (20ms) → Emit token
瓶颈在于每次都要花 70ms 加载参数,真正的计算只占一小部分。如果用个比喻:每次只端一盘菜上餐桌,但为了端这一盘菜,你得从仓库走到厨房再走回来——90% 的时间花在了走路上,真正做菜的时间只有 10%。
Speculative Decoding 的并行验证则不同:
Draft: q(t1) → q(t2) → q(t3) → q(t4) → q(t5) (小模型快,总耗时 ≈ 20ms)
Verify: p(t1,t2,t3,t4,t5) (一次推理,耗时 ≈ 90ms)
总耗时 ≈ 110ms,而传统方案需要 5 × 90ms = 450ms。关键在于:草稿模型参数只有目标模型的 1/10 甚至更少,每次生成 token 只需要加载约 15GB 的权重(8B 模型),而不是 140GB。
2.4 接受率:决定加速效果的核心指标
加速效果最终由接受率决定。接受率是指草稿生成的 token 被主模型验证通过的比率。
假设草稿阶段耗时 $T_d$(生成 $\gamma$ 个 token),验证阶段耗时 $T_v$(和生成 1 个 token 差不多),接受率为 $\alpha$,那么平均每个验证周期获得的有效 token 数约为:
$$
E[\text{tokens}] = 1 + \sum_{k=1}^{\gamma} \alpha^k = \frac{1 - \alpha^{\gamma+1}}{1 - \alpha}
$$
加速比就是:
$$
\text{Speedup} = \frac{E[\text{tokens}] \cdot T_{\text{baseline}}}{T_d + T_v}
$$
其中 $T_{\text{baseline}}$ 是传统方式生成 1 个 token 的时间。
来算个具体的:草稿模型速度是目标模型的 10 倍($T_d \approx T_v / 10$),接受率 $\alpha = 0.7$,$\gamma = 6$:
$$
E[\text{tokens}] = \frac{1 - 0.7^7}{1 - 0.7} = \frac{1 - 0.082}{0.3} \approx 3.06
$$
$$
\text{Speedup} = \frac{3.06 \cdot T_v}{0.1T_v + T_v} = \frac{3.06}{1.1} \approx 2.78x
$$
这就是 2-3 倍加速的来源。如果接受率降到 0.5,加速比变成约 1.8x;如果提高到 0.85,加速比可达 3.5x 以上。
三、草稿模型选择:工程上的核心决策
Speculative Decoding 好不好用,核心看草稿模型选得对不对。这是我花时间最多的地方,也是生产环境最容易踩坑的环节。
3.1 草稿模型的核心约束
草稿模型 $q$ 与目标模型 $p$ 之间有三个关键指标:
- 接受率(Acceptance Rate):草稿生成的 token 被主模型接受的概率,决定了每个验证周期内获得的有效 token 数。接受率主要取决于草稿模型和目标模型分布的对齐程度。
- 草稿速度比(Draft Speed Ratio):$r = \frac{\text{草稿模型 token/s}}{\text{目标模型 token/s}}$,决定了草稿阶段的损耗占比。
- 显存开销:额外加载一份草稿模型需要消耗额外的显存,这可能是最现实的约束。
这三个指标构成一个不可能三角——没有草稿模型能同时做到接受率高、速度快、显存占用小。需要在三者之间做出权衡。
3.2 常见的草稿模型组合策略
| 策略 | 草稿模型 | 目标模型 | 接受率 | 加速比 | 额外显存 |
|---|---|---|---|---|---|
| 同族降级 | Phi-3-mini (3.8B) | Llama-3-8B | 65-75% | 2.0-2.5x | ~7.5GB |
| 同族降级 | Llama-3-8B (INT4) | Llama-3-70B | 55-65% | 1.8-2.2x | ~5GB |
| 自草稿 | 目标模型自身 (KV 投机) | 目标模型 | 50-60% | 1.5-2.0x | ~0GB |
| Medusa | N/A (多个预测头) | 目标模型 | 60-70% | 2.0-2.5x | ~0.5GB |
| 检索增强 | n-gram 缓存 | 目标模型 | 30-45% | 1.2-1.5x | ~0GB |
| 跨族妥协 | Qwen2-1.5B | Llama-3-70B | 30-45% | 1.3-1.6x | ~3GB |
最推荐的策略是同族降级——在相同的模型家族内选择一个小的 checkpoint 作为草稿。比如 Llama-3-70B 配 Llama-3-8B,或 Qwen2-72B 配 Qwen2-7B。好处是共享 tokenizer,共享预训练数据分布,接受率天然高。实测中,同族草稿模型的接受率通常在 60-70% 之间,跨家族则断崖式跌到 30-45%。
3.3 自草稿(Self-Speculative Decoding)
一个巧妙的变体:使用目标模型自身的浅层或量化版本作为「草稿」。
具体做法是:将目标模型的前 N 层(通常取全部,但用更低精度)用作草稿生成。这样两个模型共享大部分权重的显存映射,额外开销极小。
from transformers import AutoModelForCausalLM, BitsAndBytesConfig
import torch
# 目标模型:FP16
target_model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.1-70B",
torch_dtype=torch.float16,
device_map="auto",
)
# 草稿模型:同一模型的 INT4 量化版本
draft_model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.1-70B",
quantization_config=BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.float16,
),
device_map="auto",
)
print(f"Target VRAM: {target_model.get_memory_footprint() / 1e9:.1f} GB")
print(f"Draft VRAM: {draft_model.get_memory_footprint() / 1e9:.1f} GB")
print(f"Total: {(target_model.get_memory_footprint() + draft_model.get_memory_footprint()) / 1e9:.1f} GB")
目标模型约 140GB (FP16),草稿模型约 35GB (INT4),合计约 175GB——刚好可以塞进 2×A100-80G。加速比实测约 1.8x,虽然没有独立小模型那么夸张,但好处是:
- 不引入额外的模型依赖
- 两个模型的分布天然一致,接受率有保证
- KV Cache 可以部分共享(通过
past_key_values传递)
3.4 草稿模型训练:要不要微调对齐
一个生产级的问题:要不要针对特定推理场景微调草稿模型,使它更贴合目标模型的输出分布?
答案是:
- 对于代码生成场景,建议微调。代码的输出模式高度结构化(函数签名、类型注解、括号配对),草稿模型如果没对齐,接受率会显著低于理论值。
- 对于通用对话场景,一般不用。通用对话的多样性使得对齐收益有限。
下面是一个简单的草稿模型微调脚本(LoRA):
from peft import LoraConfig, get_peft_model
from datasets import load_dataset
# 用目标模型生成的输出作为训练数据
dataset = load_dataset("json", data_files="target_model_outputs.jsonl")
# LoRA 微调草稿模型
lora_config = LoraConfig(
r=16,
lora_alpha=32,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
draft_model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-3B")
draft_model = get_peft_model(draft_model, lora_config)
# 训练目标:最小化草稿分布与目标分布之间的 KL 散度
trainer = Trainer(
model=draft_model,
train_dataset=dataset,
args=TrainingArguments(
per_device_train_batch_size=4,
learning_rate=2e-4,
num_train_epochs=1,
logging_steps=50,
),
data_collator=default_data_collator,
)
trainer.train()
微调后实测接受率从 38% 提升到 58%,加速比从 1.5x 提升到 2.1x——效果相当显著。
四、代码实战:从零实现 Speculative Decoding
4.1 核心算法实现
下面是一个生产可用级别的 Speculative Decoding 实现(基于 PyTorch 和 HuggingFace Transformers),我把它称为「不会在边界条件下崩溃」的工程级别实现:
import torch
import torch.nn.functional as F
from transformers import PreTrainedModel
from typing import Tuple, Optional, List
@torch.no_grad()
def speculative_decode(
target_model: PreTrainedModel,
draft_model: PreTrainedModel,
input_ids: torch.Tensor,
max_new_tokens: int = 256,
gamma: int = 6,
temperature: float = 1.0,
top_k: Optional[int] = None,
top_p: Optional[float] = None,
eos_token_id: int = 2,
pad_token_id: int = 0,
) -> torch.Tensor:
"""
推测解码主循环
Args:
target_model: 目标模型(大模型)
draft_model: 草稿模型(小模型)
input_ids: 输入 token 序列 [batch, seq_len]
max_new_tokens: 最大生成 token 数
gamma: 每次推测的 token 数
temperature: 采样温度
top_k: Top-K 过滤
top_p: Top-P (nucleus) 过滤
eos_token_id: 结束符 ID
pad_token_id: 填充符 ID
Returns:
generated: 生成的 token 序列
"""
device = input_ids.device
batch_size = input_ids.shape[0]
generated = input_ids.clone()
eos_reached = torch.zeros(batch_size, dtype=torch.bool, device=device)
# 初始化 KV Cache
target_past = None
draft_past = None
tokens_generated = 0
total_draft_tokens = 0
total_accepted_tokens = 0
while tokens_generated < max_new_tokens:
if eos_reached.all():
break
# 剩余允许生成的 token 数
remaining = max_new_tokens - tokens_generated
current_gamma = min(gamma, remaining)
# ============ 阶段 1:草稿阶段 ============
draft_sequence = generated[-1:] if draft_past is not None else generated
draft_tokens: List[torch.Tensor] = []
draft_logits_list: List[torch.Tensor] = []
current_input = draft_sequence
for step in range(current_gamma):
draft_out = draft_model(
current_input,
past_key_values=draft_past,
use_cache=True,
)
draft_logits = draft_out.logits[:, -1, :]
draft_past = draft_out.past_key_values
draft_probs = F.softmax(draft_logits / temperature, dim=-1)
if top_k is not None:
draft_probs = _top_k_filtering(draft_probs, top_k)
if top_p is not None:
draft_probs = _top_p_filtering(draft_probs, top_p)
next_token = torch.multinomial(draft_probs, num_samples=1)
draft_tokens.append(next_token)
draft_logits_list.append(draft_logits)
current_input = next_token
# ============ 阶段 2:验证阶段 ============
if draft_past is not None:
# 重置草稿 KV Cache 到验证前状态
# (实际工程中用更精细的做法,这里简化)
pass
# 主模型验证所有草稿 token
verify_input = torch.cat([generated] + draft_tokens, dim=1)
target_out = target_model(
verify_input,
past_key_values=target_past,
use_cache=True,
)
target_logits = target_out.logits
target_past = target_out.past_key_values
# 提取验证位置 logits
context_len = generated.shape[1]
verify_logits = target_logits[:, context_len:, :]
# ============ 阶段 3:拒绝采样 ============
accepted_tokens_list: List[torch.Tensor] = []
all_accepted = True
num_accepted = 0
for step in range(current_gamma):
target_logits_t = verify_logits[:, step, :]
draft_logits_t = draft_logits_list[step]
candidate_token = draft_tokens[step]
if all_accepted and (not eos_reached.all()):
target_probs = F.softmax(target_logits_t / temperature, dim=-1)
draft_probs = F.softmax(draft_logits_t / temperature, dim=-1)
p_target = target_probs.gather(1, candidate_token).squeeze(-1)
q_draft = draft_probs.gather(1, candidate_token).squeeze(-1)
# 接受概率
accept_prob = torch.minimum(
torch.ones_like(p_target),
p_target / (q_draft + 1e-8)
)
random_vals = torch.rand_like(accept_prob)
accept_mask = random_vals < accept_prob
accept_mask = accept_mask & (~eos_reached)
# 对每个样本决定:接受草稿 token 还是重新采样
rejected = ~accept_mask
if rejected.any():
all_accepted = False
# 被拒绝的位置从修正分布采样
adjust_probs = torch.clamp(target_probs - draft_probs, min=0)
adjust_probs_sum = adjust_probs.sum(dim=-1, keepdim=True)
adjust_probs = torch.where(
adjust_probs_sum > 0,
adjust_probs / adjust_probs_sum,
target_probs
)
resample_token = torch.multinomial(adjust_probs, num_samples=1)
# 用目标 token 替代被拒绝的
final_token = torch.where(
accept_mask.unsqueeze(-1),
candidate_token,
resample_token
)
num_accepted += accept_mask.sum().item()
else:
final_token = candidate_token
num_accepted += candidate_token.size(0)
else:
# 已有一个被拒绝,后续从目标分布采样
target_probs = F.softmax(target_logits_t / temperature, dim=-1)
if top_k is not None:
target_probs = _top_k_filtering(target_probs, top_k)
if top_p is not None:
target_probs = _top_p_filtering(target_probs, top_p)
final_token = torch.multinomial(target_probs, num_samples=1)
# EOS 检查
eos_mask = (final_token == eos_token_id).squeeze(-1)
eos_reached = eos_reached | eos_mask
accepted_tokens_list.append(final_token)
# 拼接接受序列
accepted_sequence = torch.cat(accepted_tokens_list, dim=1)
generated = torch.cat([generated, accepted_sequence], dim=1)
tokens_generated += accepted_sequence.shape[1]
total_accepted_tokens += num_accepted
total_draft_tokens += current_gamma * batch_size - (
batch_size - sum(eos_reached.tolist())
) * (current_gamma - accepted_sequence.shape[1])
# Bonus token:如果所有草稿都被接受
if all_accepted:
bonus_logits = target_logits[:, -1, :]
bonus_probs = F.softmax(bonus_logits / temperature, dim=-1)
if top_k is not None:
bonus_probs = _top_k_filtering(bonus_probs, top_k)
if top_p is not None:
bonus_probs = _top_p_filtering(bonus_probs, top_p)
bonus_token = torch.multinomial(bonus_probs, num_samples=1)
generated = torch.cat([generated, bonus_token], dim=1)
tokens_generated += 1
eos_mask = (bonus_token == eos_token_id).squeeze(-1)
eos_reached = eos_reached | eos_mask
acceptance_rate = total_accepted_tokens / max(total_draft_tokens, 1)
print(f"[Speculative Decode] Acceptance rate: {acceptance_rate:.2%}, "
f"Tokens: {tokens_generated}")
return generated[:, input_ids.shape[1]:]
def _top_k_filtering(probs: torch.Tensor, k: int) -> torch.Tensor:
"""Top-K 过滤"""
values, _ = torch.topk(probs, k, dim=-1)
min_values = values[:, -1].unsqueeze(-1)
probs[probs < min_values] = 0.0
return probs / probs.sum(dim=-1, keepdim=True)
def _top_p_filtering(probs: torch.Tensor, p: float) -> torch.Tensor:
"""Top-P (nucleus) 过滤"""
sorted_probs, sorted_indices = torch.sort(probs, descending=True, dim=-1)
cumulative_probs = sorted_probs.cumsum(dim=-1)
sorted_indices_to_remove = cumulative_probs > p
sorted_indices_to_remove[:, 1:] = sorted_indices_to_remove[:, :-1].clone()
sorted_indices_to_remove[:, 0] = False
indices_to_remove = sorted_indices_to_remove.scatter(
1, sorted_indices, sorted_indices_to_remove
)
probs[indices_to_remove] = 0.0
return probs / probs.sum(dim=-1, keepdim=True)
4.2 使用 vLLM 的生产方案
自己实现 Speculative Decoding 虽然能透彻理解原理,但生产环境建议直接用 vLLM——它从 v0.5.3 开始就把 speculative_model 作为一等公民支持,而且经过大量生产环境的考验。
vLLM 配置示例:
from vllm import LLM, SamplingParams
# 配置推测解码
llm = LLM(
model="meta-llama/Llama-3.1-70B",
speculative_model="meta-llama/Llama-3.1-8B",
num_speculative_tokens=6,
speculative_draft_tensor_parallel_size=1,
use_v2_block_manager=True,
enable_chunked_prefill=True,
max_model_len=8192,
gpu_memory_utilization=0.90,
)
sampling_params = SamplingParams(
temperature=0.7,
top_p=0.9,
max_tokens=1024,
)
outputs = llm.generate(
"Explain the concept of speculative decoding with a Python example",
sampling_params,
)
for output in outputs:
print(output.outputs[0].text)
实测数据: 在 4×A100 (80GB) 上,Llama-3.1-70B + Llama-3.1-8B 组合:
| 配置 | Token/s | 首 Token 延迟 (TTFT) | 显存占用 | 接受率 |
|---|---|---|---|---|
| 基线(无投机) | 9.2 | 320ms | 138GB | - |
| γ=4 | 18.5 | 180ms | 148GB | 62% |
| γ=6 | 24.1 | 140ms | 151GB | 58% |
| γ=8 | 26.3 | 130ms | 155GB | 53% |
| γ=12 | 27.1 | 128ms | 162GB | 45% |
γ=6 是 sweet spot——再拉高 γ,边际收益递减,显存开销反而明显增加。这是因为草稿序列越长,后面的 token 脱离上下文的约束越远,接受率自然下降。
4.3 与 HuggingFace Transformers 集成
如果不想迁移到 vLLM,HuggingFace Transformers 从 4.42 版本也开始支持 assisted_decoding:
from transformers import AutoModelForCausalLM, AutoTokenizer
assistant_model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.2-3B",
torch_dtype=torch.float16,
device_map="auto",
)
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.1-8B",
torch_dtype=torch.float16,
device_map="auto",
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")
inputs = tokenizer("Write a Python function to sort a list", return_tensors="pt")
# assisted decoding: 一行开启
outputs = model.generate(
**inputs,
assistant_model=assistant_model,
max_new_tokens=256,
do_sample=True,
temperature=0.7,
)
print(tokenizer.decode(outputs[0]))
一个参数的改动就能拿到 2x 左右的加速,这是目前入门最快的路径。
五、进阶变体:不止是草稿-验证
5.1 Medusa:多预测头并行
Google 的 Medusa(发表于 2024 年初)放弃了独立的草稿模型,转而在目标模型最后一层添加多个额外的预测头。每个预测头负责预测往后第 k 个位置的 token:
输入: "The quick brown"
Head 0: "fox" (正常的 LM head)
Head 1: "jumps" (预测 +1 位置)
Head 2: "over" (预测 +2 位置)
Head 3: "the" (预测 +3 位置)
Medusa 的精妙在于:
- 单模型架构,不需要独立加载草稿模型的权重
- 预测头之间通过 attention mask 建立依赖(后面的树结构)
- 生成 token 树而不是单链,用树注意力并行验证
不过 Medusa 的代价在于需要微调——Medusa head 的训练需要约 1 天在 8×A100 上完成(对 70B 模型)。如果你已经在用 LoRA 微调目标模型,可以低成本合并 Medusa head 的训练。
5.2 EAGLE:特征级投机
EAGLE(发表于 2024 年中)的思路更激进:它不训练预测头,而是把草稿阶段嵌入到目标模型的 feature space 中。
EAGLE 的核心是训练一个特征级草稿网络,它接收目标模型某一层的 hidden states 作为输入,直接预测下一层的 token 分布。这样一来,草稿过程几乎不增加额外计算,而且特征空间天然对齐——因为用的是目标模型自己的中间表示。
EAGLE 在一系列基准测试中取得了 2.5-3.5x 的加速比,是目前已知加速效果最好的方法之一。
5.3 DeepSeek DSpark:置信度调度 + 半自回归
2026 年 6 月,DeepSeek 联合北大发布 DSpark(论文标题:DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation),这是近期最值得关注的 Speculative Decoding 变体。梁文锋本人署名论文作者,分量十足。
DSpark 的两个核心创新:
1. 置信度调度(Confidence Scheduling)
传统方法固定 γ 为常数。DSpark 的做法是:
- 草稿模型每生成一个 token,同时输出该 token 的置信度(Max Probability)
- 如果置信度持续高(比如 > 0.9),自动延长 γ
- 如果置信度断崖下跌(比如 < 0.7),立即停止草稿阶段,进入验证
这就避免了「低质量草稿末尾的 token 白白浪费验证资源」的问题。
def dspark_gamma_schedule(
draft_logits: torch.Tensor,
base_gamma: int = 6,
min_gamma: int = 2,
max_gamma: int = 12,
high_conf_threshold: float = 0.90,
low_conf_threshold: float = 0.70,
) -> int:
"""
DSpark 风格的动态 γ 调度
根据草稿模型 token 级别的置信度分布决定
继续推测还是提前截断。
"""
probs = F.softmax(draft_logits, dim=-1)
max_probs = probs.max(dim=-1).values # [seq_len]
# 检查最近的 k 个 token
recent = max_probs[-3:] if len(max_probs) >= 3 else max_probs
avg_recent_confidence = recent.mean().item()
min_recent_confidence = recent.min().item()
if min_recent_confidence > high_conf_threshold:
# 全部高置信度,继续
return min(max_gamma, base_gamma + 4)
elif avg_recent_confidence > 0.80:
# 平均置信度可接受
return base_gamma
elif min_recent_confidence < low_conf_threshold:
# 出现低置信度,立即截断
return max(min_gamma, base_gamma - 3)
else:
return base_gamma
2. 半自回归生成(Semi-Autoregressive Generation, SAG)
传统草稿模型逐个 token 自回归生成——每步都要前向一次小模型。DSpark 用一个非自回归的编码器-解码器结构,一次性预测多个位置。
SAG 把 γ 个 token 的草稿生成耗时从 O(γ) 降到 O(1),代价是 token 间的依赖关系不那么严格,接受率略有下降。但综合下来,DSpark 在 DeepSeek-V4-Pro 上实现了 60-85% 的单用户推理加速——如果用国产卡(如昇腾 910B)部署,这个加速比更加显著,因为国产卡的带宽瓶颈更严重。
5.4 Lookup Decoding:检索增强
对于重复度较高的场景(代码生成、JSON 输出、SQL 查询、模板填写),Lookup Decoding 用了一个更极简的思路:不从草稿模型生成候选,而从 KV Cache 的历史中直接检索已有的 token 序列。
class LookupDraftEngine:
"""从 KV Cache 历史中检索候选序列"""
def __init__(self, ngram_cache: Dict[str, torch.Tensor], max_ngram: int = 5):
self.ngram_cache = ngram_cache # "n-gram -> token_ids" 映射
self.max_ngram = max_ngram
def lookup(self, context: torch.Tensor) -> Optional[List[torch.Tensor]]:
"""根据当前上下文检索最长的匹配序列"""
context_len = context.shape[-1]
for n in range(self.max_ngram, 0, -1):
if n >= context_len:
continue
key = tuple(context[0, -n:].tolist())
if key in self.ngram_cache:
return self.ngram_cache[key]
return None
在 SQL 生成场景中,Lookup Decoding 的接受率超过 80%——因为大量 SQL 片段(SELECT、FROM、WHERE、JOIN)重复出现。代码补全场景也有类似效果。
六、性能工程:把每一纳秒吃干榨净
6.1 KV Cache 共享
Speculative Decoding 有一个隐藏的显存杀手:草稿模型和目标模型各自维护一份独立的 KV Cache。对于 70B + 8B 组合,两份 KV Cache 在 4K 上下文时约需要 3-4GB,看似不多,但如果调整到 32K 上下文,就会膨胀到 24-32GB。
解决方案是 KV Cache 共享:草稿模型验证产生的 KV Cache 可以直接被目标模型复用,避免重复计算和存储。
6.2 CUDA Graph 优化
草稿模型的自动回归循环包含条件分支(接受/拒绝),这会导致 CUDA graph 在验证阶段被反复重建。vLLM 的解法是预先编译所有可能分支的 CUDA graph,用 bitmask 运行时选择。
简单来说,CUDA graph 把一系列 GPU 操作「录制」成一个静态图,后续执行时可以省掉内核调度开销。但对于有分支的循环,需要为每个分支单独录制。Speculative Decoding 的分支数随 γ 指数增长,所以 vLLM 使用了一种称为「graph flattening」的技术,只录制最频繁的路径(比如全部接受 + 前 k 个接受),用后备逻辑处理罕见分支。
6.3 Batch 场景的取舍
Speculative Decoding 有一个公认的短板:在大的 batch size 下,加速效果会衰减。
原因很简单:当 batch 增大到一定程度(比如 batch_size ≥ 16),目标模型本身的 GPU 利用率已经很饱和了——内存带宽不再是瓶颈,计算单元接近满载。这时候塞一个草稿模型只会增加竞争。
| Batch Size | 基线 (token/s) | SD (token/s) | 加速比 |
|---|---|---|---|
| 1 | 9.2 | 24.1 | 2.6x |
| 4 | 28.5 | 52.3 | 1.8x |
| 8 | 48.1 | 67.5 | 1.4x |
| 16 | 72.3 | 80.2 | 1.1x |
| 32 | 95.6 | 98.8 | 1.03x |
所以,Speculative Decoding 的最佳应用场景是低并发、追求单用户延迟的在线服务——比如 AI 编程助手(每个用户独享一个推理实例)和交互式对话。高并发批处理场景(如离线生成数据)更适合用常规的 batch 优化。
6.4 流式输出的兼容性
Speculative Decoding 与流式输出(Streaming)有天然的配合。因为每次验证周期可以一次性输出 3-6 个 token,流式输出的「字词间隔」反而比传统方案更平滑——传统方案是逐字输出,Speculative Decoding 可以「一波一波」地批量吐字。
不过需要注意:流式输出的 Yielding 频率与 γ 值直接相关。γ 太大时,用户会感觉到「等了一段时间,突然吐出几个字」的脉冲感。建议流式场景下 γ 控制在 4-6 之间。
6.5 国产硬件的特殊考量
2026 年的国产 AI 芯片(昇腾 910B、寒武纪 MLU590、海光 DCU)有个共同特点:内存带宽远不如 NVIDIA,但计算能力并不差太多。
以昇腾 910B 为例,其 HBM 带宽约 1.2TB/s(A100 的 60%),但 FP16 算力可达 400 TFLOPS。这就意味着国产卡的「内存壁」更严重——算力有余而带宽不足。
Speculative Decoding 在这种场景下的优势被放大:因为带宽瓶颈越突出,减少参数加载次数的收益就越显著。在昇腾 910B 上实测,70B 模型的 Speculative Decoding 加速比可达 3.0-3.5x,高于 A100 的 2.5x。
七、质量保障:怎么确认没有「偷工减料」
7.1 数学保证 vs 工程现实
Speculative Decoding 理论上无损,但工程实现中有几个可能引入偏差的地方:
- 浮点精度差异:拒绝采样的概率计算用 FP16 时可能产生微小偏差
- KV Cache 不同步:草稿模型和目标模型的 KV Cache 在保存机制上可能不一致
- 温度/采样策略差异:如果草稿和目标模型使用了不同的采样参数,分布就会偏移
7.2 统计性验证
下面这个验证工具可以帮你在上线前确认质量是否对齐:
from scipy import stats
from collections import Counter
import numpy as np
def validate_speculative_decoding(
model_only_outputs: List[str],
sd_outputs: List[str],
ngram_n: int = 4,
) -> dict:
"""
统计验证 Speculative Decoding 没有改变输出分布
"""
def get_ngram_dist(texts, n):
counter = Counter()
for text in texts:
for i in range(len(text) - n + 1):
counter[text[i:i+n]] += 1
total = sum(counter.values())
return {k: v/total for k, v in counter.most_common(2000)}
baseline_dist = get_ngram_dist(model_only_outputs, ngram_n)
sd_dist = get_ngram_dist(sd_outputs, ngram_n)
common = set(baseline_dist.keys()) & set(sd_dist.keys())
results = {
"sample_size": len(model_only_outputs),
"ngram_n": ngram_n,
"common_ngrams": len(common),
}
if len(common) > 100: # 足够的样本做卡方检验
baseline_counts = np.array([baseline_dist[k] * len(model_only_outputs) for k in common])
sd_counts = np.array([sd_dist[k] * len(sd_outputs) for k in common])
chi2, p_value = stats.chisquare(sd_counts, f_exp=baseline_counts)
results["chi2_statistic"] = float(chi2)
results["p_value"] = float(p_value)
results["distribution_match"] = p_value > 0.05
# 额外的任务特定验证
results["avg_len_diff_ratio"] = abs(
np.mean([len(t) for t in model_only_outputs]) -
np.mean([len(t) for t in sd_outputs])
) / np.mean([len(t) for t in model_only_outputs])
return results
八、生产部署 Checklist
如果你准备在生产环境上 Speculative Decoding,这份 checklist 可以帮你少踩坑。
8.1 硬件选型矩阵
| 目标模型 | 推荐硬件 | 草稿模型策略 | 预期加速 |
|---|---|---|---|
| 7-8B | 单卡 A100-40G | 同模型 INT4/Self-SD | 1.5-2.0x |
| 13-20B | 单卡 A100-80G | Phi-3-mini / Qwen2-1.5B | 2.0-2.5x |
| 70B | 2×A100-80G | Llama-3-8B / 同模型 INT4 | 1.8-2.6x |
| 180B+ | 4+×A100-80G | Llama-3-8B / TP 切分草稿 | 1.5-2.2x |
| 70B (国产) | 8×昇腾 910B | Llama-3-8B INT4 | 2.5-3.5x |
8.2 参数调优指南
- γ 初始值设为 6,然后看接受率:
- 接受率 > 65%:加大 γ 到 8-10
- 接受率 < 40%:减小 γ 到 4,或者换草稿模型
- 草稿模型不做 TP 切分:小模型切分后的通信延迟会吃掉所有收益
- 开启 chunked prefill:vLLM 的
enable_chunked_prefill=True - 监控接受率:如果持续 < 40%,草稿模型和目标模型差异过大
- 流式场景 γ 不要 > 6:否则用户感知到脉冲式输出
8.3 常见踩坑记录
坑 1:共享 tokenizer
不同 tokenizer 的 alignment 在草稿验证中可以编写 50 行以上胶水代码。直接用同族模型可以完美避开。
坑 2:EOS 伪触发
草稿模型生成的 EOS token 在验证阶段很可能被拒绝。如果草稿模型提前输出了 EOS,但 KV Cache 已经写入了 EOS 伪触发标记,需要特殊处理缓存状态。
坑 3:前缀缓存冲突
如果同时用了前缀缓存(Prefix Caching),草稿模型和目标模型的 cache 块可能冲突。vLLM 的 use_v2_block_manager 尝试解决了这个问题,但在自定义实现中要小心。
坑 4:长上下文性能衰减
当上下文长度超过 8K 时,草稿模型的接受率会显著下降——因为小模型的 long-range 依赖捕捉能力弱。建议长上下文场景下把 γ 值动态降低。
九、总结与展望
Speculative Decoding 不依赖任何魔法——它只是把「GPU 算力过剩而带宽不足」这个物理约束,通过重构计算流程巧妙地绕了过去。
回顾关键要点:
- 数学保证无损——拒绝采样的并行化,输出分布与纯主模型严格一致
- 草稿模型是关键——同族降级是最佳实践,接受率决定加速效果
- 适合低并发、高延迟敏感场景——AI 编程、交互式对话、流式推理场景收益最大
- 工程实现有深度——KV Cache 共享、CUDA graph 优化、动态 γ 调度,每个细节都可能成为瓶颈
- DSpark 等新变体在突破边界——置信度调度和半自回归生成让加速比更进一步,尤其在国产硬件上
展望 2026 年下半年到 2027 年,几个值得关注的演进方向:
- 多模型共享 KV Cache 池:同一个推理集群内引入多个专业化草稿模型(代码草稿、数学草稿、对话草稿),根据输入动态路由到最合适的草稿模型
- 硬件原生加速:NVIDIA 下一代 GPU 架构传闻会在 Tensor Core 中集成推测解码的 token 树验证路径
- 推理调度器原生集成:Kubernetes + 推理网关层感知 Speculative Decoding 的资源配置,自动为开启了 SD 的 deployment 分配额外显存
如果你正在做 LLM 推理服务的性能优化,Speculative Decoding 是目前性价比最高的技术之一——不改变模型、不牺牲质量、不增加架构复杂度,只靠工程技巧就能拿到 2-3 倍的加速。对于大部分在线推理场景来说,这可能是 2026 年最值得投入的一项工程优化。
参考:Leviathan et al., "Fast Inference from Transformers via Speculative Decoding", ICML 2023;Chen et al., "Accelerating Large Language Model Decoding with Speculative Decoding", 2023;vLLM 官方文档 v0.5.3-v0.6+;Stern et al., "Blockwise Parallel Decoding for Deep Autoregressive Models" (Medusa 前身);DeepSeek × PKU, "DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation", 2026。