Speculative Decoding:让大模型生成速度翻倍的秘密
·
LLaMA-13B 解码一 Token 约 1.2ms。一秒钟只能出 800 个字——不够快。为什么不能一次出多个 Token?因为解码是自回归的——下一个 Token 依赖上一个的输出。
推测解码(Speculative Decoding)打破了自回归瓶颈。它用一个超小的 Draft Model 提前"猜"一串 Token,然后用大模型并行验证。猜对了就一次吐出多个 Token。
为什么大模型输出速度慢
# 标准 Decode——逐 Token 自回归
def standard_decode(model, prompt, max_tokens=128):
output = prompt[:]
for step in range(max_tokens):
# 每次只能算 1 个 Token
logits = model.infer(output) # 1 次推理 → 1 个 Token
next_token = sample(logits[-1])
output.append(next_token)
if next_token == EOS:
break
return output
# 128 个 Token = 128 次推理 → 128 × 1.2ms = 154ms
自回归限制:第 5 个 Token 的推理必须等第 1-4 个算完。NPU 的 24 TFLOPS 算力在解码阶段只用了 35%。
Draft Model 如何提前生成 Token
# 推测解码——Draft Model 猜一串,大模型批量验证
def speculative_decode(large_model, draft_model, prompt, gamma=5):
"""
large_model: LLaMA-13B(目标模型)
draft_model: 一个小模型(如 LLaMA-68M,只有目标模型的 0.5% 参数)
gamma: 每次猜测几 Token
"""
output = prompt[:]
while len(output) < max_tokens:
# Step 1: Draft Model 快速猜 gamma 个 Token
draft_tokens = []
draft_state = draft_model.init_state(output)
for i in range(gamma):
logits = draft_model.infer(draft_state) # 0.01ms(68M 模型)
next_tok = sample(logits[-1])
draft_tokens.append(next_tok)
draft_state = draft_model.append_state(draft_state, next_tok)
# Step 2: 大模型并行验证所有 Draft Token
candidate_seq = output + draft_tokens
# 一次推理算完所有 candidate 的 logits!
large_logits = large_model.infer(candidate_seq)
# 不是 gamma 次推理——只推理 1 次
# Step 3: 逐个验证——保留通过验证的前 k 个 Token
accepted = 0
for i in range(gamma):
# 大模型的 predicted 分布 vs Draft Model 的分布
large_probs = softmax(large_logits[len(output) + i - 1])
draft_probs = softmax(draft_logits[i])
if draft_tokens[i] == sample(large_probs):
# Draft 猜对了 → 接受
accepted += 1
else:
# Draft 猜错了 → 从大模型的正确分布重新采样
output.append(sample(large_probs))
break
# 如果全部猜对——一次验证拿到了 gamma+1 个 Token
if accepted == gamma:
output.extend(draft_tokens)
output.append(sample(large_logits[-1]))
return output
# 如果 gamma=5、接受率 80%——每步平均产出 4 Token
# 128 Token 从 128 次推理降到 32 次 → 延迟从 154ms 降到 38ms
Draft Model 的关键:它算得快但不是必须准——猜错了只是浪费一次验证,不会降低输出质量。大模型的验证保证最终输出跟不用推测解码完全一致。
昇腾NPU如何并行验证
// CANN Runtime 上的并行验证——Draft Token 在同一个 Batch 中推理
// draft_tokens = [a, b, c, d, e]
// 不推测时:5 次解码推理
for (int i = 0; i < 5; i++) {
aclmdlExecute(model, single_token_input, output);
}
// 推测解码:拼接 draft_tokens 做一次推理
// 输入序列 = prompt + [a, b, c, d, e]
// GE 会执行 Prefill 路径(输入多 Token)——比 Decode 路径快 N 倍
int16_t candidate[] = {45, 892, 312, 67, 1289, 34, a, b, c, d, e};
aclmdlExecute(model, candidate_input, output);
// 一次推理产出所有 candidate 位置的 logits
// 时间 = 一次 Prefill ≈ 多次 Decode
// 验证在 CPU 上做——跟推理异步
for (int i = 0; i < gamma; i++) {
if (reject(output[i], draft_logits[i])) {
output[output_len] = resample(output[i]);
break;
}
output[output_len + i] = draft_tokens[i];
}
关键点:大模型用一次 Prefill 推理验证了 gamma 个 Token。Prefill 的 Cube 利用率(80%)远高于 Decode(35%)——GPU 算力的浪费减少了。
性能提升
# 推测解码的实测性能(LLaMA-13B, Ascend 910)
results = {
"gamma_1": {"tokens/s": 800, "accept_rate": 1.0}, # 等价于标准解码
"gamma_3": {"tokens/s": 1800, "accept_rate": 0.75}, # 猜 3 个接受 2.25 个
"gamma_5": {"tokens/s": 2400, "accept_rate": 0.65}, # 猜 5 个接受 3.25 个
"gamma_8": {"tokens/s": 2800, "accept_rate": 0.55}, # 接受率下降后收益饱和
}
# 最佳 gamma≈5:速度提升约 3 倍
Draft Model 的参数量=目标模型的 0.5%-1%。太大就没有速度优势。太小接受率太低。LLaMA-13B 的常用 Draft Model 是 LLaMA-68M——0.5% 的参数量,推理快了 100 倍,接受率约 65%。
// Draft Model 的 CANN 部署——跟大模型共享 Runtime 实例
// 两个模型用一个 Device、不同 Context
aclrtContext draft_ctx, main_ctx;
aclrtCreateContext(&draft_ctx, 0);
aclrtCreateContext(&main_ctx, 0);
uint32_t draft_model_id, main_model_id;
aclmdlLoadFromFile("draft_68m.om", &draft_model_id);
aclmdlLoadFromFile("llama_13b.om", &main_model_id);
// 交替推理:Draft → 大模型 → Draft → 大模型
for (int step = 0; step < gamma; step++) {
aclrtSetContext(draft_ctx);
aclmdlExecute(draft_model_id, draft_input, draft_output); // 0.01ms
}
aclrtSetContext(main_ctx);
aclmdlExecute(main_model_id, verify_input, verify_output); // 1 次替代 5 次
参考仓库
更多推荐


所有评论(0)