ARTICLE DETAIL

资讯详情

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

Muon优化器在Stiefel流形上的闭式更新:从迭代近似到精确投影

Muon优化器在Stiefel流形上的闭式更新:从迭代近似到精确投影 看到Muon on the Stiefel Manifold Admits an Exact Closed-Form Update这个主题时我第一反应是这说的不就是给 Muon 优化器换一个更精确的正交化步骤吗后来在一组小实验里把 Newton-Schulz 迭代换成一次 SVD 投影我才意识到这个“闭式更新”不只是实现细节上的优化它背后是两套完全不同的思考方式一个是把正交化当成数值问题一个是把正交化当成约束优化的解析结果。我更愿意用一句话概括这件事的价值Muon 在 Stiefel 流形上的更新并不只能靠迭代逼近它存在一个精确的闭式解。我们需要讨论的是这个解到底意味着什么工程上应该怎么落地以及它会不会取代我们习惯的 Newton-Schulz 方案。1. 当优化器开始约束矩阵事情就不只是“调学习率”这么简单1.1 为什么优化器要把参数限制在 Stiefel 流形上Stiefel 流形是所有列正交矩阵的集合。用数学语言说一个n × p矩阵X属于 Stiefel 流形当且仅当X^T X I_p。这里的I_p是p × p单位阵。也就是说矩阵的每一列都和其他列正交而且每一列的范数都是 1。这个约束在深度学习中经常出现。Transformer 结构里有些权重矩阵如果偏离正交性太远信息在多层之间流动时会发生协方差偏移训练会变得不稳定。循环神经网络里正交权重能帮助缓解梯度消失和梯度爆炸。子空间学习方法中要求一组基是正交的实际上就是要求矩阵落在 Stiefel 流形上。问题在于普通优化器并不认识这个流形。SGD、Adam 这些常见优化器的更新规则是纯粹在欧氏空间里工作的。每一步都是X ← X - lr * G其中G是梯度。这个更新没有任何机制保证更新后的X仍然满足X^T X I。一次更新可能偏离很小但训练动辄成千上万步偏差会累积。于是你会看到正交性误差越来越大网络的条件数越来越差最后表现为 Loss 曲线震荡、梯度范数异常、甚至训练直接发散。因此如果我们要优化一个必须落在 Stiefel 流形上的参数就不能只是“调学习率”。需要优化器本身具备一种能力在每步更新后把参数“拉”回流形上。Muon 这类优化器的价值就在这里。1.2 普通优化器直接更新会破坏正交性很多第一次接触 Stiefel 流形优化的人会有一个误解只要学习率足够小参数就一直在正交矩阵附近就算有点偏离也没关系。这个说法在非常短的训练里可能成立。但只要训练步数变长问题就会出现。假设当前矩阵X满足X^T X I梯度G包含了一些与当前子空间共振的分量那么X - lr * G的正交性偏差通常和lr * ||G||同阶。单看一步偏差很小但优化器没有纠偏机制偏差会带着方向性地持续累积。一旦正交性被破坏后续的梯度更新会在一个扭曲的坐标系里进行。这不是数值误差的问题而是优化轨迹本身偏离了约束流形。所以Muon 这类方案做的关键一步是在欧氏更新之后把结果投影回 Stiefel 流形。真正值得关注的不是“有没有投影”而是“投影这一步怎么算”。2. Muon 到底做了什么先把它放到 Stiefel 流形的语境里看2.1 Muon 优化器的基本直觉Muon 优化器并不是一个只调学习率的工具。它的核心思想是对于矩阵参数应该在保持某种正交结构的条件下更新而不是把矩阵当成普通向量来处理。常见的实现可以拆成两个阶段首先按照某种动量或梯度规则在欧氏空间里更新矩阵。然后把更新后的矩阵投影回 Stiefel 流形。第二个阶段是关键。它等价于回答一个问题有一个矩阵Z它可能已经偏离了流形我要找一个在流形上的矩阵Y让它尽可能接近Z。这就是一个投影问题。Muon 的很多实现在这个投影步骤里使用 Newton-Schulz 迭代而不是直接做 SVD。为什么因为在 GPU 上矩阵乘法非常快而 SVD 的数值实现相对重批量小矩阵时反而不划算。Newton-Schulz 迭代只需要矩阵乘法可以用很少的迭代次数逼近正交化结果。但这里有一个容易混淆的地方Newton-Schulz 只是计算投影的一种数值工具它不是 Muon 本身。Muon 背后的约束逻辑是“更新后要回到 Stiefel 流形”这个逻辑可以被不同的数值方法实现。2.2 为什么 Muon 会和 Newton-Schulz 绑定出现如果你读过一些 Muon 相关的开源实现会发现代码里经常出现一个函数对输入矩阵做若干次 Newton-Schulz 迭代然后输出一个接近正交的矩阵。这个函数在名字上往往写着orthogonalize或polar之类。Newton-Schulz 迭代的基本形式是Y (3/2) * X - (1/2) * X X.T X重复应用可以让X的正交部分被保留非正交部分逐渐衰减。这是一种非常简洁的迭代格式每一步只涉及矩阵乘法。对于大矩阵这种操作模式在 GPU 上非常高效所以它成为很多实现的默认选择。但“迭代”意味着需要选择迭代次数意味着结果是一个近似值。不同实现里迭代次数可能不同这会带来实验复现上的细微差别。如果你训练一个模型用 3 次 Newton-Schulz 迭代和用 5 次迭代最后得到的训练曲线可能几乎一样也可能有微妙差异。于是问题就来了如果 Muon 在 Stiefel 流形上的更新本身存在一个精确闭式解我们为什么还要忍受迭代带来的近似和不确定性这正是这次主题要回答的问题。3. 闭式更新不是玄学从极分解到 SVD一条清晰的推导路径3.1 一个约束优化问题的标准答案我们先把问题抽象成数学形式。假设当前参数是X_k梯度是G。常见的 Muon 更新先做一步欧氏移动Z X_k - lr * G之后我们希望找到一个列正交矩阵Y使得Y尽量接近Z。用 Frobenius 范数来衡量“接近”就得到投影问题min_Y || Y - Z ||_F^2 subject to Y^T Y I_p这个问题有非常成熟的解。设Z的奇异值分解为Z U Σ V^T那么最优投影就是Y_star U V^T如果Z是方阵且可逆这个结果也等于极分解里取正交因子。也就是说U V^T是一个精确的闭式更新不需要通过迭代去逼近。这一步的推导并不复杂但它改变了我们对 Muon 更新的理解所谓“正交化”其实就是在解一个最小化投影距离的问题。而这个问题的最优解可以直接写出来。3.2 为什么说这恰好是 Muon 的精确更新标题里说“admits an exact closed-form update”核心信息就在这里Muon 在 Stiefel 流形上的投影步骤理论上没有必要依赖 Newton-Schulz 迭代。只要先做欧氏更新再用 SVD 或极分解求出正交因子就能得到精确的流形投影结果。这个说法把 Muon 的更新路径从“近似计算”变成了“解析表达”。工程上这意味着正交化误差可以降到机器精度级别。不再需要调节 Newton-Schulz 迭代次数。实验结果可复现性更强。每一步更新都对应一个明确的最优解。但请注意这个“精确”指的是投影步骤的精确不是整个训练过程的最优。它没有改变优化器的动量策略没有改变学习率调度也没有改变损失函数。它只是把“如何回到流形”这件事变成了一个可验证的精确解。3.3 一个容易混淆的地方精确投影不等于黎曼梯度下降很多读者看到“Stiefel 流形 优化器”时会立刻联想到黎曼优化。黎曼优化是在流形的切空间上定义梯度然后通过指数映射或收缩映射把更新映射回流形。这是一套完整且优雅的数学框架。但 Muon 这类“先欧氏更新再投影”的做法和严格的黎曼梯度下降并不等价。它更像投影梯度法在欧氏空间走一步再投影回约束集。闭式解给出的是“投影算子”的精确结果而不是黎曼梯度的精确结果。这两者的训练动态可能会有差异。不能因为在 Stiefel 流形上找到了闭式投影解就认为 Muon 已经是标准的黎曼优化器。4. 工程上怎么用从最小验证到接入训练循环4.1 一个最小可运行的投影实现如果你已经在用 PyTorch想验证闭式更新代码可以非常短。核心是用torch.linalg.svd实现投影import torch def project_to_stiefel(X): U, _, Vh torch.linalg.svd(X) return U Vh这段代码接受一个n × p矩阵返回一个n × p矩阵。只要输入矩阵的秩不低于p输出就会满足近似列正交。注意SVD 返回的U是n × n正方形矩阵Vh是p × p正方形矩阵。U Vh是n × p刚好覆盖整个 Stiefel 流形。如果你只想要“窄版”的左奇异向量可以用full_matricesFalse但上面的写法在大多数情况下足够。然后一个简化版 Muon 更新可以写成这样def muon_step(X, grad, lr0.01, momentum0.9, bufNone): if buf is None: buf torch.zeros_like(grad) buf momentum * buf grad Z X - lr * buf return project_to_stiefel(Z), buf这个实现只展示结构不代表完整优化器。实际接入训练时你需要按参数管理优化器的状态字典处理梯度裁剪、权重衰减、学习率分组等问题。4.2 验证正交性误差替换实现之前先不要跑大规模训练。先做一个最小实验生成一个随机矩阵分别用 Newton-Schulz 迭代和 SVD 投影看看两者输出的正交性误差有多大。def orth_error(X): p X.shape[1] I torch.eye(p, deviceX.device) return torch.linalg.norm(X.T X - I) / p对同一个ZSVD 投影的正交性误差通常能到1e-6甚至更低。Newton-Schulz 迭代的误差则取决于迭代次数。迭代 3 次和迭代 10 次结果会有明显差距。这个实验能直观告诉你闭式解在数值上确实更精确。但“更精确”不等同于“更好”。你还需要看训练曲线。4.3 从单步验证到批量训练建议的路径是先在小规模随机矩阵上验证投影结果。在单个训练 batch 上对比新旧实现的 loss。在一个很小的 Transformer 或 MLP 上跑几百步观察 loss 曲线和正交性误差。最后才考虑在正式模型里替换实现。我遇到过一种情况闭式版本的正交误差更低但训练 loss 反而比 Newton-Schulz 版本略高。原因可能不是投影精度而是 SVD 在每步引入的方向符号变化或额外噪声。所以必须用训练结果来决策而不是只看正交误差。5. 误差、稳定性和排查链路闭式解不是免死金牌5.1 闭式更新也有坑SVD 投影理论上漂亮工程上却有一些实际问题。第一SVD 的反向传播会消耗更多内存。PyTorch 对 SVD 的自动微分是支持的但中间变量比矩阵乘法多。在训练大型模型时这可能导致显存占用上升。第二半精度矩阵下的 SVD 不够稳定。很多训练框架使用fp16或bf16如果直接对这些精度的矩阵做 SVD可能会遇到收敛失败、NaN 或正交性误差不降的情况。第三SVD 的结果在奇异值有重根或接近重根时左右奇异向量的符号可能不唯一。也就是说U V^T可能在某些步突然改变方向。这种符号不稳定性会在自动微分中造成梯度跳跃进而影响训练稳定性。第四SVD 的计算代价不是固定的。对于很小的矩阵它可能比 Newton-Schulz 还快对于几千乘几千的矩阵它可能明显更慢。这需要实际测不能凭空判断。5.2 如果训练出了问题按什么顺序排查闭式更新并不能让你避开常见训练问题。我建议按这个顺序排查先确认输入数据没有异常。检查Z中是否有 NaN 或 Inf。如果原版 Muon 正常闭式版本出现 NaN先看 SVD 输入。再看学习率。闭式投影本身不会造成梯度爆炸但它不负责让更新量变小。学习率过大时任何投影策略都救不回来。再看正交性误差。如果闭式版本的正交性误差也很大说明 SVD 没有正常返回或者矩阵在截断后秩不足。再看训练曲线。Loss 抖动通常不是投影精度导致的更多是动量、学习率调度和权重重分配的问题。对比两个实现。固定随机种子让 Newton-Schulz 版本和闭式版本跑同一组实验看差异出现在第几步。如果差异从一开始就很大说明算法行为本身不同如果差异在几十步后才出现可能是累积误差导致。检查确定性。SVD 在 CUDA 上默认可能具有不确定性。如果实验需要可复现开启torch.use_deterministic_algorithms(True)或者固定设备、关闭 TF32。5.3 一个折中策略什么时候用闭式解什么时候继续用迭代从工程经验看不需要在所有场景里二选一。可以设一个阈值矩阵维度较小例如p 512时用 SVD 闭式投影。矩阵维度较大例如p 512时继续用 Newton-Schulz 迭代。训练前期可以用闭式解验证网络结构是正确的后期切换到迭代来省显存。如果你正在做优化器方向的实验闭式解更合理因为它省掉了迭代次数这个超参数。如果你在做大规模预训练Newton-Schulz 仍然是更省资源的投影策略。6. 适用边界别把“闭式”当作“免费”6.1 适合什么人和什么场景闭式更新最适合下面这些情况你在研究优化器本身需要精确控制每一步的更新路径。你的模型里有少量矩阵需要严格保持正交性例如某些归一化层或嵌入层。你的矩阵维度不大SVD 计算成本可以接受。你希望实验稳定可复现不愿意让 Newton-Schulz 迭代次数成为另一个超参数。对小规模实验和原型验证来说闭式更新是一个非常值得尝试的替代方案。它让思路更清晰每次更新后参数确实在流形上。6.2 不适合什么场景闭式更新不适合作为所有矩阵参数的默认替换。原因很简单SVD 的计算成本和稳定性问题可能抵消掉它带来的精度优势。超大模型的全参数训练是最典型的不适用场景。模型参数量达到十亿或百亿规模时每一步都要对所有矩阵做 SVD代价非常高。即使你只约束部分矩阵SVD 的反向传播也可能让显存和耗时变得不可控。另外如果你的训练框架依赖半精度和并行策略闭式更新需要仔细适配。要么把 SVD 放在高精度下计算要么在低精度下做额外保护。否则闭式解带来的“精确”反而可能变成“精确地传播数值噪声”。6.3 回到主判断它改变了什么没有改变什么精确闭式更新并没有让 Muon 变成一种全新的优化器。它改变的只是“如何把参数投影回 Stiefel 流形”这一步的数值实现方式。这个改变让投影过程更精确、更可解释、更容易复现也让我们能更清楚地看到 Muon 更新背后的约束优化本质。但它没有改变的是优化器仍然需要合理的学习率仍然需要处理动量、权重衰减和训练动态仍然不能保证每一个问题都收敛。闭式解是一个更好的工具不是一个神奇的答案。如果让我给一个行动建议我会说先别急着全局替换。打开编辑器写一个 20 行的脚本生成几个随机矩阵在 CPU 或 GPU 上比较 SVD 投影和 Newton-Schulz 投影的正交性误差和耗时。把这个小实验做完你对“闭式更新”的理解会比读十篇文章都深。然后再决定你的模型到底需不需要这个精确解。
返回列表