
这次我们来看 MALT。论文标题是 “MALT: Lightweight Curvature-Aware Muon via Diagonal Preconditioning”一句话概括它给 Muon 优化器做了一次轻量化改造用对角预条件器去近似曲率信息绕开 SVD、Newton-Schulz 迭代这类重计算目标是让 Muon 路线的优化器在训练大模型时更省显存、更省算力。这个方向最近关注度很高。现在不少前沿模型训练开始从 AdamW 转向 Muon 类优化器但 Muon 的预条件器构造和正交化步骤开销并不小。MALT 解决的问题很明确在保留 Muon 类收敛优势的前提下把额外开销压到最低。如果你关心大模型训练优化器选型正在复现 1B 以内的小规模预训练实验或者想在老显卡上跑更大的 batch这篇文章值得看完。这篇文章会做四件事先讲清楚 Muon 为什么有效、瓶颈在哪再拆解 MALT 的对角预条件思路然后给出 PyTorch 环境下的接入方式和伪代码示意最后给出一套可执行的验证实验和排错清单。MALT 不是一个带 WebUI 的应用它没有启动页面也不监听端口而是一个需要集成进训练脚本的优化器实现所以整篇文章的实操重点会放在算法理解、代码接入和实验设计上。1. 核心能力速览能力项说明项目类型深度学习优化器曲率感知优化算法技术方向Muon 优化器的轻量化改进对角预条件器替代矩阵正交化主要功能大模型训练优化、梯度缩放、显存与计算开销降低依赖框架通常是 PyTorch / JAX 生态具体以官方源码为准显存需求不直接决定总占用额外预条件器开销低于矩阵类方法启动方式无需独立服务作为 optimizer 集成进训练脚本接口能力暴露 Python 优化器接口接入 PyTorch Optimizer 或 JAX Optax批量任务面向大批量、多卡分布式训练场景适合场景预训练、大 batch 微调、资源受限的模型训练实验上手难度中等难点在于超参数对齐和实验验证需要强调一点MALT 的额外显存和计算开销到底比 Muon 低多少必须以论文实验表格和本机测试为准。下面所有分析都是基于论文标题、Muon 类优化器的公开讨论和优化器设计常识做的定性拆解不替代官方源码和原文数据。2. 为什么 Muon 有效又为什么沉重要理解 MALT先要理解 Muon 优化器。Muon 是最近在开源社区和前沿模型训练中频繁出现的一类优化器核心设计可以拆成三步第一步对梯度做零中心化。Transformer 的权重大多是矩阵形状Muon 会去掉每个权重矩阵行方向上的均值让梯度在矩阵乘法意义上保持“居中”。第二步构造右乘预条件器。这一步不是像 AdamW 那样把每个参数元素单独缩放而是试图捕捉输入维度之间的相关性用一个矩阵形式的预条件器去缩放梯度。第三步做正交化。常用做法是 Newton-Schulz 迭代或者通过 SVD、QR 分解把预条件结果拉回正交附近。这个操作本质上是让更新方向保持在该矩阵几何下的合理尺度。这套设计的直觉是Transformer 中有大量线性层每一层都是一个矩阵乘法。AdamW 的逐元素缩放忽略了维度间的相关性而 Muon 类方法显式建模了这种相关性所以在大模型上经常出现“同样的 token 预算下Muon 比 AdamW 更快到达目标 loss”的观察。问题出在开销上。Muon 的实际部署成本来自三个方面第一预条件器构造需要计算矩阵统计量。如果要保存一个和权重矩阵尺寸匹配的预条件矩阵显存会随矩阵维度平方增长。对一个 4096×4096 的矩阵额外保存一个同类矩阵就是 128MB 起步这还没有算中间计算量。第二Newton-Schulz 迭代需要多轮矩阵乘法。每轮迭代都是与权重形状相同的大矩阵乘法在深层模型中会显著拉长单 step 时间。第三分布式场景下通信和计算需要额外同步。Muon 类优化器的预条件器和正交化步骤在张量并行或流水线并行场景里都需要额外的通信设计工程复杂度比 AdamW 高不少。这正是 MALT 这类“轻量级 Muon 变体”的价值空间如果对角预条件器就能近似出值得用的曲率信息那就不需要完整矩阵预条件器也不需要显式正交化整体开销可以降一个量级。3. MALT 核心思路用对角预条件器做曲率感知从论文标题可以拆出三个关键词Lightweight、Curvature-Aware、Diagonal Preconditioning。先说 curvature-aware。曲率感知的意思是优化器应该知道损失曲面在当前参数位置各个方向的弯曲程度。高曲率方向说明损失函数在这个方向变化剧烈步长要小低曲率方向说明损失变化平缓步长可以大。AdamW 里的二阶矩本质上就是一种逐元素的曲率估计但它只知道每个参数自身的历史梯度量级不知道参数之间的相关性。再说 diagonal preconditioning。Muon 用完整矩阵预条件器MALT 把它限制成对角矩阵。也就是说MALT 不去计算完整的维度间相关矩阵而是只估计一个按元素缩放的对角向量。这个对角向量可以来自当前梯度的二阶统计量也可以来自历史统计量的移动平均。最后是 lightweight。因为没有矩阵形式的相关性估计MALT 不需要 Newton-Schulz 迭代不需要 SVD也不需要 QR 分解。额外保存的状态只有与参数形状相同的预条件向量内存增量大幅收敛。从论文的命名习惯推断MALT 至少有两种可能的变体设计MALT-S。S 可能代表 Single即使用单步梯度的统计量直接构造对角预条件器。这种变体开销最低每一步只有一次对角线统计计算但单步梯度噪声较大曲率估计的稳定性会受影响。MALT-M。M 可能代表 Multi-step 或 Moving Average即在多个迭代步上对对角统计量做指数移动平均。这样得到的曲率估计更稳定代价是需要多保存一份状态向量额外的显存占用仍然是对角级的。具体公式和变体定义需要对照论文原文但可以确定的是MALT 的目标是在 Muon 的“强曲率感知”和 AdamW 的“低额外开销”之间取一个折中点。下面给出一个定性对比帮助你快速理解 MALT 的定位优化器预条件方式额外计算量额外显存曲率感知强度AdamW一阶矩 二阶矩逐元素 EMA低低弱逐元素Muon矩阵预条件器 Newton-Schulz 正交化高中高强能捕捉维度相关性MALT对角预条件器低低中等近似曲率信息这张表是定性判断不是实测数据。MALT 的对角预条件器能不能在特定模型上逼近 Muon 的效果要跑实验说话。4. 环境准备与项目集成方式MALT 是一个优化器没有 WebUI不需要一键启动脚本也不需要关心端口占用。它的“启动”就是被你的训练脚本 import 进优化器位置。所以环境准备的核心是保证训练代码能正常跑并且把 MALT 当作普通优化器接入。前置检查清单如下操作系统Linux 或 macOS 均可大规模训练建议 Linux。Python 版本以官方仓库要求为准通常是 Python 3.10 或更高。深度学习框架PyTorch 2.x 或 JAX需要确认 MALT 官方实现依赖哪个框架。CUDA 环境如果只有 CPU 环境也可以用极小模型验证算法逻辑但无法得到有意义的性能结论。显卡优化器本身不挑显卡型号只要 PyTorch 能正常调用 CUDA 即可。老显卡也能跑瓶颈在模型规模和 batch size。磁盘空间源码和依赖占几 GB 以内模型权重和训练数据集另算。安装方式通常是先克隆官方仓库再以可编辑模式安装git clone 官方仓库地址 cd malt-optimizer pip install -e .注意这里的仓库地址、包名和安装命令需要用官方 README 替换。不同实现可能叫malt、malt-optimizer或别的名字不要照抄命令就执行。验证安装是否成功可以直接在 Python 里确认 import 是否通过import malt_optimizer # 如果类名不确定打印模块下的公开属性 print([name for name in dir(malt_optimizer) if not name.startswith(_)])如果看不到 MALT 相关类名说明安装不完整需要回头检查依赖版本。5. Python 接口实现与训练循环接入MALT 的价值最终要体现在训练循环里。作为一个优化器它需要实现标准的optimizer.step()接口。下面的代码是一般性示意用来展示“零中心化 对角预条件器”的核心框架不是论文逐行源码也不保证与官方实现完全一致。接入前应以官方实现为准。import torch class MALT(torch.optim.Optimizer): def __init__(self, params, lr1e-3, betas(0.9, 0.95), eps1e-8, use_centeringTrue): defaults dict(lrlr, betasbetas, epseps, use_centeringuse_centering) super().__init__(params, defaults) def step(self, closureNone): loss None if closure is not None: loss closure() for group in self.param_groups: beta1, beta2 group[betas] lr group[lr] eps group[eps] for p in group[params]: if p.grad is None: continue grad p.grad.data state self.state[p] if len(state) 0: state[step] 0 state[exp_avg_sq] torch.zeros_like(grad) state[step] 1 exp_avg_sq state[exp_avg_sq] # 零中心化对二维以上权重做行均值消除 if grad.dim() 2 and group[use_centering]: grad grad - grad.mean(dim-1, keepdimTrue) # 对角二阶矩 EMA exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value1 - beta2) # 对角预条件更新 denom exp_avg_sq.sqrt().add_(eps) p.data.addcdiv_(grad, denom, value-lr) return loss这段代码非常直观对矩阵梯度做零中心化维护一个与参数同形状的二阶矩 EMA然后按元素缩放梯度更新。它具备了对角预条件器的基本形态但 MALT 论文中的预条件器构造可能更精细可能包含对梯度范数的归一化、特定的曲率估计器或不同的统计方式。接入训练循环的模板如下import torch model get_model() optimizer MALT(model.parameters(), lr1e-3) for step, batch in enumerate(train_loader): optimizer.zero_grad() loss model(batch, labelsbatch) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() if step % 100 0: print(fstep {step}, loss {loss.item():.4f})这里需要注意三个容易踩坑的地方第一个zero_grad()一定要在 loss 计算之前调用否则梯度会累计。第二个梯度裁剪顺序必须在loss.backward()之后、optimizer.step()之前。先裁剪再更新避免异常梯度直接修改权重。第三个学习率不要直接照搬 AdamW 的配置。Muon 类优化器通常对 lr 更敏感建议从1e-3起步做小范围扫描而不是一上来就试3e-4或1e-4。如果你想继续扩展也可以顺手封装一个 callback在训练过程中记录参数更新范数和梯度范数方便定位收敛问题def log_grad_and_update_norm(model, prev_params, step): total_grad_norm 0.0 total_update_norm 0.0 for name, p in model.named_parameters(): if p.grad is not None: total_grad_norm p.grad.norm().item() ** 2 update p.data - prev_params[name] total_update_norm update.norm().item() ** 2 print(fstep {step} | grad_norm {total_grad_norm ** 0.5:.6f} | update_norm {total_update_norm ** 0.5:.6f})6. 功能测试与效果验证方案优化器不是一个能“跑出张图”或“识别一段文字”的模型应用它的功能验证必须通过训练实验完成。下面是一套可复现的验证流程比较适合第一次接触 MALT 的时候使用。第一步固定模型和数据。建议选择一个 50M 到 200M 参数的 Transformer 模型数据集固定为一个可复现的中型语料。不要同时更换模型架构和数据否则优化器之间的差异会被其他变量淹没。第二步固定训练配置。batch size、学习率调度、warmup、梯度裁剪、种子、token 预算全部保持相同。只有优化器本身不同。第三步跑三个实验python train.py --optimizer adamw --lr 3e-4 --batch_size 32 python train.py --optimizer muon --lr 1e-3 --batch_size 32 python train.py --optimizer malt --lr 1e-3 --batch_size 32数据要记录至少四个指标训练 loss 曲线重点看下降速度。验证集 loss重点看最终水平。训练吞吐例如 tokens/s 或 samples/s。峰值显存用torch.cuda.max_memory_allocated()统计。第五步判断 MALT 是否值得用。判断标准不要只看 loss 最终值要看“在相同 wall-clock 时间内 MALT 是否更快到达目标 loss”。如果 MALT 每个 step 更快但需要明显更多的 step 才能收敛那总时间不一定占优。如果 MALT 在更短时间内到达相同或更低的 loss同时额外显存不高于 AdamW就说明这个优化器在你的场景下有效。测量显存的代码可以参考import torch torch.cuda.reset_peak_memory_stats() # 跑一个完整的 training step loss.backward() optimizer.step() peak_memory torch.cuda.max_memory_allocated() / 1024**3 print(fPeak GPU memory: {peak_memory:.2f} GB)建议把峰值显存记录脚本化和自动保存。手动记录容易漏掉中间的峰值而且很难在不同实验中保持一致。判断成功的标准建议写成表格指标AdamW 基线Muon 基线MALT 实验结论到达 3.0 loss 所需 step 数待测待测待测对比下降速度最终验证 loss待测待测待测对比收敛质量吞吐 tokens/s待测待测待测对比单 step 开销峰值显存待测待测待测对比资源占用7. 资源占用与性能观察方法优化器的资源开销看起来隐蔽但大模型训练中非常关键。AdamW 之所以长期是默认选择就是因为它的额外状态只有一阶矩和二阶矩两个张量随参数量线性增长。Muon 的问题是预条件器可能引入平方级开销而 MALT 的设计目标就是把这部分拉回线性级。实操中观察优化器开销有几个简单方法。第一个是看 optimizer state 的显存占用。在训练脚本里统计每个参数的 state 张量形状def count_optimizer_state(optimizer): total_bytes 0 for state_dict in optimizer.state.values(): for key, tensor in state_dict.items(): if torch.is_tensor(tensor): total_bytes tensor.numel() * tensor.element_size() return total_bytes / 1024**3 print(fOptimizer state memory: {count_optimizer_state(optimizer):.3f} GB)如果 MALT 实现确实是对角预条件state 里应该只有一阶统计量或二阶统计量形状与参数一致。如果出现矩阵形状的 state那就说明实现里引入了完整矩阵预条件器不能算是严格意义上的对角化。第二个是看单 step 时间。不要用整个训练流程的平均时间因为数据加载和日志打印会干扰。可以在训练循环里单独计时import time start time.time() loss.backward() optimizer.step() torch.cuda.synchronize() elapsed time.time() - start第三个是观察峰值显存。优化器的额外显存可能不是峰值来源峰值往往来自激活值和 optimizer state 叠加后的最大值。所以要在训练开始前先reset_peak_memory_stats()在若干个 step 后再读取。第四个是注意 DDP 场景下的通信开销。Muon 类优化器在分布式训练中可能涉及额外通信MALT 如果保持对角化通信量与 AdamW 接近更容易套用现有 DDP 逻辑。但这一点取决于官方实现是否做了特殊通信设计要看源码。性能观察的核心思路是分别测量“收敛质量”和“单位时间的训练进度”不要只盯着 loss 曲线也不要只盯着 step 速度。8. 常见问题与排查方法MALT 作为新优化器实际使用时最容易遇到的问题主要有这些问题现象可能原因排查方式解决方案训练 loss 不下降学习率设置不当看梯度范数和参数更新范数从 1e-3 起步做 lr 扫描显存高于 AdamW 很多预条件器实现不是对角级检查 optimizer state 张量形状确认官方实现是否纯对角化训练出现 NaN数值不稳定或学习率过高检查 grad norm、缩放因子增加 eps降低 lr加 grad clipDDP 同步后不收敛梯度同步顺序或随机种子问题核对 zero_grad 与 DDP hook 顺序固定 seed按 PyTorch 标准顺序写与论文结果差距大超参数和 token 预算不同对比 batch、lr schedule、tokens统一实验配置再对比单 step 时间比预期慢预条件器计算未充分向量化profile step 时间检查是否调用了循环逐层计算性能忽高忽低显存碎片或数据加载波动观察多 step 时间曲线固定数据加载线程增加缓存逐个展开说。loss 不下降是最常见的。新优化器往往对 lr 范围敏感不要直接照搬 AdamW 的 lr。建议用 1e-4、3e-4、1e-3、3e-3 四组做一次小规模 lr 扫描每组只跑几百步画出 loss 曲线。出现 NaN 时先看梯度范数。如果梯度范数在爆炸前就非常大优先降低 lr 并开启梯度裁剪。MALT 这类带有曲率感知的优化器默认 eps 如果太小可以在数值不稳定时把 eps 从 1e-8 提高到 1e-6 或 1e-4。显存异常偏高时不要只看总显存要拆解成模型参数、optimizer state、激活值三部分分别统计。如果 optimizer state 中出现矩阵形状的张量说明当前实现的预条件器不是对角级的。想排查实现逻辑问题最直接的办法是把 MALT 和 AdamW 在同一个极小模型、同一份数据上做几次前向反向对比每一步更新方向的余弦相似度。这个对比能快速验证 MALT 是否真的在按预期方式缩放梯度。9. 最佳实践与合规注意MALT 这类优化器是否值得引入核心是实验成本和收益的权衡。下面几条实践建议可以直接用。第一次接触先在小模型上验证。建议用 50M 到 200M 参数的模型batch size 可以大一些跑一个固定 token 预算的实验。不要一开始就上 1B 以上模型新优化器的超参坑还没有摸清时大规模实验只会浪费时间和算力。保留 AdamW 基线。每个实验都可以在相同的模型和数据配置下跑一个 AdamW 对照。没有基线的优化器实验结果很难判断好坏。固定 token 预算而不是固定 epoch 数。同一个数据集上不同优化器的收敛步数不一样固定 epoch 会掩盖吞吐差异。用固定 token 数更公平。记录每一步的梯度范数和更新范数。这是判断优化器是否正常工作的最直接指标。如果梯度范数震荡剧烈或者更新范数突然跳变大概率是 lr 或预条件器参出了问题。学习率和 schedule 分开调。先把恒定 lr 下的行为摸清楚再加 warmup 和 cosine schedule。同时调整多个超参数时troubleshooting 会非常困难。如果是分布式训练先在单卡上跑通 MALT再切 DDP。批量切换优化器和分布式训练方式会让问题难以定位。MALT 如果保持对角预条件通信开销应该与 AdamW 接近但要确认官方实现的 DDP 兼容性。合规方面也要注意。使用论文和开源实现时要检查许可证确认是否允许商用和修改。在技术分享中引用 MALT 的算法思想要标明论文作者和出处。如果是用于内部项目或对外产品需要评估优化器实现是否包含第三方依赖以及这些依赖的许可证限制。另外任何训练实验都要注意数据版权和隐私合规。不要使用未经授权的数据做训练不要在自己的业务数据上直接跑来源不明的代码。优化器代码本身不直接涉及数据风险但训练流程中的数据合规问题不会因为换了优化器而消失。公开分享对比结果时只对同一个固定配置下的实验做结论不要夸大某个优化器的优势。不同的模型规模、batch size、lr schedule 下优化器表现可能完全不同。10. 总结与下一步MALT 这个项目最值得尝试的地方是它给出了一个很轻的“Muon 替代品”方向。如果 Muon 的正交化开销在你当前任务上显存或计算不堪重负MALT 的对角预条件思路值得认真测一测。它解决的是真实痛点大模型训练里优化器额外开销的每一分都可能转化为更小的 batch、更低的吞吐和更长的实验周期。建议第一次使用时优先验证三个方面第一确认官方实现是否真的做到对角预条件。检查 optimizer state 的形状如果出现矩阵张量说明实现和论文描述可能不一致。第二跑一组 100M 参数级别的小模型对比。用固定 token 预算对比 AdamW、Muon、MALT 的 loss 曲线、吞吐和峰值显存。这一步能最快判断 MALT 在你场景下是否值得替代现有优化器。第三从1e-3附近开始做 lr 扫描。新优化器最常见的坑就是直接套用 AdamW 的 lr导致 loss 不下降或震荡。调超参数时先固定其他变量一次只改一个数字。在正式任务中使用 MALT 之前还需要确认官方实现是否支持 DDP、FSDP 或你正在使用的分布式框架以及是否经过了大规模训练的稳定性验证。如果官方仓库提供了与 Muon 或 AdamW 的对比脚本建议直接复用减少自己搭建实验的成本。后续可以继续关注 MALT 在长上下文训练、多卡大规模预训练和低显存微调场景下的表现。先把小实验跑通再决定要不要让 MALT 进入正式训练管线。