大模型推理加速复盘:投机解码将首 Token 延迟从 1.8s 压至 300ms
一、首 Token 延迟的致命瓶颈:为什么用户觉得“卡”
在承接一个实时对话场景的 LLM 推理需求时,业务方给出了硬性指标:首 Token 生成延迟(TTFT)不超过 500ms。初版部署的 Llama-3-8B 模型在 A100 上测得 TTFT 为 1.8s,远高于要求线。这意味着用户发出消息后,需要等待近两秒才能看到模型开始回复,在实时对话场景中几乎不可接受。
传统的 KV Cache 预填充优化(Prefix Caching)将 TTFT 从 1.8s 压到了 1.2s,但距离 500ms 还有显著差距。需要引入一套更激进的加速方案——投机解码(Speculative Decoding)。其核心思想是:用一个轻量级的小模型(Draft Model)快速生成多个候选 token,再用大模型(Target Model)并行验证这些候选,从而将大模型的自回归瓶颈从逐 token 串行变为批量并行。
二、草稿模型的选择:不是越小越好
投机解码的加速效果取决于草稿模型的"接受率"——即草稿模型生成的 token 被目标模型验证认可的比例。接受率每提升 10%,整体加速比约提升 0.3~0.5 倍。
实验了三种草稿模型配置:
| 草稿模型 | 参数量 | 单 Token 延迟 | 接受率 | 等效加速比 |
|---|---|---|---|---|
| Llama-3-8B(同模型浅层) | ~1.2B | 18ms | 72% | 3.4x |
| TinyLlama-1.1B | 1.1B | 15ms | 58% | 2.9x |
| GPT-2-Small-124M | 124M | 3ms | 31% | 1.6x |
关键的意外发现是:使用目标模型的浅层网络(前 8 层)作为草稿模型的方案在 8B 规模上表现最优。原因在于同一模型的词表分布和隐藏表示完全对齐,浅层的输出分布与深层的输出分布在 semantics 上高度相关。而 TinyLlama 虽然参数接近,但训练数据和词表的差异导致接受率偏低。
# 投机解码核心逻辑 —— 草稿生成 + 批量验证 import torch import torch.nn.functional as F class SpeculativeDecoder: def __init__(self, target_model, draft_model, gamma: int = 4): """ gamma: 草稿模型每次生成的候选 token 数量 经验值:70B 模型 gamma=4, 8B 模型 gamma=5 """ self.target = target_model self.draft = draft_model self.gamma = gamma @torch.inference_mode() def generate(self, input_ids: torch.Tensor, max_new_tokens: int = 512): generated = [] current_ids = input_ids while len(generated) < max_new_tokens: # 阶段 1:草稿模型生成 gamma 个候选 token draft_ids = current_ids.clone() draft_tokens = [] for _ in range(self.gamma): logits = self.draft(draft_ids).logits[:, -1, :] # 从草稿模型的分布中采样,而非贪心选择 # 目的是在与目标模型分布不完全对齐时提高接受概率 next_token = torch.multinomial( F.softmax(logits / 0.6, dim=-1), num_samples=1 ) draft_tokens.append(next_token) draft_ids = torch.cat([draft_ids, next_token], dim=-1) # 阶段 2:目标模型并行验证整个候选序列 # 一次前向传播处理 gamma+1 个位置(包含原始 prompt 位置) target_logits = self.target(draft_ids).logits # [1, seq_len+gamma, vocab] # 逐 token 对比决定接受/拒绝 accepted = 0 for i in range(self.gamma): pos = current_ids.shape[1] + i # 当前位置在序列中的索引 p_target = F.softmax(target_logits[:, pos, :], dim=-1) # 概率接受规则: # P_accept = min(1, P_target(token) / P_draft(token)) draft_token_prob = p_target[0, draft_tokens[i].item()].item() draft_dist = F.softmax( self.draft(draft_ids[:, :pos+1]).logits[:, -1, :], dim=-1 ) draft_prob = draft_dist[0, draft_tokens[i].item()].item() accept_prob = min(1.0, draft_token_prob / (draft_prob + 1e-8)) if torch.rand(1).item() < accept_prob: accepted += 1 else: # 拒绝:从修正后的分布中重新采样 corrected_prob = F.relu(p_target - draft_dist) corrected_prob /= corrected_prob.sum() bonus_token = torch.multinomial(corrected_prob, 1) generated.append(bonus_token.item()) current_ids = torch.cat( [current_ids, torch.tensor([draft_tokens[:accepted] + [bonus_token]])], dim=-1 ) break else: # 全部接受:从目标模型分布中额外采样一个 token logits = target_logits[:, current_ids.shape[1]+self.gamma, :] bonus_token = torch.multinomial(F.softmax(logits, dim=-1), 1) for t in draft_tokens: generated.append(t.item()) generated.append(bonus_token.item()) current_ids = torch.cat( [current_ids, torch.stack(draft_tokens + [bonus_token], dim=1)], dim=-1 ) # 更新循环条件 if accepted < self.gamma: continue return generated三、投机解码的工程化权衡
投机解码的加速效果在某些边界场景下会退化:
- 低熵场景(如模板补全):草稿模型的输出与目标模型高度一致,接受率超 85%,加速比可达 4~5 倍;
- 高熵场景(如创意写作):草稿模型与目标模型的分歧增大,接受率降至 40%
50%,加速比缩至 1.52 倍; - Batch 场景:批量推理时,投机解码的并行验证优势与 Continuous Batching 叠加,加速比可进一步提升。
综合实测数据:
| 指标 | 无优化 | Prefix Cache only | + 投机解码 |
|---|---|---|---|
| TTFT(7B 模型) | 1.8s | 1.2s | 310ms |
| TTFT(70B 模型) | 12.5s | 8.2s | 1.9s |
| 每秒生成 Token 数 | 42 | 42 | 126 |
| GPU 利用率 | 38% | 42% | 78% |
四、与量化方案的协同效应
投机解码与模型量化的组合使用产生了 1+1>2 的效果。INT8 量化后的目标模型前向传播速度提升 1.8 倍,但草稿模型也同时受益。然而需要注意:
- 草稿模型不建议量化:草稿模型的精度对接受率有放大效应,INT8 量化后接受率下降了 12 个百分点(72%→62%),抵消了延迟降低的收益;
- 目标模型的 INT8 量化对验证阶段的精度影响在 0.5% 以内,对接受率的判断基本无影响。
五、总结
投机解码在推理加速中的核心结论:
- 草稿模型选择"同模型浅层"优于"异模型小模型":词表和隐藏表示的共享带来的接受率增益远超参数量的微小差异;
- gamma 参数需按模型规模调整:8B 级 gamma=5, 70B 级 gamma=4。gamma 越大理论上加速潜力越高,但接受率会边际递减;
- 投机解码对低熵任务加速效果最显著:代码生成、模板补全等确定性高的任务加速比可达 4
5 倍;创意写作等高熵任务加速效果收敛在 1.52 倍; - 与 Continuous Batching 天然兼容:投机解码产出的 token 批量与调度器的批量合并机制正交,部署时无需额外改造。
适用边界:投机解码对第一 token 生成延迟(TTFT)的改善有限——TTFT 受 Prompt 编码阶段的计算量主导,投机解码只能加速后续 token 的生成流速。改善 TTFT 仍需要 Prefix Caching 或 KV Cache 优化。