推理加速算法
大模型推理的延迟主要来自两个方面: Prefill 阶段(处理输入 prompt )和 Decode 阶段(逐个生成token)。由于Decode阶段是内存带宽受限的,每生成一个token都需要加载完整的模型参数和KV Cache。本节介绍通过投机执行来突破这一瓶颈的系列算法。
7.2.1 Speculative Decoding (投机解码)
问题背景
在自回归生成中,每个token的生成都需要: 1. 加载整个模型到GPU 2. 执行一次完整的前向传播 3. 采样得到下一个token
这个过程是顺序的,无法并行化。然而,语言模型生成的 token往往具有较强的可预测性,特别是常见短语和模板化内容。
核心思想
投机解码采用小模型起草,大模型验证的策略:
- 起草阶段(Draft):使用轻量级小模型(draft model)快速生成 \(\gamma\) 个候选token
- 验证阶段(Verify):大模型(target model)并行验证这些候选token
- 接受/拒绝:根据概率分布决定是否接受候选token
数学原理
设:
- \(M_p\):大模型( target model )的概率分布 \(p(x)\)
- \(M_q\):小模型( draftmodel)的概率分布 \(q(x)\)
对于候选token\(x\),接受概率为:
\(P(\text{accept}) = \min\left(1, \frac{p(x)}{q(x)}\right)\)
。修正采样
如果候选token\(x\) 被拒绝,从修正分布中采样:
\(p'(x) = \frac{p(x) - q(x)}{1 - \sum_{x'} \min(p(x'), q(x'))}\)
或者使用更简单的形式:
\(p'(x) = \text{normalize}(\max(0, p(x) - q(x)))\)
算法伪代码
def speculative_decoding(prompt, M_p, M_q, gamma, max_tokens):
"""
投机解码算法
M_p: 大模型(target model)
M_q: 小模型(draft model)
gamma: 每次起草的token数
"""
tokens = tokenize(prompt)
while len(tokens) < max_tokens:
# === 起草阶段 ===
draft_tokens = []
draft_probs = []
for _ in range(gamma):
q_dist = M_q(tokens + draft_tokens)
x = sample(q_dist)
draft_tokens.append(x)
draft_probs.append(q_dist)
# ===验证阶段 ===
# 大模型并行计算所有候选token的概率
all_tokens = tokens + draft_tokens
p_dists = M_p(all_tokens, return_all_probs=True)
# === 接受/拒绝 ===
accepted = 0
for i, (x, q_dist) in enumerate(zip(draft_tokens, draft_probs)):
p_dist = p_dists[len(tokens) + i]
p_x = p_dist[x]
q_x = q_dist[x]
# 计算接受概率
accept_prob = min(1.0, p_x / q_x)
if random() < accept_prob:
# 接受token
tokens.append(x)
accepted += 1
else:
# 拒绝token,从修正分布采样
adjusted_dist = normalize(max(0, p_dist - q_dist))
tokens.append(sample(adjusted_dist))
break
# 如果全部接受,从p采样一个额外token
if accepted == gamma:
tokens.append(sample(p_dists[-1]))
return detokenize(tokens)加速比分析
理论加速比取决于: 1. 接受率 \(\alpha\):候选 token 被接受的比例 2. 起草成本比 \(\beta\):小模型与大模型的推理时间比 3. 起草数量 \(\gamma\)
理想加速比(假设 \(\beta \ll 1\)):
\(\text{Speedup} \approx \frac{1}{1 - \alpha^{\gamma+1}} \cdot \frac{\gamma + 1}{\gamma \beta + 1}\)
典型场景( \(\alpha = 0.8, \gamma = 4, \beta = 0.1\) ): \(\text{Speedup} \approx 2.5-3.0x\)
实际应用与限制
优势:
- 无损加速:保持大模型的输出分布
- 通用性强:适用于任何自回归模型
- 实现相对简单
限制:
- 小模型选择:需要与大模型行为相似的小模型
- 内存开销:需要同时加载两个模型
- 接受率依赖:在低可预测性内容上效果差
- 通信开销:如果大小模型在不同设备上
典型配置:
- LLaMA-70B + LLaMA-7B作为draft model
- \(\gamma = 4-8\)
- 实测加速比:1.5-2.5x
7.2.2 Medusa
问题背景
投机解码需要维护两个独立的模型,带来额外的内存和工程复杂度。Medusa提出了一种无需辅助模型的投机解码方法,通过在原始模型上添加轻量级头来实现。
核心思想
Medusa的核心创新: 1. 多头解码:在原始模型的顶层添加多个解码头 2. 每个头预测未来不同位置的token:头1预测下一个token,头2预测下下个token,依此类推 3. 树状注意力:高效组织多个候选序列的验证
多头结构
每个Medusa Head是一个轻量级MLP:
\(\text{Medusa}_i(h) = \text{Softmax}(W_i \cdot \text{MLP}_i(h))\)
训练方法
训练数据
使用原始模型生成的数据作为监督信号: 1. 从训练集中采样prompt 2. 用原始模型生成完整序列 3. 将生成的token作为每个头的目标
损失函数
\(\mathcal{L} = \mathcal{L}_{\text{base}} + \lambda \sum_{i=1}^{k}\mathcal{L}_{\text{Medusa}_i}\)
其中 \(\mathcal{L}_{\text{base}}\) 是原始语言模型头的交叉熵损失, \(\mathcal{L}_{\text{Medusa}_i}\) 是第 \(i\) 个Medusa头的损失。
训练策略
- 冻结主干:保持原始模型参数不变
- 只训练Medusa Heads:通常只有几百万参数
- 典型训练时间:1-2个epoch,数小时到一天
树状注意力验证
Medusa提出树状注意力来高效验证多个候选序列
候选序列树:
t+1
/ | \
t+2a t+2b t+2c
/ | |
t+3a ... ...通过树状注意力,可以在一次前向传播中验证所有候选路径。Medusa 伪代码
class MedusaModel(nn.Module):
def __init__(self, base_model, num_heads, num_candidates):
super().__init__()
self.base_model = base_model
self.num_heads = num_heads
self.num_candidates = num_candidates
# Medusa Heads: 每个头预测未来第i个token
self.medusa_heads = nn.ModuleList([
MedusaHead(base_model.config.hidden_size, base_model.vocab_size) # 补全右括号
for _ in range(num_heads)
])
def forward(self, input_ids):
# 基础模型前向
outputs = self.base_model(input_ids, output_hidden_states=True)
hidden = outputs.hidden_states[-1]
# 基础预测
base_logits = outputs.logits[:, -1, :]
# Medusa 预测
medusa_logits = []
for head in self.medusa_heads:
medusa_logits.append(head(hidden[:, -1, :]))
return base_logits, medusa_logits
def generate(self, prompt, max_new_tokens):
tokens = tokenize(prompt)
for _ in range(max_new_tokens // (self.num_heads + 1)):
# 获取所有预测
base_logits, medusa_logits = self.forward(tokens)
# 构建候选序列树
candidates = build_tree(base_logits, medusa_logits, self.num_candidates)
# 树状注意力验证
verified_tokens = self.verify_tree(tokens, candidates)
tokens.extend(verified_tokens)
return detokenize(tokens)
def verify_tree(self, prefix, candidates):
"""使用树状注意力验证候选序列"""
# 构建树状注意力掩码
tree_mask = build_tree_mask(candidates)
# 一次前向验证所有候选
all_sequences = [prefix + cand for cand in candidates]
logits = self.base_model(all_sequences, attention_mask=tree_mask).logits
# 选择接受的token序列
accepted = []
for i, cand in enumerate(candidates):
if verify_acceptance(logits[i], cand):
accepted.extend(cand)
else:
# 从修正分布采样
accepted.append(sample_adjusted(logits[i]))
break
return accepted与投机解码的对比
| 特性 | Speculative Decoding | Medusa |
|---|---|---|
| 辅助模型 | 需要独立的小模型 | 无需辅助模型 |
| 内存开销 | 2x模型大小 | 1.01-1.05x模型大小 |
| 训练需求 | 无需训练 | 需要训练Medusa Heads |
| 加速比 | 1.5-2.5x | 1.8-2.8x |
| 实现复杂度 | 中等 | 较高 |
性能分析
在Vicuna-7B上的实测结果:
- 平均加速比:2.0-2.5x
- Medusa Heads数量:4-8个
- 训练时间:约4-8小时
- 内存增加:<5%
7.2.3 EAGLE
问题背景
Medusa在特征空间进行预测,但直接预测未来 token的分布可能面临信息不足的问题。EAGLE ( Extrapolation Algorithm for Greater Language-model Efficiency )提出在特征级别进行投机,利用更丰富的上下文信息。
核心思想
EAGLE的核心创新: 1. 特征级投机:预测未来token的特征(而非token本身) 2. 自回归头(Auto-regressive Head):使用轻量级自回归模型预测特征序列 3. 特征到 token映射:通过LM头将预测特征转换为token分布
LM Head(Language Model Head)是位于主干网络之后的一个线性层(Linear Layer),它的核心作用是:将模型的隐藏状态(Hidden State)映射到词汇表大小的概率分布上。
输入 Token → Embedding → Transformer Layers → [LM Head] → Logits → Softmax → 概率分布
↑
Linear(768, vocab_size)特征级预测的优势
相比token级预测,特征级预测:
- 包含更丰富的语义信息
- 更容易捕获长距离依赖
- 预测任务更简单(连续空间vs离散空间)
自回归头架构
EAGLE的自回归头是一个小型Transformer:
输入: [h_t, h_t-1, ..., h_t-k] (历史特征序列)
|
v
+-----------------------+
| 小型Transformer | (2-4层)
+-----------------------+
|
v
输出: [ĥ_t+1, ĥ_t+2, ..., ĥ_t+m] (预测的未来特征)自回归头的规模通常只有原模型的1-10%。
算法流程
阶段1: 特征投机
- 使用自回归头预测未来m个特征
Ĥ = AR_Head([h_t, h_t-1, ...])
阶段2: Token生成
- 通过LM头将特征转换为token分布
- 对每个
ĥ_t+i: p_t+i = Softmax(W_lm · ĥ_t+i)
阶段3: 大模型验证
- 大模型并行验证候选token
- 接受/拒绝机制同投机解码
EAGLE 伪代码
class EAGLEModel(nn.Module):
def __init__(self, base_model, ar_layers=2, ar_hidden=None):
super().__init__()
self.base_model = base_model
hidden_size = base_model.config.hidden_size
# 自回归头:小型Transformer
self.ar_head = nn.TransformerDecoder(
nn.TransformerDecoderLayer(
d_model=hidden_size,
nhead=hidden_size // 64,
dim_feedforward=ar_hidden or hidden_size * 2,
batch_first=True
),
num_layers=ar_layers
)
# 特征投影
self.feature_proj = nn.Linear(hidden_size, hidden_size)
def forward(self, input_ids):
# 获取基础模型特征
outputs = self.base_model(input_ids, output_hidden_states=True)
features = outputs.hidden_states[-1] # [batch, seq, hidden]
return features
def speculate_features(self, features, num_steps):
""" 使用自回归头投机未来特征"""
speculated = []
current = features
for _ in range(num_steps):
# 自回归预测下一个特征
next_feature = self.ar_head(current, current)
next_feature = self.feature_proj(next_feature[:, -1:, :])
speculated.append(next_feature)
current = torch.cat([current, next_feature], dim=1)
return torch.cat(speculated, dim=1)
def features_to_tokens(self, features):
""" 将特征转换为token分布"""
logits = self.base_model.lm_head(features)
return F.softmax(logits, dim=-1)
def generate(self, prompt, max_tokens, gamma=5):
tokens = tokenize(prompt)
while len(tokens) < max_tokens:
# 获取当前特征
features = self.forward(tokens)
# 投机未来特征
spec_features = self.speculate_features(features[:, -4:, :], ga
# 转换为token分布
draft_dists = self.features_to_tokens(spec_features)
draft_tokens = [sample(dist) for dist in draft_dists[0]]
# 大模型验证
verified = self.verify(tokens, draft_tokens, draft_dists[0])
tokens.extend(verified)
return detokenize(tokens)
def verify(self, prefix, draft_tokens, draft_dists):
"""验证候选token"""
all_tokens = prefix + draft_tokens
p_dists = self.base_model(all_tokens, return_all_probs=True)
accepted = []
for i, (x, q_dist) in enumerate(zip(draft_tokens, draft_dists)):
p_dist = p_dists[len(prefix) + i]
p_x = p_dist[x]
q_x = q_dist[x]
if random() < min(1.0, p_x / q_x):
accepted.append(x)
else:
adjusted = normalize(max(0, p_dist - q_dist))
accepted.append(sample(adjusted))
break
return accepted训练方法
EAGLE的训练分为两个阶段
- 特征收集阶段:
- 用基础模型处理训练数据
- 保存每层的隐藏状态
- 自回归头训练:
- 输入:历史特征序列 \([h_{t-k}, ..., h_t]\)
- 目标:未来特征序列 \([h_{t+1}, ..., h_{t+m}]\)
- 损失:MSE损失 \(\mathcal{L} = \sum_{i=1}^{m} ||\hat{h}_{t+i} - h_{t+i}||^2\).
性能分析
EAGLE相比 Medusa的优势:
- 更高的接受率(特征级预测更准确)
- 更好的长距离依赖建模 - 在复杂任务上表现更稳定
实测加速比:
- LLaMA-2-7B: 2.5-3.0x
- LLaMA-2-70B: 2.0-2.5x
7.2.4 Lookahead Decoding
问题背景
投机解码、Medusa和EAGLE都需要额外的模型或训练。Lookahead Decoding提出了一种无需任何额外模型或训练的并行解码方法。
核心思想
Lookahead Decoding基于 Jacobi迭代的思想: 1. 打破顺序依赖:将自回归生成视为求解方程组 2. 并行猜测:基于n-gram匹配生成多个候选token 3. 并行验证:一次前向传播验证所有猜测
- n-gram:长度为 n 的连续子序列(子串或子词)。
- n-gram 匹配:在两个序列中寻找共同的 n-gram,并以此衡量它们的相似度或重叠程度。
Jacobi迭代视角
自回归生成可以表示为:
\(x_{t+1} = f(x_1, x_2, ..., x_t)\)
Jacobi 迭代同时更新所有位置:
\(x_i^{(k+1)} = f(x_1^{(k)}, x_2^{(k)}, ..., x_{i-1}^{(k)})\)
对于语言模型,这意味着可以并行猜测多个未来token。
\(x_1\), \(x_2\), …, \(x_t\):已经生成的 token id 序列(或更严格地说,是这些 token 对应的离散变量 / 隐表示)\(x_i\)就是“第 \(i\) 个位置上的 token”。
这里的 \(k\)是猜测/修正轮次(iteration index),Jacobi 做法是“用上一轮的全部值,同时算下一轮的全部值”。
n-gram匹配策略
Lookahead Decoding维护一个n-gram窗口:
已生成序列: [The, cat, sat, on, the, ...]
|___| (2-gram: "cat sat")
在窗口中查找匹配的n-gram:
位置3: "cat sat" → 下一个token是 "on"
位置10: "cat sat" → 下一个token是 "under"
候选token: ["on", "under", ...]Lookahead Decoding 伪代码
class LookaheadDecoding:
def __init__(self, model, window_size=5, ngram_size=3, num_candidates=1
self.model = model
self.window_size = window_size
self.ngram_size = ngram_size
self.num_candidates = num_candidates
self.ngram_cache = {} # n-gram -> 后续token列表
def update_cache(self, tokens):
""" 更新n-gram缓存"""
for i in range(len(tokens) - self.ngram_size):
ngram = tuple(tokens[i:i + self.ngram_size])
next_token = tokens[i + self.ngram_size]
if ngram not in self.ngram_cache:
self.ngram_cache[ngram] = []
if next_token not in self.ngram_cache[ngram]:
self.ngram_cache[ngram].append(next_token)
def generate_candidates(self, tokens):
""" 基于n-gram匹配生成候选token"""
candidates = []
# 使用最近的n-gram查找
for n in range(self.ngram_size, 1, -1):
if len(tokens) < n:
continue
ngram = tuple(tokens[-n:])
if ngram in self.ngram_cache:
candidates.extend(self.ngram_cache[ngram])
# 去重并限制数量
candidates = list(dict.fromkeys(candidates))[:self.num_candidates]
return candidates
def verify_candidates(self, tokens, candidates):
""" 并行验证候选token"""
if not candidates:
return [self.model.generate_next(tokens)]
# 构建验证批次
sequences = [tokens + [cand] for cand in candidates]
# 并行前向
logits = self.model.batch_forward(sequences)
# 检查哪些猜测是正确的
accepted = []
for i, cand in enumerate(candidates):
# 如果模型预测的下一个token与猜测一致
predicted = logits[i].argmax(dim=-1)
if predicted == cand:
accepted.append(cand)
else:
# 添加正确的token并停止
accepted.append(predicted.item())
break
return accepted
def generate(self, prompt, max_tokens):
tokens = tokenize(prompt)
while len(tokens) < max_tokens:
# 生成候选token
candidates = self.generate_candidates(tokens)
# 验证候选
verified = self.verify_candidates(tokens, candidates)
# 更新序列和缓存
tokens.extend(verified)
self.update_cache(tokens)
return detokenize(tokens)复杂度分析
| 操作 | 复杂度 | 说明 |
|---|---|---|
| n-gram查找 | \(O(1)\) | 哈希表实现 |
| 候选生成 | \(O(k)\) | k为候选数 |
| 并行验证 | \(O(1)\) | 单次batch前向 |
| 缓存更新 | \(O(w)\) | w为窗口大小 |
实际应用与限制
优势:
- 无需额外模型或训练
- 实现简单
- 适用于任何自回归模型
限制: 1. n-gram稀疏性:长n-gram匹配率低 2. 内容依赖:在模板化内容上效果好,在创造性内容上效果差 3. 缓存开销:需要维护n-gram缓存
典型加速比:
- 代码生成:1.5-2.0x
- 对话:1.2-1.5x
- 创意写作:<1.2x