MALT 这种“轻量级曲率感知 Muon 优化器”最值得关注的地方不是名字本身而是它想同时解决两个问题让更新方向携带二阶信息同时保持普通自适应优化器的低计算开销。实现里的关键落点就是对角预条件。如果你之前主要用 Adam、AdamW或者正在训练中小规模模型这篇文章可以帮你理解 MALT 这类优化器适合什么场景以及自己动手复现时最容易在哪里翻车。我先说一个总体判断这类优化器实际测试起来比想象中麻烦因为“能跑”和“能收敛”是两件事。我习惯先把单条训练任务跑通再替换优化器观察 loss 曲线最后才看具体指标。下面按这个顺序展开。1. MALT 到底解决什么问题1.1 从名字拆一遍MALT 这个名字看起来更像一个组合标签而不是一个完整缩写。把它拆开看每个词都指向一个具体设计点。Lightweight 是最容易理解的目标计算和内存开销不能像二阶优化器那样高。Curvature-Aware 表示更新时希望感知损失曲面的弯折程度而不是只拿梯度方向硬走。Muon 则代表这个优化器可以归入 Muon 这类更新范式也就是不把原始梯度直接当作更新方向而是先做一些变换或缩放再更新参数。Diagonal Preconditioning 是落地方案用对角矩阵去近似曲率信息从而绕开完整 Hessian 矩阵的计算。把这几层串起来MALT 的核心问题就清楚了用对角预条件来逼近每个参数应该被放大的倍数以感知局部曲率同时控制在训练大模型时能接受的开销。这种设计思路并不算激进。Adam 其实也使用二阶矩做逐参数缩放本质上也是一种对角预条件。MALT 的特殊之处在于更强调“曲率感知”这一层并且希望保留 Muon 式更新的优势而不是简单复制 Adam 的公式。1.2 和 Adam 相比多了什么Adam 的基本更新规则很好理解维护梯度的一阶矩和二阶矩然后用梯度的一阶矩除以二阶矩的平方根再乘以学习率。这个操作让每个参数都有一个独立的学习率缩放所以 Adam 在稀疏梯度和非平稳目标上表现不错。MALT 从结构上很像 Adam但它在预条件上的解释会更靠近“曲率”而不是“梯度归一化”。也就是说它不只是为了让更新步长均匀而是把分母看作对损失函数局部曲率的近似。曲率大的方向更新的步子要小曲率小的方向更新的步子可以大一些。从实现层面看这种差异会体现在几个地方二阶矩的估计方式、一阶动量的处理、是否对矩阵参数使用额外的分组策略以及是否加入类似 scale 的超参数来控制对角预条件的强度。我不是说 MALT 一定会在这些地方全部和 Adam 不同而是说当你去读一个 MALT 风格优化器代码时应该重点看这些设计点而不是只看名字。1.3 和二阶优化器相比少了什么真正意义上的二阶优化器比如 K-FAC、牛顿法变体会构造并求逆一个大矩阵。这个矩阵保存了参数之间的大量相关性所以更新方向更准但内存和计算成本往往高得离谱。大模型训练基本扛不住。MALT 只保留对角线自然失去参数之间的相关性信息。它的预条件矩阵只能描述单个参数自身方向和幅度无法表达两个参数之间的联合曲率。这是对角预条件的老问题也是它“轻量”的代价。所以 MALT 更适合作为“比 Adam 更懂曲率、比完整二阶优化器更便宜”的中间方案。如果某个任务确实需要精细捕捉参数相关性那 MALT 这类方案可能不够需要回到更重的二阶方法。2. 跑通 MALT 之前环境和基准怎么准备2.1 硬件和依赖MALT 的硬件要求不会比 Adam 高太多。因为它本质上还是逐参数维护状态不会出现一个巨大的预条件矩阵。最基础的环境只需要一个支持自动求导的深度学习框架比如 PyTorch。CPU 也能验证单步更新是否正确但只要涉及真实训练我还是建议用 GPU。显存不需要很高很多 CV 分类任务用一块普通显卡就能跑。依赖方面除了框架本身最好准备一个日志工具。TensorBoard、WandB 都行没有的话也可以手动打印 loss。关键是你能同时看到梯度范数、更新范数和 loss 曲线否则很难判断优化器是不是正常工作。还需要确认 PyTorch 的版本。不同版本在混合精度、优化器 step 行为上有细微差别如果代码里出现了版本相关的 API最好先固定环境。2.2 用一个小模型先验证我第一次测试一个陌生优化器时不会直接上大模型。更稳妥的方法是先在简单任务上验证数据集用 CIFAR-10 或类似的小规模分类数据集。模型用一个小型 CNN 或两三层 MLP。先只跑 20 到 50 个 batch看 loss 能不能稳定下降。这个阶段不是看最终精度而是确认优化器在真实梯度上不会直接崩掉。如果 loss 一路变成 NaN说明更新规则或数值处理有问题不值得继续训练大模型。验证流程可以分成三步建一个最小模型随机初始化。用默认学习率跑少量 step打印 loss、梯度范数、更新范数。如果一切正常再替换成真实数据集和稍大模型。2.3 观察哪些指标观察优化器效果不能只看 loss 曲线。我会同时记录下面几个指标loss 是否稳定下降有没有周期性震荡。梯度范数是否正常不会出现从 1e-4 突然跳到 1e4。参数更新范数是否有界更新量过大说明学习率偏大。训练吞吐是否比 Adam 低太多如果低 30% 以上就要考虑实现方式是否存在额外开销。如果你能把这些指标记录下来后面排查问题会比较容易。否则当训练崩掉时你连“是学习率问题还是预条件问题”都说不清。3. 一个 MALT 风格优化器的可执行实现3.1 更新规则拆解这里给出的是一个简化版实现目的是帮助你理解 MALT 风格优化器的核心公式不代表官方实现。假设第 t 步参数梯度为 g_t。首先维护一阶动量 m_t 和二阶矩 v_tm_t beta1 * m_{t-1} (1 - beta1) * g_tv_t beta2 * v_{t-1} (1 - beta2) * g_t^2在真正更新参数之前需要做偏差校正让初期估计更准确m_hat_t m_t / (1 - beta1^t)v_hat_t v_t / (1 - beta2^t)最终更新theta_t theta_{t-1} - lr * m_hat_t / (sqrt(v_hat_t) eps)这里的 1 / (sqrt(v_hat_t) eps) 就是对角预条件。它逐参数调整梯度缩放曲率大的方向自然会被压制。如果要更贴近 Muon 式更新可以考虑增加一个 curvature_scale 参数对预条件强度做整体缩放。这个参数的作用是控制我们在多大程度上信任二阶信息。scale 越大更新越依赖曲率缩放scale 趋近于 0则退化成带动量的 SGD。3.2 PyTorch 代码下面是一个简化版本。它继承了 PyTorch 的 Optimizer适合直接替换到现有训练循环里。import math import torch from torch.optim import Optimizer class MALT(Optimizer): def __init__( self, params, lr1e-3, betas(0.9, 0.999), eps1e-8, weight_decay0.0, curvature_scale1.0, ): if lr 0: raise ValueError(lr must be positive) if not 0.0 betas[0] 1.0 or not 0.0 betas[1] 1.0: raise ValueError(betas must be in [0, 1)) defaults dict( lrlr, betasbetas, epseps, weight_decayweight_decay, curvature_scalecurvature_scale, ) 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: beta1, beta2 group[betas] lr group[lr] eps group[eps] weight_decay group[weight_decay] scale group[curvature_scale] for p in group[params]: if p.grad is None: continue grad p.grad if weight_decay ! 0: grad grad.add(p, alphaweight_decay) 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_t state[step] exp_avg.mul_(beta1).add_(grad, alpha1 - beta1) exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value1 - beta2) bias_correction1 1 - beta1 ** step_t bias_correction2 1 - beta2 ** step_t denom (exp_avg_sq.sqrt() / math.sqrt(bias_correction2)).add_(eps) step_size lr / bias_correction1 * scale p.addcdiv_(exp_avg, denom, value-step_size) return loss这段代码里有两个重点。第一个是addcmul_它用来累积梯度平方也就是二阶矩。第二个是addcdiv_它执行最终更新用一阶动量除以对角预条件。3.3 参数解释学习率 lr 是最敏感的参数不能直接照搬 Adam 的设置。由于预条件缩放可能让部分参数更新更激进MALT 风格优化器通常需要更保守的学习率。beta1 控制一阶动量通常取 0.9。beta2 控制二阶矩通常取 0.999。beta2 越大对角线曲率估计越稳定但反应越慢。如果在某些任务上 loss 震荡严重可以试着把 beta2 调小到 0.99。eps 是防止分母为零的常数同时还起到限制最大更新幅度的作用。eps 太小会让预条件在平坦区域产生极大步长容易导致训练不稳定。weight_decay 在示例代码里以 L2 正则方式加入梯度。这种做法简单但不完全等同于 AdamW 的 decoupled weight decay。如果要严格对比我建议实现一个 decoupled 版本把权重衰减从梯度计算中分离出来。curvature_scale 是本示例额外加的缩放超参。它让你能够整体控制对角预条件的影响。当 scale 为 1 时更新规则接近 Adam当 scale 减小时更新更接近 SGD。这个参数可以作为调参入口用来测试曲率信息在当前任务上是否有正面作用。3.4 单步更新验证安装好 PyTorch 后可以用下面这个小脚本验证优化器能否正常更新参数import torch from torch import nn model nn.Linear(4, 1) opt MALT(model.parameters(), lr1e-3) x torch.randn(8, 4) y torch.randn(8, 1) loss nn.functional.mse_loss(model(x), y) opt.zero_grad() loss.backward() opt.step() print(loss:, loss.item())如果优化器状态和参数更新正常这一步不会报错并且 loss 会是一个有效数值。你可以在 step 前后分别打印model.weight的 norm确认参数确实发生了变化。跑这个小脚本只需要几秒。它能帮你排除最基础的 API 使用错误再进入真实训练测试。4. 训练中的关键参数和效果判断4.1 学习率怎么给我给 MALT 风格优化器设置学习率时一般不会直接取 Adam 的默认值。原因是它对角预条件之后更新方向会被缩放实际步长可能和 Adam 差别很大。比较稳妥的做法是先取一个较小的学习率比如 1e-4。跑到 100 到 200 个 step观察 loss 是否下降。如果 loss 不变再逐步调大学习率但每次只调整 2 倍左右。如果 loss 出现发散或大幅震荡立刻降回来。不要一上来就用超大学习率测试因为很多二阶风格优化器在初期估计不准时更新量会非常大。如果你在某个任务上已经知道 AdamW 的最佳学习率可以先用它的 1/3 到 1/2 作为起点。这个经验不能保证适用所有情况但比瞎猜要快。4.2 权重衰减怎么处理示例代码里的 weight_decay 是直接加到梯度上的这是传统 L2 正则方式。但近年主流训练更多使用 decoupled weight decay也就是 AdamW 的做法在更新参数时把权重衰减单独从梯度里分离出来。如果你想严格复现某个实验建议改成分离式权重衰减。可以这样处理if weight_decay ! 0: p.mul_(1 - lr * weight_decay)然后在梯度计算里不再加入 weight_decay。这样权重衰减不会受到预条件缩放影响超参数更容易迁移。注意不同任务对 weight_decay 的敏感度差别很大。CV 分类任务常用 5e-4 到 1e-2Transformer 训练常用 0.01 到 0.1。先用任务原本的推荐值再根据验证集指标微调。4.3 曲线和成功率怎么看训练一个模型跑完几十个 epoch 后光看最终 loss 不够。我会额外关注两个东西优化器替代是否带来了“早期收敛加速”以及最终效果是否不低于基线。如果 MALT 在早期几个 epoch 里 loss 下降明显更快但后期曲线趋于平缓甚至被 AdamW 追平那么它更适合用于训练前期加速。如果最终精度比 AdamW 低很多那即使收敛快也不值得在生产环境中替换。你可以在同一份数据和同一套超参数下固定随机种子分别跑 AdamW 和 MALT并记录每 20 个 step 的 loss。我一般会跑 15 到 20 个 epoch 再看结论而不是只跑 1 个 epoch 就下判断。4.4 Batch size 和梯度裁剪对角预条件优化器对 batch size 的敏感度不一定和 Adam 一样。小 batch 时梯度噪声大预条件估计可能不稳定大 batch 时曲率估计更准但单步更新也可能更激进。如果训练 Transformer 或大模型我强烈建议保留梯度裁剪。你可以先设 max_grad_norm 1.0 或 5.0。梯度裁剪不会解决所有问题但能防止个别异常梯度把预条件状态污染。如果你在测试不同 batch size不要同时改学习率和梯度裁剪。一次只改一个变量否则你很难判断到底是哪个因素导致效果变化。5. 常见问题排查5.1 Loss 爆炸或变成 NaN这个问题最容易出现也最容易误判。很多人第一反应是优化器代码写错了但实际更多是学习率太大、输入里有 NaN或者预条件分母太小。我的排查顺序是检查输入数据确认没有 NaN 或 Inf。打印梯度看梯度本身是否已经出现异常。降低学习率至少降 10 倍看问题是否消失。检查混合精度配置确认优化器更新是在 fp32 下完成的。最后再检查优化器实现比如二阶矩是否被错误清零或者 eps 是否设置得太小。不要跳过前两步。如果输入本身有问题换任何优化器都会崩。5.2 收敛过慢如果你把学习率设得很小确实可能不崩但收敛会很慢。这时候不要急着把学习率提高先看一下梯度范数和更新范数的比例。如果梯度范数很正常但参数更新范数偏小说明预条件缩放把更新压得太狠。可以把 curvature_scale 调大一点或者减小 eps。如果更新范数和学习率变化对不上再检查代码里是否有额外的缩放逻辑。还有一点容易忽略bias correction 是否生效。如果你的实现没有做偏差校正前几百步的预估会偏小可能表现为早期更新迟缓。5.3 混合精度下不稳定混合精度训练时如果用 AMP梯度会先缩放到 fp16再反缩放回 fp32。如果二阶矩在 fp16 下累积可能因为下溢或上溢导致数值不稳定。我建议把优化器状态保持在 fp32全部更新操作都在 fp32 下完成。PyTorch 的 AMP 通常不会改变优化器内部状态精度但如果你自己写了 kernel 或自定义逻辑需要格外小心。如果你在混合精度下跑出 NaN可以先关闭 AMP 测试。如果关闭后一切正常那大概率是梯度缩放或参数更新精度问题。5.4 什么时候不要用这个方案MALT 这类轻量曲率感知优化器并不适合所有场景。如果你只有一个非常短的训练任务比如只跑 20 个 step 来测试 pipeline直接用现成 AdamW 更省事没有必要引入新优化器。如果你的任务里大多数参数是稀疏的比如超大 embedding 表用简化的对角预条件可能效果一般。此时更适合使用针对稀疏梯度设计的状态管理方式。如果你的训练目标是复现某个官方 baseline别为了新鲜感换优化器。先跑通原配方再单独做消融实验。6. 进阶落地建议6.1 单卡到多卡MALT 这类优化器在多卡训练下没有额外的通信成本。DDP 只会同步参数梯度优化器在每个 rank 上维护各自的状态。只要梯度同步正确优化器本身不需要做特殊处理。需要注意的是如果你在多卡下使用不同 batch size实际梯度噪声会不同。最优学习率也可能随之变化。跨卡对比时要保证 batch size 全局一致或者明确记录单卡 batch 和梯度累积步数。如果使用 FSDP 或 Sharded Optimizer还要注意优化器状态的分片保存和恢复。恢复 checkpoint 时如果 step 计数不对偏差校正会错得很离谱。6.2 和现有训练框架结合比较合适的接入方式是把 MALT 封装成标准 Optimizer再配合 scheduler、AMP、DDP 一起使用。这样可以最小化改动。例如model nn.Linear(10, 2) optimizer MALT(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50)训练循环里调用optimizer.step()和scheduler.step()的顺序要保持固定。常见的顺序是 loss.backward - optimizer.step - scheduler.step或者按框架惯例来。如果你做的是研究型实验建议把优化器单独写成一个文件不要混在训练循环里。这样更换参数和复现实验都更方便。6.3 稳定性和可重复性我实际跑下来的感受是这类优化器的结果可重复性还不错但前提是固定好随机种子和数据顺序。需要保存的优化器状态至少包括一阶动量、二阶矩和 step 计数。断点续训时只保存 model 而忘记 optimizer statebias correction 会从 0 重新开始导致恢复后的前几百步更新异常。如果你要发论文或做对比实验最好把每次实验的优化器参数、学习率 schedule、数据增强方式一起记录下来。否则两个实验之间产生差异时你根本没法判断是随机噪声还是优化器改动导致的。6.4 我个人更建议的做法如果你只是想快速验证某个任务有没有用最简单的办法是写一个对比脚本一份用当前任务的默认 AdamW一份换成 MALT其余完全一致。先跑短时间看 loss 和验证指标。如果 MALT 在早期收敛或最终指标上有任何一方优势再深入调试参数。如果没有明显优势也不要强行调参。优化器只有在合适任务和合适超参下才有价值不是所有问题都需要换新方法。我自己更倾向于把这类轻量级曲率感知优化器当成训练工具箱里的一个备选而不是默认选择。真正让它发挥价值的地方通常是那些存在明显方向差异、曲率不均匀的任务比如 Transformer 或某些多层异质网络。遇到这类场景MALT 的思路值得你花时间跑一轮对比。实际落地时最该盯住的不是功能列表而是输入格式、资源占用、数值稳定性和失败重试。优化器代码再漂亮只要训练日志不明、checkpoint 不完整最终都很难用于生产。