状态空间模型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(dim0, keepdimTrue)这一步移除梯度中的“整体平移分量”让更新更关注矩阵行与行之间的差异。第二步对中心化后的梯度做正交化。这里使用的是 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_{k1} \frac{3}{2} X_k - \frac{1}{2} X_k X_k^T X_k$$迭代若干次后$X$ 会逼近 $G$ 的正交因子。PyTorch 实现通常写成这样def newton_schulz(g, steps5): 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 编译、版本匹配等问题干扰主线。推荐环境如下组件推荐配置说明Python3.10 或 3.11依赖torch2.xPyTorch2.1 或更高需要支持torch.optim.Optimizer自定义硬件CPU 可跑通有 GPU 更快本文示例规模较小CPU 也可完成额外包无不需要mamba_ssm不需要额外 CUDA 扩展如果本机已经安装了 Anaconda 或 Miniconda可以直接创建虚拟环境conda create -n muon-mamba python3.10 -y conda activate muon-mamba conda install pytorch pytorch-cuda12.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_input2, d_model32, d_state8, d_delta4): 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, devicex.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, keepdimTrue) outputs.append(output) out torch.stack(outputs, dim1) 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, lr0.02, momentum0.95, nesterovTrue, ns_steps5, adamw_lr1e-3, adamw_betas(0.9, 0.95), adamw_eps1e-8, wd0.01, ): defaults dict( lrlr, momentummomentum, nesterovnesterov, ns_stepsns_steps, adamw_lradamw_lr, adamw_betasadamw_betas, adamw_epsadamw_eps, wdwd, muonTrue, ) super().__init__(params, defaults)这里muon标志会在参数组中传递。之后创建优化器时可以人为指定哪些参数走 Muon哪些走 AdamWoptimizer 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中需要区分两类参数muonTrue梯度先零中心化再应用动量缓冲再经过 Newton-Schulz 正交化最后用lr更新。muonFalse使用 AdamW 的动量与二阶动量估计更新。def step(self, closureNone): 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, alpha1 - momentum) g newton_schulz(buf, stepsns_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, alpha1 - beta1) exp_avg_sq.mul_(beta2).addcmul_(g, g, value1 - 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 隐藏层权重Muon0.01 - 0.05默认从 0.02 开始1D bias、norm 权重AdamW1e-3不参与正交化embedding / headAdamW1e-3 或更低按实际任务调整在Muon构造器中lr控制 Muon 部分的学习率adamw_lr控制 AdamW 部分的学习率。两个值可以分开调整。实验时建议先固定adamw_lr1e-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_t1$ 时模型需要把当前信号累加到状态中当 $g_t0$ 时可以忽略。目标是在每个时间步输出当前累计和。def generate_batch(batch_size, seq_len, devicecpu, gate_prob0.15): v torch.randn(batch_size, seq_len, devicedevice) g (torch.rand(batch_size, seq_len, devicedevice) gate_prob).float() x torch.stack([v, g], dim-1) y torch.cumsum(v * g, dim1).unsqueeze(-1) return x, y这个任务的优点是必须依赖状态跨时间步传递能测试模型长程记忆能力。有明确的“选择”语义只有门控为 1 的时间步需要写入状态。目标序列与输入序列等长方便用 MSE 评估。5.2 训练循环与验证训练循环不复杂关键是控制随机种子让两个优化器在相同初始化下对比。def train_one_model(model, optimizer, steps1200, batch_size64, seq_len64): model.train() for step in range(steps): xb, yb generate_batch(batch_size, seq_len, devicenext(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(fstep{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_muonTrue): 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(), lr1e-3, weight_decay0.01)注意这里 AdamW 对比实验把所有参数放在一组是公平的对照。如果希望更细致也可以让 AdamW 对一维参数保持lr1e-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 迭代前没有对梯度