Muon 优化器最近在深度学习训练里讨论度不低尤其是矩阵形状参数的更新场景。它最关键的做法是让动量重新变成正交列矩阵而这个“正交列矩阵”的集合就是 Stiefel 流形。标题里的“Muon on the Stiefel Manifold Admits an Exact Closed-Form Update”直接点题在 Stiefel 流形上Muon 的每一步更新可以写成精确的闭式更新不需要每步都去跑 QR 或 SVD。这篇文章适合两类读者一类是想读懂 Muon 实现细节、准备把它替换进自己训练脚本的人另一类是刚开始接触流形优化、想知道“闭式更新”到底比 QR/SVD 多解决什么问题的同学。下面按背景、意义、实现、验证、边界、排查的顺序拆尽量把每一步的“为什么”也讲清楚。1. 先理解 Muon 为什么必须面对流形约束1.1 Muon 到底改了普通动量的哪一步普通动量的写法非常直接momentum mu * momentum grad weight - lr * momentumAdam 会在上面基础上再除以二阶矩估计并做偏置修正。Muon 的不同点在于当权重形状是[out, in]这种矩阵时它不直接拿原始 momentum 去做权重减法而是先把 momentum 重新处理成“列与列正交”的矩阵再把它当成更新方向。为什么这样做因为矩阵形状的参数学到的特征方向是有几何结构的。如果 momentum 里混入了让某些列向量“压扁”或“倾斜”的成分下一步更新就会把整个参数矩阵带偏。先做一次正交化相当于把更新方向洗干净让它反映的是一组互相独立的单位方向而不是被梯度噪声污染的任意矩阵。所以 Muon 的核心动作不是“普通动量 夹断”而是“普通动量合成 流形重投影”。这一步重投影才是它和 Adam 在几何上最本质的区别。1.2 正交列矩阵的集合就是 Stiefel 流形Stiefel 流形这个名字听起来抽象定义其实很朴素所有满足X^T X I_p的n x p矩阵组成的集合。换句话说取一个 n 维空间在里面挑 p 个互相垂直的单位向量把它们排成一个矩阵这个矩阵就在 Stiefel 流形上。Muon 更新后的 momentum恰好就是一个这样的矩阵。它不是一个普通向量而是一个“正交列矩阵”。如果动量初始化时满足正交条件那么后续每一步更新都应该尽量让这个条件继续保持。这里的关键在于这个集合不是欧氏空间而是一个流形。欧氏空间里x y永远还在同一个空间里但在 Stiefel 流形上两个正交列矩阵相加之后基本不会再满足正交条件。你从流形上取一个点往任意方向走一步大概率就离开流形了。所以 Muon 每一步都面临同一个问题如何在尽量保持正交列结构的前提下让 momentum 按照梯度给出的方向前进。1.3 流形上的更新为什么不能直接套向量加法最容易犯的错误是把 momentum 当成普通矩阵直接做momentum mu * momentum grad这一步做完动量通常已经不再满足M^T M I。如果你继续拿这个非正交矩阵去更新权重Muon 里“正交化动量”的核心意义就丢了。因此必须多一步把合成后的动量重新拉回流形。这一步在工程上最常见的是 QR 分解、SVD 或者施密特正交化。把动量拉回流形之后再乘上学习率去更新权重。但问题也随之出现拉回流形的方式并不唯一不同的正交化路径会给出不同的更新结果。这就像从流形上的 A 点走到邻近的 B 点你可以沿着流形的“测地线”走也可以先在旁边转一圈再挤回流形。路径不同更新的几何意义也不同。闭式更新的价值正是在于它给出了一条确定而且精确的路径。2. “精确闭式更新”到底解决了什么问题2.1 常见的正交化路线QR、SVD 和投影目前多数 Muon 实现里正交化这一步用的是 QR 分解momentum mu * momentum grad momentum, _ torch.linalg.qr(momentum)QR 有个特点它把矩阵分解成一个正交矩阵和一个上三角矩阵取前面的正交矩阵作为更新后的动量。实现简单效果通常也不差。但 QR 存在几个值得注意的工程问题对列向量逐列做 Gram-Schmidt 时数值敏感性会随着列数增加而累积。QR 的正交矩阵在符号上不唯一前后两步的列符号可能翻转。QR 和 SVD 都是一次“分解 重新拼装”的过程它保证了结果在流形上但没有回答“从几何上看这一步到底是不是合理的流形更新”。SVD 比 QR 更稳定但计算量更大。还有一类做法是把动量投影到流形上再做一次迭代修正这需要额外循环。这些方法都能让动量回到流形上但本质上都是“先把动量弄走再强行拉回来”。2.2 闭式更新到底指什么闭式更新closed-form update的意思是给定当前流形点M和梯度G通过一个明确的表达式直接算出下一步M而且M精确落在 Stiefel 流形上。它不依赖内层迭代不需要每步都做 QR、SVD也不需要设置投影循环的终止条件。这种闭式更新通常长成这样先由当前流形点和梯度构造一个反对称矩阵skew-symmetric matrix再对这个反对称矩阵做矩阵指数matrix exponential最后用指数结果去乘当前流形点。M_new M exp( -lr * A )其中A是一个反对称矩阵满足A^T -A。反对称矩阵的指数有一个很好的性质算出来的结果一定是正交矩阵。一个正交矩阵乘以一个正交列矩阵结果仍然是正交列矩阵。也就是说只要你从满足M^T M I的M出发按这个结构更新下一步在形式上天然不会离开流形。闭式更新不意味着完全没有浮点误差。矩阵指数本身也是数值计算但它的误差来源明确、可控不依赖内层迭代的收敛情况。这是它和 QR/SVD 路线最大的区别。2.3 这件事为什么值得关注闭式更新的意义可以分三层看。第一层是计算层面。训练循环里参数多、步数多每一步都做 QR 或 SVD 是有额外开销的。用闭式更新替代后正交化过程变成一次矩阵指数计算省掉一轮分解和重新组合。第二层是数值稳定层面。QR 的列符号翻转、Gram-Schmidt 的误差累积、SVD 在接近退化矩阵时的行为都会影响训练过程的复现性。闭式更新把不确定性收敛到一个固定表达式上数值行为更可预期。第三层是几何层面。闭式更新通常对应着精确的黎曼梯度步或测地线步它能告诉你“这一步在流形上到底走了多远、沿哪个方向走”。这种一致性是工程实现和理论研究都能受益的。下面用一张表来对比两条路线判断维度QR/SVD 路线闭式更新路线每步主操作分解 重新组合构造反对称矩阵 矩阵指数是否依赖内层迭代Gram-Schmidt 需要逐列迭代否一次性表达式正交保持依赖分解精度结构上由矩阵指数保证是否对应测地线一般是投影/回缩可以对应精确黎曼步实现复杂度低一行qr需要expm略高调试难度中间状态容易理解需要理解反对称和指数结构3. 从动量更新到闭式更新的实现路径3.1 最小可运行的结构示例下面给一个示意代码。注意这段代码的目的是讲清结构不是对某篇论文公式的逐字复刻。实际落地前请以原始论文或作者仓库里的公式为准。import torch def expm_stable(A): # 反对称矩阵的矩阵指数float32 下建议先升精度 orig_dtype A.dtype E torch.linalg.matrix_exp(A.double()) return E.to(orig_dtype) def closed_form_muon_step(M, G, lr, mu): # M: [n, p]满足 M.T M I # G: [n, p]和 M 同形状的梯度 # 1. 普通动量合成合成后暂时离开流形 M_mom mu * M G # 2. 提取反对称方向示意写法 S M.T M_mom - M_mom.T M # 3. 矩阵指数给出正交变换 M_new M expm_stable(-lr * S) return M_new解释一下这段代码里的关键点。第一步M_mom mu * M G和普通动量没有区别这一步会让动量离开流形。第二步S M.T M_mom - M_mom.T M构造的是一个p x p矩阵。手动验证一下就知道(M.T M_mom).T M_mom.T M两者相减之后天然满足反对称。反对称结构是后面“指数结果一定是正交矩阵”的前提。第三步M_new M exp_stable(-lr * S)是更新的主体。exp_stable(-lr * S)相当于在p维空间里作用一个正交变换M与它相乘后列之间的正交关系保持不变。如果你已经在使用 Muon可以把原来的 QR 替换成这类结构先在小样本上对比再逐步扩大规模。3.2 矩阵指数实现时的数值选择矩阵指数在 PyTorch 里可以直接用torch.linalg.matrix_exp在 CPU 上用scipy.linalg.expm也很方便。真正要注意的是精度。torch.linalg.matrix_exp在 float16 下容易出现精度损失因为矩阵指数对输入矩阵的异常值和尺度很敏感。我的习惯是先把反对称矩阵转成 float64 计算指数再转回原精度。这样多一次类型转换但换来的稳定性值得。尤其当p较大、矩阵接近病态时float32 下直接算 expm 的结果可能明显偏离正交条件。如果p很大比如几千expm 的开销也不小。这时候可以先想想是不是所有参数矩阵都需要走闭式更新。如果只是对一部分关键的矩阵参数使用 Muon其他参数继续用 AdamW整体收益可能更好。3.3 一个可复现的小验证脚本在动手训练之前先用随机矩阵跑一个单步验证torch.manual_seed(0) n, p 128, 16 M0, _ torch.linalg.qr(torch.randn(n, p)) # 生成一个正交列矩阵 G torch.randn(n, p) M_new closed_form_muon_step(M0, G, lr0.01, mu0.9) err (M_new.T M_new - torch.eye(p)).abs().max().item() print(err)如果能正常输出err通常在1e-5到1e-6这个量级。如果err到了1e-2甚至更大说明你的 expm 实现或反对称矩阵构造有问题先不要进训练循环。这一步很便宜但能提前拦住大部分低级错误。4. 验证闭式更新真的“正确”吗4.1 第一类检查正交性误差闭式更新最基本的要求是更新后还在流形上。所以第一个检查指标就是正交性误差def orth_error(M): I torch.eye(M.shape[1], dtypeM.dtype, deviceM.device) return (M.T M - I).norm().item()判断标准float32 下正交性误差在1e-5以下说明实现基本正确。float64 下应该能到1e-12量级。如果误差在1e-2量级说明闭式结构没有真正闭合先检查反对称矩阵是否构造正确、expm 是否在低精度下被调用。这个检查可以放到每一步更新之后做。刚开始调试时每隔几步打印一次确认误差没有随着步数累积上升。4.2 第二类检查和 QR 路线对比把闭式更新和 QR 更新放在同一个随机输入上对比M_qr, _ torch.linalg.qr(mu * M0 G) M_cf closed_form_muon_step(M0, G, lr0.01, mu0.9) print(torch.norm(M_qr - M_cf).item())这里要注意两个结果大概率不一致。不要把这个当成“谁对谁错”的依据。QR 是一种投影路径闭式更新是流形上的精确更新路径两者本来就不应该落在同一个点。你需要确认的是两个结果都满足正交性。在简单目标函数上两者的 loss 都在下降。闭式版本的行为符合论文或作者的描述。如果闭式版本的单步更新和 QR 在 loss 方向上完全相反那就要警惕了可能你的梯度符号反了也可能反对称矩阵的符号取反了。4.3 第三类检查小规模收敛行为最实用的验证方法是做一个小规模训练实验。比如训练一个线性层y Wx用闭式更新优化损失函数随机生成输入X和真值Y。用普通 SGD、Adam、QR-Muon、闭式 Muon 四种优化器各跑一遍。对比 loss 下降曲线。重点不是看谁收敛速度最快而是确认闭式更新不会在几十步内就出现 NaN 或发散。如果 loss 曲线震荡剧烈先把学习率调小一个数量级再试。闭式更新在流形上的步长感受和普通欧氏空间不一样Adam 里