拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Muon优化器在Stiefel Manifold上的闭环投影更新详解

Muon 优化器最近在开源社区和模型训练圈子里讨论度上升得很快。很多人第一次见到这个词是因为某些大模型训练日志里出现了它的名字接着就看到了“Newton-Schulz 迭代”和“正交化”这些略显劝退的术语。但真正让 Muon 区别于普通优化器的地方其实藏在它的参数更新方式里每一轮更新之后参数矩阵都要被拉回一类特殊的几何约束面上。这个约束面就是 Stiefel Manifold而把参数拉回去的过程传统上依赖迭代近似算法。这篇文章想讲清楚的正是这个关键问题Muon 在 Stiefel Manifold 上的投影更新为什么可以写成精确的闭环形式closed-form update而不是只能用迭代逼近我会先解释 Muon 的核心设计动机再用通俗的语言讲明白 Stiefel Manifold 是什么然后推导闭环更新公式的来源最后给出一份 PyTorch 参考实现和工程落地建议。如果你正在做大模型训练、LoRA 微调或者对优化器原理感兴趣这篇文章值得收藏。读完你会明白什么时候可以用闭环更新替代 Newton-Schulz 迭代闭环更新的代价在哪里以及在实际代码里怎么落地。1. 这篇文章真正要解决的问题先聊一个实际问题。训练一个 Transformer 或 LLM 时参数张量里有很多是矩阵形式比如注意力层的 Q/K/V 投影矩阵、MLP 层的权重矩阵。常规优化器如 AdamW对每个参数元素独立地做一阶矩和二阶矩估计更新方向是逐元素的根本不关心矩阵内部列与列之间的关系。Muon 的设计思路完全不同。它先像 SGD 一样对动量项做更新然后把动量项对应的矩阵“投影”到一个约束集合上让更新后的矩阵保持某种结构通常是列正交结构。这样做的好处是参数更新的几何路径被限制在一个更“规矩”的流形上训练过程往往更稳定收敛曲线也更平滑。但问题来了将矩阵投影到 Stiefel Manifold 上传统做法是拿 Newton-Schulz 迭代反复逼近极分解polar decomposition中的正交因子。迭代就要设迭代次数就要调超参数而且每一步都是矩阵乘法维度高的时候计算量并不低。于是有人会问能不能不求近似直接写一个精确的数学表达式一步算出投影结果答案是可以。从矩阵分解的角度看极分解的正交因子本身就有精确表达式只是直接算矩阵平方根逆或 SVD在工程上开销太大。所谓 “Exact Closed-Form Update”不是说能在常数时间内算完而是说在数学上可以一次性写出投影后的精确结果不需要多层迭代逼近。对于小到中等规模的矩阵特别是 LoRA 场景下的低秩矩阵这个精确解是可以直接算的。这篇文章要解决的问题就是四件事Muon 为什么要约束到 Stiefel ManifoldStiefel Manifold 上的投影和闭环更新到底是什么关系闭环更新的数学公式如何推导、如何理解工程上什么时候用闭环更新、什么时候继续用迭代法2. Muon 优化器的核心思路2.1 动量与正交分解Muon 这个名字在社区里的解释通常是 “MomentUm Orthogonalized by Newton-schulz” 的缩写。从名字就能看出来它的核心链条是动量更新 —— 正交化 —— 更新参数具体来说维护一个动量变量 M每一轮迭代按照类似 SGD-Momentum 的方式更新M - β M (1 - β) G其中 G 是当前梯度。得到动量矩阵 M 之后Muon 的下一步不是直接用 M 去更新权重 W而是先对 M 做一次“正交化”再将正交化后的结果乘上一个学习率W - W - lr * orthogonalize(M)这里的 orthogonalize(M) 就是让输出尽量接近一个正交矩阵。如果 M 本身已经是方阵或者近正方矩阵这个操作等价于求 M 的极分解中的正交因子。2.2 为什么要正交化看到这里很多人会问我只想训练模型为什么要多此一举做正交化从优化几何的角度看普通的逐元素更新在矩阵参数空间中走的是“逐坐标”路径而正交矩阵构成的流形是弯曲的。如果权重矩阵天然被期望具有正交性比如某些归一化层或者特征映射层那么每次更新后把参数拉回正交约束面相当于保证模型始终在一个合理的参数子空间内移动。直观类比在地球表面行走时你希望每一步都贴在地球表面而不是从一个点直接跳到球体外再从外面硬拉回来。在实际训练中Muon 风格的更新往往能减少 loss 的剧烈震荡对学习率的敏感度也比 AdamW 低一些。这也是为什么一些大模型训练实验会尝试这类“正交化动量优化器”。3. Stiefel Manifold 到底是什么3.1 从定义说起Stiefel Manifold 的数学定义是St(n, p) { X ∈ R^(n×p) | X^T X I_p }翻译成人话所有满足“列向量两两正交、且每一列都是单位向量”的 n×p 矩阵组成一个流形。当 n p 时这个集合就是正交群 O(n)即所有正交方阵。这里要区分两个容易混淆的概念正交矩阵X^T X I列向量彼此正交且范数为 1。Stiefel Manifold当矩阵不是方阵而是 n×p通常 n ≥ p时列正交但不一定是方阵这就是 Stiefel Manifold 的元素。3.2 为什么优化问题会跑到 Stiefel Manifold 上Muon 更新里的“动量矩阵” M 通常是一个普通矩阵可能行数大于列数和权重矩阵形状一致。当我们要把它变成一个列正交矩阵时就是在 Stiefel Manifold 上找一点使得这一点离 M 最近。这个问题在数学上写为minimize || X - M ||_F^2 subject to X^T X I_p也就是说Muon 的正交化步骤本质上是一个“投影到 Stiefel Manifold”的优化问题。3.3 极分解的视角任何一个满列秩矩阵 M都可以分解成M Q P其中 Q 是一个列正交矩阵Q^T Q IP 是一个对称半正定矩阵。这个分解叫极分解polar decomposition。这里的 Q 恰好就是上述投影问题的解。所以“把 M 投影到 Stiefel Manifold”和“求 M 的极分解中的正交因子 Q”是同一个问题。这一步是整个理解的关键。因为极分解不是只能靠迭代近似它和 SVD 之间有明确关系如果 M 的 SVD 是 M U Σ V^T那么Q U V^T这就是极分解里的正交因子。而“闭环更新”本质上就是顺着这条 SVD 的路径直接算 Q而不是用 Newton-Schulz 一遍遍逼近。4. 从迭代逼近到闭环更新4.1 Newton-Schulz 迭代在做什么Newton-Schulz 迭代是数值线性代数里求极分解的经典办法。它的迭代式可以写成Y - (3/2) Y - (1/2) Y (Y^T Y)或更高阶的变体Y - (15/8) Y - (5/4) Y (Y^T Y) (3/8) Y (Y^T Y)^2这个迭代的出发点是如果 Y 已经是正交矩阵那么 Y^T Y I右边会保持 Y 不变如果 Y 偏离正交状态迭代会把 Y 往正交方向拉。实际使用前还要对 M 做谱范数归一化除以最大奇异值否则迭代可能不收敛。在 Muon 的实现里常见做法是对 M 先除以其 Frobenius 范数或更精细的谱范数估计再做若干次 Newton-Schulz 迭代最后乘回一个范数补偿因子。这个方案的优点是全是矩阵乘法GPU 上很容易并行不需要做 SVD速度可控。缺点是近似程度受迭代次数影响迭代太少正交性不足迭代太多浪费计算而且归一化方式也会影响最终行为。4.2 闭环更新意味着什么所谓 “closed-form update”指的是我们可以直接写出投影结果的解析表达式不依赖迭代逼近。对于极分解来说这个表达式就是基于 SVD 的Q U V^T如果 M 是对称矩阵则表达式还可以进一步写成矩阵函数形式Q M (M^T M)^(-1/2)或者等价地Q (M M^T)^(-1/2) M选择哪个形式取决于哪一侧的维度更小。对列数小于行数的矩阵通常计算 (M^T M)^(-1/2) 更省。4.3 为什么“精确”不等于“免费”这里必须强调一个关键判断闭环更新是数学精确的但它依赖 SVD 或矩阵平方根逆的计算复杂度通常是 O(n p^2) 级别对 n×p 矩阵在大规模稠密矩阵上未必比有固定迭代次数的 Newton-Schulz 更快。所以这篇文章说的“exact closed-form update”真正的价值主要体现在中低维矩阵场景比如 LoRA 的 low-rank 矩阵。对正交性有严格要求的场景迭代误差不可接受。希望简化超参数、去掉迭代次数和范数归一化因子的实验场景。科研对比实验用精确投影做基准衡量迭代近似的误差与速度。5. 闭环更新公式的推导框架这一节给出推导的数学主干。虽然 CSDN 读者不一定都关心完整推导但理解推导能帮你避免实现时用错公式。5.1 问题重述给定矩阵 M ∈ R^(n×p)满列秩。目标是求Q argmin_{X^T X I_p} || X - M ||_F^25.2 利用 SVD 分解对 M 做 SVDM U Σ V^T其中 U 是 n×p 列正交矩阵V 是 p×p 正交矩阵Σ 是 p×p 对角矩阵。极分解的标准形式为 M Q P其中Q U V^T P V Σ V^T验证Q^T Q V U^T U V^T V V^T I_p所以 Q 确实属于 Stiefel Manifold。5.3 从 SVD 到闭式公式进一步利用 M^T MM^T M V Σ^2 V^T因此(M^T M)^(-1/2) V Σ^(-1) V^T从而M (M^T M)^(-1/2) U Σ V^T V Σ^(-1) V^T U V^T Q这就是常见实现里写u v.T的来源。5.4 对对称矩阵的特殊情况如果 M 是对称的U 与 V 在符号选择上可能一致于是Q U U^T但工程上一般不依赖这种特殊情况直接走通用路径更稳。实际实现时还要注意 SVD 的数值稳定性比如使用torch.linalg.svd时选择full_matricesFalse并考虑奇异值截断或加小常数防止除零。6. PyTorch 中的 Muon 参考实现下面给出一个完整可运行的 PyTorch 参考实现包含两种正交化方式默认的 Newton-Schulz 迭代以及可选的 SVD 闭环更新。代码统一使用类组织方便插入训练脚本。6.1 核心实现# 文件路径muon.py import torch import torch.nn as nn import torch.optim as optim def newton_schulz_orthogonalize(M, iterations5, order5): 使用 Newton-Schulz 迭代将矩阵投影到 Stiefel Manifold 附近。 参数: M: 需要正交化的矩阵 iterations: 迭代次数 order: 迭代阶数5 对应五次收敛成本更高但更精确 返回: 正交化后的矩阵 # 使用 Frobenius 范数做缩放保证迭代稳定 norm M.norm() 1e-16 X M / norm if order 3: a, b, c 1.5, -0.5, 0.0 elif order 5: a, b, c 15.0 / 8.0, -5.0 / 4.0, 3.0 / 8.0 else: raise ValueError(order 参数只支持 3 或 5) for _ in range(iterations): XT X.transpose(-2, -1) X a * X b * X (XT X) c * X (XT X) (XT X) return X * norm def closed_form_orthogonalize(M): 使用 SVD 计算极分解的正交因子得到精确闭环投影。 返回: Q: 满足 Q^T Q I 的列正交矩阵 U, _, Vh torch.linalg.svd(M, full_matricesFalse) Q U Vh return Q class Muon(optim.Optimizer): Muon 优化器动量更新 Stiefel Manifold 投影。 参数: params: 需要优化的参数 lr: 学习率 momentum: 动量系数 use_closed_form: True 使用 SVD 闭环投影False 使用 Newton-Schulz 迭代 ns_iterations: Newton-Schulz 迭代次数仅在 use_closed_formFalse 时生效 def __init__(self, params, lr0.02, momentum0.95, use_closed_formFalse, ns_iterations5): if momentum 0.0 or momentum 1.0: raise ValueError(momentum 应在 [0, 1] 范围内) defaults dict(lrlr, momentummomentum, use_closed_formuse_closed_form, ns_iterationsns_iterations) super().__init__(params, defaults) torch.no_grad() def step(self, closureNone): loss None if closure is not None: with torch.enable_grad(): loss closure() for group in self.param_groups: lr group[lr] momentum group[momentum] use_closed_form group[use_closed_form] ns_iterations group[ns_iterations] for p in group[params]: if p.grad is None: continue grad p.grad state self.state[p] if momentum_buffer not in state: state[momentum_buffer] torch.zeros_like(p) buf state[momentum_buffer] buf.mul_(momentum).add_(grad) if use_closed_form: projected closed_form_orthogonalize(buf) else: projected newton_schulz_orthogonalize( buf, iterationsns_iterations ) p.sub_(projected, alphalr) return loss代码说明newton_schulz_orthogonalize先按 Frobenius 范数缩放迭代结束后再乘回缩放系数这和常见的 Muon 开源实现思路一致。closed_form_orthogonalize直接调用torch.linalg.svd用full_matricesFalse避免计算无用的大矩阵。Muon类维护动量 buffer更新后对动量矩阵做正交投影再乘学习率更新参数。6.2 快速验证代码下面的脚本用一个 64x32 的随机参数验证优化器的基本流程并检查投影后的正交性。# 文件路径demo_train.py import torch import torch.nn as nn from muon import Muon def check_orthogonality(matrix): 返回 ||Q^T Q - I||_F越小说明正交性越好。 Q matrix I torch.eye(Q.shape[1], deviceQ.device) diff Q.transpose(-2, -1) Q - I return diff.norm().item() def run_demo(use_closed_form): torch.manual_seed(42) model nn.Linear(64, 32, biasFalse) optimizer Muon( model.parameters(), lr0.02, momentum0.95, use_closed_formuse_closed_form, ns_iterations5, ) inputs torch.randn(16, 64) targets torch.randn(16, 32) loss_fn nn.MSELoss() print(f use_closed_form{use_closed_form} ) for step in range(5): optimizer.zero_grad() outputs model(inputs) loss loss_fn(outputs, targets) loss.backward() optimizer.step() # 检查更新后的权重矩阵是否近似列正交 orth_error check_orthogonality(model.weight) print(fstep{step} loss{loss.item():.6f} orth_error{orth_error:.6f}) if __name__ __main__: run_demo(use_closed_formFalse) print() run_demo(use_closed_formTrue)6.3 运行与验证在项目目录下执行python demo_train.py预期的观察结果是loss 会随训练步数下降。两种模式下都能完成更新流程。闭环模式的orth_error应该非常小接近机器精度因为 SVD 的投影结果满足正交约束到数值误差级别。Newton-Schulz 模式的orth_error取决于迭代次数通常比闭环模式大一点但训练效果不一定更差。如果运行报错优先检查PyTorch 版本是否支持torch.linalg.svd建议 1.13 以上。参数是否是二维矩阵。Muon 只对二维参数做正交投影更合理多维参数最好展平或按最后一维处理。是否混入了 requires_gradFalse 的参数代码里已经用p.grad is None做了跳过。7. 闭环更新与 Newton-Schulz 迭代的对比理解了两种实现方式后最自然的问题是到底该用哪一种对比维度Newton-Schulz 迭代SVD 闭环更新数学精度近似受迭代次数影响精确计算复杂度迭代次数 × 若干次矩阵乘法SVD 分解O(n p^2) 级别超参数迭代次数、迭代阶数、归一化方式几乎无额外超参数GPU 友好度高纯矩阵乘中等SVD 在大矩阵上可能成为瓶颈小矩阵/低秩场景迭代可能浪费更划算大规模稠密矩阵主流选择通常过慢可解释性偏黑盒数学上清晰我的判断是如果你在训练千亿参数大模型矩阵动辄上万维Newton-Schulz 迭代仍然是务实的选择。如果你在做 LoRA 或轻量微调低秩矩阵只有几十到几百维闭环更新的速度和精度都更好。如果你在写论文或做对照实验希望排除“正交化不精确”的影响闭环更新是最可靠的 baseline。如果你的训练脚本里已经有成熟的 Newton-Schulz 实现但每次调迭代次数都很痛苦不如试着把迭代次数设大一点或者直接切到闭环模式跑一次对照。8. 闭环更新的数学优雅与实际边界8.1 为什么这个公式值得关注从数学角度看闭环更新提供了一条关键路径优化问题不再依赖逐步逼近而是直接通过矩阵分解给出全局最优解。在 Stiefel Manifold 投影的语境下这个问题有非常干净的闭式解就是因为极分解与 SVD 之间存在一一对应的映射关系。这个结果属于数值线性代数里的经典结论但在优化器设计中被大规模使用是 Muon 带来的新关注点。8.2 闭环更新的限制闭环更新并不是银弹。它有一个隐性前提矩阵必须满列秩否则 SVD 时会出现零奇异值(M^T M)^(-1/2) 无法直接计算。在实际训练中动量矩阵的秩基本是满的但极端情况下如果梯度在某几个方向长期为零动量矩阵也可能退化。工程上可以考虑加一个小的正则项例如用 (M^T M ε I)^(-1/2) 代替。另一个限制是 SVD 在反向传播中的梯度。Muon 的投影步骤通常放在torch.no_grad()下投影本身不参与梯度计算所以不需要关心 SVD 的反向传播。如果你要把它嵌入到可微分模块里就要注意torch.linalg.svd的反向传播在奇异值重复或为零时可能出现数值不稳定。8.3 闭环更新的适用场景总结推荐使用LoRA 微调、小规模全参训练、正交性验证实验、数学对照实验。不推荐使用超大隐藏维度矩阵、每次 step 都做 SVD 且矩阵超过 2048x2048 的场景、对训练吞吐极度敏感的生产环境。如果你身处第二种场景但确实想要更高精度可以考虑混合策略先用 Newton-Schulz 迭代得到近似 Q再用一步闭环校正比如用 SVD 修正到最接近的正交矩阵。这在实践中比较少见但值得了解。9. 常见误区与排查思路围绕 Muon 和 Stiefel Manifold有些概念特别容易混淆。下面用表格说明问题现象、可能原因和排查方式。问题现象可能原因排查方式解决方案训练 loss 不下降反而震荡学习率过大或正交化破坏了动量方向调低 lr检查 loss 曲线对比去掉正交化步骤将 lr 调小 3-10 倍或改用 warmupNewton-Schulz 迭代不收敛没有对矩阵做范数归一化或迭代次数过大导致数值膨胀打印迭代过程中 X 的范数变化先按 Frobenius 范数归一化迭代结束后再缩放回来闭环更新计算很慢矩阵维度较大SVD 成为瓶颈统计 step 耗时对比 Newton-Schulz 模式换用迭代法或先用低秩投影压缩维度使用闭环更新后 orth_error 依然偏大SVD 数值误差或矩阵退化检查奇异值是否有接近 0 的值加小常数 ε 到 (M^T M)^(-1/2)或对奇异值做截断更关心训练效果而非正交精度正交化只是辅助手段不是目标比较两种模式在同一个任务上的 loss/acc优先选速度快的模式不必无限追求精确投影误把 Stiefel Manifold 当作单位球面概念混淆检查定义Stiefel 是列正交矩阵集合与向量范数约束不同对照维度n×p 矩阵约束是 n×p 而非 p×1其中第 4 条最容易踩坑。很多人以为 SVD 出来的 Q 一定完美满足 Q^T Q I但浮点运算下 U 和 Vh 本身会有数值误差而且如果矩阵接近退化误差会被放大。解决办法是在测试中直接输出 orth_error不要凭感觉判断。10. 最佳实践与工程建议10.1 在训练脚本中如何接入 Muon接入方式比想象中简单。只需要把原来的 AdamW 实例替换成 Muon 实例# 文件路径train.py from muon import Muon optimizer Muon( model.parameters(), lr0.02, momentum0.95, use_closed_formFalse, ns_iterations5, )替换时注意三点Muon 当前实现只对二维参数做正交投影因此卷积层、embedding 层等形状不匹配的参数会按全要素动量更新不做投影。实际中更精细的做法是只对指定的大矩阵启用投影通过参数分组实现。学习率策略建议从较小值开始观察 loss 曲线后逐步调大。Muon 对 lr 的敏感度低于 AdamW但并非完全免疫。如果和 AdamW 混用比如同一份代码在 Transformer 的不同模块使用不同优化器可以使用 PyTorch 的 param_groups 或自定义层名匹配。10.2 参数分组示例# 文件路径train_with_groups.py import torch.nn as nn from muon import Muon orthogonal_module_names {attn.q_proj, attn.k_proj, attn.v_proj} ortho_params [] other_params [] for name, param in model.named_parameters(): if param.requires_grad and any(key in name for key in orthogonal_module_names): ortho_params.append(param) elif param.requires_grad: other_params.append(param) optimizer Muon([ {params: ortho_params, lr: 0.01, use_closed_form: True}, {params: other_params, lr: 0.01, use_closed_form: False}, ])这种做法的好处是对维度较小的 attention 投影矩阵使用闭环投影对数量庞大的其他参数使用 Newton-Schulz 迭代或普通动量更新兼顾精度和吞吐。10.3 安全与可复现性提醒任何涉及优化器替换的实验都应该先在单卡小模型上跑通再扩展到多卡。不要在大规模训练中突然替换优化器避免训练曲线无法解释。SVD 闭环更新虽然数学精确但每次 SVD 的结果符号可能存在不确定性。如果追求绝对可复现建议固定随机种子并在日志中记录每次迭代的矩阵范数和正交误差。涉及生产环境变更时先保存 checkpoint确认新一轮 loss 和梯度范数都在合理区间再决定是否继续。10.4 性能优化建议如果决定使用闭环更新可以尝试以下优化方向只在预训练的前若干个 epoch 开启闭环投影后续切回 Newton-Schulz 迭代节省时间。使用torch.compile编译优化器 step 函数里的 SVD 调用某些 GPU 组合下能获得可感知的加速。对极低秩矩阵如 rank 小于 64可以显式计算 M^T M 的 Cholesky 分解或特征分解代替通用 SVD。11. 总结与后续学习方向这篇文章从 Muon 优化器的设计动机出发讲清楚了三个层次的问题第一Muon 为什么要把动量矩阵投影到 Stiefel Manifold第二Stiefel Manifold 上的投影本质上对应极分解而极分解的正交因子可以通过 SVD 一步得到这就是闭环更新的数学基础第三工程上闭环更新和 Newton-Schulz 迭代各有利弊低维低秩场景推荐闭环大维度稠密场景推荐迭代。想要进一步深入可以依次看这几个方向极分解与 SVD 的完整证明弄懂 Q U V^T 为什么是距离 M 最近的正交矩阵。Newton-Schulz 迭代的收敛性分析为什么三次和五次迭代式的系数是这样设计的。Stiefel Manifold 上的 Riemannian 优化从“投影到流形”升级到“沿流形测地线走”这是更高级的优化器设计方向。LoRA 与 Muon 的结合低秩矩阵做闭环投影可能是效果与速度兼顾的折中方案。如果你打算在自己的项目里实验 Muon建议先跑一遍 demo_train.py观察 loss 和 orth_error 的变化再决定选用哪种投影方式。把闭环更新当作精确 baseline把 Newton-Schulz 迭代当作速度优化手段这样对比会非常清晰。建议收藏这篇文章后面写代码或调参时可以直接翻到第 6 节的参考实现和第 9 节的排查表。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门