- AI 技能
- 人工智能
- 大模型
- 深度学习
【免费下载链接】AI-Research-SKILLs
Comprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.
本指南围绕开源仓库 AI-Research-SKILLs 中 19-emerging-techniques/speculative-decoding 技能包的 Medusa 参考资料展开,系统讲解 Medusa(arXiv 2401.10774,2024)这一"多解码头"LLM 推理加速框架的架构原理、Medusa-1/Medusa-2 两种训练方法、树形注意力验证算法、超参数调优与生产部署全流程。读完本文,你将掌握如何在不引入独立草稿模型的前提下,为现有 LLM 附加预测头并在单次前向中验证整棵候选树,获得 2.2–3.6× 的推理加速,并能在 transformers 生态与 vLLM 等部署场景中正确选型与落地。
Medusa 解决的问题:绕开"双模型"推测解码
传统推测解码(Speculative Decoding)用一个小型草稿模型(draft model)快速生成 K 个候选 token,再由大目标模型并行验证,通常能获得 1.5–2× 加速。但它的部署代价是必须同时维护两套模型:草稿模型 + 目标模型,二者占据双份显存,且草稿模型与目标模型的分布差异会直接影响接受率。
Medusa 的核心创新是把"草稿能力"内化进模型本身:不再外挂草稿模型,而是在现有 LLM 的隐藏状态之上并联添加多个解码头(Medusa heads),每个头专门预测未来第 t+1、t+2、t+3、t+4 个位置的 token。这一思路来自论文MEDUSA: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads(arXiv 2401.10774,2024),仓库中的 medusa.md 将其定位为"简单"(Simple):不动主干、不改注意力、只加轻量头,即可获得 2.2–3.6× 加速且不损失质量。
核心架构:多解码头 + 树形注意力
多解码头设计
在 medusa.md 架构小节 中,Medusa 把原有 LLM 视为"骨干 + 原始输出头"的组合:
Input → Base LLM (frozen or fine-tuned) → Hidden State ├→ Head 0 (original, predicts t+1) ├→ Head 1 (predicts t+2) ├→ Head 2 (predicts t+3) └→ Head 3 (predicts t+4)- Head 0 是模型原有的输出头,负责预测下一个 token(t+1);
- Head 1~Head N 是新增的轻量预测头(通常为单层线性层或小型 MLP),分别负责预测 t+2、t+3、t+4 等更远位置的 token;
- 所有头共享骨干网络输出的同一份隐藏状态,因此推理时骨干只需前向一次,多头的额外计算量极小(论文与仓库资料称整体显存开销约为基础模型的 1–2%)。
从该技能包 SKILL.md 的 Medusa 章节 可以看到同样的结构描述,且强调"无需独立草稿模型、仅需极少量训练(只训头)、可与任意 LLM 兼容"三大优点。
树形注意力:一次前向验证所有候选路径
多头分别给出不同位置的预测后,需要把它们组织起来一次性验证。Medusa 的做法是构造候选树(candidate tree),再用树形注意力掩码在单次前向传播中并行打分整棵树。
以 2 个头、每个头取 top-2 候选为例(见 medusa.md 的树结构示例):
Root (current token) / \ Candidate 1a Candidate 1b (Head 1: 2 options) / \ / \ C2a C2b C2c C2d (Head 2: 4 total paths)- 第一层由 Head 1 产生 2 个候选(1a、1b);
- 第二层由 Head 2 在每个候选下再分支,共 4 条完整路径;
- 树形注意力掩码让所有路径共享前缀 token 的 KV 缓存,只有分叉处需要额外计算;
- 最终一条前向即可对 4 个候选路径并行打分,从中挑选可接受的最长路径,一次生成多个 token。
这与传统逐 token 自回归(每步只生成 1 个 token)形成本质区别,也是加速的来源。
训练方法:Medusa-1 与 Medusa-2
Medusa 提供两种训练模式,取舍核心在于"骨干是否参与训练"(见 medusa.md 训练章节)。
Medusa-1:冻结骨干,只训头
思路:保持基础 LLM 完全冻结(torch.no_grad()前向取隐藏状态),仅训练新增的 Medusa 头。
优点:
- 无损:基础模型参数不变,原始能力零损失;
- 训练快:约数小时(8 张 GPU 规模);
- 数据需求小:约 1000 万 token 即可。
性能:2.2× 加速。
# Training loop for Medusa-1 for batch in dataloader: # Frozen base model with torch.no_grad(): hidden_states = base_model(**batch, output_hidden_states=True).hidden_states[-1] # Train Medusa heads for i, head in enumerate(medusa_heads): logits = head(hidden_states) # Target: tokens shifted by (i+1) positions targets = batch['input_ids'][:, i+1:] loss += F.cross_entropy(logits[:, :-i-1], targets) loss.backward() optimizer.step()训练数据的构造逻辑很直观:第 i 个头以当前位置隐藏状态预测未来第 i+1 个 token,因此目标序列就是原始input_ids向右偏移 (i+1) 个位置;配合logits[:, :-i-1]截断末尾,保证对齐。训练数据可以是任意文本语料(Wikipedia、C4 等)。
Medusa-2:骨干与头联合微调
思路:解冻基础模型,把骨干和 Medusa 头一起微调,让各头与骨干分布对齐。
优点:
- 预测准确率更高(头与骨干联合优化);
- 加速上限更高(2.3–3.6×)。
挑战:联合微调可能破坏基础模型原有能力。论文给出的对策是特殊的训练配方(见 medusa.md L85-L90):
- 从预训练基础模型出发;
- 添加 Medusa 头;
- 骨干与头联合微调,采用精细的学习率调度;
- 使用高质量数据避免能力退化。
# Medusa-2 training # All parameters trainable for param in base_model.parameters(): param.requires_grad = True # Unfreeze base for param in medusa_heads.parameters(): param.requires_grad = True # Different learning rates optimizer = torch.optim.AdamW([ {'params': base_model.parameters(), 'lr': 1e-5}, # Lower for base {'params': medusa_heads.parameters(), 'lr': 1e-3}, # Higher for heads ])关键细节在于分组学习率:骨干用较小的1e-5(保护已有能力),头用较大的1e-3(快速收敛)。SKILL.md 的进阶模式 还展示了与之互补的代码形态——先以nn.Linear(hidden_size, vocab_size, bias=False)逐个构造头,再冻结骨干参数只优化头,同样可用于 Medusa-1 场景。
推理算法:生成、验证、接受三步走
候选生成
生成阶段由骨干给出基础 token,各 Medusa 头各自取 top-k 预测,再以笛卡尔积组合成候选序列(medusa.md L113-L134):
def medusa_generate_candidates(base_logits, medusa_head_logits, top_k=10): """Generate candidate sequences using tree structure.""" candidates = [] # Base token (original LLM output) base_token = torch.argmax(base_logits, dim=-1) # For each Medusa head, get top-k predictions medusa_candidates = [] for head_logits in medusa_head_logits: top_k_tokens = torch.topk(head_logits, k=top_k, dim=-1).indices medusa_candidates.append(top_k_tokens) # Build candidate tree (all combinations) # With 4 heads, top-2 each: 2^4 = 16 candidates for combo in itertools.product(*medusa_candidates): candidate = [base_token] + list(combo) candidates.append(candidate) return candidates # Shape: (num_candidates, seq_len)注意候选数量随头数与 top-k 指数增长:4 个头各取 top-2 即 2^4 = 16 条路径。
树形验证
候选树通过特殊构造的注意力掩码打包进同一 batch,一次前向完成打分与挑选(medusa.md L139-L160):
def medusa_verify_candidates(model, candidates, past_key_values): """Verify all candidates in single forward pass using tree attention.""" # Construct tree attention mask # All candidates share prefix, diverge at different points attention_mask = build_tree_attention_mask(candidates) # Single forward pass for all candidates outputs = model( input_ids=candidates, attention_mask=attention_mask, past_key_values=past_key_values, use_cache=True ) # Score each candidate scores = compute_acceptance_scores(outputs.logits, candidates) # Accept longest valid candidate best_candidate = select_best(candidates, scores) return best_candidate接受准则:后验概率阈值
Medusa 采用后验阈值(posterior threshold)判定是否接受某个候选 token:当该 token 的概率超过阈值即接受(medusa.md L167-L174):
def should_accept(token, token_prob, threshold=0.09): """Medusa acceptance criterion.""" return token_prob >= threshold # Typical thresholds: # - 0.09: Standard (from paper) # - 0.05: Conservative (fewer rejections, slower) # - 0.15: Aggressive (more rejections, faster when works)阈值语义:0.05更保守(拒绝少、更慢但更稳),0.15更激进(接受门槛高,命中时更快,但拒绝多、可能影响质量)。
性能结果与关键结论
论文在Vicuna-7B + MT-Bench上报告的加速与质量数据(medusa.md 性能表):
| Configuration | Speedup | Quality (MT-Bench score) |
|---|---|---|
| Baseline | 1.0× | 6.57 |
| Medusa-1 (frozen) | 2.2× | 6.57 (lossless) |
| Medusa-2 (joint) | 2.3× | 6.60 (+0.03) |
| Medusa-2 (optimized) | 3.6× | 6.55 (-0.02) |
关键结论:
- Medusa-1 因骨干冻结,质量完全无损(6.57 持平);
- Medusa-2 联合微调反而可能带来轻微质量提升(6.60,+0.03);
- 极端优化配置(3.6×)以微幅质量波动(-0.02)为代价,印证"越激进越快、但可能轻微影响质量"的权衡。
超参数调优指南
解码头数量(Number of Heads)
# Typical configurations: num_heads = 2 # Conservative (2× speedup) num_heads = 3 # Balanced (2.5× speedup) num_heads = 4 # Standard (3× speedup, from paper) num_heads = 5 # Aggressive (3.5×+ speedup) # Rule: More heads = more candidates but also more computation # Optimal: 3-4 heads for most models头越多 → 候选路径指数增多 → 潜在加速更大,但验证开销也同步上升;对多数模型,3–4 个头是经验最优区间。
每头 Top-K
# Candidates per head top_k = 2 # Standard (2^num_heads total candidates) top_k = 3 # More candidates (3^num_heads) top_k = 5 # Many candidates (5^num_heads) # Example with 4 heads: # top_k=2: 16 candidates (fast) # top_k=3: 81 candidates (slower verification)树结构:medusa_choices
medusa_choices显式指定要探索的候选路径,比全笛卡尔积更可控(medusa.md L226-L240):
# Standard configuration (from paper) medusa_choices = [ [0], # Only head 0 [0, 0], # Head 0, then head 1 (first candidate) [0, 1], # Head 0, then head 1 (second candidate) [0, 0, 0], # All heads (first path) ] # Aggressive configuration (more paths) medusa_choices = [ [0], [0, 0], [0, 1], [0, 0, 0], [0, 0, 1], [0, 1, 0], [0, 1, 1], ]列表中的每个子列表代表一条路径的"步进选择序列",例如[0, 1]表示先走 Head 0 的第一个候选、再走 Head 1 的第二个候选。标准配置收敛于较短路径,激进配置覆盖更多分支组合。SKILL.md 超参数小节 亦给出同样的[[0], [0, 0], [0, 1], [0, 0, 0]]作为深度 3 的典型设置。
端到端训练配方
数据需求
| 模式 | 数据量 | 数据质量 | 训练时间(8× A100) |
|---|---|---|---|
| Medusa-1 | 1000 万–1 亿 token | 任意文本语料即可 | 2–8 小时 |
| Medusa-2 | 1 亿–10 亿 token | 高质量、与目标场景同域 | 1–3 天 |
训练脚本
仓库 medusa.md 训练脚本 给出可直接套用的命令行:
# Clone Medusa repo git clone https://github.com/FasterDecoding/Medusa cd Medusa # Train Medusa-1 (frozen base) python medusa/train/train.py \ --model_name_or_path lmsys/vicuna-7b-v1.3 \ --data_path ShareGPT_Vicuna_unfiltered/ShareGPT_V4.3_unfiltered_cleaned_split.json \ --bf16 True \ --output_dir medusa-vicuna-7b-v1.3 \ --num_train_epochs 3 \ --per_device_train_batch_size 4 \ --gradient_accumulation_steps 8 \ --learning_rate 1e-3 \ --medusa_num_heads 4 \ --medusa_num_layers 1 \ --freeze_base_model True # Medusa-1 # Train Medusa-2 (joint fine-tuning) python medusa/train/train.py \ --model_name_or_path lmsys/vicuna-7b-v1.3 \ --data_path high_quality_data.json \ --bf16 True \ --output_dir medusa-vicuna-7b-v1.3-joint \ --num_train_epochs 1 \ --per_device_train_batch_size 4 \ --gradient_accumulation_steps 8 \ --learning_rate 1e-5 \ # Lower LR for base model --medusa_num_heads 4 \ --freeze_base_model False # Medusa-2 (joint)参数要点:
--medusa_num_heads:新增预测头数量(上文经验值 3–4);--medusa_num_layers:每个头的层数,轻量场景取 1;--freeze_base_model:切换 Medusa-1/Medusa-2 的总开关;--learning_rate:Medusa-1 取1e-3(只训头,可大胆些),Medusa-2 取1e-5(联合微调,保护骨干)。
部署与生成
加载 Medusa 模型
medusa.md 部署小节 提供两种加载方式:
from medusa.model.medusa_model import MedusaModel # Load pre-trained Medusa model model = MedusaModel.from_pretrained( "FasterDecoding/medusa-vicuna-7b-v1.3", torch_dtype=torch.float16, device_map="auto" ) # Or load base + Medusa heads separately base_model = AutoModelForCausalLM.from_pretrained("lmsys/vicuna-7b-v1.3") medusa_heads = torch.load("medusa_heads.pt") model = MedusaModel(base_model, medusa_heads)第二种方式适合"已有基础模型 + 离线训练好的头文件"的复用场景。SKILL.md 快速上手 中同样演示了MedusaModel.from_pretrained(...)配合medusa_generate的完整调用链。
生成调用
# Generate with Medusa outputs = model.medusa_generate( input_ids, max_new_tokens=256, temperature=0.7, posterior_threshold=0.09, # Acceptance threshold posterior_alpha=0.3, # Tree construction parameter medusa_choices=medusa_choices, # Candidate paths )posterior_threshold:后验接受阈值(0.09 为论文标准值);posterior_alpha:树构造参数(0.3);medusa_choices:候选路径配置,与上文调优小节对应。
与推测解码(Draft Model)的对比与选型
medusa.md 对比表 系统对比了两条技术路线:
| Aspect | Medusa | Speculative Decoding |
|---|---|---|
| Draft Model | Built-in (heads) | External (separate model) |
| Training | Minimal (heads only) | None (use existing small model) |
| Memory | Base + heads (~1-2% overhead) | Base + draft (can be large) |
| Speedup | 2-3.6× | 1.5-2× |
| Deployment | Single model | Two models |
何时选 Medusa:
- 希望单模型部署(少一套模型、少一份显存);
- 可以承受极少量训练(只训头);
- 需要最大加速(3× 以上)。
何时选推测解码:
- 手头已有现成的小模型可当草稿;
- 零训练预算;
- 追求最简配置。
该技能包还给出了混合模式(SKILL.md 进阶模式):把训练好的 Medusa 模型当作推测解码的草稿模型,assistant_model=draft_medusa传入generate,让更大的目标模型验证——同时获得 Medusa 的多 token 草稿能力与大模型的生成质量。此外 SKILL.md 的选型建议 给出了更细的决策树:新部署优先 Medusa;已有小版本模型优先草稿式推测解码;要求零训练即插即用则用 Lookahead Decoding(Jacobi 迭代方案,详见同目录 lookahead.md)。
在本仓库中的使用方式
本技能位于仓库的 19-emerging-techniques/speculative-decoding/ 目录下,包含三份核心材料,可配合查阅:
- SKILL.md:技能总纲,含安装(
pip install transformers accelerate、Medusa 仓库pip install -e .)、三种方法的快速上手、进阶模式(训练 Medusa 头、混合推测解码、vLLM 部署speculative_model参数)与最佳实践; - references/medusa.md:本文的主题来源,Medusa 架构、训练与推理算法的完整参考;
- references/lookahead.md:Lookahead Decoding(Jacobi 迭代)的互补方案,适合零训练、即插即用场景。
按上述流程,你可以从"冻结骨干训练 Medusa-1 快速验证效果",再到"联合微调 Medusa-2 榨取 2.3–3.6× 加速",最后以medusa_generate单模型部署上线;若追求更激进的吞吐,还可通过 vLLM 的speculative_model参数将 Medusa 作为草稿模型接入生产服务。
- AI 技能
- 人工智能
- 大模型
- 深度学习
【免费下载链接】AI-Research-SKILLs
Comprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.
相关推荐
Medusa:加速LLM生成的多解码头简单框架
Medusa:加速LLM生成的多解码头简单框架 项目介绍 Medusa是一个旨在通过多解码头技术加速大型语言模型(LLM)生成的简单框架。该项目通过在同一模型上
人工智能大模型微调本地部署终极LLMA加速指南:如何实现大语言模型2-3倍无损推理加速
终极LLMA加速指南:如何实现大语言模型2 3倍无损推理加速 LMOps(GitHub加速计划)中的LLMA(Large Language Model Acce
大模型深度学习NLPRAGAI Agent如何用LyricsX打造macOS终极歌词体验:完整配置指南
如何用LyricsX打造macOS终极歌词体验:完整配置指南 LyricsX是一款专为macOS设计的终极歌词应用程序,能够自动搜索并显示当前播放歌曲的歌词,为
桌面应用音视频
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考