【大模型】- 推测解码:EAGLE3

算法

推测解码:EAGLE3

使用小模型加速大模型生成

类型: 学习 | 语言: Python | 🏷 前置:《开源模型架构》(本系列第 13 篇)

学习目标

  • 理解推测解码的原理
  • 实现 EAGLE3 架构
  • 训练草稿模型
  • 优化推测解码性能
  • 评估加速效果

推测解码概述

推测解码(Speculative Decoding)使用小模型(草稿模型)猜测多个 token,然后用大模型(目标模型)验证。如果猜测正确,可以跳过多次前向传播。

基本原理

1
2
传统生成:大模型 → token1 → 大模型 → token2 → 大模型 → token3
推测解码:小模型 → [token1, token2, token3] → 大模型验证 → 接受/拒绝

加速原理

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

EAGLE3 架构

EAGLE3(Extrapolation Algorithm for Greater Language-model Efficiency)是一种先进的推测解码方法。

架构设计

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
class EAGLE3(nn.Module):
def __init__(self, target_model, draft_model):
super().__init__()
self.target_model = target_model # 大模型
self.draft_model = draft_model # 小模型

# 特征提取器
self.feature_extractor = FeatureExtractor(target_model)

# 预测头
self.prediction_head = PredictionHead(
input_dim=target_model.config.hidden_size,
output_dim=target_model.config.vocab_size
)

def forward(self, input_ids, features):
"""前向传播"""
# 草稿模型预测
draft_logits = self.draft_model(input_ids)

# 特征融合
fused_features = self.fuse_features(features, draft_logits)

# 预测头
logits = self.prediction_head(fused_features)

return logits

def fuse_features(self, target_features, draft_logits):
"""融合目标模型和草稿模型的特征"""
# 线性融合
fused = torch.cat([target_features, draft_logits], dim=-1)
return fused

特征提取器

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
class FeatureExtractor(nn.Module):
def __init__(self, model):
super().__init__()
self.model = model

# 提取中间层特征
self.layer_indices = [i for i in range(len(model.layers))]

def forward(self, input_ids):
"""提取特征"""
features = []

# 逐层提取
hidden_states = self.model.embed_tokens(input_ids)

for layer_idx in self.layer_indices:
layer = self.model.layers[layer_idx]
hidden_states = layer(hidden_states)
features.append(hidden_states)

return features

预测头

1
2
3
4
5
6
7
8
9
10
11
class PredictionHead(nn.Module):
def __init__(self, input_dim, output_dim):
super().__init__()
self.linear = nn.Linear(input_dim, output_dim)
self.layer_norm = nn.LayerNorm(output_dim)

def forward(self, x):
"""预测下一个 token"""
logits = self.linear(x)
logits = self.layer_norm(logits)
return logits

训练草稿模型

训练数据

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
def prepare_training_data(target_model, tokenizer, dataset):
"""准备训练数据"""
training_data = []

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

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

# 提取特征
features = extract_features(target_model, inputs["input_ids"])

# 草稿模型输入
draft_input = inputs["input_ids"][:, :-1]

# 标签
labels = inputs["input_ids"][:, 1:]

training_data.append({
"input_ids": draft_input,
"features": features,
"labels": labels,
"target_logits": target_logits[:, :-1]
})

return training_data

训练循环

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
def train_eagle3(eagle3_model, training_data, config):
"""训练 EAGLE3"""
optimizer = torch.optim.AdamW(
eagle3_model.parameters(),
lr=config["lr"]
)

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

for batch in training_data:
# 前向传播
logits = eagle3_model(
batch["input_ids"],
batch["features"]
)

# 计算损失
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
batch["labels"].view(-1)
)

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

total_loss += loss.item()

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

return eagle3_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
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 SpeculativeDecoder:
def __init__(self, target_model, eagle3_model, tokenizer, k=5):
self.target_model = target_model
self.eagle3_model = eagle3_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:
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

return generated, acceptance_rate

def draft_generate(self, input_ids):
"""草稿模型生成"""
# 提取特征
features = extract_features(self.target_model, input_ids)

tokens = []
probs = []

for _ in range(self.k):
with torch.no_grad():
logits = self.eagle3_model(input_ids, features)
next_probs = F.softmax(logits[:, -1, :], 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 = []
last_accepted_pos = -1

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))
last_accepted_pos = i
else:
# 拒绝:从目标模型重新采样
# 修正概率分布
修正_probs = torch.clamp(target_probs - draft_probs, min=0)
修正_probs = 修正_probs / (修正_probs.sum() + 1e-8)

new_token = torch.multinomial(修正_probs, 1)
return torch.cat(accepted_tokens, dim=1) if accepted_tokens else new_token, 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
class BatchSpeculativeDecoder:
def __init__(self, target_model, eagle3_model, tokenizer, k=5, batch_size=32):
self.target_model = target_model
self.eagle3_model = eagle3_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()

# KV 缓存
kv_cache = self.init_kv_cache(batch_size)

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

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

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

return generated

def batch_draft_generate(self, input_ids, kv_cache):
"""批量草稿生成"""
batch_size = input_ids.shape[0]

all_tokens = []
all_probs = []

for _ in range(self.k):
with torch.no_grad():
logits = self.eagle3_model(input_ids)
probs = F.softmax(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)

自适应猜测长度

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
class AdaptiveSpeculativeDecoder:
def __init__(self, target_model, eagle3_model, tokenizer):
self.target_model = target_model
self.eagle3_model = eagle3_model
self.tokenizer = tokenizer

# 自适应参数
self.min_k = 1
self.max_k = 10
self.current_k = 5
self.acceptance_history = []

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

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)

评估加速效果

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, eagle3_model, tokenizer, test_prompts):
"""评估加速效果"""
# 传统生成
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, eagle3_model, tokenizer)

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

总结

EAGLE3 通过融合目标模型和草稿模型的特征,实现了高效的推测解码。关键组件包括特征提取器、预测头和自适应猜测长度。推测解码可以显著加速 LLM 推理,特别是在长序列生成中。

下一步

下一课将介绍差分注意力,改进注意力机制。

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

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