ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Muon优化器的几何本质:Stiefel流形上的精确闭式更新

Muon优化器的几何本质:Stiefel流形上的精确闭式更新 Muon 优化器从 2025 年初开始持续引发关注。它不是又一个 Learning Rate Scheduler 的变体而是把“对矩阵做正交化”这个几何操作重新提升到了优化器设计的核心位置。很多人在复现 Muon 时注意力都放在了 Newton-Schulz 迭代上认为这是 Muon 收敛效果好的关键。但如果把问题放到 Stiefel 流形上重新看一遍会得到一个更本质、也更容易被忽略的结论Muon 在 Stiefel 流形上需要的那个投影其实存在一个精确的闭式解理论上根本不需要靠迭代去逼近。这篇文章想解决的不只是“Muon 怎么用”而是三个更实际的问题Stiefel 流形为什么会出现在优化器里Muon 的正交化步骤到底在算什么以及当你需要精确正交约束时怎样用 SVD 给出一个 closed-form 的替代方案。读完这篇文章你能理解 Muon 的几何本质也能在 PyTorch 中亲手把迭代版和闭式版跑通并知道它们各自适合什么场景。1. Muon 为什么和 Stiefel 流形扯上了关系先回顾一下 Muon 的核心思想。Muon 的出发点非常朴素SGD 加上动量的更新方向存在“各向异性”某些方向被过度放大导致隐藏层的更新质量下降。Adam 用逐元素二阶矩归一化来解决这个问题但代价是引入了大量额外状态并且对学习率的敏感度很高。Muon 选择的路径完全不同——它把梯度或动量矩阵按块划分然后做一次正交化让更新方向在矩阵的奇异值尺度上变得更均匀。这里的关键概念是 blockwise orthogonalization。Muon 每次更新时会取出动量缓冲矩阵 M对它做一次正交化得到 O然后执行参数更新 W ← W - lr * O。也就是说Muon 默认你关心的是更新方向的“整体形状”而不是单个元素的二阶统计量。那么 Stiefel 流形为什么会出现因为“正交化”这四个字背后对应的几何对象就是 Stiefel 流形。一个 n × k 矩阵 X如果满足 XᵀX Iₖ那么它就落在 Stiefel 流形 Vₖ(ℝⁿ) 上。Muon 对动量矩阵做正交化本质上就是把这个矩阵“拉”到 Stiefel 流形上。它不是优化器作者凭空发明的操作而是所有带正交约束的矩阵都必须经过的几何投影。读者经常遇到的困惑是一边看 Muon 的伪代码一边看 Stiefel 流形的数学定义觉得两者毫无关系。实际上Muon 的 orth 操作就是流形投影的一次工程近似。理解了这一点你就能理解为什么标题说“Muon on the Stiefel Manifold Admits an Exact Closed-Form Update”——它的意思是这个投影操作本身不需要迭代近似存在解析公式。2. Stiefel 流形与正交约束概念先讲清楚Stiefel 流形可以理解为“所有满足列正交的矩阵”的集合。它的数学定义并不难[ V_k(\mathbb{R}^n) { X \in \mathbb{R}^{n \times k} \mid X^T X I_k } ]直观来说这个集合里的每个元素都是一组相互正交、长度归一化的列向量。当你训练一个带正交约束的神经网络层时你希望这一层的权重矩阵始终停留在 Stiefel 流形上不偏离到附近的普通欧氏空间去。几个常见特例有助于建立直觉k 1 时V₁(ℝⁿ) 是 n 维单位球面。k n 时Vₙ(ℝⁿ) 是正交群 O(n)也就是所有正交方阵的集合。介于两者之间时Vₖ(ℝⁿ) 是“半正交矩阵”的集合。流形最重要的结构是切空间。在欧氏空间里一个点的切空间就是整个空间本身在 Stiefel 流形上切空间是受限的。如果你站在流形上的点 X 处想要移动一个微小步长 Δ这个 Δ 必须满足[ X^T \Delta \Delta^T X 0 ]也就是说XᵀΔ 必须是反对称矩阵。这个概念对理解 Muon 很关键因为它意味着普通的梯度下降更新 W ← W - lr * g 会立刻破坏正交约束正确的更新要么走投影收缩projected retraction路线要么显式地把更新方向限制在切空间里。下面这张表可以帮你快速区分几个容易混的概念对象定义常见用途欧氏空间 ℝⁿˣᵏ所有 n×k 矩阵无约束线性层参数单位球面满足 ‖x‖ 1 的向量归一化嵌入、方向向量正交群 O(n)满足 XᵀX I 的方阵正交初值、旋转矩阵Stiefel 流形 Vₖ(ℝⁿ)满足 XᵀX Iₖ 的矩阵正交权重、子空间基、多通道滤波器为什么 Stiefel 流形比正交群更常用因为神经网络权重矩阵几乎都是矩形。卷积核、注意力矩阵、嵌入矩阵都不是方阵。Stiefel 流形允许列数 k 小于行数 n这正好匹配真实模型的需求。3. Muon 的正交化到底在算什么标准的 Muon 更新流程可以写成下面的伪代码形式。为了方便讨论这里只展示单矩阵版本实际实现中会对模型每个二维参数块分别执行。# Muon 的标准更新逻辑单块 # M动量缓冲矩阵n×k # g当前梯度n×k # lr学习率 # momentum动量系数 M momentum * M (1 - momentum) * g # 1. 动量更新 O orth(M) # 2. 正交化 W W - lr * O # 3. 参数更新这三个步骤里第 1 步是常规动量第 3 步是常规下降。真正的 Muon 特色在第 2 步。orth(M) 的目标是给定任意矩阵 M返回一个正交矩阵 O使得 O 尽可能“接近” M。这是一个典型的正交 Procrustes 问题[ \min_{O \in V_k(\mathbb{R}^n)} | M - O |_F^2 ]在数学上这个问题的闭式解非常干净。对 M 做奇异值分解[ M U \Sigma V^T ]那么最优的 O 就是[ O U V^T ]这就是标题里说的 exact closed-form update。它一步到位没有任何迭代也不需要设置牛顿迭代的步数。但 Muon 的原始实现并没有直接使用这个闭式解。它用的是 Newton-Schulz 迭代。为什么关键在计算效率。SVD 的精确分解在 GPU 上并不便宜特别是当模型参数量很大时每一轮训练都要对大量矩阵做 SVD计算开销会显著上升。而 Newton-Schulz 迭代只需要若干次矩阵乘法对 GPU 的并行架构更友好虽然它得到的是近似解但经过 5 次左右迭代后正交误差通常已经足够小不影响训练收敛。用一句话概括Muon 实际想做的操作是 Stiefel 流形上的最近点投影闭式解存在于欧氏空间Newton-Schulz 只是在工程上以更低的算力成本逼近这个解。4. 核心话题为什么存在精确的闭式更新现在进入本文的核心部分。为什么说 Muon on the Stiefel Manifold 存在精确的闭式更新关键在于极分解和正交 Procrustes 问题的关系。任意矩阵 M 都可以极分解为[ M Q P ]其中 Q 是一个满足 QᵀQ I 的“正交因子”P 是半正定对称矩阵。在机器学习里Q 通常被称为 M 的“酉部分”。如果 M 满秩这个分解是唯一的。这天然就是 Stiefel 流形上的一个元素。而奇异值分解直接给出了极分解的构造方法。对 M UΣVᵀ可以写成[ M (U V^T)(V \Sigma V^T) ]这里 UVᵀ 就是 QVΣVᵀ 就是 P。换句话说SVD 的左右奇异向量相乘得到的就是极分解中的正交因子。这正是闭式更新公式 O UVᵀ 的来源。为了说明这一点可以做一个代数验证。假设你有一个不含 SVD 闭式解的近似算法比如随机初始化一个正交矩阵 R₀然后用梯度下降去最小化 ‖M - R‖²那么每一步迭代都只能朝最优解缓慢逼近。而 SVD 是一步到位的解析解不是通过优化过程逼近得到的。这就是“闭式”closed-form这个说法的数学含义。把这个结论放到 Muon 的场景里当我们说“Muon on the Stiefel Manifold admits an exact closed-form update”时我们实际是在说Muon 的正交化步骤不再需要 Newton-Schulz 迭代而是直接通过 SVD 的 UVᵀ 一步完成。从理论角度看这是最干净的版本从工程角度看SVD 版本的更新路径与 Newton-Schulz 逼近的路径不同——前者严格收敛到最优投影后者只有迭代次数趋向无穷时才收敛到同一结果。有一个常见误解需要澄清很多人以为 Newton-Schulz 迭代是 Muon 数学推导的一部分因此不可替换。实际不是。Muon 论文选择 Newton-Schulz是因为它满足“GPU 友好、自动微分友好”这两个工程要求而不是因为它唯一的理论投影方式。如果你换用 SVD UVᵀ你得到的是同一类优化器的“精确投影版本”。5. 闭式更新 vs Newton-Schulz工程上的真实差异从数学上讲SVD UVᵀ 是准确解Newton-Schulz 是逼近解。但从工程上讲选择并不是“越精确越好”而是要看计算成本、数值稳定性、反向传播效率和扩展性。第一个差异是计算复杂度。对 n×k 矩阵做截断 SVD 的复杂度大约是 O(nk²)。对于 k 很小的矩阵比如逐块、逐层参数这个开销可以接受但对于 k 较大的稠密矩阵SVD 会明显变慢。Newton-Schulz 迭代每一步都是矩阵乘法复杂度 O(nk²)做 5 步就是 5O(nk²)理论上和 SVD 同量级但常数更小而且矩阵乘法在 GPU 上的峰值效率远高于 SVD 这样的分解算法。第二个差异是精确度。Newton-Schulz 的误差取决于迭代步数和初始矩阵条件数。如果初始矩阵的奇异值偏离 1 太远5 步迭代后正交误差可能仍然达到 1e-3 量级。SVD 闭式解的正交误差只受浮点精度限制通常可以达到 1e-6 以下。对于严格要求正交约束的任务比如某些正交循环神经网络、正交卷积、特征分解模块SVD 版本更可靠。第三个差异是反向传播。SVD 的反向传播需要在最优解 U、V、Σ 附近求解一个线性系统计算代价高且存在奇异值接近 0 时导数爆炸的风险。Newton-Schulz 完全由矩阵乘法和加法组成反向传播天然稳定这也是它被大量深度学习框架选中的原因。我用下面的表格做一个快速比较维度SVD 闭式更新Newton-Schulz 迭代准确性精确到浮点精度取决于迭代步数通常 1e-3 到 1e-6GPU 效率分解算法常数较大纯矩阵乘法更适合大规模并行反向传播需要求解线性系统复杂自动微分天然稳定数值稳定性奇异值接近 0 时有风险奇异值偏离 1 时收敛变慢适用场景小矩阵、强正交约束、离线验证大规模训练、Muon 默认路线有一个工程细节值得注意SVD 反向传播的不稳定性可以通过一个技巧规避——训练时用 Newton-Schulz 版本前向计算验证时用 SVD 版本检查正交误差。两种版本可以同时存在于同一份代码里互不冲突。6. 完整代码实现从闭式正交化到 Muon这一节给出可以直接运行的 PyTorch 代码。环境要求是 PyTorch 2.xPython 3.9 以上GPU 可选。核心依赖只有 torch重点看 torch.linalg.svd 的用法和 Newton-Schulz 迭代的实现。6.1 SVD 闭式正交化函数import torch def svd_closed_form_orthogonalize(M: torch.Tensor) - torch.Tensor: 通过 SVD 极分解将任意矩阵 M 投影到 Stiefel 流形 V_k(R^n)。 返回 O U V^T满足 O^T O ≈ I_k。 M: (..., n, k) 或 (n, k) U, _, Vt torch.linalg.svd(M, full_matricesFalse) return U Vt这段代码就是闭式更新的核心。对任意形状的矩阵调用它得到的结果就是正交 Procrustes 问题的最优解。注意 full_matricesFalse 很关键它避免计算多余的 U 列。6.2 标准 Muon 的 Newton-Schulz 正交化Muon 论文中的 Newton-Schulz 采用三阶迭代格式常用系数是 a3.4445、b-4.7750、c2.0315。这组系数经过专门设计可以让迭代快速收敛到正交矩阵。def newton_schulz_orthogonalize(M: torch.Tensor, steps: int 5) - torch.Tensor: 近似正交 Procrustes 的 Newton-Schulz 迭代。 输入 M 的奇异值应尽量接近 1否则建议先做归一化或增加 steps。 a, b, c 3.4445, -4.7750, 2.0315 X M for _ in range(steps): XTX X X.mT X a * X b * (XTX X) c * (XTX (XTX X)) return X每次迭代只包含矩阵乘法和线性组合非常适合在 GPU 上批量执行。训练中一般取 steps5既保证正交误差可控又不会带来过多开销。6.3 支持两种模式的 Muon 优化器下面实现一个最小可用的 Muon 优化器支持 svd 和 newton_schulz 两种正交化模式。这里的重点是结构演示加入自定义 optimizer 时需要按实际框架的规范调整。class Muon(torch.optim.Optimizer): def __init__(self, params, lr0.02, momentum0.95, modenewton_schulz, ns_steps5): defaults dict(lrlr, momentummomentum, modemode, ns_stepsns_steps) super().__init__(params, defaults) torch.no_grad() def step(self, closureNone): loss None if closure is not None: loss closure() for group in self.param_groups: lr group[lr] momentum group[momentum] mode group[mode] for p in group[params]: if p.grad is None: continue g 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_(g) # 统一按二维矩阵处理 if p.dim() 2: if mode svd: O svd_closed_form_orthogonalize(buf) else: O newton_schulz_orthogonalize(buf, stepsgroup[ns_steps]) else: # 非二维参数保持简单动量更新不做正交化 O buf p.add_(O, alpha-lr) return loss这个实现把两种正交化方式放在同一个优化器里方便打对比实验。实际工程中你可以在初始化时通过 mode 参数切换。6.4 运行验证比较正交误差写一个独立脚本用同一随机矩阵分别调用两种方法打印正交误差。def check_orthogonality(X: torch.Tensor) - float: 计算 max |X^T X - I|值越小说明越接近 Stiefel 流形。 k X.shape[-1] I torch.eye(k, dtypeX.dtype, deviceX.device) return (X.mT X - I).abs().max().item() torch.manual_seed(42) M torch.randn(64, 16) # 随机矩阵不易直接满足正交约束 R_svd svd_closed_form_orthogonalize(M) R_ns newton_schulz_orthogonalize(M, steps5) print(SVD closed-form orth error:, check_orthogonality(R_svd)) print(Newton-Schulz orth error: , check_orthogonality(R_ns)) print(Optimality gap tr(R^T M): , torch.trace(R_svd.mT M).item(), torch.trace(R_ns.mT M).item())你可以直接运行这段代码。预期结果是SVD 版本的正交误差在 1e-6 甚至更小Newton-Schulz 版本在 1e-3 到 1e-5 之间。如果两者接近说明你的矩阵初始条件较好如果 Newton-Schulz 误差偏大可以把 steps 从 5 提高到 8 或 10但训练速度会下降。7. 常见问题与排查思路在实际使用 SVD 闭式更新或 Muon 时会遇到一些典型问题。把这些排查经验提前整理出来可以节省大量调试时间。问题现象可能原因排查方式解决方案训练 Loss 不降甚至发散学习率过大或正交化后更新方向被过度压缩打印每层更新范数对比 SVD 与 NS 的更新幅度降低学习率或对梯度做裁剪SVD 反向传播报错或 NaN矩阵接近亏秩奇异值非常小检查奇异值分布 torch.linalg.svdvals加入微小 epsilon 正则或改用 Newton-SchulzNewton-Schulz 正交误差始终偏大矩阵奇异值偏离 1 太远迭代步数不足打印 M 的奇异值范围增加 steps或先做 Frobenius 范数归一化SVD 版训练速度明显变慢参数矩阵 k 过大SVD 开销高用 torch.profiler 查看算子耗时对 k 较大的矩阵回退到 Newton-Schulz梯度裁剪后效果变差正交化发生在裁剪前两者顺序不匹配检查裁剪位置与 orth 位置统一先裁剪、再做动量更新与正交化自定义层中正交误差缓慢累积只在初始化时做一次正交化后续更新未约束检查前向是否包含正交投影算子在每个前向中显式调用投影或使用参数化层一个容易忽略的坑是如果把 SVD 版 Muon 用在 k 很小但 n 很大的矩阵上SVD 本身的复杂度并不高真正昂贵的是 PyTorch autograd 对 SVD 线性系统的反向传播。这种情况下使用自定义 forward、将 SVD 替换为两次 Householder 反射的闭式实现可以进一步减少开销但需要更复杂的自定义算子。另外在分布式训练中Muon 的正交化是逐参数块执行的不涉及跨卡同步因此数据和模型并行不影响正交化逻辑。唯一需要注意的是如果你在优化器 step 中加入了 SVD并且使用了梯度累积那么 momentum buffer 的更新频率必须和梯度累积步数对齐否则动量语义会发生偏移。8. 最佳实践与工程建议基于前面的分析可以给出几条明确的工程建议帮助你在实际项目中选型。第一默认场景建议继续使用 Newton-Schulz 版本。Muon 设计这套迭代方案是为了在 GPU 上获得最佳吞吐量。只要你的训练任务不强制要求“精确正交”NS 版本通常更划算。一个推荐的起点是 steps5观察验证集 Loss 后可以微调。第二当任务要求严格的正交约束时优先使用 SVD 闭式更新。典型例子包括正交循环网络、正交卷积核、子空间基学习、基于正交矩阵的注意力投影层。这些场景里残差正交误差不仅影响数值稳定性还可能直接影响模型表达能力。第三将两种实现做成可配置项而不是二选一。在模型较小、矩阵规模不夸张的开发阶段可以用 SVD 模式跑通全流程因为它的结果是精确的方便你判断问题是否出在正交化环节。模型变大后再切回 NS 模式做大规模训练。第四对参数矩阵做二维 reshape 时要确保 reshape 后的语义清晰。Muon 通常把权重按最后一维对齐做正交化因此一个 4 维卷积张量 (out, in, h, w) 需要先 reshape 成 (out, inhw) 再正交化。这个 reshape 前后的梯度对应关系必须一致否则正交化方向会与参数布局错位。第五学习率设置建议参考 SGD 的经验而不是 Adam。Muon 的正交化让更新方向的尺度被约束在较低水平因此它对学习率的敏感度介于 SGD 与 Adam 之间。实践上可以从 0.01 到 0.05 起步用两三次小规模实验确定合理范围。第六配合 spectral normalization 使用时注意顺序。Spectral Normalization 本身也在限制权重矩阵的谱范数与 Muon 的正交化目标部分重叠。如果两者叠加建议在实验组里分别记录基线、单独 Muon、单独 SN、MuonSN 四组效果避免正交互补变相抑制模型表达能力。最后要认识到“闭式”不等于“免费”。SVD 闭式更新虽然在数学上简洁但它的反向传播复杂度和数值敏感性并不低。真正能体现闭式更新优势的场景是参数矩阵规模可控、正交性要求高、并且你可以接受略微降低吞吐量的任务。9. 总结与下一步实践方向这篇文章想讲清楚的核心事实是Muon 的正交化步骤不是黑盒魔法它的数学本质是 Stiefel 流形上的最近点投影而该投影在理论上存在一个精确的闭式解即通过对动量矩阵做 SVD 后取 UVᵀ 得到。Newton-Schulz 迭代只是工程上更便宜的近似方案它并不能推翻闭式解的存在性。如果你正在复现 Muon建议动手做三件事。第一用本文的代码在同一组随机矩阵上对比 SVD 版和 NS 版的正交误差建立对两种实现的直觉。第二在一个小模型上分别用两种模式训练对比 Loss 曲线和训练吞吐从而判断你的场景更适合哪种正交化方式。第三尝试把 SVD 闭式更新用在带正交约束的模型层中比如正交循环单元或正交注意力投影观察它是否比逐次投影更稳定。下一步值得深入的方向有两个一是阅读 Muon 原始论文中关于 blockwise orthogonalization 的推导理解正交化的尺寸选择如何影响优化轨迹二是研究极分解在深度学习的不同入口比如权重归一化、白化与去相关、网络正则化它们背后共享同一套矩阵分解数学。理解 Stiefel 流形上的闭式更新不只是为了调好一个优化器更是为了在面对任何“需要保持正交结构”的问题时知道该从哪里找到解析解。
返回列表