状态空间模型(SSM)在长序列建模上提供了一条不同于 Transformer 的技术路线。Mamba 通过选择性扫描机制让模型能够根据输入决定“记住什么、忘掉什么”,并在推理阶段保持线性复杂度,这是它被广泛讨论的核心原因。真正把 Mamba 从论文变成可用模型时,难点往往不在网络结构本身,而在训练稳定性:状态转移矩阵 A 的谱性质、投影矩阵的更新方向、以及长序列下梯度如何传播,这些都会直接影响模型能否记住远程信息。2025 年开源社区讨论较多的 Muon 优化器,把方向更新和正交化放到了优化器层面,用 Newton-Schulz 迭代保持权重更新在正交意义下的良好形态,正好适合这类对矩阵结构敏感的模型。这篇文章沿着“Muon meets Mamba”的路线,讲清楚状态空间模型为什么需要谱优化、Muon 的核心机制是什么,以及如何用最小 PyTorch 实现把两者组合起来,在一个确定性记忆任务上完成训练和对比。
这里不打算停留在公式层面,而是给出可运行的简化 Mamba 模块、可复用的 Muon 优化器实现,以及一组针对训练现象的判断方法。读者可以在 CPU 或单张 GPU 上直接跑通,再决定是否把同一套优化思路迁移到官方 Mamba、Vision Mamba 或更复杂的生产项目中。
1. 先理解 Mamba 的状态传递与谱性质问题
1.1 从连续状态空间到选择性扫描
经典状态空间模型描述的是一个连续系统:
$$x'(t) = A x(t) + B u(t)$$
$$y(t) = C x(t) + D u(t)$$
其中 $x(t)$ 是隐藏状态,$u(t)$ 是输入,$y(t)$ 是输出,$A$ 是状态转移矩阵,$B$ 和 $C$ 分别是输入矩阵和输出矩阵,$D$ 是直通项。实际处理序列时,需要把连续系统离散化。常见方式是零阶保持(ZOH):
$$A_bar = \exp(\Delta A)$$
$$B_bar = (\Delta A)^{-1}(\exp(\Delta A) - I) \Delta B$$
离散化后,每一步的状态更新可以写成:
$$h_t = A_bar h_{t-1} + B_bar u_t$$
$$y_t = C h_t + D u_t$$
Mamba 的关键变化是让 $\Delta$、$B$、$C$ 都依赖当前输入。也就是说,模型不再用同一组转移参数处理所有 token,而是根据输入选择当前时刻更关注哪些信息、遗忘哪些信息。这个机制被称为选择性扫描。它让 SSM 能处理 Transformer 中依赖输入内容的建模需求,同时保持推理时的循环结构。
需要明确一点:Mamba 并不只是“在序列维度上做循环”。论文中为了充分并行化,训练阶段还使用了硬件感知的并行扫描算法,把选择后的离散参数组织成可以并行计算的关联扫描。教学用的简化实现可以不追求并行效率,但必须保留“状态随输入变化”这个核心语义。
1.2 状态转移矩阵 A 为什么是训练难点
在状态空间模型里,长期记忆能力很大程度上取决于 $A_bar$ 的谱性质。直观理解是:如果 $A_bar$ 的谱半径远小于 1,那么状态 $h_t$ 会指数级地遗忘过去信息;如果谱半径大于 1,状态又容易发散。理想状态是让 $A_bar$ 的谱半径接近 1,这样信息可以在长序列中缓慢衰减,同时通过输入依赖的 $B$ 和 $\Delta$ 控制写入多少信息。
Mamba 对 $A$ 采用的是对数参数化:
$$A = -\exp(A_log)$$
训练时优化的是 $A_log$,而不是直接优化 $A$。这样能够保证离散化前的 $A$ 始终是负对角线矩阵,从而在理论上倾向于稳定。但这里存在一个容易被忽略的问题:最终决定“记住多少”的是 $\exp(\Delta A)$,而 $\Delta$ 又是输入依赖的。也就是说,即使 $A$ 本身稳定,训练过程中 $\Delta$ 和 $A_log$ 的耦合也可能导致状态在部分时间步上被放大或快速清空。这种动态平衡很难用一个固定更新策略去处理。
另一个难点是状态维度。真实 Mamba 中每个特征通道都有自己的一组 $A$、$B$、$C$ 参数。状态维度 $d_state$ 往往远小于模型维度,模型需要在高维输入和低维状态之间反复投影。如果投影矩阵的训练不稳定,状态空间就学不出有意义的记忆表示。
1.3 AdamW 更新与矩阵结构之间的冲突
AdamW 是训练 Transformer 的主流优化器,但它并不天然适合所有矩阵参数。AdamW 会对每个梯度元素做归一化:
$$m_t = \beta_1 m_{t-1} + (1 - \beta_1) g_t$$
$$v_t = \beta_2 v_{t-1} + (1 - \beta_2) g_t^2$$
$$\theta_t = \theta_{t-1} - \eta \frac{m_t}{\sqrt{v_t} + \epsilon}$$
这里每个元素独立缩放,意味着更新方向被“逐元素”地拉平了。对于 embedding 这类稀疏参数,这种特性很好;但对于依赖矩阵乘法结构的权重,比如 SSM 的 $A$、$B$、$C$ 投影矩阵、以及 Mamba 中的线性投影层,逐元素归一化可能破坏矩阵奇异值之间的相对关系。
更具体地说,AdamW 会让梯度方向中的“大值分量”被压缩,而“小值分量”被放大。当模型需要保持某个矩阵的低秩结构、正交方向或谱分布时,这种更新并不理想。Muon 优化器的思路则是:对二维权重矩阵,先做零中心化,再用 Newton-Schulz 迭代把梯度方向投影到正交矩阵附近,最后用较大的学习率更新。这与 SSM 对矩阵结构的要求更为契合。
2. Muon 优化器:更尊重矩阵结构的更新方式
2.1 Muon 在做什么:零中心化、正交化、大学习率
Muon 由 Keller Jordan 等人在论文《Muon: An optimizer for hidden layers in neural networks》中提出。论文的核心观察是:神经网络隐藏层的二维权重矩阵,在训练时更适合使用“方向更新”,而不是像 AdamW 那样逐元素缩放。
Muon 对二维权重 $W$ 的处理可以拆成三个动作:
第一步,对梯度做零中心化。常见做法是对梯度矩阵的某一维做均值减法,例如:
def zerocenter(g): return g - g.mean(dim=0, keepdim=True)这一步移除梯度中的“整体平移分量”,让更新更关注矩阵行与行之间的差异。
第二步,对中心化后的梯度做正交化。这里使用的是 Newton-Schulz 迭代,把梯度矩阵投影到正交矩阵附近。直观理解是:更新量不再是一个任意方向的矩阵,而更像是一个“旋转方向”。
第三步,以较大的学习率更新参数。Muon 论文中使用的学习率通常是 AdamW 的数十倍。例如在常见实现中,Muon 的学习率可能在0.01到0.05范围,而同一模型中使用 AdamW 时学习率可能是1e-3。因为更新方向被限制在正交流形附近,模型对学习率过大的敏感度会下降。
需要注意,Muon 论文建议把 Muon 用于隐藏层二维参数,而 embedding 和输出 head 仍然使用 AdamW。同时,一维参数如 bias、LayerNorm 的 scale 和 shift 也不适合用正交化更新。因此在实现 Muon 时,通常需要按参数形状和用途分组。
2.2 Newton-Schulz 迭代与极分解
对于一个矩阵 $G$,我们希望找到它对应的正交因子。这本质上是在做极分解:把 $G$ 分解成 $G = U P$,其中 $U$ 是正交矩阵,$P$ 是对称半正定矩阵。完整计算极分解可以使用 SVD,但 SVD 在训练循环中计算量太大,且梯度回传成本高。
Newton-Schulz 迭代提供了一种轻量近似。先对 $G$ 做 Frobenius 范数缩放,保证迭代初始矩阵的谱范数有界,然后迭代:
$$X_{k+1} = \frac{3}{2} X_k - \frac{1}{2} X_k X_k^T X_k$$
迭代若干次后,$X$ 会逼近 $G$ 的正交因子。PyTorch 实现通常写成这样:
def newton_schulz(g, steps=5): a, b = g.shape g = g / (g.norm() + 1e-12) if a < b: g = g.T x = g for _ in range(steps): x = 1.5 * x - 0.5 * x @ x.T @ x if a < b: x = x.T return x关键点是g = g / (g.norm() + 1e-12)这一步。如果不先缩放,Newton-Schulz 迭代可能发散。steps控制了近似精度,常用的取值范围是 3 到 6。步数越多,结果越接近正交矩阵,但计算开销也越大。对于嵌入层、中间层和头部分类层,有些实现会使用更少的迭代步数,以减少训练开销。
实际使用中,x @ x.T @ x的矩阵乘法会占用额外显存。对于超大隐藏层,这一步会成为性能瓶颈。可以在实现中限制 Newton-Schulz 迭代只作用于隐藏层二维权重,而不作用于 embedding 和高维 head。
2.3 为什么谱优化思路适合 Mamba
把 Muon 与 Mamba 放在一起,不是简单的“换一个优化器”。两者的结合点在于“谱结构”这个词。
Mamba 的状态空间模型最核心的参数是 $A$,它的谱性质直接决定模型能记忆多长的信息。训练时,如果 $A_log$、$\Delta$、$B$、$C$ 的更新方向不稳定,模型就难以在“记住过去”和“写入当前”之间找到平衡。Muon 的零中心化和正交化让二维权重参数的更新更像“旋转 + 方向移动”,而不是“逐元素拉伸”,这有助于保持投影矩阵的几何结构。
第二层关系是梯度传播。SSM 在长序列上训练时,梯度需要穿过很多时间步。状态转移矩阵的谱半径如果偏离 1 太远,梯度会出现指数级衰减或爆炸。Muon 不直接约束 $A$ 的谱半径,但它通过更稳定的更新方向,减少了训练过程中参数剧烈变化导致的状态发散。实际项目中,通常还会配合梯度裁剪、$A_log$ 初始化范围限制等策略。
可以这样理解:Muon 优化的是“参数更新时矩阵结构的保持问题”,而 Mamba 需要解决的是“状态转移矩阵在长序列上的稳定性问题”。前者从优化器层面提供了更好的几何更新方向,后者决定模型能否学到长程记忆。两者结合,是一种“谱优化”思路:不仅关注 loss 下降,还关注参数矩阵在谱意义下是否健康。
3. 最小可运行环境与简化 Mamba 模块
3.1 环境与依赖
本文所有示例代码使用纯 PyTorch 实现,不依赖官方mamba_ssm包。这样做的目的是先跑通优化器与状态空间模型之间的协作关系,避免因为 CUDA kernel 编译、版本匹配等问题干扰主线。
推荐环境如下:
| 组件 | 推荐配置 | 说明 |
|---|---|---|
| Python | 3.10 或 3.11 | 依赖torch2.x |
| PyTorch | 2.1 或更高 | 需要支持torch.optim.Optimizer自定义 |
| 硬件 | CPU 可跑通,有 GPU 更快 | 本文示例规模较小,CPU 也可完成 |
| 额外包 | 无 | 不需要mamba_ssm,不需要额外 CUDA 扩展 |
如果本机已经安装了 Anaconda 或 Miniconda,可以直接创建虚拟环境:
conda create -n muon-mamba python=3.10 -y conda activate muon-mamba conda install pytorch pytorch-cuda=12.1 -c pytorch -c nvidia需要说明的是,这里创建环境只是为了隔离依赖。如果本机已经装好 PyTorch,可以跳过 conda 步骤直接用。若之后想跑官方 Mamba,需要按照mamba_ssm官方 README 的说明安装,并提前确认 PyTorch 版本、CUDA 版本和编译工具链,否则容易出现 C++/CUDA 编译报错。
3.2 简化版选择性 SSM 模块设计
教学用的 Mamba 实现没有必要完整复刻官方S6模块的所有细节。真正的 Mamba 里,每个特征通道会维护独立的隐藏状态,训练时用并行关联扫描加速。这里为了把“状态更新、选择性、优化器”讲清楚,写一个共享状态的简化版本。它保留了 SSM 的核心递推关系,但状态维度与模型维度不直接绑定。
import math import torch import torch.nn as nn import torch.nn.functional as F class SelectiveSSM(nn.Module): def __init__(self, d_input=2, d_model=32, d_state=8, d_delta=4): super().__init__() self.d_input = d_input self.d_model = d_model self.d_state = d_state self.d_delta = d_delta self.encoder = nn.Linear(d_input, d_model) self.u_proj = nn.Linear(d_model, 1) self.A_log = nn.Parameter(torch.randn(d_state)) self.D = nn.Parameter(torch.randn(d_model)) self.x_proj = nn.Linear(d_model, d_delta + 2 * d_state) self.dt_proj = nn.Linear(d_delta, 1) self.decoder = nn.Linear(d_model, 1) self._init_weights() def _init_weights(self): self.dt_proj.weight.data.zero_() self.dt_proj.bias.data.fill_(0.1) self.A_log.data.uniform_(-5.0, -1.0) def forward(self, x): B, L, _ = x.shape u = F.silu(self.encoder(x)) A = -torch.exp(self.A_log) h = torch.zeros(B, self.d_state, device=x.device) outputs = [] for t in range(L): xt = u[:, t, :] u_scalar = self.u_proj(xt).squeeze(-1) xb = self.x_proj(xt) dt = F.softplus(self.dt_proj(xb[:, :self.d_delta])).squeeze(-1) Bt = xb[:, self.d_delta:self.d_delta + self.d_state] Ct = xb[:, self.d_delta + self.d_state:] dA = torch.exp(dt.unsqueeze(-1) * A) dB = ((dt.unsqueeze(-1) * A).expm1() / A) * Bt dB = dB * u_scalar.unsqueeze(-1) h = dA * h + dB output = (Ct * h).sum(dim=-1, keepdim=True) outputs.append(output) out = torch.stack(outputs, dim=1) out = self.decoder(F.silu(out)) return out这个模块的关键设计如下:
encoder把原始输入映射到隐藏维度,u_proj再压缩成标量,作为当前时刻写入状态的“标量输入”。A_log初始化为[-5, -1]之间的均匀分布,再取负指数,保证初始 $A$ 是负值,状态更新不会一上来就发散。x_proj输出被切成三段:前d_delta维用于计算步长 $\Delta$,接下来d_state维是输入依赖的 $B$,最后d_state维是输入依赖的 $C$。- 循环中的
h = dA * h + dB是离散化状态更新。dB使用expm1计算 $\frac{\exp(\Delta A) - 1}{A}$,同时对 $B$ 乘以当前输入的标量表示。 - 输出是当前状态与 $C$ 的内积,经
decoder映射成最终预测。
需要注意,这里没有使用官方 Mamba 的 1D 卷积、通道扩张和并行扫描。它是一个“能体现选择性状态更新”的最小版本,适合调试优化器和状态动态。实际生产中替换成官方 Mamba 时,优化器部分逻辑仍然可以复用。
3.3 参数与形状对照
| 参数名 | 含义 | 示例值 | 形状说明 |
|---|---|---|---|
d_input | 输入特征维度 | 2 | 任务输入是(B, L, 2) |
d_model | 隐藏维度 | 32 | 内部编码维度 |
d_state | 状态空间维度 | 8 | 每个时刻维护的隐藏状态长度 |
d_delta | 控制 $\Delta$ 的中间维度 | 4 | 由输入映射得到 |
L | 序列长度 | 64 | 任务生成的序列长度 |
d_state越大,模型记忆容量越高,但训练成本也越高。d_delta只是用于生成 $\Delta$ 的中间表示,可以把它理解为步长控制分支的一个隐藏层。
4. 用 PyTorch 实现 Muon 优化器
4.1 优化器整体结构
Muon 优化器可以基于torch.optim.Optimizer自定义实现。整体结构是:把模型参数分成两组,一组是“使用 Muon 的二维隐藏层参数”,另一组是“使用 AdamW 的一维参数和 embedding/head 参数”。
class Muon(torch.optim.Optimizer): def __init__( self, params, lr=0.02, momentum=0.95, nesterov=True, ns_steps=5, adamw_lr=1e-3, adamw_betas=(0.9, 0.95), adamw_eps=1e-8, wd=0.01, ): defaults = dict( lr=lr, momentum=momentum, nesterov=nesterov, ns_steps=ns_steps, adamw_lr=adamw_lr, adamw_betas=adamw_betas, adamw_eps=adamw_eps, wd=wd, muon=True, ) super().__init__(params, defaults)这里muon标志会在参数组中传递。之后创建优化器时,可以人为指定哪些参数走 Muon,哪些走 AdamW:
optimizer = Muon([ {"params": [p for name, p in model.named_parameters() if p.ndim >= 2]}, {"params": [p for name, p in model.named_parameters() if p.ndim < 2], "muon": False}, ])4.2 Newton-Schulz 正交化实现
在step中,需要区分两类参数:
muon=True:梯度先零中心化,再应用动量缓冲,再经过 Newton-Schulz 正交化,最后用lr更新。muon=False:使用 AdamW 的动量与二阶动量估计更新。
def step(self, closure=None): with torch.no_grad(): for group in self.param_groups: if group.get("muon", True): self._step_muon(group) else: self._step_adamw(group) return None_step_muon的实现如下:
def _step_muon(self, group): lr = group["lr"] momentum = group["momentum"] ns_steps = group["ns_steps"] wd = group["wd"] for p in group["params"]: if p.grad is None: continue g = p.grad.detach() g = zerocenter(g) state = self.state[p] if "momentum_buffer" not in state: state["momentum_buffer"] = torch.zeros_like(g) buf = state["momentum_buffer"] buf.mul_(momentum).add_(g, alpha=1 - momentum) g = newton_schulz(buf, steps=ns_steps) if wd != 0: p.mul_(1 - lr * wd) p.add_(g, alpha=-lr)这段代码中有一个容易混淆的细节:先动量和后动量的顺序。这里选择的顺序是“零中心化 -> 动量缓冲 -> Newton-Schulz -> 更新”。如果改成“零中心化 -> Newton-Schulz -> 动量缓冲”,更新结果会不同。不同开源实现有两种顺序,落地时应该先在自己的小任务上验证,再固定下来。本文为了演示,采用了先动量后正交化的顺序。
_step_adamw是标准 AdamW 更新:
def _step_adamw(self, group): lr = group["adamw_lr"] beta1, beta2 = group["adamw_betas"] eps = group["adamw_eps"] wd = group["wd"] for p in group["params"]: if p.grad is None: continue g = p.grad.detach() if wd != 0: p.mul_(1 - lr * wd) state = self.state[p] if "step" not in state: state["step"] = 0 state["exp_avg"] = torch.zeros_like(g) state["exp_avg_sq"] = torch.zeros_like(g) exp_avg = state["exp_avg"] exp_avg_sq = state["exp_avg_sq"] state["step"] += 1 exp_avg.mul_(beta1).add_(g, alpha=1 - beta1) exp_avg_sq.mul_(beta2).addcmul_(g, g, value=1 - beta2) bias_corr1 = 1 - beta1 ** state["step"] bias_corr2 = 1 - beta2 ** state["step"] denom = (exp_avg_sq.sqrt() / math.sqrt(bias_corr2)).add_(eps) step_size = lr / bias_corr1 p.addcdiv_(exp_avg, denom, value=-step_size)4.3 学习率与参数分组
Muon 对隐藏层二维参数使用较大的学习率,对一维参数和 AdamW 组使用较小的学习率。参考配置可以用下表:
| 参数分组 | 优化器 | 学习率 | 说明 |
|---|---|---|---|
| 2D 隐藏层权重 | Muon | 0.01 - 0.05 | 默认从 0.02 开始 |
| 1D bias、norm 权重 | AdamW | 1e-3 | 不参与正交化 |
| embedding / head | AdamW | 1e-3 或更低 | 按实际任务调整 |
在Muon构造器中,lr控制 Muon 部分的学习率,adamw_lr控制 AdamW 部分的学习率。两个值可以分开调整。实验时建议先固定adamw_lr=1e-3,只调整 Muon 的lr,观察 loss 是否出现剧烈波动。
需要注意,模型中的A_log、D、x_proj、dt_proj等参数形态不同。A_log和D是一维参数,走 AdamW 分支;encoder.weight、decoder.weight等二维参数走 Muon 分支。这符合 Muon 论文中“hidden layers 用 Muon,其他参数用 AdamW”的原则。
5. 训练一个确定性记忆任务并对比优化器
5.1 任务设计:延迟选择性求和
为了验证状态空间模型和优化器的组合,需要一个能明确考察“选择性记忆”的任务。这里选择延迟选择性求和:输入序列中每个时刻有两个值,一个随机信号 $v_t$ 和一个门控信号 $g_t$。当 $g_t=1$ 时,模型需要把当前信号累加到状态中;当 $g_t=0$ 时,可以忽略。目标是在每个时间步输出当前累计和。
def generate_batch(batch_size, seq_len, device="cpu", gate_prob=0.15): v = torch.randn(batch_size, seq_len, device=device) g = (torch.rand(batch_size, seq_len, device=device) < gate_prob).float() x = torch.stack([v, g], dim=-1) y = torch.cumsum(v * g, dim=1).unsqueeze(-1) return x, y这个任务的优点是:
- 必须依赖状态跨时间步传递,能测试模型长程记忆能力。
- 有明确的“选择”语义:只有门控为 1 的时间步需要写入状态。
- 目标序列与输入序列等长,方便用 MSE 评估。
5.2 训练循环与验证
训练循环不复杂,关键是控制随机种子,让两个优化器在相同初始化下对比。
def train_one_model(model, optimizer, steps=1200, batch_size=64, seq_len=64): model.train() for step in range(steps): xb, yb = generate_batch(batch_size, seq_len, device=next(model.parameters()).device) pred = model(xb) loss = F.mse_loss(pred, yb) optimizer.zero_grad() loss.backward() optimizer.step() if step % 200 == 0 or step == steps - 1: print(f"step={step:04d} loss={loss.item():.6f}")使用两个模型分别训练:
torch.manual_seed(0) model_adam = SelectiveSSM() model_muon = SelectiveSSM()为了让两个模型初始参数一致,可以先构造一个模型,再把参数复制给另一个:
def copy_model(src, dst): dst.load_state_dict(src.state_dict()) model_base = SelectiveSSM() model_adam = SelectiveSSM() model_muon = SelectiveSSM() copy_model(model_base, model_adam) copy_model(model_base, model_muon)然后分别构建优化器:
def build_optimizer(model, use_muon=True): if use_muon: return Muon([ {"params": [p for name, p in model.named_parameters() if p.ndim >= 2]}, {"params": [p for name, p in model.named_parameters() if p.ndim < 2], "muon": False}, ]) else: return torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)注意这里 AdamW 对比实验把所有参数放在一组,是公平的对照。如果希望更细致,也可以让 AdamW 对一维参数保持lr=1e-3,对二维参数使用同样的学习率。由于 Muon 本身对隐藏层使用0.02,直接对比时学习率差异很大,这正是两种优化器的特性差异,而不是 bug。
5.3 AdamW 与 Muon 的收敛对比观察
训练完成后,重点看两个指标:loss 曲线的下降速度和最终 MSE 水平。
在小规模任务上,Muon 通常能更快把 loss 压下去,尤其是训练初期。原因是初始化阶段模型的状态转移能力较弱,AdamW 逐元素归一化更新较保守,而 Muon 的大学习率配合正交化能在保持结构稳定的前提下更快探索参数方向。
这里要强调一点:不要把这个结果解读成“Muon 一定优于 AdamW”。在 embedding、attention 类模型或某些非矩阵结构任务上,Muon 未必比 AdamW 好。本文的对比只在“简化 SSM + 延迟选择性求和任务”这个范围内成立。
如果 Muon 组出现 loss 不降或波动过大,优先检查:
- 是否混入了 embedding 参数。Muon 应该只用于隐藏层二维权重。
lr是否过大。可以降到0.005再试。- Newton-Schulz 迭代步数是否足够。默认 5 步通常够用。
- 模型是否有状态爆炸。打印训练中的
dA均值,看是否远大于 1。
6. 常见问题与排查路径
6.1 Muon 更新后 loss 出现 NaN
现象:训练几十步后 loss 变成nan或inf。
可能原因:
- Muon 学习率过大,更新步长越过了稳定区域。
- Newton-Schulz 迭代前没有对梯度