MALT优化器:融合Adam自适应与Muon正交化的轻量级实现

📅 发布时间:2026/8/27 22:55:31
MALT优化器:融合Adam自适应与Muon正交化的轻量级实现 训练神经网络时优化器选择的本质是在更新方向、步长与计算成本之间做权衡。Adam 擅长给每个坐标提供自适应步长但它在更新时没有利用权重矩阵本身的结构Muon 这类优化器通过矩阵正交化让更新方向更接近正交却对梯度的绝对尺度非常敏感。MALTMuon with Adaptive Lightweight diagonal preconditioning是一种把两者结合起来的轻量级做法在牛顿-舒尔茨迭代做正交化之前先给梯度加一层对角预条件。这篇文章会从 Muon 的原理开始讲清楚为什么对角预条件能补足 Muon 的短板再给出 MALT 的算法定义、PyTorch 最小实现、运行验证方法和工程排错清单。读者不需要提前了解优化器前沿研究只需要熟悉 PyTorch 的Optimizer基础用法就能跟下来。1. Muon 优化器为什么值得关注1.1 Muon 的出发点矩阵层需要结构感知更新深度学习模型里全连接层、卷积层、注意力层的权重在内存中通常被组织成二维矩阵。常见的 Adam 类优化器把每个元素当作独立标量来更新逐元素地维护一阶矩和二阶矩。这样做的好处是简单、稳定坏处是忽略了同一矩阵内部的坐标关系。Muon 的核心想法是如果当前参数是一个矩阵那么更新方向也应该尽量保留矩阵层面的结构性质。具体来说它希望在梯度方向上叠加一个动量后让最终更新方向接近“正交矩阵”的方向。正交在这里有两种理解方向一是列向量之间接近标准正交二是行向量之间接近标准正交具体取决于矩阵形状。这个性质在某些深层网络中被认为有助于控制隐藏状态的规模漂移因为一个接近正交的更新方向不会在反复叠加后无限膨胀。可以把 Muon 理解成一个“中间路线”它不像 K-FAC 那样显式建模二阶曲率也不像 Adam 那样完全无视矩阵结构。它只额外做一步矩阵层面的正交化因此实现成本比 K-FAC 低表达能力又比纯逐元素方法强。1.2 Muon 的典型更新流程一个常见的 Muon 风格更新流程可以拆成三步。第一步对梯度做动量累积。这一步和 SGD 动量类似目的不是为了自适应而是为了稳定更新方向减少随机采样带来的抖动。第二步对动量矩阵做正交化。工程实现里通常不直接做 SVD因为 SVD 在训练高频迭代中太贵。一般用固定步数的牛顿-舒尔茨迭代逼近“极分解”中的正交因子。迭代会保留梯度的大致方向同时让结果的列向量或行向量更接近正交。第三步用学习率缩放后更新参数。整个过程中没有逐坐标二阶矩因此 Muon 的内存占用通常明显低于 Adam。下面是 Muon 风格流程的伪代码。输入当前权重 W梯度 G动量缓冲 M学习率 lr动量系数 mu M mu * M G O orthogonalize(M) W W - lr * O这里的关键点是“正交化”作用在哪个对象上。如果把正交化直接作用在原始梯度上那么梯度范数剧烈变化时正交化后的方向会非常不稳定。如果先动量后正交化相当于在一个平滑过的梯度方向上做几何修正稳定性更好。1.3 Muon 的遗留问题梯度尺度敏感Muon 虽然利用了矩阵结构但没有引入逐元素的自适应缩放。这带来一个实际问题当一个矩阵参数的不同列、不同行或者不同矩阵之间的梯度尺度差异很大时Muon 的更新幅度只由统一的学习率和动量决定。结果是某些坐标的梯度可能过小几乎不更新另一些坐标的梯度可能过大网络训练早期就出现异常 spike。传统 Adam 解法是维护二阶矩然后对每个元素做归一化。于是很自然的思路就是能不能把 Adam 的对角预条件拿过来给 Muon 的梯度先做一次逐元素缩放再做正交化这正是 MALT 要解决的问题。2. 对角预条件如何给 Muon 补上尺度信息2.1 Adam 的自适应缩放是在做什么Adam 的每一次更新可以拆成两个部分方向归一化以及幅度归一化。方向归一化体现在它用m / sqrt(v)替代原始梯度。m是梯度的指数移动平均v是梯度平方的指数移动平均。这个比值可以粗略理解为“带符号的信噪比”某个坐标梯度长期为正它就会得到一个稳定的正更新某个坐标梯度不断正负跳跃它的更新会被分母抑制。幅度归一化体现在整个更新步长最终由学习率控制。即使某个参数的梯度绝对值非常大除以sqrt(v)之后也会被压回一个相对稳定的范围。因此 Adam 对学习率和初始梯度尺度没有那么敏感。但这种逐元素缩放丢掉了矩阵结构。两个相邻参数元素可能被缩放成完全不同的更新幅度矩阵的正交性、谱结构完全不在考虑范围内。2.2 MALT 的轻量对角预条件MALT 选择一个折中方案保留 Adam 的逐元素二阶矩估计但对二维权重参数在预条件之后继续做一轮矩阵正交化。也就是说它既知道每个坐标的尺度差异又保留矩阵层面的几何约束。预条件公式可以写成g_pre g / (sqrt(v_hat) eps)其中g是当前梯度v_hat是去偏后的二阶矩估计eps是数值稳定项。这一步是“轻量”的关键。它不需要构造完整的曲率矩阵也不需要计算 Hessian 的逆。每个参数只需要额外维护一个和自身形状相同的二阶矩缓冲以及一个动量缓冲。相比对完整矩阵做预条件普通 GPU 显存也能接受。2.3 预条件与正交化的先后顺序工程上最自然的顺序是先逐元素预条件再做正交化。如果先做正交化再按二阶矩逐元素缩放那么正交化带来的矩阵几何结构很可能被逐元素缩放破坏。因为逐元素缩放是非线性、非等距的变换它会把原本正交的方向扭曲掉。如果先逐元素预条件再做正交化相当于用“尺度修正后的梯度”参与矩阵几何修正。正交化会把方向重新拉回近正交流形同时它对全局范数有归一化作用因此预条件的绝对尺度不会影响最终更新量只影响矩阵内部每个元素的方向权重。这个顺序有一个需要注意的副作用牛顿-舒尔茨迭代一旦对矩阵做了归一化预条件的全局缩放就会被抹掉。最终起作用的只有预条件的方向信息而不是幅度。这是符合预期的因为步长应该由学习率控制而不应该由历史梯度规模控制。3. MALT 算法设计与超参语义3.1 矩阵参数分支的更新公式对于形状为二维的参数pMALT 每个 step 执行以下过程。首先更新二阶矩估计v beta2 * v (1 - beta2) * g^2 v_hat v / (1 - beta2^t)然后计算预条件梯度g_pre g / (sqrt(v_hat) eps)接着对g_pre执行牛顿-舒尔茨正交化o orthogonalize(g_pre)最后更新动量并让动量自己承担长期记忆m beta1 * m (1 - beta1) * o p p - lr * m这里刻意没有像 Adam 那样对m做去偏。因为第二步已经把o归一化到接近等范数的状态即使早期m偏小也不会造成入口阶段的异常大更新。如果项目希望和 Adam 行为更接近也可以对m去偏但建议先用默认不做去偏的形式跑通。3.2 非矩阵参数的回退策略并不是所有参数都是二维权重。偏置项、LayerNorm 的 scale、Embedding 的某些一维参数形状不是二维或者某个维度等于 1勉强套正交化没有意义。MALT 的工程实现通常对这部分参数回退到 AdamW 风格的更新。回退分支用标准的 AdamW 公式m beta1 * m (1 - beta1) * g v beta2 * v (1 - beta2) * g^2 m_hat m / (1 - beta1^t) v_hat v / (1 - beta2^t) p p - lr * (m_hat / (sqrt(v_hat) eps) weight_decay * p)这里weight_decay * p是解耦权重衰减不是 L2 正则化。两者名字经常混用但在实现上有一个显著区别解耦权重衰减不参与动量也不会被自适应缩放而是直接以固定比例缩小参数。3.3 超参表与内存分析MALT 的超参定义和 Adam 非常接近含义也基本一致。超参常用默认值作用调整建议lr1e-3 左右全局学习率比 SGD 更敏感建议配合 warmupbeta10.95一阶动量指数衰减系数小 batch 用 0.9大 batch 可调高到 0.98beta20.999二阶矩指数衰减系数数据非平稳时可降到 0.99eps1e-8数值稳定项如果 loss 出现 NaN可提高到 1e-6weight_decay0.01解耦权重衰减过大容易欠拟合从 0 开始调ns_steps5牛顿-舒尔茨迭代步数3 到 7 之间通常够用chunk_sizeNone正交化分块大小大矩阵建议 128 或 256内存方面MALT 对每个参与矩阵分支的参数维护m和v两个缓冲对每个非矩阵参数同样维护两个缓冲因此整体内存大约是参数量的 3 倍和 AdamW 基本一致。Muon 只有动量一个缓冲内存是参数量的 2 倍。MALT 多出的这部分内存就是对角预条件带来的成本。4. PyTorch 最小实现4.1 工具函数Newton-Schulz 迭代先写一个独立的牛顿-舒尔茨函数。输入是二维梯度矩阵输出是尽量靠近正交因子的矩阵。为了保证数值稳定迭代前先对矩阵做 Frobenius 范数归一化。import torch def newton_schulz(grad_matrix: torch.Tensor, steps: int 5, eps: float 1e-8): if grad_matrix.dim() ! 2: raise ValueError(newton_schulz only supports 2D tensors) transposed False if grad_matrix.shape[0] grad_matrix.shape[1]: grad_matrix grad_matrix.t() transposed True grad_matrix grad_matrix / (grad_matrix.norm() eps) eye torch.eye( grad_matrix.shape[1], dtypegrad_matrix.dtype, devicegrad_matrix.device, ) for _ in range(steps): grad_matrix 0.5 * grad_matrix.mm( 3.0 * eye - grad_matrix.t().mm(grad_matrix) ) return grad_matrix.t() if transposed else grad_matrix这个实现的核心公式是X - 0.5 * X (3I - X^T X)。它是求解极分解中正交因子的经典牛顿迭代。注意迭代前必须归一化否则矩阵范数过大时迭代会发散。steps5在大多数场景是性能和精度的折中。4.2 MALT 优化器主体下面实现一个继承torch.optim.Optimizer的 MALT 优化器。为了方便切换矩阵参数走 Muon 风格非矩阵参数走 AdamW 风格。import torch from torch.optim import Optimizer class MALT(Optimizer): def __init__( self, params, lr1e-3, betas(0.95, 0.999), eps1e-8, weight_decay0.0, ns_steps5, chunk_sizeNone, ): if not 0.0 lr: raise ValueError(fInvalid lr: {lr}) if not 0.0 eps: raise ValueError(fInvalid eps: {eps}) if not 0.0 betas[0] 1.0: raise ValueError(fInvalid beta1: {betas[0]}) if not 0.0 betas[1] 1.0: raise ValueError(fInvalid beta2: {betas[1]}) if weight_decay 0.0: raise ValueError(fInvalid weight_decay: {weight_decay}) defaults dict( lrlr, betasbetas, epseps, weight_decayweight_decay, ns_stepsns_steps, chunk_sizechunk_size, ) super().__init__(params, defaults) def _orthogonalize(self, grad_matrix: torch.Tensor) - torch.Tensor: chunk_size self.defaults[chunk_size] ns_steps self.defaults[ns_steps] if chunk_size is None or grad_matrix.shape[1] chunk_size: return newton_schulz(grad_matrix, stepsns_steps) chunks [] for start in range(0, grad_matrix.shape[1], chunk_size): chunk grad_matrix[:, start : start chunk_size] chunks.append(newton_schulz(chunk, stepsns_steps)) return torch.cat(chunks, dim1) 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: beta1, beta2 group[betas] lr group[lr] eps group[eps] weight_decay group[weight_decay] for p in group[params]: if p.grad is None: continue grad p.grad if grad.is_sparse: raise NotImplementedError(MALT does not support sparse gradients) state self.state[p] if len(state) 0: state[step] 0 state[exp_avg] torch.zeros_like(p) state[exp_avg_sq] torch.zeros_like(p) exp_avg state[exp_avg] exp_avg_sq state[exp_avg_sq] state[step] 1 step state[step] # 更新二阶矩 exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value1 - beta2) bias_correction2 1 - beta2**step v_hat exp_avg_sq / bias_correction2 is_matrix p.dim() 2 and p.shape[0] 1 and p.shape[1] 1 if is_matrix: # 对角预条件 g_pre grad / (v_hat.sqrt() eps) # 矩阵正交化 g_pre self._orthogonalize(g_pre) # 动量 exp_avg.mul_(beta1).add_(g_pre, alpha1 - beta1) update exp_avg else: # AdamW 回退分支 exp_avg.mul_(beta1).add_(grad, alpha1 - beta1) bias_correction1 1 - beta1**step m_hat exp_avg / bias_correction1 denom v_hat.sqrt().add_(eps) update m_hat / denom if weight_decay ! 0.0: p.mul_(1 - lr * weight_decay) p.add_(update, alpha-lr) return loss这里有一个细节值得解释矩阵分支里exp_avg_sq是用原始梯度维护的而不是用预条件后的梯度。因为二阶矩的作用是估计原始梯度的尺度一旦把预条件后的梯度也纳入二阶矩整个缩放会进入循环依赖行为更难预测。chunk_size的作用是控制正交化过程中的矩阵乘法规模。一个4196 x 4196的矩阵做完整牛顿-舒尔茨迭代时单次矩阵乘法就是千万级别元素显存和计算压力都很大。按列切块后每个块独立做正交化可以大幅降低单次矩阵乘法的峰值开销。4.3 用一个三层 MLP 跑通训练为了让验证足够简单这里用一个三层 MLP 做示例。数据源可以使用 MNIST 或自己生成随机数据下面是网络定义和优化器接入方式。import torch from torch import nn class MLP(nn.Module): def __init__(self, in_dim784, hidden512, num_classes10): super().__init__() self.net nn.Sequential( nn.Linear(in_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, num_classes), ) def forward(self, x): return self.net(x) model MLP() optimizer MALT( model.parameters(), lr1e-3, betas(0.95, 0.999), eps1e-8, weight_decay0.01, ns_steps5, chunk_size256, ) criterion nn.CrossEntropyLoss()训练循环保持和普通 PyTorch 代码一致。没有特殊 API只需要在loss.backward()之后调用optimizer.step()。for epoch in range(20): for images, labels in train_loader: images images.view(images.size(0), -1) logits model(images) loss criterion(logits, labels) optimizer.zero_grad(set_to_noneTrue) loss.backward() optimizer.step() print(fepoch {epoch}: loss {loss.item():.4f})如果原始数据集的类别数、输入维度不同只需要调整MLP的in_dim和num_classes。跑通的标志是 loss 能稳定下降并且前几个 epoch 不出现 NaN。5. 运行验证与异常定位5.1 从数值上确认正交化生效写代码容易但很难保证牛顿-舒尔茨迭代的真实效果符合预期。建议在训练脚本里单独加一个检查函数每个 epoch 检查一次二维参数的更新方向是否接近正交。def check_orthogonality(g: torch.Tensor): if g.dim() ! 2: return None if g.shape[0] g.shape[1]: matrix g.t().mm(g) target torch.eye(g.shape[1], dtypeg.dtype, deviceg.device) else: matrix g.mm(g.t()) target torch.eye(g.shape[0], dtypeg.dtype, deviceg.device) return (matrix - target).norm().item()理想情况下check_orthogonality的值应该小于 1并且不随训练明显增大。如果该值很大说明正交化没有生效需要检查newton_schulz的归一化逻辑或者确认参数是否确实被传进了矩阵分支。5.2 三路对比实验的设计要验证 MALT 是否真的同时具备“自适应”和“结构感知”能力最直接的做法是跑三组对照实验AdamW、Muon、MALT。每组实验固定相同的网络结构、数据切分、随机种子、batch size 和 epoch 数。对比时不要只比较最终 loss还要记录训练 loss 曲线。验证集准确率。二维梯度正交化误差。每步更新范数的稳定程度。单位时间吞吐量。MALT 相对 Muon 多出来的计算量应该主要来自逐元素除法和二阶矩更新这部分非常少。相对 AdamW 多出来的计算量来自牛顿-舒尔茨迭代这部分才是真实开销。如果ns_steps5并且权重矩阵很大吞吐量下降会很明显此时应该开启chunk_size。5.3 从日志与梯度范数定位 NaN 和震荡训练早期出现 NaN不要只盯着学习率。排查顺序可以这样走。首先看梯度范数。在optimizer.step()之前手动打印torch.norm(p.grad).item()。如果梯度本身就出现 NaN问题在网络、数据或 loss 计算不在优化器。其次看二阶矩。打印exp_avg_sq的min()。如果某个坐标的二阶矩长时间接近 0预条件会把梯度放大到巨大值导致更新溢出。此时可以把eps从 1e-8 提高到 1e-6。最后看预处理后的梯度范数。检查在newton_schulz之前和之后的范数变化。如果正交化前范数已经超过 1e8大概率是预条件分母太小不是正交化的问题。6. 工程落地常见问题排查6.1 常见问题速查表问题现象常见原因检查方式处理建议训练前几步 loss 直接变成 NaN学习率过大或预条件分母过小打印梯度和exp_avg_sq.min()降低 lr把 eps 调到 1e-6检查数据归一化损失曲线长期不下降正交化导致梯度方向失真用check_orthogonality查看梯度方向误差降低ns_steps检查是否错误地对 bias 执行了正交化更新方向变得异常小预条件后newton_schulz做了范数归一化打印torch.norm(update)这是预期行为步长应完全由 lr 控制大矩阵训练时显存暴涨牛顿-舒尔茨完整矩阵乘法峰值过高查看torch.cuda.max_memory_allocated()设置chunk_size128或chunk_size256和 AdamW 对比时 loss 略高正交化的几何约束限制了更新方向检查是否所有超参一致适当调大 lr或做学习率网格搜索换 batch size 后训练不稳定beta1、beta2 与 batch size 不匹配记录梯度跳变幅度大 batch 提高 beta1 到 0.98非平稳数据降低 beta2权重衰减没生效实现成 L2 正则而不是解耦衰减检查参数是否在做动量前被放大使用p.mul_(1 - lr * weight_decay)6.2 生产环境需要注意的额外事项生产环境不能只验证训练曲线。需要额外配上梯度裁剪、日志、checkpoint 和回滚机制。梯度裁剪建议加在optimizer.step()之前。MALT 的预条件已经能抑制大多数尺度问题但极端 batch、异常数据样本仍可能产生超大梯度。用torch.nn.utils.clip_grad_norm_把整体梯度范数限制在 1.0 左右即可。checkpoint 保存时建议把优化器的state_dict一起保存。优化器状态里保存了exp_avg和exp_avg_sq如果只保存模型权重中断后重启的效果会出现明显波动。恢复训练时要注意step也被恢复否则 bias correction 会重新从 0 计算。学习率调度建议使用 warmup。MALT 的更新方向在最初几百步内还不稳定直接使用大学习率容易破坏正交化的收敛过程。warmup 步数可以参考总训练步数的 1% 到 5%之后接 cosine 衰减。7. 最佳实践与扩展方向7.1 参数配置与学习率调度建议MALT 的默认参数适合作为起点但不适合直接上线。建议按下面顺序调参。先固定ns_steps5、chunk_size256。用一个小数据集跑 500 步观察 loss 是否稳定下降。如果 loss 震荡严重把lr降为原来的五分之一。如果下降太慢再逐渐提高 lr。调整beta1时要结合 batch size。batch size 越大每步梯度越接近全量梯度beta1可以设得大一些。小 batch 训练时建议从 0.9 起步。beta2控制二阶矩的滞后程度。默认 0.999 适合相对平稳的数据分布。如果模型需要持续适应新数据分布可以降到 0.99这样预条件能更快响应尺度变化。weight_decay建议从 0 开始测试需要正则化时再逐步提高到 0.01。Muon 系优化器对权重衰减通常比较敏感一次提高太多容易导致欠拟合。7.2 发布前检查清单上线前可以按下面清单逐项检查避免在长训练后被一个底层小问题浪费大量时间。[ ] 确认所有二维权重确实进入了矩阵分支bias 和归一化层参数进入 AdamW 分支。[ ] 确认newton_schulz的输入和输出形状一致。[ ] 确认exp_avg_sq没有更新到预条件之后的梯度。[ ] 确认weight_decay是解耦衰减而不是把weight_decay * w直接写进 loss。[ ] 确认 checkpoint 保存了优化器state_dict。[ ] 确认每个 step 后exp_avg和exp_avg_sq没有 NaN。[ ] 确认学习率 warmup 生效。[ ] 确认大矩阵场景下chunk_size已启用并用显存统计验证峰值可控。7.3 进一步扩展的三个方向第一个方向是面向更大矩阵的块对角预条件。当前实现只维护逐元素对角预条件后续可以考虑对同一行或同一列的梯度做分组归一化让预条件感知到局部结构。第二个方向是动态调整正交化步数。训练前期梯度方向变化剧烈可以适当提高ns_steps训练中后期模型基本稳定可以把步数降到 3 或 2节省计算时间。第三个方向是与低精度训练结合。MALT 的预条件会改变梯度尺度低精度混合精度训练时需要仔细检查二阶矩的分辨率是否足够。如果使用 bf16 训练eps可能需要调高否则极小的二阶矩值容易被舍入误差清零。MALT 给训练优化带来的核心价值不是“替换 Adam”而是提供了一条新的组合思路用对角预条件保留 Adam 的自适应能力用矩阵正交化保留 Muon 的结构感知能力。相比完整二阶方法它足够轻量相比普通逐元素优化器它多了一层矩阵几何约束。在遇到深层网络训练不稳定、需要更多结构信息又不想引入显式二阶矩阵时MALT 是一个值得放进实验清单的选项。