它减少的是 target 串行步数,不是 target 计算凭空消失
普通自回归 decode 每个 token 都要一次 target model 串行前向。投机解码让 draft model 从当前 prefix 连续提出 个候选 ,target model 对这整个候选块做一次并行 verification,随后按协议提交一个连续接受前缀。
如果 draft 足够便宜、与 target 分布足够接近,单次 target call 可以提交多个 token,减少串行关键路径;但 target 仍要计算候选块各位置的 logits,draft 也有成本,还会发生 rejected suffix 的浪费与 KV rollback/copy-on-write。
精确接受与校正分布
在位置 ,draft 条件分布为 ,target 条件分布为 ,候选 。接受概率:
从左到右测试;第一次拒绝后,丢弃该候选及其后缀,并从校正分布采样 replacement:
其中 。直觉上,draft 已通过接受候选贡献了 的概率质量,拒绝时补上 target 比 draft 多出的正差。若 个候选全部接受,再从 target 在完整候选之后的分布采样一个 bonus token,因此一次 verification 最多提交 个 token。
边界情况:若 ,该 token 不可能由 draft 提出;公式只需在实际提出的 上求比值。若数值误差使 总和接近 0,实现需要稳定 fallback,不能除以 0。
INTERACTIVE EXPLAINER
proposal、并行 verification、接受前缀与拒绝校正
展示 k=4 的 exact sampling 协议;第 4 步是假设 y₃ 被拒绝,第 5 步展示全部接受时的 bonus 分支。
小 draft model 从当前 prefix 依次采样 y₁…y₄,并记录每个位置的 proposal distribution qᵢ。
图中把“一次 target verification”画成一个框,但真实模型会对 prefix 后 个位置产生 logits,并管理 candidate KV、accepted prefix 与 rejected suffix 的状态。不同 runtime 可能重算、暂存或用 copy-on-write blocks;这些是调度/缓存实现,不改变接受分布。
可运行的经验分布验证
实验使用 vocab=4 的一阶 autoregressive categorical target/draft。普通 target sampling 与 exact speculative sampling 各生成 10 万 token,比较每个前一 token 条件下的转移频率。它还报告平均接受 draft tokens/call 与 target calls/token。
import platform,torch
torch.manual_seed(101); DEVICE=torch.device("cpu"); DTYPE=torch.float64
TARGET=torch.tensor([[.55,.25,.15,.05],[.10,.60,.20,.10],
[.15,.15,.55,.15],[.25,.20,.15,.40]],dtype=DTYPE)
DRAFT=torch.tensor([[.50,.28,.17,.05],[.12,.55,.23,.10],
[.18,.14,.50,.18],[.28,.18,.17,.37]],dtype=DTYPE)
def sample(probs,g): return torch.multinomial(probs,1,generator=g).item()
def target_stream(length,seed):
g=torch.Generator().manual_seed(seed); current=0; output=[]
for _ in range(length): current=sample(TARGET[current],g); output.append(current)
return output
def speculative_step(current,k,g):
candidates=[]; q_rows=[]; state=current
for _ in range(k):
q=DRAFT[state]; token=sample(q,g)
candidates.append(token); q_rows.append(q); state=token
p_rows=[]; state=current
for token in candidates: p_rows.append(TARGET[state]); state=token
bonus_p=TARGET[state]; accepted=[]
for token,p,q in zip(candidates,p_rows,q_rows):
a=min(1.,(p[token]/q[token]).item())
if torch.rand((),generator=g).item()<=a:
accepted.append(token); continue
correction=torch.clamp(p-q,min=0); correction/=correction.sum()
accepted.append(sample(correction,g))
return accepted,len(accepted)-1
accepted.append(sample(bonus_p,g)); return accepted,k
def speculative_stream(length,k,seed):
g=torch.Generator().manual_seed(seed); current=0; output=[]
calls=accepted_draft=0
while len(output)<length:
emitted,accepted=speculative_step(current,k,g)
calls+=1; accepted_draft+=accepted
output.extend(emitted[:length-len(output)]); current=output[-1]
return output,calls,accepted_draft
def frequencies(tokens):
counts=torch.zeros_like(TARGET); previous=0
for token in tokens: counts[previous,token]+=1; previous=token
return counts/counts.sum(-1,keepdim=True)
def main():
length,k=100_000,4
baseline=target_stream(length,202)
speculative,calls,accepted=speculative_stream(length,k,303)
f1,f2=frequencies(baseline),frequencies(speculative)
e1=(f1-TARGET).abs().max().item(); e2=(f2-TARGET).abs().max().item()
cross=(f1-f2).abs().max().item()
assert e1<.02 and e2<.02 and cross<.02
print(f"python={platform.python_version()} torch={torch.__version__} device={DEVICE} dtype={DTYPE}")
print(f"tokens={length} vocab=4 proposal_len={k}")
print(f"baseline_vs_target_max_error={e1:.5f}")
print(f"speculative_vs_target_max_error={e2:.5f}")
print(f"baseline_vs_speculative_max_error={cross:.5f}")
print(f"mean_accepted_draft_tokens_per_call={accepted/calls:.3f}")
print(f"target_calls_per_emitted_token={calls/length:.3f}")
if __name__=="__main__": main()
完整文件位于 examples/speculative_decoding_exact.py。实际结果:baseline 对 target 最大误差 0.00635,speculative 对 target 0.00498,二者最大差 0.01029;平均接受 3.512 个 draft tokens/target call,target calls/token 为 0.222。
经验频率只能在容差内支持分布一致,不能数学证明实现无 bug;这里用已知有限状态 target 做可重复单元测试。Python 循环速度也不是生产 kernel 性能。
三类指标必须分开
- 接受质量:平均接受 draft tokens、全接受率、拒绝位置分布。
- target 串行步数:target calls/token 或 tokens/verification call。
- 端到端性能:TTFT、TPOT、吞吐,包括 draft 成本、target verification shape、调度、KV 管理、采样与通信。
高接受率不保证加速:draft 太大、verification kernel 低效、batch 动态性差、proposal length 过大或 rollback 成本高时,端到端可能变慢。反之,小 draft 与高 target/draft latency ratio、相近分布、适合并行 verification 的硬件更可能受益。
与 MTP、树形投机的边界
MTP 可以在训练时加入多个未来 token 预测目标,也可以为推理提供 candidate heads/blocks;但“训练时预测多个 token”不等于“一次接受多个 token”。只要候选要保持 target sampling 分布,仍需明确 verifier 与接受/校正协议。
Medusa/EAGLE 等树形候选把单链 proposal 扩展为 candidate tree,并用 tree attention/verification 复用计算;它们改变候选结构和调度,不自动取消 target correctness 要求。
解码扩展总览用静态 candidate tree 与可运行 rollback 例子区分 tentative nodes、accepted path 和 committed KV;本页则专注单链 exact acceptance/correction 分布。
上下游关系
- Logits 与基础采样定义 target 与 draft 的 categorical 分布;本页在此基础上保持精确 target 分布。
- Prefill/Decode给出普通串行基线与 TPOT;投机解码尝试让一次 target decode/verification 提交多个 token。
- KV Cache 与 PagedAttention 承载候选状态、accepted prefix 与 rejected suffix 的回滚/共享。
- 后续 MTP 页面必须区分训练目标、候选生成模块和本页 verifier 协议。