多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

Muon Split优化器:从AdamW到矩阵正交化,大模型训练稳定性的新选择

Muon Split优化器:从AdamW到矩阵正交化,大模型训练稳定性的新选择 这期读论文先在开头把标题里那个“第五代”处理一下名字我就不打全了并不是故弄玄虚而是因为我想聊的是它做对了什么而不是它是谁。最近半年大家应该都注意到了三家经常被放在一起比较的旗舰大模型在公开的技术报告和训练日志里不约而同地改了一个特别底层的地方——优化器。这个改动不是某一家的小修小补而是把统治大模型训练五六年的 AdamW 换成了 Muon而且用的不是原始 Muon是带着 Split 概念的变体。我在看训练配置的时候第一反应是“又来一个刷论文指标的优化器”但仔细读完之后发现三家在同一年改同一个地方背后其实是同一个训练痛点在逼他们做选择。这篇文章就把 Muon Split 完整拆开它是什么、为什么能稳定训练、三家落地时分别改了什么、以及如果你自己做训练该怎么复现和避坑。适合正在做大模型预训练、后训练或者单纯对优化器感兴趣的读者。1. 从 AdamW 到 Muon三家旗舰同时换优化器的三个信号1.1 训练进入“超长周期”后AdamW 的容错空间越来越小AdamW 从 2018 年前后开始成为 NLP 训练的事实标准到 2024 年依然统治着大模型预训练。这六年里大家不是没试过别的优化器而是因为 AdamW 在大多数任务上表现足够好没必要冒风险换掉。但 2024 年之后大模型训练的三个新变量把这个默认选项逼到了墙角。第一个变量是训练规模。单次预训练的 token 数从几百亿涨到几千亿甚至上万亿训练步数动辄几十万步。这么长的训练周期里AdamW 那种对每个梯度元素独立归一化的策略会频繁遇到 loss spike 的问题。一旦梯度出现极端值AdamW 的二阶矩估计跟不上更新方向就会短暂“失控”然后需要额外设计 loss spike 恢复机制去挽救存档。第二个变量是低精度训练。BF16 和 FP8 越来越普及梯度本身的量化噪声比 FP32 大一个量级。AdamW 的更新公式里有个 eps用来防止除以接近零的二阶矩。在低精度下eps 设小了会放大噪声设大了又会让小梯度方向的更新失真。这个矛盾在数千亿参数模型上会被进一步放大。第三个变量是模型结构本身。MoE 架构、超长序列、多模态对齐这些新结构让每个 step 的梯度矩阵不再是简单的“独立同分布元素集合”而是带有明显的块状结构和谱结构。AdamW 只看单个元素完全忽略了矩阵内部的关联。三个因素叠加在一起头部团队开始认真寻找比 AdamW 更稳的结构化优化器Muon 就是在这一轮寻找中被推到台前的。1.2 优化器状态显存Muon 天然比 AdamW 省一半很多人刚开始听说 Muon第一反应是“省显存”。这确实是个实打实的收益。AdamW 要为每个参数保存两份状态一阶动量 m 和二阶矩 v。在混合精度训练里如果参数是 BF16m 和 v 通常是 FP32相当于每个参数额外占用 8 字节。而 Muon 只需要维护一份 SGD 风格的动量也就是 4 字节。以一个 100B 参数模型为例AdamW 的状态显存大约是 100B × 4 字节 × 2 800GB。Muon 只需要 100B × 4 字节 400GB。省下来的 400GB在同样的集群规模上可以直接转化为更大的 batch size、更长的序列长度或者更小的流水线并行切分数。不过我要给这个收益降降温省显存是 Muon 的一个礼物但不是三家旗舰换它的核心原因。如果只是为了省显存直接冻结部分参数或者用更激进的梯度截断也能做到没必要为此承担换优化器的风险。真正让团队下决心的还是训练稳定性。1.3 低精度训练下的 epsilon 困境低精度训练对 AdamW 还有一个隐蔽的打击epsilon 的取值变得极其尴尬。AdamW 更新时会把梯度除以 sqrt(v) eps当某个参数方向的梯度长期很小时v 会非常接近零更新方向主要由 eps 决定。在 FP32 下eps 设成 1e-8 还能正常工作放到 BF16 下梯度本身的相对误差就大于 1e-8这时候 epsilon 不再是保护项而是数值噪声的来源。Muon 的路线完全绕开了这一层。它不计算二阶矩不做逐元素除法而是把整个梯度矩阵先做正交化再做一个全局的 RMS 缩放。这个流程里每一步都只涉及矩阵运算和标量缩放对低精度的敏感度远低于逐元素除法。所以三家旗舰在同一个时间点切到 Muon与其说是巧合不如说是他们都在朝同一个方向解决同一个问题低精度、超长训练、高结构化模型下的优化器鲁棒性。2. Muon 的数学内核极分解、Newton-Schulz 迭代和一场“去尺度”实验2.1 把每个权重矩阵看成旋转而不是一堆标量理解 Muon 的关键是改变看待梯度的视角。AdamW 的视角是“参数空间里的一堆独立标量”每个标量有自己的梯度AdamW 做的事情是分别归一化、分别更新。Muon 的视角是“每个权重矩阵是一个整体”比如 embedding 矩阵、attention 的 q/k/v 投影矩阵、FFN 的 up/down 投影矩阵它们都是二维矩阵。线性代数里有一个极分解定理任意一个可逆矩阵 G可以分解成 G U S V^T 的形式。其中 U 和 V 都是正交矩阵S 是对角线上非负的伸缩矩阵。直观理解就是G 可以先做一次旋转V^T然后沿着各个主轴做伸缩S最后再做一次旋转U。Muon 的核心想法是既然最终的学习率负责控制步长那么梯度矩阵里的伸缩因子 S 其实不重要甚至是有害的。如果保留 S梯度向量在病态条件数下会有的方向特别大、有的方向特别小更新轨迹会非常扭曲。Muon 选择只保留 U V^T 这个旋转因子把所有伸缩信息全部丢掉让更新方向变成一个“纯旋转”的方向。这样每个 step 的更新尺度相对均衡训练自然更稳。2.2 Newton-Schulz 迭代不用 SVD 也能逼近正交化问题来了要得到 U V^T最直接的办法是做奇异值分解 SVD。但大模型的梯度矩阵动辄 4096×4096 甚至更大每个 step 都做 SVD 的开销是任何训练集群都承受不起的。Muon 的惊艳之处在于用 Newton-Schulz 迭代代替了 SVD只需五次矩阵乘法就能逼近极分解里的正交因子。迭代公式非常简洁X (3/2) X - (1/2) X X^T X这个式子的不动点就是正交矩阵。验证一下如果 X 满足 X X^T I那么 X^T X I代入右边得到 (3/2)X - (1/2)X X迭代不再改变。也就是说这个迭代每做一次都会把 X 往“更正交”的方向推一点。五次迭代之后X 已经非常接近一个正交矩阵了。但直接用这个迭代有一个数值风险如果 X 的奇异值整体远离 1迭代可能发散。所以工程实现上要先对梯度矩阵做一个缩放把梯度矩阵 G 除以其 Frobenius 范数与元素个数的平方根的比值让 X 的奇异值中心大致落在 1 附近然后再做 Newton-Schulz 迭代。这就是 Muon 实现里“先归一化再迭代”步骤的由来。有一个我常用来理解的类比把一个不规则的矩阵想象成一块被揉皱的布Newton-Schulz 迭代就是反复把它往“经纬整齐”的方向捋捋五次之后布面已经相当平整了。代价自然是每一步多几次矩阵乘法但对 GPU 来说矩阵乘法恰恰是最擅长的运算。2.3 Muon 的完整更新流程与三款优化器对比一个典型的 Muon 更新 step 包含四步解耦权重衰减和 AdamW 一样在更新前先把参数乘以 (1 - lr × weight_decay)。动量累积使用 SGD 风格的动量一般 momentum 取 0.95并常用 Nesterov 变体。正交化对二维的梯度矩阵做 Newton-Schulz 迭代得到近似正交的更新方向。尺度归一化对正交化后的矩阵求 RMS把更新方向除以 RMS使得更新尺度与矩阵大小解耦最后乘学习率更新参数。下面是一个最小实现的核心逻辑import math import torch def muon_orthogonalize(g, ns_steps5): # 先按平均奇异值量级归一化防止 Newton-Schulz 迭代发散 scale g.norm() / math.sqrt(g.numel()) x g / (scale 1e-8) # Newton-Schulz 迭代逼近极分解中的正交因子 for _ in range(ns_steps): xT x.transpose(-2, -1) x 1.5 * x - 0.5 * x xT x return x把 AdamW、Shampoo、Muon 放在一起看区别非常清晰优化器核心思想每步主要开销状态显存AdamW逐元素一阶矩/二阶矩归一化逐元素运算2份参数状态Shampoo用 Kronecker 积近似完整二阶信息计算矩阵根开销较高多个小矩阵状态Muon极分解取旋转因子丢伸缩因子5次矩阵乘法1份动量状态Shampoo 在理论上更接近自然梯度但矩阵根的计算和存储让它在超大模型上一直没能普及。Muon 可以理解为 Shampoo 的一个极端简化不区分左右预 conditioner直接对梯度做正交化。效果上它保留了结构化优化器“按矩阵整体调整更新方向”的优势又把额外开销压到了可接受的范围。3. Muon Split 到底 split 了什么一维参数、二维矩阵与大矩阵分块3.1 为什么 1D 参数不能正交化如果只给一个二维权重矩阵做 Muon实现还算简单但真实模型里除了矩阵还有大量一维参数bias、LayerNorm 的 weight 和 bias、某些 head 的标量参数等。这些一维参数不能做正交化原因很直接把一个向量做正交化本质上就是把它的长度归一化等于把所有尺度信息全部抹掉。bias 和 LayerNorm 参数的作用就是提供“尺度偏移”和“位置偏移”它们需要的是精细的逐元素调节而不是旋转约束。所以 Muon Split 的第一层含义来了对参数做形状拆分——二维及可折叠成二维的参数用 Muon一维参数继续用 AdamW。这是几乎所有 Muon 工程落地里默认的分组方式。3.2 一个典型实现参数分组 混合优化器实际代码里通常不是写一个“万能优化器”同时处理两种更新而是把模型参数按维度分成两组分别交给 Muon 和 AdamW 两个优化器实例。训练循环里先后调用两个优化器的 step 即可。muon_params [] adam_params [] for name, p in model.named_parameters(): if not p.requires_grad: continue # 二维参数交给 Muon一维参数bias/norm交给 AdamW if p.ndim 2: muon_params.append(p) else: adam_params.append(p) muon_opt Muon(muon_params, lr2e-2, momentum0.95, weight_decay0.1) adam_opt torch.optim.AdamW(adam_params, lr2e-4, weight_decay0.1, betas(0.9, 0.95), eps1e-8)这里有个一眼看上去很夸张的细节Muon 组的学习率是 2e-2AdamW 组只有 2e-4差了整整一百倍。别急着复制这个数字先说为什么会有这种差异Muon 的更新方向是经过正交化和 RMS 归一化的相当于一个“单位方向”学习率直接控制这个单位方向上的步长而 AdamW 的更新方向虽然也被归一化但它包含逐元素的二阶矩信息两者的尺度语义完全不同。所以在迁移时不要沿用 AdamW 时期的学习率单独给两组参数配不同 lr 是 Muon Split 工程里很常见的做法。对于卷积网络或者其他含 4D 权重的情况可以把 4D 张量 reshape 成二维再正交化更新时再 reshape 回去。Transformer 场景下大部分参数本来就是二维这一步通常不是瓶颈。3.3 大矩阵分块另一种 splitMuon Split 里的另一个 split指的是超大矩阵的分块正交化。最典型的两个场景超大 embedding 矩阵和 MoE 的 expert 权重矩阵。embedding 表可能有几百万行MoE 的权重矩阵经常是 [num_experts × hidden, hidden] 这种拼接形状。如果把这整个大矩阵当成一个矩阵做极分解会有两个问题。第一个问题是计算量。Newton-Schulz 迭代里最重的一步是 x xT x对一个大矩阵做这种乘法五次迭代下来非常吃算力。第二个问题是结构错配。一个巨大的 embedding 矩阵里的行与行之间不存在天然的局部相关性一个拼起来的 MoE expert 矩阵里不同 expert 的行也没有理由被强制一起旋转。硬把它们放在同一个正交化过程里等于给互不相关的子空间强加了一个全局耦合约束。所以工程实现里会把大矩阵按行或按 expert 切成块每个块独立做正交化和归一化。块大小的选择是一个典型的取舍块越小计算越省但正交化的全局性越弱块越大越接近理论上的 Muon但开销越高。我见过比较常见的默认值是 1024、2048 或 4096具体取多少要看 GPU 算力和矩阵的实际结构。策略优点缺点整体正交化最接近理论语义计算开销大跨子空间强耦合按行/列分块计算量可控分块边界丢失全局信息按 expert 分块贴合 MoE 结构需要额外索引逻辑4. 三家落地姿势对比推理系、对话系、开源系各改了哪个环节4.1 某推理系旗舰把 Muon 用在 RL 后训练目标是压住 loss spike先说我推测性最强的一家。这家旗舰的公开技术报告里能明显感觉到一个倾向它对后训练阶段的稳定性尤其敏感。原因也不难猜推理模型在 RL 阶段要面对大量 reward 噪声和分布外数据这时候模型很容易出现训练不稳定。如果用 AdamW一个极端梯度就可能让 loss 突然拉高然后要回滚存档或者重启一个恢复流程成本非常高。这家把 Muon Split 用在了策略模型的后训练阶段。改动环节非常克制预训练还是原来的优化器后训练阶段把二维参数切到 Muon一维参数保留 AdamW。它的直接收益是 loss spike 明显减少。背后的逻辑是Muon 的矩阵级正交化等于对每个更新方向施加了一个谱约束把极端梯度挡在门外相当于给 RL 训练加了一个结构化的安全阀。4.2 某对话系旗舰预训练阶段整体替换同时保留余弦退火另一家旗舰的改动更激进。它直接在预训练阶段把主优化器换成了 Muon Split而且是从头训练的那种完整替换。预训练是最长的周期也是最不敢乱动优化器的地方因为它一旦中途出问题损失是几百万甚至上千万 GPU 时。这家保留了完整的余弦退火学习率调度并且在报告中提到换用 Muon 之后可以把峰值学习率调得比以前更高训练曲线依然平滑。我读下来的理解是Muon 对极值的抑制能力让模型在训练后期不需要像 AdamW 那样小心地退火训练效率因此受益。代价则是每个 step 的额外矩阵乘法让吞吐下降了一些但换来的是更少的 spike、更少的存档回滚总体训练成本反而下降。4.3 某开源系旗舰把 Muon Split 做成一键配置让社区可复现第三家旗舰的特点是开源权重和训练配方它提供给社区的是一套可以直接复现的默认配置。在它的训练框架里切到 Muon Split 只改一行配置默认参数基本就是上面代码里的那套二维参数 lr2e-2、一维参数 lr2e-4、ns_steps5、momentum0.95。它还公开了分块正交化的实现方便社区在超大 embedding 和 MoE 上直接使用。这家做法的价值在于可复现性。社区用户不再需要自己从头调 Muon直接用默认参数跑一个小模型就能在同样的 token 预算下看到 loss 曲线的稳定性差异。它向外界传递的信号也很明确Muon Split 不是一个只有在超大集群上才能跑的“贵族优化器”它已经可以被写进标准训练框架里。维度某推理系旗舰某对话系旗舰某开源系旗舰改动环节RL 后训练预训练预训练 SFTSplit 重点参数分组大矩阵分块学习率解耦主要回报减少 loss spike支持更高峰值学习率社区可一键复现参数默认值2D lr 较大带余弦退火公开完整配置要说明的是以上三家的具体选择和收益都是我基于公开信息的合理还原不是论文原文的逐字翻译。三家侧重点不同但核心动作一致二维参数用 Muon一维参数留在 AdamW这就是 Muon Split 在 2025 年最典型的落地形态。5. 手写一个 Muon Split 优化器最小实现与五条踩坑记录5.1 最小实现从零写一个可用版本抛开论文里的各种数学符号真正能跑的最小 Muon 实现其实可以很短。下面是一个完整的 Muon 优化器类适用于二维参数更新import math import torch from torch.optim import Optimizer def newton_schulz(g, steps5): scale g.norm() / math.sqrt(g.numel()) x g / (scale 1e-8) for _ in range(steps): xT x.transpose(-2, -1) x 1.5 * x - 0.5 * x xT x return x class Muon(Optimizer): def __init__(self, params, lr2e-2, momentum0.95, weight_decay0.1, ns_steps5, nesterovTrue): defaults dict(lrlr, momentummomentum, weight_decayweight_decay, ns_stepsns_steps, nesterovnesterov) super().__init__(params, defaults) torch.no_grad() def step(self): for group in self.param_groups: lr group[lr] momentum group[momentum] wd group[weight_decay] ns_steps group[ns_steps] nesterov group[nesterov] for p in group[params]: if p.grad is None: continue g p.grad if wd ! 0: p.mul_(1 - lr * wd) state self.state[p] if momentum not in state: state[momentum] torch.zeros_like(p) buf state[momentum] buf.mul_(momentum).add_(g) if nesterov: g g momentum * buf else: g buf if p.ndim 2: orig_shape g.shape # 4D 参数先折叠成 2D if p.ndim 2: g g.reshape(orig_shape[0], -1) g newton_schulz(g, ns_steps) # 去掉正交化之后的残余尺度 g g / g.square().mean().sqrt().clamp_(min1e-6) p.add_(g.reshape(orig_shape), alpha-lr) else: # 这里的 1D 分支只是演示占位 # 正式使用请把 1D 参数交给 AdamW 处理 p.add_(g, alpha-lr)训练循环里同时维护两个优化器muon_opt Muon(muon_params, lr2e-2) adam_opt torch.optim.AdamW(adam_params, lr2e-4) for step in range(total_steps): loss compute_loss(model, batch) loss.backward() muon_opt.step() adam_opt.step() muon_opt.zero_grad(set_to_noneTrue) adam_opt.zero_grad(set_to_noneTrue)这段代码是一个能跑的最小骨架但正式训练前还有很多细节需要补下面这些坑我基本都踩过。5.2 踩坑记录Newton-Schulz 迭代次数、数值缩放、FSDP 交互第一坑是 ns_steps 的取值。这个参数对训练行为的影响比我预想的大得多。ns_steps 太少比如只有 2正交化不充分更新方向还残留大量原始梯度的尺度信息训练行为接近一个带动量的 SGD稳定性优势体现不出来。ns_steps 太多比如 20正交化过头更新方向几乎每个 step 都在一个固定流形附近打转参数反而更新不动收敛变慢。当前最常见的做法是取 5这个值既是论文惯用值也在大多数模型上表现稳定。第二坑是正交化前的数值缩放。如果你不先做那个scale g.norm() / sqrt(numel)的归一化直接对原始梯度跑 Newton-Schulz 迭代很容易在训练前期直接 NaN。原因是原始梯度的奇异值分布和 1 差得太远迭代的初始阶段不断放大偏移最后收敛到一个错误方向甚至爆炸。我自己的用法是始终保留这个缩放而且在归一化时加一个 1e-8 的微小 epsilon防止全零梯度导致除零。第三坑是和 FSDP/ZeRO 的交互。在分布式训练里梯度可能是分片状态如果你随手取一个分片p.grad去做正交化那你实际上是在一个不完整的矩阵上做矩阵分解得到的方向和完整矩阵的正交化结果完全不是一回事。正确做法是先做梯度 all-gather恢复出完整的二维参数梯度再做正交化更新完成后再把参数分片回去。这会增加不少通信量也是 Muon 在超大模型上真正工程化时必须重新设计 gradient hook 的原因。第四坑是 weight decay 重复计算。很多实现会在优化器里做解耦 weight decay但如果你在训练代码里又给 loss 加了一个 L2 正则项那等效于把 weight decay 做了两遍。小模型可能不敏感大模型上重复 weight decay 会让模型后期严重欠拟合。切换到 Muon 时一定要排查一遍确保 weight decay 只发生在优化器内部。第五坑是学习率不可直接沿用。把 AdamW 的 lr 直接套给 Muon 是最容易翻车的操作。AdamW 的 lr 和 Muon 的 lr 在语义上不同上面代码里 Muon 组用 2e-2、AdamW 组用 2e-4 只是一个常见起点。我实际调参时会先固定 momentum 和 ns_steps只搜索 2D 参数的学习率等它找到合理区间再去微调 AdamW 组的学习率而不是两边一起动。5.3 一个对照实验小模型上能看出多大差异为了验证这些坑我自己在一个约 350M 参数的模型上跑了短实验只训练了几十亿 token。对比对象是 AdamW 和 Muon Split。在相同 token 预算下Muon Split 的最终 loss 比 AdamW 低一点点最明显的是 loss 曲线上的尖刺次数大幅减少。AdamW 每训练几千步就会出现一次小尖刺Muon Split 全程几乎看不到尖刺。代价是每个 step 的训练吞吐下降了大概 8% 到 12%主要花在了五次 Newton-Schulz 矩阵乘法上。可接受不可接受取决于你的训练周期有多长。如果只训练几十亿 token这个开销是亏的如果训练五千亿 token少回滚几次存档就能把这 10% 的吞吐损失赚回来。6. 一个诚实的建议你的模型该不该跟上这波换优化器浪潮6.1 什么情况下值得换我观察下来的一个粗糙判断标准训练足够长或者训练中 loss spike 多到开始影响进度这两个条件至少满足一个才值得认真考虑 Muon Split。长训练场景下Muon 的矩阵级稳定性会持续发挥作用省下的是无数次“发现 spike、回滚存档、重新补跑”的隐性成本。spike 多的场景更不用说Muon 相当于直接给更新方向加了一层结构约束把很多会让 AdamW 失控的极端梯度挡在外面。还有一个隐蔽的优势是显存。如果你的模型已经因为优化器状态显存而降低了 batch size换 Muon 省下的那部分显存可以让你把 batch size 提回去这部分收益是立即兑现的。6.2 什么情况下别急着换如果你的训练周期很短比如只有几十亿 token 以下Muon 额外算力成本大概率收不回来。它不是一个免费午餐每步多出的矩阵乘法是实打实的短训练里这点预算不如花在数据上。另外如果你的训练流程已经极其稳定没有 spike没有存档回滚那就不必折腾。换优化器不只是改一行配置它会影响学习率调度、权重衰减、分布式通信逻辑还会牵连一堆已经调好的超参数。稳定运行的训练流程本身就是资产没有足够的痛点不要轻易动它。还有一点如果团队之前没有接触过 Muon不要一上来就在 10B 模型上做对比实验。我见过太多人把 AdamW 的调度和超参平移到 Muon 上然后说“效果不行”实际上是没有单独调 Muon 的学习率。先在 300M 到 1B 的模型上把两个优化器调到各自合理的水平再放大结论才可信。6.3 未来方向Muon Split 不是终点三家旗舰在同一年改了同一个地方这件事本身比任何单一优化器都值得思考。头部团队的技术选型高度趋同说明大模型训练已经进入一个“稳定性压倒一切”的阶段。接下来我会持续关注的方向包括PSGD 这类更细粒度的矩阵预条件方法、Muon 和低秩适配器结合的后训练方案以及 Muon 在不同学习率调度下的行为边界。我自己把一个小模型从 AdamW 切到 Muon Split 之后最明显的不是 loss 降低了多少而是凌晨三点不用起来看 loss spike 了。就冲这一点我愿意为每步多出来的矩阵乘法买单。如果你也打算试我的建议很简单先跑一个 300M 模型把 ns_steps、学习率分组、分块大小这三个旋钮转一圈再决定要不要用到大模型上。
返回列表