训练神经网络时,优化器选择的本质是在更新方向、步长与计算成本之间做权衡。Adam 擅长给每个坐标提供自适应步长,但它在更新时没有利用权重矩阵本身的结构;Muon 这类优化器通过矩阵正交化让更新方向更接近正交,却对梯度的绝对尺度非常敏感。MALT(Muon 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 矩阵参数分支的更新公式
对于形状为二维的参数p,MALT 每个 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 非常接近,含义也基本一致。
| 超参 | 常用默认值 | 作用 | 调整建议 |
|---|---|---|---|
lr | 1e-3 左右 | 全局学习率 | 比 SGD 更敏感,建议配合 warmup |
beta1 | 0.95 | 一阶动量指数衰减系数 | 小 batch 用 0.9,大 batch 可调高到 0.98 |
beta2 | 0.999 | 二阶矩指数衰减系数 | 数据非平稳时可降到 0.99 |
eps | 1e-8 | 数值稳定项 | 如果 loss 出现 NaN,可提高到 1e-6 |
weight_decay | 0.01 | 解耦权重衰减 | 过大容易欠拟合,从 0 开始调 |
ns_steps | 5 | 牛顿-舒尔茨迭代步数 | 3 到 7 之间通常够用 |
chunk_size | None | 正交化分块大小 | 大矩阵建议 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], dtype=grad_matrix.dtype, device=grad_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)。它是求解极分解中正交因子的经典牛顿迭代。注意迭代前必须归一化,否则矩阵范数过大时,迭代会发散。steps=5在大多数场景是性能和精度的折中。
4.2 MALT 优化器主体
下面实现一个继承torch.optim.Optimizer的 MALT 优化器。为了方便切换,矩阵参数走 Muon 风格,非矩阵参数走 AdamW 风格。
import torch from torch.optim import Optimizer class MALT(Optimizer): def __init__( self, params, lr=1e-3, betas=(0.95, 0.999), eps=1e-8, weight_decay=0.0, ns_steps=5, chunk_size=None, ): if not 0.0 <= lr: raise ValueError(f"Invalid lr: {lr}") if not 0.0 <= eps: raise ValueError(f"Invalid eps: {eps}") if not 0.0 <= betas[0] < 1.0: raise ValueError(f"Invalid beta1: {betas[0]}") if not 0.0 <= betas[1] < 1.0: raise ValueError(f"Invalid beta2: {betas[1]}") if weight_decay < 0.0: raise ValueError(f"Invalid weight_decay: {weight_decay}") defaults = dict( lr=lr, betas=betas, eps=eps, weight_decay=weight_decay, ns_steps=ns_steps, chunk_size=chunk_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, steps=ns_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, steps=ns_steps)) return torch.cat(chunks, dim=1) @torch.no_grad() def step(self, closure=None): 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, value=1 - 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, alpha=1 - beta1) update = exp_avg else: # AdamW 回退分支 exp_avg.mul_(beta1).add_(grad, alpha=1 - 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_dim=784, hidden=512, num_classes=10): 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(), lr=1e-3, betas=(0.95, 0.999), eps=1e-8, weight_decay=0.01, ns_steps=5, chunk_size=256, ) 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_none=True) loss.backward() optimizer.step() print(f"epoch {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], dtype=g.dtype, device=g.device) else: matrix = g.mm(g.t()) target = torch.eye(g.shape[0], dtype=g.dtype, device=g.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_steps=5并且权重矩阵很大,吞吐量下降会很明显,此时应该开启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_size=128或chunk_size=256 |
| 和 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_steps=5、chunk_size=256。用一个小数据集跑 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 是一个值得放进实验清单的选项。