全部术语

GLOSSARY ENTRY

投机解码

  • Speculative Decoding
  • Speculative Sampling

让较便宜的 draft 模型提出多个候选,再由 target 模型一次并行验证;精确采样用接受概率与拒绝校正分布保持 target 分布,同时尝试减少 target 的串行调用次数。

它减少的是 target 串行步数,不是 target 计算凭空消失

普通自回归 decode 每个 token 都要一次 target model 串行前向。投机解码让 draft model 从当前 prefix 连续提出 kk 个候选 y1,,yky_1,\ldots,y_k,target model 对这整个候选块做一次并行 verification,随后按协议提交一个连续接受前缀。

如果 draft 足够便宜、与 target 分布足够接近,单次 target call 可以提交多个 token,减少串行关键路径;但 target 仍要计算候选块各位置的 logits,draft 也有成本,还会发生 rejected suffix 的浪费与 KV rollback/copy-on-write。

精确接受与校正分布

在位置 ii,draft 条件分布为 qiq_i,target 条件分布为 pip_i,候选 yiqiy_i\sim q_i。接受概率:

ai(yi)=min(1,pi(yi)qi(yi))a_i(y_i)=\min\left(1,\frac{p_i(y_i)}{q_i(y_i)}\right)

从左到右测试;第一次拒绝后,丢弃该候选及其后缀,并从校正分布采样 replacement:

ri(x)=[pi(x)qi(x)]+z[pi(z)qi(z)]+r_i(x)=\frac{[p_i(x)-q_i(x)]_+}{\sum_z[p_i(z)-q_i(z)]_+}

其中 [u]+=max(u,0)[u]_+=\max(u,0)。直觉上,draft 已通过接受候选贡献了 min(p,q)\min(p,q) 的概率质量,拒绝时补上 target 比 draft 多出的正差。若 kk 个候选全部接受,再从 target 在完整候选之后的分布采样一个 bonus token,因此一次 verification 最多提交 k+1k+1 个 token。

边界情况:若 qi(y)=0q_i(y)=0,该 token 不可能由 draft 提出;公式只需在实际提出的 qi(y)>0q_i(y)>0 上求比值。若数值误差使 [pq]+[p-q]_+ 总和接近 0,实现需要稳定 fallback,不能除以 0。

INTERACTIVE EXPLAINER

proposal、并行 verification、接受前缀与拒绝校正

展示 k=4 的 exact sampling 协议;第 4 步是假设 y₃ 被拒绝,第 5 步展示全部接受时的 bonus 分支。

prefix state · target p · draft q · proposal length k=4accepted prefixKV state before drafty₁draft q1y₂draft q2y₃draft q3y₄draft q4one target verification forwardp₁(y₁), p₂(y₂), p₃(y₃), p₄(y₄) + p_bonus(·)left-to-right acceptanceaᵢ = min(1, pᵢ(yᵢ) / qᵢ(yᵢ)) · stop at first rejectionaccept y₁,y₂,y₃,y₄bonus ~ p(· | prefix,y₁,y₂,y₃,y₄)bonustarget-correct token
STEP 01 / 05Draft 串行提出 4 个候选

小 draft model 从当前 prefix 依次采样 y₁…y₄,并记录每个位置的 proposal distribution qᵢ。

图中把“一次 target verification”画成一个框,但真实模型会对 prefix 后 kk 个位置产生 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 性能。

三类指标必须分开

  1. 接受质量:平均接受 draft tokens、全接受率、拒绝位置分布。
  2. target 串行步数:target calls/token 或 tokens/verification call。
  3. 端到端性能: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 pp 与 draft qq 的 categorical 分布;本页在此基础上保持精确 target 分布。
  • Prefill/Decode给出普通串行基线与 TPOT;投机解码尝试让一次 target decode/verification 提交多个 token。
  • KV CachePagedAttention 承载候选状态、accepted prefix 与 rejected suffix 的回滚/共享。
  • 后续 MTP 页面必须区分训练目标、候选生成模块和本页 verifier 协议。

参考资料