【大模型】- 推测解码

算法

推测解码

加速自回归生成的通用方法

类型: 学习 | 语言: Python | 🏷 前置:《推理优化》(本系列第 11 篇)

学习目标

  • 理解推测解码的通用原理
  • 实现基础推测解码器
  • 训练草稿模型
  • 优化推测解码参数
  • 评估加速效果

推测解码概述

推测解码(Speculative Decoding)是一种加速自回归生成的通用方法,使用小模型猜测多个 token,大模型验证。

核心思想

  • 草稿模型:小而快的模型,用于生成候选 token
  • 目标模型:大而准确的模型,用于验证候选 token
  • 接受/拒绝:根据概率决定是否接受候选 token

加速原理

  1. 并行验证:大模型一次验证多个 token
  2. 减少前向传播:如果猜测正确,跳过多次计算
  3. 批量处理:提高 GPU 利用率

实现

基础推测解码器

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
import torch
import torch.nn.functional as F

class SpeculativeDecoder:
def __init__(self, target_model, draft_model, tokenizer, k=5):
self.target_model = target_model
self.draft_model = draft_model
self.tokenizer = tokenizer
self.k = k # 每次猜测的 token 数

def generate(self, prompt, max_new_tokens):
"""推测解码生成"""
input_ids = self.tokenizer.encode(prompt, return_tensors="pt")

generated = input_ids.clone()
total_tokens = 0
accepted_tokens = 0

while generated.shape[1] < max_new_tokens:
# 步骤 1:草稿模型猜测
draft_tokens, draft_probs = self.draft_generate(generated)

# 步骤 2:目标模型验证
accepted, new_token = self.target_verify(
generated, draft_tokens, draft_probs
)

# 步骤 3:更新生成序列
if accepted is not None:
generated = torch.cat([generated, accepted], dim=1)
accepted_tokens += accepted.shape[1]
else:
generated = torch.cat([generated, new_token], dim=1)

total_tokens += self.k + 1

# 检查结束条件
if self.tokenizer.eos_token_id in generated:
break

# 计算接受率
acceptance_rate = accepted_tokens / total_tokens if total_tokens > 0 else 0

return generated, acceptance_rate

def draft_generate(self, input_ids):
"""草稿模型生成"""
tokens = []
probs = []

for _ in range(self.k):
with torch.no_grad():
outputs = self.draft_model(input_ids)
logits = outputs.logits[:, -1, :]
next_probs = F.softmax(logits, dim=-1)

# 采样
next_token = torch.multinomial(next_probs, 1)

tokens.append(next_token)
probs.append(next_probs.gather(1, next_token))

# 更新输入
input_ids = torch.cat([input_ids, next_token], dim=1)

return torch.cat(tokens, dim=1), torch.cat(probs, dim=1)

def target_verify(self, input_ids, draft_tokens, draft_probs):
"""目标模型验证"""
# 拼接输入
full_input = torch.cat([input_ids, draft_tokens], dim=1)

with torch.no_grad():
outputs = self.target_model(full_input)
target_logits = outputs.logits[:, input_ids.shape[1]-1:-1, :]

# 逐个验证
accepted_tokens = []

for i in range(draft_tokens.shape[1]):
# 目标模型概率
target_probs = F.softmax(target_logits[:, i, :], dim=-1)

# 草稿模型概率
draft_prob = draft_probs[:, i]

# 计算接受概率
token = draft_tokens[:, i]
target_prob = target_probs.gather(1, token.unsqueeze(1))

accept_ratio = target_prob / (draft_prob + 1e-8)
accept_prob = torch.min(torch.ones_like(accept_ratio), accept_ratio)

# 接受或拒绝
if torch.rand(1).item() < accept_prob.item():
accepted_tokens.append(token.unsqueeze(1))
else:
# 拒绝:从目标模型重新采样
# 修正概率分布
corrected_probs = torch.clamp(target_probs - draft_probs, min=0)
corrected_probs = corrected_probs / (corrected_probs.sum() + 1e-8)

new_token = torch.multinomial(corrected_probs, 1)

if accepted_tokens:
return torch.cat(accepted_tokens, dim=1), new_token
else:
return None, new_token

# 所有 token 都被接受
if accepted_tokens:
return torch.cat(accepted_tokens, dim=1), None
else:
return None, draft_tokens[:, :1]

自适应推测解码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
class AdaptiveSpeculativeDecoder:
def __init__(self, target_model, draft_model, tokenizer, min_k=1, max_k=10):
self.target_model = target_model
self.draft_model = draft_model
self.tokenizer = tokenizer
self.min_k = min_k
self.max_k = max_k
self.current_k = (min_k + max_k) // 2

# 自适应参数
self.acceptance_history = []
self.history_size = 100

def adapt_k(self, acceptance_rate):
"""根据接受率调整猜测长度"""
self.acceptance_history.append(acceptance_rate)

if len(self.acceptance_history) > self.history_size:
self.acceptance_history.pop(0)

if len(self.acceptance_history) >= 10:
avg_acceptance = sum(self.acceptance_history[-10:]) / 10

if avg_acceptance > 0.8:
# 接受率高,增加猜测长度
self.current_k = min(self.current_k + 1, self.max_k)
elif avg_acceptance < 0.5:
# 接受率低,减少猜测长度
self.current_k = max(self.current_k - 1, self.min_k)

def generate(self, prompt, max_new_tokens):
"""自适应推测解码生成"""
input_ids = self.tokenizer.encode(prompt, return_tensors="pt")

generated = input_ids.clone()

while generated.shape[1] < max_new_tokens:
# 使用当前 k 值
draft_tokens, draft_probs = self.draft_generate(generated, self.current_k)

accepted, new_token = self.target_verify(
generated, draft_tokens, draft_probs
)

if accepted is not None:
generated = torch.cat([generated, accepted], dim=1)
acceptance_rate = accepted.shape[1] / self.current_k
else:
generated = torch.cat([generated, new_token], dim=1)
acceptance_rate = 0

# 自适应调整
self.adapt_k(acceptance_rate)

if self.tokenizer.eos_token_id in generated:
break

return generated

def draft_generate(self, input_ids, k):
"""草稿模型生成 k 个 token"""
tokens = []
probs = []

for _ in range(k):
with torch.no_grad():
outputs = self.draft_model(input_ids)
logits = outputs.logits[:, -1, :]
next_probs = F.softmax(logits, dim=-1)

next_token = torch.multinomial(next_probs, 1)

tokens.append(next_token)
probs.append(next_probs.gather(1, next_token))

input_ids = torch.cat([input_ids, next_token], dim=1)

return torch.cat(tokens, dim=1), torch.cat(probs, dim=1)

批量推测解码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
class BatchSpeculativeDecoder:
def __init__(self, target_model, draft_model, tokenizer, k=5, batch_size=32):
self.target_model = target_model
self.draft_model = draft_model
self.tokenizer = tokenizer
self.k = k
self.batch_size = batch_size

def generate_batch(self, prompts, max_new_tokens):
"""批量推测解码"""
# 分词
input_ids = self.tokenizer(
prompts,
return_tensors="pt",
padding=True,
truncation=True
).input_ids

batch_size = input_ids.shape[0]
generated = input_ids.clone()

while generated.shape[1] < max_new_tokens:
# 批量草稿生成
draft_tokens, draft_probs = self.batch_draft_generate(generated)

# 批量目标验证
accepted, new_tokens = self.batch_target_verify(
generated, draft_tokens, draft_probs
)

# 更新
generated = self.update_generated(generated, accepted, new_tokens)

return generated

def batch_draft_generate(self, input_ids):
"""批量草稿生成"""
all_tokens = []
all_probs = []

for _ in range(self.k):
with torch.no_grad():
outputs = self.draft_model(input_ids)
probs = F.softmax(outputs.logits[:, -1, :], dim=-1)

tokens = torch.multinomial(probs, 1)

all_tokens.append(tokens)
all_probs.append(probs.gather(1, tokens))

input_ids = torch.cat([input_ids, tokens], dim=1)

return torch.cat(all_tokens, dim=1), torch.cat(all_probs, dim=1)

def batch_target_verify(self, input_ids, draft_tokens, draft_probs):
"""批量目标验证"""
full_input = torch.cat([input_ids, draft_tokens], dim=1)

with torch.no_grad():
outputs = self.target_model(full_input)
target_logits = outputs.logits[:, input_ids.shape[1]-1:-1, :]

# 批量验证
accepted_list = []
new_tokens_list = []

for i in range(draft_tokens.shape[0]):
accepted, new_token = self.verify_single(
input_ids[i:i+1],
draft_tokens[i:i+1],
draft_probs[i:i+1],
target_logits[i:i+1]
)

if accepted is not None:
accepted_list.append(accepted)
else:
new_tokens_list.append(new_token)

return accepted_list, new_tokens_list

def verify_single(self, input_ids, draft_tokens, draft_probs, target_logits):
"""验证单个序列"""
accepted_tokens = []

for i in range(draft_tokens.shape[1]):
target_probs = F.softmax(target_logits[:, i, :], dim=-1)
draft_prob = draft_probs[:, i]

token = draft_tokens[:, i]
target_prob = target_probs.gather(1, token.unsqueeze(1))

accept_ratio = target_prob / (draft_prob + 1e-8)
accept_prob = torch.min(torch.ones_like(accept_ratio), accept_ratio)

if torch.rand(1).item() < accept_prob.item():
accepted_tokens.append(token.unsqueeze(1))
else:
corrected_probs = torch.clamp(target_probs - draft_probs, min=0)
corrected_probs = corrected_probs / (corrected_probs.sum() + 1e-8)
new_token = torch.multinomial(corrected_probs, 1)

if accepted_tokens:
return torch.cat(accepted_tokens, dim=1), new_token
else:
return None, new_token

if accepted_tokens:
return torch.cat(accepted_tokens, dim=1), None
else:
return None, draft_tokens[:, :1]

训练草稿模型

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
def train_draft_model(target_model, draft_model, tokenizer, train_data, config):
"""训练草稿模型"""
optimizer = torch.optim.AdamW(draft_model.parameters(), lr=config["lr"])

for epoch in range(config["epochs"]):
total_loss = 0

for batch in train_data:
# 目标模型生成
inputs = tokenizer(batch["text"], return_tensors="pt")

with torch.no_grad():
target_outputs = target_model(**inputs)
target_logits = target_outputs.logits

# 草稿模型预测
draft_outputs = draft_model(**inputs)
draft_logits = draft_outputs.logits

# 计算损失
loss = F.cross_entropy(
draft_logits.view(-1, draft_logits.size(-1)),
inputs["input_ids"].view(-1)
)

# 反向传播
loss.backward()
optimizer.step()
optimizer.zero_grad()

total_loss += loss.item()

print(f"Epoch {epoch + 1}: Loss = {total_loss / len(train_data):.4f}")

return draft_model

评估加速效果

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
def evaluate_speedup(target_model, draft_model, tokenizer, test_prompts, k=5):
"""评估加速效果"""
# 传统生成
start_time = time.time()
for prompt in test_prompts:
inputs = tokenizer(prompt, return_tensors="pt")
target_model.generate(**inputs, max_new_tokens=100)
traditional_time = time.time() - start_time

# 推测解码
decoder = SpeculativeDecoder(target_model, draft_model, tokenizer, k)

start_time = time.time()
total_acceptance = 0
for prompt in test_prompts:
_, acceptance_rate = decoder.generate(prompt, max_new_tokens=100)
total_acceptance += acceptance_rate
speculative_time = time.time() - start_time

avg_acceptance = total_acceptance / len(test_prompts)
speedup = traditional_time / speculative_time

print(f"传统生成时间: {traditional_time:.2f}s")
print(f"推测解码时间: {speculative_time:.2f}s")
print(f"加速比: {speedup:.2f}x")
print(f"平均接受率: {avg_acceptance:.2%}")

return speedup, avg_acceptance

参数优化

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
def optimize_parameters(target_model, draft_model, tokenizer, eval_data):
"""优化推测解码参数"""

k_values = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
best_k = 5
best_speedup = 0

for k in k_values:
decoder = SpeculativeDecoder(target_model, draft_model, tokenizer, k)

speedup, acceptance = evaluate_speedup(
target_model, draft_model, tokenizer, eval_data[:10], k
)

if speedup > best_speedup:
best_speedup = speedup
best_k = k

print(f"k={k}: 加速比={speedup:.2f}x, 接受率={acceptance:.2%}")

print(f"\n最佳 k={best_k}, 加速比={best_speedup:.2f}x")

return best_k

总结

推测解码是一种通用的加速自回归生成的方法,通过草稿模型猜测和目标模型验证来减少计算。关键参数包括猜测长度 k 和草稿模型的选择。自适应推测解码可以根据接受率动态调整参数。

下一步

下一课将介绍梯度检查点,优化内存使用。

📚 本文改编自 AI Engineering from Scratch(MIT License · 作者 Rohit Ghumare),中文内容来自官方中文镜像。原课程共 503 课 · 20 阶段 · 免费开源,教程网站见 aiengineeringfromscratch.com

  • 标题: 【大模型】- 推测解码
  • 作者:
  • 创建于 : 2026-08-19 09:23:00
  • 更新于 : 2026-08-21 16:20:12
  • 链接: https://sxl-space.tk/2026/08/19/010_LLM/010_LLM-23-SpeculativeDecoding/
  • 版权声明: 版权所有 © 宋,禁止转载。