先记住一句话

混合 optimizer 的关键不是类名,而是每个 trainable parameter 恰好路由到一条明确的 update rule,并拥有正确的 LR、decay、state 与 checkpoint 语义。

枚举参数按 shape/type 分类分 decay/no-decay建立两套 state同步 scheduleassert 无重漏
Trainable paramsname, shape, type
2D hiddenMuon
embed/head/1DAdamW
Shared stepupdate + log
同一次 training step 可有不同 update geometry,但 parameter ownership 必须互斥且完备。

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_params

is_embedding_or_output_head 最好依据 module type 或显式 registry,而不是只靠字符串包含关系。Tied embedding/head 必须按 parameter identity 去重,否则同一个 storage 可能被更新两次。

3. 两条 update path

核心 state方向变换典型参数
MuonmomentumNewton–Schulz orthogonalizationattention/MLP 2D weights
Aux AdamWfirst/second moments逐坐标 adaptive scalingembedding、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

  1. 从相同 model weights、data order 与 global batch 开始;
  2. 为两种 optimizer 分别 sweep 合理 LR/decay,不沿用同一数值;
  3. 同时报告 loss/token、达到目标 loss 的 wall-clock、step time 与 peak memory;
  4. 记录各 group numel、update/weight ratio 和 nonfinite;
  5. 至少复现实验 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 与实现接口。