大模型推理的速度瓶颈集中在自回归生成阶段——每生成一个token都要跑一遍完整前向传播。投机采样(Speculative Decoding)用一个小模型快速草拟多个候选token,再由大模型一次性验证,在不损失精度的前提下把推理速度提升2到4倍。本文拆解投机采样的数学原理、Medusa多头方案的实现方式,以及在vLLM中的部署配置。
投机采样的基本原理:草稿模型加验证模型
自回归生成的串行性是推理慢的根本原因。投机采样的思路是:用一个参数量小得多的draft model快速猜测接下来k个token,然后把这k个token和原prompt拼接,交给target model一次前向传播完成验证和修正。如果draft猜对了,等于用一次前向拿到了k个token;如果猜错了,从错误位置截断,保留正确前缀,后续token由target model自己生成。
关键在于接受概率的计算。设draft model分布为q(x),target model分布为p(x),接受概率r = min(1, p(x)/q(x))。当p(x)大于q(x)时必定接受;当p(x)小于q(x)时按概率p(x)/q(x)接受,拒绝时从修正分布norm(max(p(x)-q(x),0))重新采样。这个采样过程保证最终输出与target model独力生成的分布完全一致,不会引入精度损失。
Medusa头:不引入额外模型的投机采样
传统投机采样需要两个模型,draft model选型与部署增加工程复杂度。Medusa方案在target model的隐藏层上直接添加多个解码头,每个头独立预测不同位置的token,省去独立draft model。
import torch
import torch.nn as nn
class MedusaHead(nn.Module):
def __init__(self, hidden_size, vocab_size, num_heads=4):
super().__init__()
self.heads = nn.ModuleList([
nn.Sequential(
nn.Linear(hidden_size, hidden_size),
nn.SiLU(),
nn.Linear(hidden_size, vocab_size)
) for _ in range(num_heads)
])
def forward(self, hidden_states):
# 每个头预测当前位置之后第i+1个token
return [head(hidden_states) for head in self.heads]
Medusa头在训练时冻结base model参数,只训练头的权重。推理时base model一次前向传播,4个头同时预测接下来4个token,通过tree attention验证候选序列,接受率高的token链直接保留。相比双模型方案,Medusa只需维护一个模型实例,显存开销仅增加几个轻量线性层。
vLLM中的投机采样部署配置
vLLM从0.4.x版本开始支持speculative decoding,内置了draft model方案和Medusa方案两种路径。配置方式:
# 方式一:使用独立draft model
python -m vllm.entrypoints.openai.api_server \
--model meta-llama/Llama-3-70B-Instruct \
--speculative-model meta-llama/Llama-3-8B-Instruct \
--num-speculative-tokens 5 \
--gpu-memory-utilization 0.9
# 方式二:使用Medusa头(需先训练Medusa权重)
python -m vllm.entrypoints.openai.api_server \
--model meta-llama/Llama-3-70B-Instruct \
--speculative-model medusa-llama-3-70b \
--num-speculative-tokens 4
num-speculative-tokens控制每次草拟的token数量,通常设为4到8。值越大,猜中时吞吐提升越明显,但猜错时的计算浪费也越多,需要根据draft model与target model的相似度调参。
加速效果与适用场景
投机采样的加速比取决于draft model与target model的分布吻合度。实测数据参考:
- Llama-3-8B作为Llama-3-70B的draft model,accept rate约70%,吞吐提升约2.3倍;
- Medusa-4头方案在Llama-3-70B上,accept rate约55%,吞吐提升约1.8倍;
- 同系列模型做draft时效果最好,跨架构(如用GPT-2做Llama的draft)accept rate骤降。
适用场景:单请求延迟敏感的交互式服务(如ChatBot),投机采样能显著降低TTFT后的生成延迟。对于高并发批量推理场景,投机采样的收益较小,因为批处理本身已经充分利用GPU并行度。选择方案时先测accept rate,低于40%时加速效果不如直接增大batch size。
原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/da-mo-xing-tou-ji-cai-yang-tui-li-jia-su-shi-zhan/