MoE 模型预训练完整 Recipe¶
更新日期:2026-04-15
一、MoE Recipe 与 Dense 的差异¶
| 维度 | Dense Recipe | MoE Recipe | 为什么存在差异 | 实践注意 |
|---|---|---|---|---|
| 并行策略 | TP + PP + DP | TP + PP + DP + EP (Expert Parallelism) | MoE 专家数量多 (64-256),单卡放不下所有专家,需 EP 将专家分布到多卡 | EP 必须放在同一 rack 内 (All-to-All 通信密集);EP 和 DP 的乘积 = 数据维度总并行度 |
| 路由调试 | 不需要 | 需要实时监控:负载 CV、死专家数、有效专家数、router bias 范围 | 路由器是可学习模块,容易出现"赢者通吃"坍缩——少数专家垄断所有 token | 建议每步都 log 负载分布;坍缩一旦发生很难恢复,不如预防 |
| Loss 组成 | CE loss (纯交叉熵) | CE + aux_loss (可选,鼓励均衡路由) + z_loss (稳定 logits 尺度) | aux_loss 惩罚负载不均,z_loss 防止 router logits 爆炸;但 aux_loss 会损害模型质量 | DeepSeek-V3 用 Aux-Loss-Free (动态 bias 替代 aux_loss),仅保留 z_loss=0.001;这是当前最优实践 |
| 训练稳定性 | 较稳定:梯度流路径固定 | 更容易崩:路由决策引入离散性,专家负载动态变化 | 路由是离散选择 (top-K),梯度需通过 straight-through estimator 近似,天然不如连续网络稳定 | 前 2-3 层用 Dense 可大幅改善初期稳定性;warmup 阶段监控尤为重要 |
| Warmup 需求 | 2000 steps 标准 warmup | 更长 warmup (2000-5000 steps):路由器需要更多时间收敛 | 路由器初始随机,需要足够步数才能学会合理分配 token;过早上大 lr 会固化坏路由模式 | 可对 router 用独立的更小 lr (0.1× 主 lr),或增加 warmup 步数到 5000 |
| 数据配比 | 标准 | 类似,但对数据多样性更敏感 | 专家分化依赖数据多样性:如果数据太同质,专家无法形成有意义的分工 | 确保每个 batch 内数据类型足够多样,避免连续 batch 都是同一类数据 |
| 超参调节难度 | 中:主要调 lr/batch size | 高:额外需调 aux_loss 权重/z_loss/router lr/bias lr/top-K/专家数量 | MoE 引入的超参维度多、相互耦合,且对初始化更敏感 | 建议先在小规模 (8 专家, top-2) 上调通超参,再 scale 到大规模 (256 专家, top-8) |
二、参考 Recipe:DeepSeek-V3¶
2.1 模型配置¶
2.2 训练超参¶
2.3 并行配置¶
DeepSeek-V3 完全不用 TP!这是反直觉的设计。因为 MLA 的 KV Cache 已经很小,且 EP 提供了足够的并行。不用 TP 避免了 AllReduce 开销。
三、MoE 训练的关键 Tricks¶
3.1 前几层用 Dense¶
# DeepSeek-V3 前 3 层用 Dense FFN, 后 58 层用 MoE
# 原因:
# 1. 前几层处理底层特征, 无需专家分化
# 2. 稳定早期训练 (MoE 容易在早期崩塌)
# 3. 减少 warmup 阶段的路由噪声
3.2 Shared Expert¶
# 每个 MoE 层保留 1 个始终激活的专家
# 作用:
# - 学习"基础知识" (不需要路由就能访问)
# - 兜底: 即使路由坍缩, shared expert 仍然工作
# - 减少路由专家的冗余 (专业化)
3.3 Aux-Loss-Free 负载均衡¶
# 关键创新: 不用 auxiliary loss, 而是动态调整 router bias
for step in training_loop:
# 正常路由
logits = router.gate(x) + router.bias
top_k_ids = topk(logits, k=8).indices
# 计算负载
load = compute_expert_load(top_k_ids)
target = 1.0 / n_experts
# 动态调整偏置
router.bias += bias_lr * (target - load)
# bias 不参与梯度! 只是控制路由分布
3.4 MTP (Multi-Token Prediction)¶
DeepSeek-V3 的另一个创新:在每个位置预测多个未来 token (不只是下一个)。参考 Multi-Token Prediction (Meta, 2024)。
# 传统: 预测下一个 token
# MTP: 预测接下来的 4 个 token
class MTPHead:
def __init__(self, d_model, vocab_size, n_future=4):
self.shared_embed = Embedding(vocab_size, d_model) # 与 main embed 共享
self.mtp_layers = [TransformerBlock(d_model) for _ in range(n_future)]
self.lm_head = Linear(d_model, vocab_size)
def forward(self, hidden, targets):
losses = []
for k in range(n_future):
# 用第 k 个 MTP 层预测第 k+1 个未来 token
mtp_input = concat([hidden, embed(targets[:, :-k-1])])
h_k = self.mtp_layers[k](mtp_input)
logits = self.lm_head(h_k)
loss_k = cross_entropy(logits, targets[:, k+1:])
losses.append(loss_k)
return sum(losses) * lambda_mtp # 加权求和
# 效果:
# 1. 训练时作为辅助 loss → 提高数据效率
# 2. 推理时可用于投机解码 → 加速
四、多阶段训练¶
flowchart LR
s1["S1: Dense warmup<br/>0-3 layers dense<br/>(MoE 后期不稳)"]
s2["S2: MoE main<br/>所有 MoE 层激活<br/>aux-loss-free LB"]
s3["S3: Annealing<br/>+ MTP head<br/>+ 高质量数据"]
s4["S4: Long context<br/>YaRN extend"]
s1 --> s2 --> s3 --> s4
classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
class s1,s2,s3,s4 stage
4.1 各阶段配置¶
五、监控指标¶
5.1 MoE 专属指标¶
class MoEMonitor:
def log_per_step(self, moe_layers):
for i, layer in enumerate(moe_layers):
counts = layer.last_expert_counts # [n_experts]
# 1. 负载变异系数
load = counts / counts.sum()
cv = load.std() / load.mean()
wandb.log({f'layer{i}/load_cv': cv})
# 2. 死专家数
dead = (load < 1e-4).sum()
wandb.log({f'layer{i}/dead_experts': dead})
# 3. 专家利用率熵 (有效专家数)
entropy = -(load * torch.log(load + 1e-10)).sum()
effective = torch.exp(entropy)
wandb.log({f'layer{i}/effective_experts': effective})
# 4. 路由器 bias 的范围 (Aux-Loss-Free 特有)
bias_range = layer.router.bias.max() - layer.router.bias.min()
wandb.log({f'layer{i}/bias_range': bias_range})
# 5. 专家间梯度 norm 差异
grad_norms = [exp.weight.grad.norm() for exp in layer.experts]
wandb.log({f'layer{i}/grad_norm_cv': std(grad_norms) / mean(grad_norms)})
5.2 告警阈值¶
六、MoE Recipe 选择¶
flowchart TB
start["要训 MoE"]
q1{"从哪开始"}
start --> q1
q1 -->|"全新训练"| scratch["从零训练<br/>(推荐, 稳定)"]
q1 -->|"已有 Dense"| up["Upcycling<br/>(Dense FFN 复制 N 份)"]
scratch --> q2{"参数配置"}
up --> q2
q2 -->|"小规模 7B-30B"| small["8-32 expert × top-2<br/>(Mixtral / Qwen3-MoE)"]
q2 -->|"大规模 100B+"| large["128-256 expert × top-8<br/>(DeepSeek-V3 / Kimi K2)"]
q2 -->|"极致稀疏"| sparse["384+ expert × top-8<br/>(K2 路线)"]
classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
classDef decision fill:#f5f3eb,stroke:#bdb9ab,color:#1a1a1a;
class start,scratch,up,small,large,sparse stage
class q1,q2 decision
6.1 从 Dense 到 MoE 的升级¶
# 方法 1: 从零训练 MoE (推荐)
# 更稳定, 路由器从零学习
# 方法 2: Upcycling (Dense → MoE)
# 将 Dense 的 FFN 复制 N 份作为 N 个专家初始化
# 优点: 利用 Dense 模型已有能力
# 缺点: 初始专家同质, 需要更多训练才能分化
# 参考: "Sparse Upcycling" (Komatsuzaki et al., 2023)
def upcycle_dense_to_moe(dense_model, n_experts=8):
for layer in dense_model.layers:
old_ffn = layer.ffn
# 复制 N 份
layer.moe = MoELayer(
experts=[copy(old_ffn) for _ in range(n_experts)],
router=random_init()
)
layer.ffn = None
return dense_model
七、常见错误¶
| 错误 | 后果 | 解决 | 为什么会出问题 | 预防措施 |
|---|---|---|---|---|
| 没有 warmup 或 warmup 太短 | 早期路由崩塌:少数专家垄断,不可逆 | 至少 2000 步 warmup,MoE 可加到 5000 | 路由器随机初始化时 logits 分布不稳定,大 lr 会放大初始偏差并通过正反馈锁定坏模式 | 训练前 5K 步每步 log 负载分布,观察路由是否在收敛到均匀 |
| EP 跨机架 | 训练慢 2-5×:All-to-All 通信受限于跨 rack 带宽 | EP 放同一 rack (IB 直连) | All-to-All 是 MoE 特有的通信模式,每 token 需发到其路由专家所在 GPU;跨 rack 延迟 5-10× 于 rack 内 | 硬件拓扑感知的进程分配;EP 维度一定映射到物理上最近的 GPU 组 |
| aux loss 权重太大 | 损害模型质量:模型"学会"均匀分配但不会正确路由 | 用 Aux-Loss-Free (动态 bias 替代) | aux_loss 梯度与主 CE loss 梯度竞争,大权重时均衡目标 dominate 路由学习 | 如必须用 aux_loss,权重建议 0.001-0.01;更好的选择是 Aux-Loss-Free |
| router 学习率与主模型相同 | 路由不稳定:router 参数少但梯度大,容易震荡 | router lr = 0.1 × 主 lr | router 是 d_model → n_experts 的小矩阵,参数量远少于专家 FFN,per-param 梯度相对更大 | 用独立的 lr group 管理 router 参数;也可用更强的 weight decay 约束 |
| 前几层就用 MoE | 训练崩塌:底层特征还未稳定时路由决策噪声极大 | 前 2-3 层用 Dense FFN | 底层处理 token embedding/位置编码等基础变换,所有 token 需要相似处理,不适合专家分化 | DeepSeek-V3 用 3 层 Dense;Mixtral 用 0 层 (但仅 8 专家风险低);专家数越多越需要 Dense 前缀层 |
| 不监控专家负载 | 悄悄崩塌:loss 可能看不出异常,但路由已退化 | 每步监控负载 CV/死专家数/有效专家数 | 路由坍缩时 loss 可能仅微升 (shared expert 兜底),但模型容量实际大幅下降 | 设置自动告警:死专家 >5 或有效专家 <150 时发通知,不要等到训练结束才发现 |
参考文献¶
-
[1] DeepSeek-V3 Technical Report. 2024. 论文
-
[2] DeepSeek-AI. DeepSeekMoE. 2024. 论文
-
[3] Aux-Loss-Free Load Balancing. 2024. 论文
-
[4] Meta AI. Multi-Token Prediction. 2024. 论文
-
[5] Jiang et al. Mixtral. 2024. 论文
-
[6] Fedus et al. Switch Transformer. 2022. 论文
-
[7] Komatsuzaki et al. Sparse Upcycling. 2023. 论文
↑ 上级 · D. 预训练 Recipe