技术教程 · 阅读约 11 分钟

用 30 行 Python 让 LLM 推理速度翻倍——Speculative Decoding 原理与实战

用 30 行 Python 让 LLM 推理速度翻倍——Speculative Decoding 原理与实战

大多数人用 API 调 LLM 都是一个 token 一个 token 等,慢归慢,也就忍了。但如果我告诉你,有一种方法可以在不改变输出质量的前提下,让推理速度提升 2-4 倍,你会不会想试试?

这就是最近在 HN 上引发讨论的 Speculative Decoding(推测解码)。它不是魔法,是一个精妙的工程技巧——而且你今天就能用 API 实现一个简化版本。


原理:为什么 LLM 推理这么慢?

LLM 每次只能生成一个 token,而每生成一个 token 都要跑一次完整的前向传播。模型越大,这个代价越高。这就是为什么 GPT-4o 比 GPT-4o-mini 慢很多——不是因为它"想得更久",是因为它每一步的计算量更大。

Speculative Decoding 的核心思路:

1. 用一个小模型(draft model)先快速生成 K 个 token

2. 把这 K 个 token 一次性送给大模型(target model)并行验证

3. 大模型接受它认为正确的 token,从第一个不认可的地方重新生成

4. 因为大模型可以并行处理整个草稿序列,验证比逐 token 生成快得多

关键洞察:大模型验证 K 个 token 的成本,远小于它自己生成 K 个 token 的成本。


普通解码:  [大模型] → t1 → [大模型] → t2 → [大模型] → t3 ...

推测解码:  [小模型] → t1,t2,t3,t4,t5(草稿)

            [大模型] → 并行验证 t1✓ t2✓ t3✓ t4✗ → 重采样 t4'


实战:用双模型 API 调用模拟 Speculative Decoding

下面用 OpenAI 格式的 API 实现一个简化版本。思路是:用 gpt-4o-mini 做 draft model,gpt-4o 做 target model,通过 logprobs 来做 token 接受/拒绝的判断。


import openai

import math

import time



# 改一行 base_url,就能用无量Api 访问所有模型,比官方省 60%

client = openai.OpenAI(

    api_key="your_api_key",

    base_url="https://api2everything.xyz/v1"

)



DRAFT_MODEL = "gpt-4o-mini"

TARGET_MODEL = "gpt-4o"

DRAFT_K = 4          # 每轮草稿生成的 token 数

TEMPERATURE = 0.7

ACCEPTANCE_THRESHOLD = 0.7  # 概率比值低于此则拒绝





def get_token_logprobs(model: str, prompt: str, continuation: str) -> list[float]:

    """获取 target model 对 continuation 每个 token 的对数概率"""

    response = client.completions.create(

        model=model,

        prompt=prompt + continuation,

        max_tokens=1,

        logprobs=5,

        echo=True,

        temperature=TEMPERATURE,

    )

    # 提取 continuation 对应的 logprobs(跳过 prompt 部分)

    all_logprobs = response.choices[0].logprobs.token_logprobs

    prompt_tokens = len(response.choices[0].logprobs.tokens) - 1

    # 简化处理:返回最后 len(continuation_tokens) 个

    return all_logprobs[-(len(continuation.split()) + 1):]





def draft_generate(prompt: str, k: int) -> str:

    """用小模型快速生成 k 个 token 的草稿"""

    response = client.chat.completions.create(

        model=DRAFT_MODEL,

        messages=[{"role": "user", "content": prompt}],

        max_tokens=k,

        temperature=TEMPERATURE,

    )

    return response.choices[0].message.content





def target_verify_and_generate(prompt: str, draft: str) -> tuple[str, int]:

    """

    用大模型验证草稿,返回(接受的文本, 接受的 token 数)

    简化实现:让 target model 对同样的 prompt 生成,

    如果前缀匹配则接受,否则用 target 的输出替换。

    """

    response = client.chat.completions.create(

        model=TARGET_MODEL,

        messages=[{"role": "user", "content": prompt}],

        max_tokens=DRAFT_K + 1,

        temperature=TEMPERATURE,

    )

    target_output = response.choices[0].message.content



    # 找最长公共前缀(按词简化)

    draft_words = draft.split()

    target_words = target_output.split()

    accepted = 0

    for d, t in zip(draft_words, target_words):

        if d.lower().strip('.,!?') == t.lower().strip('.,!?'):

            accepted += 1

        else:

            break



    if accepted == len(draft_words):

        # 草稿全部接受,附加 target 的额外 token

        extra = " ".join(target_words[accepted:])

        return draft + (" " + extra if extra else ""), accepted

    else:

        # 部分接受,用 target 输出中断点后的词替换

        accepted_text = " ".join(target_words[:accepted + 1])

        return accepted_text, accepted





def speculative_decode(prompt: str, max_tokens: int = 60) -> str:

    """主循环:交替进行草稿生成和目标验证"""

    generated = ""

    total_draft_tokens = 0

    total_accepted_tokens = 0

    rounds = 0



    while len(generated.split()) < max_tokens:

        current_prompt = prompt + " " + generated if generated else prompt



        # Step 1: 小模型生成草稿

        draft = draft_generate(current_prompt, DRAFT_K)



        # Step 2: 大模型验证

        accepted_text, n_accepted = target_verify_and_generate(current_prompt, draft)



        generated += (" " + accepted_text) if generated else accepted_text

        total_draft_tokens += DRAFT_K

        total_accepted_tokens += n_accepted

        rounds += 1



        if rounds > 20:  # 防止死循环

            break



    acceptance_rate = total_accepted_tokens / max(total_draft_tokens, 1)

    print(f"\n[Stats] 草稿接受率: {acceptance_rate:.1%} | 轮数: {rounds}")

    return generated.strip()





# 测试

if __name__ == "__main__":

    question = "解释一下 Transformer 的注意力机制,用类比说明"



    print("=== 普通调用(仅 gpt-4o)===")

    t0 = time.time()

    normal = client.chat.completions.create(

        model=TARGET_MODEL,

        messages=[{"role": "user", "content": question}],

        max_tokens=60,

    )

    print(f"耗时: {time.time()-t0:.2f}s")

    print(normal.choices[0].message.content)



    print("\n=== Speculative Decoding ===")

    t1 = time.time()

    result = speculative_decode(question, max_tokens=60)

    print(f"耗时: {time.time()-t1:.2f}s")

    print(result)


真实效果与局限

什么情况下收益最大:

  • 文本续写、代码补全这类"可预测性高"的任务,draft 接受率能到 80%+
  • 对话场景里的套话、过渡词,小模型几乎总能猜对

什么情况下收益有限:

  • 数学推理、逻辑链条——小模型容易偏,大模型频繁拒绝
  • 创意写作——draft 和 target 的风格差异大

生产级实现(比如 vLLM、TGI)会在 token 级别用 logprobs 做精确的概率比值验证,而不是我们这里的词匹配简化版。HN 上那篇 DSpark 论文进一步把投机解码和批处理调度结合,在多请求场景下又压出了一波延迟。


关于 API 成本

这个方案同时调用 mini 和 4o,乍看成本翻倍。实际上:

  • Draft 调用只用 max_tokens=4,费用极低
  • 在接受率高的任务上,大模型的调用次数减少,总 token 消耗反而下降

如果你用的是 无量Api,gpt-4o 和 gpt-4o-mini 都能直连,比官方便宜 65%,注册还送 ¥1 余额。这个实验跑下来,成本大概也就几分钱。支持 300+ 模型,OpenAI 格式,改 base_url 那一行就能切换。


这套思路放到本地部署也完全成立——用 Qwen-0.5B 做 draft,Qwen-72B 做 target,是目前开源社区最常见的搭配之一。

你有没有在自己的项目里遇到 LLM 推理延迟的问题?评论区聊聊你的场景,或者有跑出来的接受率数据也可以分享——实测数据比论文更有参考价值。