尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

Muon在Stiefel流形上的精确闭式更新解析

Muon在Stiefel流形上的精确闭式更新解析 在 2025 年初这一轮大模型预训练优化器讨论中Muon 是一个绕不开的名字。它和 AdamW 的核心差异在于对于二维权重矩阵Muon 并不是逐元素更新参数而是先构造一个动量矩阵再从这个矩阵的极分解或 SVD 中取出正交因子作为更新方向。由于正交因子天然接近 Stiefel 流形网上开始出现一个更进一步的表述Muon on the Stiefel Manifold Admits an Exact Closed-Form Update意思是当参数被约束在 Stiefel 流形上时Muon 的流形更新步骤可以写成精确闭式解而不需要像早期实现那样用 Newton-Schulz 迭代去逼近正交因子。这篇文章的目标是把这个标题讲清楚先说明 Muon 为什么和 Stiefel 流形有关再解释 Newton-Schulz 迭代的近似代价然后给出基于 Cayley 变换的精确闭式更新公式最后用 PyTorch 实现一个可运行的 Muon-Cayley 优化器并补上验证方法、常见坑和落地检查清单。适合正在复现 Muon、准备把正交约束引入大模型训练、或者只是想理解几何优化器背后原理的读者。1. 先理解 Muon 为什么要把优化器限制在 Stiefel 流形上1.1 Muon 优化器要解决什么问题AdamW 是目前大模型训练最常用的优化器之一但它有一个明显特点对每个参数独立做一阶矩和二阶矩归一化。这种逐元素更新在大多数场景下很稳定却没有利用矩阵参数本身的结构。当参数是一个完整的二维权重矩阵时逐元素更新等价于在欧氏空间里沿着坐标轴方向移动更新方向可能远离“矩阵空间里有意义的方向”。Muon 的思路是反过来的。对于二维参数它先维护一个动量矩阵然后计算这个动量矩阵的正交因子用正交因子作为更新方向。这样做的几何含义是更新方向不再被逐元素噪声主导而是被矩阵的“主要方向”主导。在 Transformer 中这类二维参数大量存在embedding 矩阵attention 里的 Q、K、V 投影MLP 的第一层和第二层权重head 层权重。对这些矩阵做结构化的更新理论上可以改善优化路径避免参数在训练中产生大量冗余的坐标变化。1.2 Stiefel 流形列正交矩阵集合Stiefel 流形是数学里一个经典集合记作 V_k(R^n)。通俗定义V_k(R^n) 是所有满足 W^T W I_k 的 n 行 k 列实矩阵 W 的集合其中 k n。用文字说就是“一个 n x k 的矩阵每一列之间相互正交并且每一列的模长为 1”。这样的矩阵称为列正交矩阵。为什么大模型训练里会关心这个集合因为很多权重矩阵在理想条件下并不希望自由漂移到任意数值。比如embedding 矩阵如果保持列正交不同 token 的向量表示不会出现严重共线attention 投影矩阵如果保持列正交特征之间的冗余会更小某些低秩适配层和循环网络中的状态矩阵也会用到正交约束。当然不是所有层都适合套 Stiefel 约束。实际落地时通常只对部分层使用或者把 Stiefel 约束当成一种正则化手段而不是对所有参数强制使用。1.3 切空间更新加回流形才是 Muon 的几何本质如果一个参数 W 被要求始终落在 Stiefel 流形上那么一次更新不能直接把 W 加上一个任意的梯度方向因为 W lr * G 很可能不满足 W^T W I。在流形优化里标准做法分两步把梯度投影到当前点 W 处的切空间在切空间走一步再用 retraction 映射回流形。Stiefel 流形在 W 处的切空间 T_W 由所有满足下式的矩阵 X 组成X^T W W^T X 0给定一个任意的梯度 G切空间投影公式是G_t G - W * sym(W^T G)其中 sym(A) (A A^T) / 2。这个投影的含义是把 G 中“垂直撞向流形”的分量去掉只保留“沿流形滑动”的分量。Muon 和 Stiefel 流形的关系就在这里Muon 使用正交因子作为更新方向本质上是希望更新方向尽量贴合流形结构而如果把参数本身约束在 Stiefel 流形上就必须处理切空间投影和回流形这两个步骤。2. Newton-Schulz 迭代解决了问题但引入了近似误差2.1 极分解与正交因子的作用给定一个 n x k 矩阵 M它的极分解可以写成M U * P其中 U 是 n x k 列正交矩阵P 是 k x k 半正定矩阵。U 可以看作 M 的“正交因子”P 是“尺度因子”。Muon 的经典做法就是先构造动量矩阵 M再计算 M 的正交因子 U最后用 U或者 U 的某种缩放作为更新方向。如果直接用 SVD可以这样算U, S, Vt torch.linalg.svd(M, full_matricesFalse) direction U Vt注意 direction U Vt 是 n x k 矩阵当 M 是方阵时它就是一个正交矩阵。这个 direction 就是极分解里的正交因子。问题在于 SVD 在大模型场景下太贵。每个 step 都对所有二维参数做一次 SVD计算量会非常高。于是早期 Muon 实现普遍采用 Newton-Schulz 迭代来近似计算这个正交因子。2.2 Newton-Schulz 迭代的代码实现Newton-Schulz 迭代的核心是用矩阵乘法不断逼近矩阵的极分解因子。下面是一个常见实现def zeropower_via_newtonschulz(G, steps10): a, b, c 0.2222, 0.8889, 0.3859 X G / (G.norm(2) 1e-8) for _ in range(steps): A X X.T X a * X b * A X c * A A X return X这里的 a、b、c 来自对矩阵函数的多项式逼近。迭代会逐步把 X 拉向 M 的正交因子。steps 越大逼近精度越高但计算量也线性增长。需要特别说明的是G.norm(2) 计算的是谱范数谱范数本身通常也要通过 SVD 计算所以很多实际实现会改用 Frobenius 范数做粗略缩放或者直接不缩放。示例代码只用于说明迭代逻辑真正落地时要根据依赖版本和性能要求调整。2.3 近似迭代的三个代价Newton-Schulz 迭代不是免费的它有三个明显代价。第一迭代次数固定精度不可控。steps 设为 5 就是 5 次矩阵乘法设为 10 就是 10 次。输入矩阵条件数很糟糕时固定步数可能不足以保证结果足够接近正交。第二数值稳定性受梯度 scale 影响。如果动量矩阵的范数非常大或者混合精度下出现溢出Newton-Schulz 迭代很容易放大误差最终产生 NaN。第三计算成本并不低。每次迭代都要做 n x k 和 k x n 的矩阵乘法在 hidden size 为 4096 的模型上反复执行会明显拖慢训练。方法精度主要成本数值风险SVD 精确分解高高低但计算慢Newton-Schulz 迭代取决于 steps中对 scale 和条件数敏感Cayley 闭式更新高取决于求解方式矩阵求逆可能不稳定这正好引出标题里的 Core 信息Muon 在 Stiefel 流形上的更新并不一定要用迭代近似它可以写成精确闭式解。3. 用 Cayley 变换得到精确闭式更新3.1 Cayley 变换为什么能保持 Stiefel 流形Cayley 变换是一个经典的矩阵变换。给定一个斜对称矩阵 A满足 A^T -A那么下面这个矩阵 Q 是正交矩阵Q (I c * A)^(-1) * (I - c * A)其中 I 是单位矩阵c 是实数标量。因为 Q 是正交矩阵所以把 Q 作用到列正交矩阵 W 上即计算W_new Q W得到的 W_new 仍然满足 W_new^T W_new I。这是 Cayley 变换用于 Stiefel 流形更新的核心原因。直观理解是Q 在欧氏空间里相当于一个旋转或反射操作它不改变列之间的夹角和长度。把 W 整体旋转一下自然不会破坏 W 的列正交性。3.2 闭式更新的公式与推导关键要让 Cayley 变换真正服务于优化需要从梯度方向构造一个斜对称矩阵 A。假设 W 是当前参数V 是经过切空间投影后的更新方向。构造A V W^T - W V^T由于 A^T -AA 天然是斜对称矩阵。然后用 Cayley 变换得到W_new (I (lr / 2) * A)^(-1) * (I - (lr / 2) * A) W其中 lr 是学习率。这就是一个精确的闭式更新不需要迭代逼近也不需要先算 SVD只要解一次线性方程组。这个更新之所以成立可以从两个角度看。第一A 是切空间方向 V 与当前点 W 生成的斜对称矩阵它编码了“沿切空间走”的信息。第二Cayley 变换本身是一个 retraction它把切空间里的移动映射回 Stiefel 流形上并且保持一阶一致性。实际代码可以这样写def project_tangent(W, G): M W.T G sym 0.5 * (M M.T) return G - W sym def cayley_retraction(W, V, lr): c lr / 2.0 A V W.T - W V.T n W.shape[0] I torch.eye(n, dtypeW.dtype, deviceW.device) L I c * A R I - c * A Q torch.linalg.solve(L, R) return Q W这里最关键的一行是 torch.linalg.solve(L, R)。它直接解线性方程组 L Q R得到正交矩阵 Q再乘到 W 上。3.3 与梯度中心化和动量的配合Muon 在标准实现中通常还会做两件事梯度中心化和动量累积。梯度中心化指的是把梯度在每一行上减去均值让更新方向不携带全局均值偏移。动量累积则是用一个指数滑动平均来平滑梯度噪声。在 Stiefel 流形的 Muon-Cayley 版本里顺序可以是对梯度做中心化更新动量缓冲 M beta * M (1 - beta) * G_centered把 M 投影到 W 的切空间得到 V用 Cayley 变换执行更新。需要注意这里并没有严格要求动量矩阵 M 本身先做极分解。流形优化的标准路径是“切空间投影 retraction”所以核心计算落在 V 和 A 上。4. 在 PyTorch 中实现 Muon-Cayley 优化器4.1 最小工程结构一个可复现的最小项目可以这样组织文件职责muon_cayley.py优化器实现包含工具函数train_example.py最小训练脚本验证优化器能跑通依赖只有 PyTorch。示例代码主要用于说明思路实际项目要结合自己的包名、路径和 torch 版本调整。4.2 基础工具函数Newton-Schulz 与 Cayley 对照为了对照两种做法先把工具函数写成独立的两个版本import torch def zeropower_via_newtonschulz(G, steps10): a, b, c 0.2222, 0.8889, 0.3859 X G / (G.norm(2) 1e-8) for _ in range(steps): A X X.T X a * X b * A X c * A A X return X def project_tangent(W, G): M W.T G sym 0.5 * (M M.T) return G - W sym def cayley_retraction(W, V, lr): c lr / 2.0 A V W.T - W V.T n W.shape[0] I torch.eye(n, dtypeW.dtype, deviceW.device) Q torch.linalg.solve(I c * A, I - c * A) return Q W这里有两个需要提醒的点。第一project_tangent 里 W.T G 是 k x k 矩阵sym 是对称化操作。这个投影公式要求 W 满足 W^T W I如果 W 在训练过程中已经偏离流形投影会不准确。第二cayley_retraction 构造的是 n x n 矩阵。当 n 很大时直接构造并求逆会非常费显存。下一小节会给出一个避免 n x n 求逆的变体。4.3 完整优化器类与训练循环下面是一个最小但完整的优化器类。它把二维参数拆出来使用 Muon-Cayley其余一维参数继续使用 AdamWclass MuonCayley(torch.optim.Optimizer): def __init__(self, params, lr3e-3, momentum0.9, weight_decay0.1, use_cayleyTrue, ns_steps10, adamw_lr1e-3): defaults dict(lrlr, momentummomentum, weight_decayweight_decay, use_cayleyuse_cayley, ns_stepsns_steps, adamw_lradamw_lr) 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: for p in group[params]: if p.grad is None: continue grad p.grad if p.ndim 2 and p.shape[0] p.shape[1]: self._update_muon_cayley(p, grad, group) else: self._update_adamw(p, grad, group) return loss def _update_muon_cayley(self, p, grad, group): state self.state[p] if momentum_buffer not in state: state[momentum_buffer] torch.zeros_like(grad) buf state[momentum_buffer] momentum group[momentum] buf.mul_(momentum).add_(grad, alpha1 - momentum) centered buf - buf.mean(dim0, keepdimTrue) v project_tangent(p, centered) if group[use_cayley]: new_p cayley_retraction(p, v, group[lr]) else: new_p p group[lr] * zeropower_via_newtonschulz(v, group[ns_steps]) p.copy_(new_p) def _update_adamw(self, p, grad, group): state self.state[p] if exp_avg not in state: state[exp_avg] torch.zeros_like(p) state[exp_avg_sq] torch.zeros_like(p) state[step] 0 exp_avg state[exp_avg] exp_avg_sq state[exp_avg_sq] state[step] 1 exp_avg.mul_(0.9).add_(grad, alpha0.1) exp_avg_sq.mul_(0.999).addcmul_(grad, grad, value0.001) bias_corr1 1 - 0.9 ** state[step] bias_corr2 1 - 0.999 ** state[step] denom (exp_avg_sq.sqrt() / (bias_corr2 ** 0.5)).add_(1e-8) step_size group[adamw_lr] / bias_corr1 if group[weight_decay] 0: p.mul_(1 - group[lr] * group[weight_decay]) p.addcdiv_(exp_avg, denom, value-step_size)这个实现有几个刻意简化的地方中心化只做了行均值实际 Muon 里可能有更精细的处理AdamW 部分用了最简单的写法没有做权重衰减和 Adam 更新的完整融合判定“哪些参数用 Muon”只用了 p.ndim 2 和形状条件。在训练循环里使用方式如下model SomeTransformer() optimizer MuonCayley(model.parameters(), lr3e-3, use_cayleyTrue) for x, y in dataloader: loss model.compute_loss(x, y) optimizer.zero_grad() loss.backward() optimizer.step()4.4 维度、参数与注意事项使用 Muon-Cayley 前必须确认参数形状满足列正交的基本要求。Stiefel 流形 V_k(R^n) 要求 n k也就是说矩阵行数不小于列数。如果参数是 embedding 这类 shape 为 (vocab_size, hidden_size) 的矩阵通常是 vocab_size 大于 hidden_size可以直接用。如果出现 n k需要转置后再做流形更新或者改用行正交的镜像版本。另一个注意点是 lr。Cayley 变换里的 c lr / 2 不是普通的逐元素缩放它直接进入矩阵求逆。lr 过大时I c * A 可能接近奇异导致求逆后的 Q 非常剧烈。实际项目中可以先从 1e-3 到 3e-3 开始不要一开始就开很大的学习率。5. 运行验证如何确认更新真的落在 Stiefel 流形上5.1 验证指标验证 Cayley 更新是否正确不能只看 loss 有没有下降还要看几何指标。指标计算公式期望正交性误差norm(W.T W - I)接近 0列均值W.mean(dim0)不出现 NaN更新前后夹角cosine(W_before, W_after)下降但不过度剧烈梯度范数grad.norm()数值稳定如果正交性误差在几百个 step 后涨到 1e-2 以上说明更新方式可能不是严格 retraction或者中间某一步把 W 拉出了流形。5.2 一个可复现的最小实验用一个人造目标函数来验证优化器。目标是把 W 优化到离某个目标矩阵更近同时让 W 始终保持列正交torch.manual_seed(0) n, k 64, 16 W torch.linalg.qr(torch.randn(n, k)).Q target torch.randn(n, k) def loss_fn(W): return (W - target).pow(2).mean() optimizer MuonCayley([W], lr1e-2, use_cayleyTrue) history [] for step in range(100): loss loss_fn(W) optimizer.zero_grad() loss.backward() optimizer.step() orth_err torch.norm(W.T W - torch.eye(k)).item() history.append((loss.item(), orth_err)) if step % 20 0: print(step, loss.item(), orth_err)预期结果loss 下降正交性误差始终接近 0不出现 NaN。如果使用 use_cayleyFalse 的 Newton-Schulz 版本正交性误差可能略微波动但整体也应在可接受范围内。这个实验只验证了几何性质不能用来证明 Muon-Cayley 一定比 AdamW 好。真实训练效果需要放到具体模型和数据集上对比。5.3 学习环境与生产环境差异学习环境里直接构造 n x n 矩阵并求解通常没问题。生产环境则要额外关注几点混合精度下torch.linalg.solve 的 dtype 是否和参数一致n 很大时n x n 矩阵的显存占用是否可接受梯度裁剪是否在 Cayley 变换之前执行动量缓冲是否写进 checkpoint恢复训练后状态是否一致是否需要为 Muon-Cayley 单独设置学习率 warmup。不要只在 toy case 里验证一次就上生产。建议先在小模型上跑若干 step对比 use_cayleyTrue 和 False 的 loss、正交性、耗时和显存再做决定。6. 常见问题与排查路径6.1 矩阵维度不匹配或参数不是二维现象训练时抛出 shape 错误或者 W.T W 的形状不是 k x k。可能原因模型里有二维参数但 shape 是 n k被错误地当成 Stiefel 参数处理。检查方式在 step 里打印每个二维参数的 shape确认是否满足 n k。处理建议对不满足条件的参数回退到 AdamW或者先做转置等更新后再转置回来。问题现象常见原因检查方式处理建议shape 报错参数 n k打印 p.shape回退 AdamW 或转置正交性误差不下降初始化不在流形上打印 W.TW - I用 QR 分解初始化 W只有部分层生效判断条件太窄检查 p.ndim 和 shape放宽或细化参数分组6.2 n x n 矩阵求逆导致显存或算力爆炸现象step 时间明显变长显存 OOM。可能原因cayley_retraction 里构造了 A V W.T - W V.T形状是 n x n。当 n 等于 hidden size 4096 时这个矩阵本身就有 4096 x 4096 个元素。检查方式打印 A.shape按 A.numel() * element_size 估算额外显存。处理建议如果 n 很大改用 Woodbury 形式的求解避免直接构造 n x n 逆矩阵。Woodbury 变体的思路是把 A 写成两个窄矩阵的乘积A [V, W] [W.T, -V.T]然后利用矩阵求逆引理把 n x n 方程组化简为 2k x 2k 方程组。k 远小于 n 时显存和耗时都可以大幅下降。def cayley_retraction_woodbury(W, V, lr): n, k W.shape c lr / 2.0 U torch.cat([V, W], dim1) # n x (2k) B torch.cat([W.T, -V.T], dim0) # (2k) x n I_2k torch.eye(2 * k, dtypeW.dtype, deviceW.device) M_inner B U # (2k) x (2k) C I_2k c * M_inner b W - c * U (B W) # (I - c*A) W z torch.linalg.solve(C, B b) # (2k) return b - c * U z这个变体只要求解一个 2k x 2k 的线性方程组适合 k 较小的场景。如果 k 也很大那么直接使用 Newton-Schulz 迭代回退会更稳妥。6.3 Cayley 数值不稳定或 NaN现象训练出现 NaNloss 变成 inf。可能原因学习率过大导致 I c * A 接近奇异输入梯度包含 NaN混合精度下累加误差。检查方式在 step 前打印 grad 的最大值和最小值打印 torch.linalg.cond(I c * A) 的条件数对 W 做一次 isfinite 检查。处理建议降低 lr在 Cayley 变换前对 V 做范数裁剪对斜对称矩阵 A 做 scale例如用 V.norm() 归一化后再乘 lr混合精度场景下先用 fp32 跑通再切 bf16。6.4 加入了其他优化器后训练异常现象只对部分层使用 Muon-Cayley其余层用 AdamW训练波动变大。可能原因参数分组不当导致不同层学习率差异过大一维参数被误判成二维参数Stiefel 参数初始化不满足列正交。检查方式打印每个参数所属的分组确认哪些参数用了 Muon-Cayley哪些用了 AdamW。处理建议先跑一个纯 Muon-Cayley 的 toy case确认正常后再混合使用。混合使用时把两组参数的学习率分开设置并单独记录日志方便定位是哪个分组引起的波动。7. 最佳实践与扩展方向7.1 参数选型清单参数常见起点调整方向lr1e-3 到 3e-3过大导致 Cayley 奇异过小收敛慢momentum0.9可以尝试 0.95但要注意延迟weight_decay0.1只在 AdamW 分组使用Muon 分组慎用use_cayleyTrue对比 False 时确认正交性差异ns_steps5 到 10只在 Newton-Schulz 路径使用这些参数都不是公式只是经验起点。落地前要先确认自己依赖的 torch 版本是否支持 torch.linalg.solve 的 batched 输入以及混合精度下的行为。7.2 落地前检查清单[ ] 确认模型哪些二维参数需要 Stiefel 约束哪些应该回退到 AdamW[ ] 确认参数初始化满足列正交例如用 QR 分解或正交初始化[ ] 确认 momentum buffer 在优化器 state dict 中checkpoint 能完整保存恢复[ ] 确认梯度中心化顺序先中心化再投影再 retraction[ ] 确认学习率不会让 Cayley 变换里的矩阵奇异[ ] 在 toy case 上同时验证 loss 下降和正交性误差[ ] 在真实模型的小配置上对比 use_cayleyTrue 和 False 的耗时、显存、loss 曲线[ ] 确认混合精度下没有 NaN并检查所有 dtype 一致性[ ] 如果显存紧张优先使用 Woodbury 变体或 Newton-Schulz 回退。7.3 扩展方向沿着这个方向继续深入有几个值得尝试的点。第一个是 4D 参数的处理。Transformer 里 attention 输出投影经常是 4D 权重最简单的做法是把它 reshape 成 2D 后再应用 Muon-Cayley更新完再 reshape 回去。第二个是和其他几何约束组合。比如在权重矩阵上同时施加低秩约束和正交约束或者把归一化层的尺度参数单独用一个一维优化器管理。第三个是混合策略。训练前期使用 Newton-Schulz 迭代节省显存训练后期切换成 Cayley 闭式更新来消除正交漂移。这个策略是否有效需要靠实验数据判断不能直接假设闭式解一定更好。第四个是继续理解 Muon 的理论动机。Muon 的正交因子更新方向、Stiefel 流形上的 Cayley retraction、以及梯度中心化三者本质上是三个独立设计分别解决“更新方向结构化”“参数保持在流形上”“更新方向去偏移”。理解清楚这三者的分工比盲目替换优化器更有价值。最后回到标题本身Muon on the Stiefel Manifold Admits an Exact Closed-Form Update核心信息不是“闭式解一定更快”而是“在 Stiefel 流形上Muon 的更新可以不再依赖 Newton-Schulz 迭代也不需要每一步都做 SVD而是用一个 Cayley 变换一次性完成”。这个判断对实现者最大的启发是当你看到一个优化器依赖近似迭代时先回到几何结构上想一想它是否本来就有精确解。很多时候近似只是工程妥协而不是唯一答案。
返回列表