Dion3 让 Muon 优化器每步快 6 倍,验证损失还比 NorMuon 更低

Dion3: Full-Stack Orthogonal Updates

Noah Amsel, Jack Zhang, Kwangjun Ahn, Ali Naeimi, Austin Feng, Berlin Chen, Tri Dao, John Langford

cs.LG, cs.AI

2026-08-12

Dion3 在四个层面重写 Muon 的正交化:在小 Gram 矩阵上迭代、对称 kernel、每步只正交化动量矩阵的少数行、合并通信。大模型每步快 6.5 倍,验证损失反而更低。

这篇在解决什么

Muon 是一类把动量矩阵「正交化」后再更新权重的优化器,等价于在谱范数意义下走最陡下降。它训练质量好,过去一年越来越多地被用在百亿级模型的预训练里。代价集中在每一步都要跑一次 Newton-Schulz 迭代来近似正交化,这是立方时间复杂度;一旦权重在多卡间分片(sharded),all-to-all 通信开销又叠加上来。论文给了一个直观的代价刻度:7B 模型在 4 张 GH200 上,Muon 的优化器步骤(不含前向反向)耗时是 AdamW 的 26 倍。

之前已经有人想把这一步变便宜。Dion 用幂迭代造一个低秩近似,只正交化近似部分;Trion 更激进,直接从离散余弦变换矩阵里挑列。两者都靠 error feedback(误差反馈,把这一步引入的近似误差记下来、在后续步骤补偿)兜住质量。Dion3 沿着同一条路再推一步:既然目标是把输入变小,最简单的「低秩近似」就是直接挑出动量矩阵的若干行,其余的根本不碰。

方法

Dion3 不是单点改进,而是在四个层面同时压这一步的开销。

第一个层面是算法本身。Newton-Schulz 本来直接在大输入矩阵 X 上迭代,Dion3 改成在小的对称 Gram 矩阵 XX⊤ 上迭代,最后再乘回 X。数学上输出和标准 Newton-Schulz 完全等价,但绝大多数计算都落在小得多的 n×n 对称矩阵上,大矩形乘法从 10 次降到 2 次。在典型设置(T=5、α=4)下,比用对称 GEMM 的标准 Newton-Schulz 省 55% FLOP,比不用对称 GEMM 的常见实现省 68%。代价是 Gram 矩阵可能出现假的负特征值导致数值不稳,作者在第 3 次迭代处加一次 restart(重启)来压住,并改用 float16。

第二个层面是 GPU kernel。对称矩阵乘 A⊤A 和 A²+B 只需要算下三角再拷到上三角,作者用 CuteDSL 写了专门的对称 GEMM kernel,在大矩阵上比 cuBLAS 快约 2 倍(Hopper、Blackwell)。这恰好和 Gram Newton-Schulz 叠加,因为后者用到的对称乘更多。

第三个层面是最反直觉的算法改动:每步只正交化动量矩阵的一部分行。按 ℓ1 范数挑出 k=⌈fn⌉ 行(f 推荐 1/4 或 1/8),只正交化这个子矩阵、只更新这些行对应的权重;error feedback 只把这些选中行的动量衰减一个乘性因子 μ,没被选中的行原样留着等下一轮。f=1 时退化回原版 Muon。因为每步更新的行变少,有效步长也变小,学习率要按 η ∝ 1/√f 放大,作者据此给出了一条可解释的缩放律。

第四个层面是通信。Transformer 里只有少数几种权重形状,megabatching 把同形状的矩阵打包成一次 all-to-all,通信轮数从 O(N/worldsize) 降到 O(1),和模型层数无关。

结果

速度上各项改进可以叠加。对称 kernel 加 Gram Newton-Schulz 合计比标准 Muon 快约 1.5 倍;再叠加行选择,f=1/2 时总加速 3.6 倍,f=1/4 时总加速 6.5 倍。回到开头的 7B 例子,标准 Muon 是 26× AdamW,叠满之后降到 4× AdamW。仍然比 AdamW 慢,但已经从「贵到离谱」回到「可以接受」。

megabatching 在通信受限时收益最大:1B 模型单机 8 分片时优化器步骤从 80.7ms 降到 52.1ms(−35%);但 32 分片时只省 4%,因为每张卡手里的矩阵太少、批大小本来就接近。

质量上的结果出乎作者预料。他们本来只想证明「行选择不损害损失」,结果 Dion3(f<1)反而比 NorMuon 更低。1B 模型在 100B token 的 ClimbMix 上,最优是 f=1/8,验证损失全程低于 NorMuon。放大到 3B–14B(受算力限制只跑 10B token),Dion3(f=1/4)在四个规模上的验证损失全部更低:

规模NorMuon 损失Dion3 损失Δ下游准确率 Δ(12 项基准)
3B2.2692.257−0.012+1.0
4B2.2432.232−0.011−0.3
7B2.2202.206−0.014+0.1
14B2.1892.162−0.027+0.7

损失在所有规模都赢,下游准确率赢三输一,最大改进出现在 14B(损失 −0.027、准确率 +0.7)。

为什么重要

Muon 受关注,是因为它在「每 token 训练效率」上可能优于 AdamW,但单步贵一直是它上规模的最大障碍。Dion3 把这个障碍削掉一大块:大模型上 6.5 倍的步骤加速、且质量不退反进,代码以 dion 包发布(github.com/microsoft/dion),可作为 Muon 的 drop-in 替换。对正在用或考虑用 Muon 训练大模型的团队,这是直接能拿去省钱的改动。

但 Dion3 没有把 Muon 拉到和 AdamW 同一档:7B 上仍是 4× AdamW。Muon 总体是否划算,取决于「单步更贵」换来的「收敛更快」能补多少,Dion3 让这个权衡更偏向 Muon,但没有终结讨论。

局限与存疑

论文自己的局限写得比较坦诚。Gram Newton-Schulz 和对称 kernel 都在半精度下跑,会引入和原版不完全一致的数值差异(实验显示不影响质量,但差异客观存在)。f=1 的 Dion3 也不等于逐位的 Muon:行排列、动量衰减时机、float32 下的 NorMuon 归一化、自定义 Triton kernel 都不同,只是收敛曲线几乎重合。restart 策略本身要额外付出 3(α−1)n³ FLOP。

此外还有几点没被充分验证。第一,质量提升只在 ClimbMix 一个数据集、密集(dense)Transformer 一个架构族上测过,混合专家(MoE)只测了速度没测质量;作者自己也写明「这一改进能推广到多广还需要进一步研究」。第二,14B 那行 −0.027 的损失提升是否站得住,只跑 10B token、规模有限,样本太薄。第三,4B 那行下游准确率反而掉了 0.3,说明损失和下游指标并不总是一致,单看损失曲线可能高估了收益。

术语

原文与代码

社区讨论

相关论文

全部论文解读