用 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 推理延迟的问题?评论区聊聊你的场景,或者有跑出来的接受率数据也可以分享——实测数据比论文更有参考价值。