先记住一句话
混合 optimizer 的关键不是类名,而是每个 trainable parameter 恰好路由到一条明确的 update rule,并拥有正确的 LR、decay、state 与 checkpoint 语义。
1. 为什么需要 parameter routing
PyTorch 的 Muon 文档把目标参数限定为 hidden-layer、维度至少为 2 的参数,并明确建议 embeddings、output heads、biases 与 gains 使用 AdamW。原因是 Muon 把 momentum reshape 成 matrix 后做 Newton–Schulz orthogonalization;1D vector 没有同样的 row/column matrix geometry,大 vocabulary table 和 output head 也有特殊 scale。
2. 一个公开可实现的分类器
for name, p in model.named_parameters():
if not p.requires_grad:
continue
if p.ndim >= 2 and not is_embedding_or_output_head(name, p):
muon_params.append(p)
else:
adamw_params.append(p)
assert disjoint(muon_params, adamw_params)
assert union(muon_params, adamw_params) == all_trainable_paramsis_embedding_or_output_head 最好依据 module type 或显式 registry,而不是只靠字符串包含关系。Tied embedding/head 必须按 parameter identity 去重,否则同一个 storage 可能被更新两次。
3. 两条 update path
| 组 | 核心 state | 方向变换 | 典型参数 |
|---|---|---|---|
| Muon | momentum | Newton–Schulz orthogonalization | attention/MLP 2D weights |
| Aux AdamW | first/second moments | 逐坐标 adaptive scaling | embedding、head、bias、norm |
两组的数值 LR 不能直接视为相等强度,因为 update norm 的生成方式不同;应分别 sweep,并按组记录 ‖ΔW‖/‖W‖。
4. Decay 还要再分组
每条 path 内通常还区分 decay 与 no-decay。Matrix weights 是否 decay、embedding 是否 decay 取决于 recipe;bias 和 normalization scale 常设 no-decay,但这是一项要显式验证的实验选择,不是优化器定理。最终至少形成 Muon-decay、Muon-no-decay、AdamW-decay、AdamW-no-decay 四个逻辑组。
5. Scheduler 与 checkpoint contract
- scheduler 应说明是乘所有 group LR,还是分别维护 base LR;
- checkpoint 保存 Muon momentum、Adam m/v、各 group step 和 scheduler state;
- 从纯 AdamW 切换时旧 m/v 不能直接解释为 Muon momentum,应明确 reset;
- 加载后先打印每组 LR、numel、state dtype,并检查第一拍 parameter delta;
- 改 module name、tied weights 或 freeze policy 后必须重新做 routing audit。
6. 四个可手算的 routing 例子
例 1:按 shape 分参数
模型含两个 4×4 hidden matrices、一个 10×4 embedding、三个 4D bias/norm vectors。Muon 组参数=2×16=32;aux 组=40+3×4=52;总数 84,检查 32+52=84 保证没有漏项。
例 2:两组 LR 的实际扰动
Muon matrix 的 update norm=2.0、LR=0.02,delta norm=0.04;aux Adam direction norm=5.0、LR=0.001,delta norm=0.005。同一个“step”内 Muon 组实际移动量是 aux 的 8 倍,不能只比较 LR 数字。
例 3:decoupled decay
Muon weight norm=10,LR=0.02、decay=0.01,则单拍 decay delta norm 近似 0.02×0.01×10=0.002;aux norm parameter 若 no-decay,则这一项严格为 0。
例 4:tied weight 重复更新
Embedding 与 output head 指向同一个 40-parameter tensor。若按名字各加入一次,账本会错误显示 80 并执行两次 step;按 id(parameter) 去重后仍是 40,总 trainable numel 才与模型一致。
7. 公平比较 checklist
- 从相同 model weights、data order 与 global batch 开始;
- 为两种 optimizer 分别 sweep 合理 LR/decay,不沿用同一数值;
- 同时报告 loss/token、达到目标 loss 的 wall-clock、step time 与 peak memory;
- 记录各 group numel、update/weight ratio 和 nonfinite;
- 至少复现实验 seed,并验证 resume 前后下一拍结果一致。
“用了 Muon”不等于模型的全部参数都用 Muon。一个正确的混合 recipe 往往刻意保留 auxiliary AdamW;判断事实要看 resolved parameter groups,而不是 optimizer wrapper 的名字。
自测
1. 为什么 1D bias 通常不走 Muon?
Muon 利用 matrix momentum 的 row/column geometry 并做正交化;1D vector 不具备这个 2D 结构。
2. 怎样证明路由没有重漏?
按 parameter identity 检查两组交集为空,并验证并集等于全部 requires-grad parameters、numel 总数一致。
3. 为什么切换 optimizer 后要重扫 LR?
Muon orthogonalized matrix direction 与 Adam 的逐坐标 normalized direction 有不同 norm,因此相同数值 LR 不表示相同 update scale。
一手与官方资料
Muon 原始说明解释 momentum orthogonalization 与 auxiliary AdamW 的参数边界;PyTorch Muon 官方文档给出目标参数、Newton–Schulz、adjust_lr_fn、weight decay 与实现接口。