[ PROMPT_NODE_22374 ]
medusa
[ SKILL_DOCUMENTATION ]
# Medusa: 多解码头推理
基于 arXiv 2401.10774 (2024) - MEDUSA: 具有多个解码头的简单 LLM 推理加速框架
## 概述
**来源**: https://arxiv.org/abs/2401.10774
**GitHub**: https://github.com/FasterDecoding/Medusa
Medusa 通过添加额外的解码头来并行预测多个后续 Token,从而增强 LLM 推理能力,在不损失质量的情况下实现 2.2-3.6 倍的加速。
## 架构
### 核心创新
不再使用单独的草稿模型,而是向现有 LLM 添加多个预测头:
输入 → 基础 LLM(冻结或微调) → 隐藏状态
├→ 头 0(原始,预测 t+1)
├→ 头 1(预测 t+2)
├→ 头 2(预测 t+3)
└→ 头 3(预测 t+4)
### 基于树的注意力机制
**关键机制**: 构建候选树,在单次前向传递中验证所有路径。
以 2 个头为例,每个头有前 2 个候选词:
根节点(当前 Token)
/
候选 1a 候选 1b (头 1:2 个选项)
/ /
C2a C2b C2c C2d (头 2:共 4 条路径)
单次前向传递即可并行评估整个树(4 个候选词)!
## 训练方法
### Medusa-1: 冻结主干模型
**方法**: 保持基础 LLM 冻结,仅训练 Medusa 解码头。
**优点**:
- 无损(基础模型未改变)
- 训练速度快(在 8 张 GPU 上约几小时)
- 所需数据量极小(约 10M Token)
**性能**: 2.2 倍加速
python
# Medusa-1 的训练循环
for batch in dataloader:
# 冻结基础模型
with torch.no_grad():
hidden_states = base_model(**batch, output_hidden_states=True).hidden_states[-1]
# 训练 Medusa 解码头
for i, head in enumerate(medusa_heads):
logits = head(hidden_states)
# 目标:偏移 (i+1) 个位置的 Token
targets = batch['input_ids'][:, i+1:]
loss += F.cross_entropy(logits[:, :-i-1], targets)
loss.backward()
optimizer.step()
**训练数据**: 任何文本语料库(Wikipedia, C4 等)
### Medusa-2: 联合微调
**方法**: 同时微调基础 LLM 和 Medusa 解码头。
**优点**:
- 更高的预测准确度(解码头与基础模型对齐)
- 更高的加速比(2.3-3.6 倍)
**挑战**: 必须保持基础模型的能力
**解决方案**