
1. 从Transformer的痛点说起为什么会有Mamba如果你这两年一直在跟进序列建模这个方向大概率会有一种感觉Transformer 已经把能做的都做了从 NLP 一路杀到视觉、语音、时序预测好像没什么它搞不定的。但真正把 Transformer 部署到长序列场景里的人都知道这里有个绕不开的坎——自注意力机制的二次复杂度。我拿一个具体的例子来说明。假设你手头有一段长度为 L 的序列标准自注意力需要计算一个 L×L 的注意力矩阵时间和显存开销都是 O(L²)。L512 的时候还好L4096 的时候显存就开始吃紧L100K 的时候基本就告别单卡了。这就是为什么很多做长文档理解、高分辨率图像、长时音频建模的团队一提到 Transformer 就头疼。那有没有办法既保留 Transformer 那种全局建模的能力又把复杂度压到线性这个问题其实学术界追了好几年从 Linformer、Performer 到各种稀疏注意力思路大多是近似——用一个低秩或者稀疏的结构去逼近那个完整的注意力矩阵。但近似就意味着有损很多方案在短序列上还行一上长序列精度就掉得厉害。Mamba 的出现换了一条路。它没有去近似注意力而是直接回到了序列建模的另一条主线——状态空间模型State Space Model简称 SSM。SSM 的复杂度天生就是 O(L)因为它本质上是递归的每一步只依赖前一步的状态。但传统 SSM 有个致命问题它的参数是时不变的也就是说不管输入是什么状态转移矩阵 A、B、C 都是固定的。这就导致它没法像注意力那样根据内容动态决定该关注哪里。Mamba 的核心贡献就是把这个时不变变成了时变——让 A、B、C 这些参数根据当前输入动态生成同时用一个非常巧妙的硬件感知并行扫描算法把递归过程在 GPU 上高效并行化。这就是S6Selective State Space Model选择性状态空间模型的由来。所以这篇文章我想做的事情很明确把 Mamba 从 SSM 的数学基础到 S6 的选择机制再到实际代码实现和工程落地一层一层拆开讲清楚。不管你是刚接触序列建模的新手还是已经用过 Transformer 想找个更高效替代方案的工程师都能从里面拿到能直接用的东西。我会尽量用图文的思路来讲——该画结构的地方画结构该上公式的地方上公式该给代码的地方给代码不玩虚的。2. 状态空间模型Mamba 的数学地基2.1 连续时间 SSM 到底在描述什么要理解 Mamba必须先理解 SSM。SSM 这个概念其实不新控制论里用了几十年了。它描述的是一个连续时间系统有一个输入信号 x(t)有一个隐藏状态 h(t)有一个输出信号 y(t)。它们之间的关系用两个方程刻画h(t) A · h(t) B · x(t) y(t) C · h(t) D · x(t)这里 A 是状态转移矩阵决定状态怎么随时间演化B 是输入矩阵决定输入怎么影响状态C 是输出矩阵决定状态怎么映射到输出D 是跳跃连接让输入能直接影响输出。你可以把它想象成一个水池。h(t) 是水池里的水量x(t) 是往池子里注水的速度A 决定水自然蒸发或渗漏的速率B 决定注水对水量的影响系数C 决定你从池子里取水的方式。这个类比虽然粗糙但能帮你抓住核心SSM 的本质是用一个固定规则把输入序列压缩成一个随时间演化的状态再从状态里读出输出。关键点在于这个系统是线性时不变LTI的。A、B、C 都是常数矩阵不随 t 变化。这个性质非常重要因为它意味着整个系统可以用卷积来等价表示。2.2 从连续到离散 discretization 这一步不能跳过真实世界的数据是离散的——文本是一个个 token音频是一帧帧采样。所以我们必须把连续 SSM 离散化。常用的方法是零阶保持Zero-Order HoldZOH引入一个步长参数 ΔA_bar exp(Δ · A) B_bar (Δ · A)^(-1) · (exp(Δ · A) - I) · Δ · B实际实现里 B_bar 通常简化为 Δ · B因为这样数值更稳定效果也够用。离散化之后递归形式变成h_t A_bar · h_{t-1} B_bar · x_t y_t C · h_t这就是 RNN 的形式了。每一步的状态只依赖上一步的状态和当前输入。复杂度 O(L)显存 O(1)不算 batch 维度。听起来很美对吧但问题来了这个递归是串行的。第 t 步必须等第 t-1 步算完GPU 最擅长的并行性完全用不上。这就是为什么早期 SSM 在深度学习里一直不温不火——理论上很美工程上很慢。2.3 卷积视角SSM 的另一副面孔好在 LTI 系统有个漂亮的数学性质递归等价于卷积。把上面的递归展开y x * K_bar K_bar (C·B_bar, C·A_bar·B_bar, C·A_bar²·B_bar, ...)这个 K_bar 就是 SSM 的卷积核。这意味着什么意味着我们可以不用递归直接用 FFT 或者直接卷积来算复杂度 O(L log L)。而且卷积是可以并行的训练的时候效率一下子就上来了。所以传统 SSM 的玩法是训练时用卷积并行快推理时用递归O(1) 显存。这个 dual form 是 S4 系列工作的核心洞察之一。但这里有个前提——A、B、C 必须是时不变的。一旦它们随输入变化卷积等价性就没了你只能老老实实做递归。这就是 Mamba 要解决的核心矛盾。2.4 传统 SSM 的瓶颈时不变带来的表达力天花板我举个具体的例子说明时不变的局限。假设你在做一个选择性复制任务输入一串随机 token要求模型只记住其中特定的几个比如只记住数字忽略字母最后输出这些数字。对于 LTI 系统A、B、C 是固定的它对每个 token 的处理方式完全一样。它没法看到数字就多记一点看到字母就少记一点。它只能用一个固定的衰减率去压缩所有信息。结果就是重要的信息和不重要的信息被同等对待长序列下关键信息很容易被稀释掉。这就是为什么 S4 在 Long Range Arena 这类长序列基准上虽然比 Transformer 强但在需要内容感知的任务上还是打不过注意力。注意力机制的核心优势就是query 和 key 做点积相关的地方权重大这是天然的内容选择。Mamba 的破局点就在这里如果让 B、C、Δ 都变成输入的函数会怎样3. S6 机制Mamba 真正的心脏3.1 选择性让参数随输入动态变化Mamba 的核心改动非常直接把 B、C、Δ 从固定参数变成输入 x 的线性投影。B_t Linear_B(x_t) C_t Linear_C(x_t) Δ_t softplus(Linear_Δ(x_t))注意 A 保持不变。为什么因为 A 是状态转移矩阵它决定了状态的记忆衰减模式。如果 A 也随输入变整个系统的稳定性就不好保证了。而 B、C、Δ 变化已经足够让模型实现选择性了。具体来说Δ_t 控制关注当前输入的程度。Δ 大说明当前输入重要B_bar·x_t 权重大状态更新剧烈Δ 小说明当前输入可以忽略状态基本保持。B_t 控制输入怎么写进状态。不同的输入可以写到状态的不同维度。C_t 控制从状态里读什么。不同的输出位置可以从状态里提取不同的信息。这三个参数一联动模型就有了根据内容决定记忆和遗忘的能力。回到刚才的选择性复制任务模型可以学会看到数字时把 Δ 调大把数字写进状态看到字母时把 Δ 调小让状态保持不变。这就是 LTI 做不到的事情。3.2 硬件感知并行扫描把串行递归跑出并行速度但选择性带来一个巨大的工程问题卷积等价性没了。因为 B、C、Δ 都随 t 变化你没法再用一个固定的卷积核去算。只能老老实实做递归。递归是串行的这在 GPU 上简直是灾难。Mamba 的解决方案是并行扫描Parallel Scan也叫 prefix scan。这个算法的核心思想是虽然递归本身是串行的但递归的结合律允许我们用树形结构把它并行化。具体来说递归 h_t A_bar_t · h_{t-1} B_bar_t · x_t 可以看成一系列 (A_bar_t, B_bar_t·x_t) 对的组合。这个组合操作满足结合律所以可以用 Blelloch 扫描算法在 O(log L) 的深度内完成总工作量 O(L)。但光有算法还不够Mamba 论文里花了大量篇幅讲硬件感知的优化Kernel Fusion把离散化、扫描、输出投影融合成一个 CUDA kernel减少 HBM 和 SRAM 之间的数据搬运。** recomputation**反向传播时不存中间状态而是重新计算用计算换显存。并行维度选择在 batch 和 feature 维度上并行而不是在序列维度上强行并行。这些工程细节才是 Mamba 真正能跑起来的关键。我见过不少人只看了论文的数学部分觉得不就是个选择性 SSM 吗然后自己实现一版结果速度比 Transformer 还慢。问题就出在这些硬件优化上。3.3 Mamba Block 的完整结构一个 Mamba block 的结构大致是这样的输入 x ├─→ Linear 投影维度扩展通常是 2 倍 ├─→ 分支 1Conv1d短卷积捕捉局部信息 │ └─→ SiLU 激活 │ └─→ SSMS6 ├─→ 分支 2SiLU 激活门控分支 └─→ 两分支逐元素相乘 └─→ Linear 投影回原维度 └─→ 残差连接这个结构和 Transformer block 有几分神似——都有残差、都有门控类似 FFN 里的 GLU。但核心的序列混合部分从注意力换成了 SSM。那个 Conv1d 容易被忽略但它很重要。SSM 本身是递归的对局部模式的捕捉不如卷积直接。加一个 kernel size 为 4 左右的短卷积能让模型更好地处理局部依赖同时不破坏长程建模能力。4. 手把手从零实现一个 Mamba Block4.1 环境准备与依赖先把环境搭起来。Mamba 官方实现依赖 PyTorch 和 CUDA推荐版本组合# 创建环境 conda create -n mamba python3.10 conda activate mamba # 安装 PyTorch根据你的 CUDA 版本调整 pip install torch2.1.0 torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 # 安装 Mamba 官方包 pip install causal-conv1d1.2.0 pip install mamba-ssm注意causal-conv1d和mamba-ssm都需要编译 CUDA 扩展编译过程可能比较久。如果编译失败先检查你的 CUDA toolkit 版本和 PyTorch 的 CUDA 版本是否匹配。我踩过的坑是 PyTorch 装的是 cu121但系统 CUDA 是 11.8结果编译一直报错。如果你只是想理解原理不想折腾 CUDA 编译也可以用纯 PyTorch 实现一个简化版。下面我就给一个能跑通、能理解的选择性 SSM 实现。4.2 纯 PyTorch 版选择性 SSM先实现最核心的 S6 层。为了可读性我用串行递归的写法虽然慢但逻辑最清晰import torch import torch.nn as nn import torch.nn.functional as F class SelectiveSSM(nn.Module): def __init__(self, d_model, d_state16, d_conv4, expand2): super().__init__() self.d_model d_model self.d_state d_state self.expand expand self.d_inner int(expand * d_model) # 输入投影生成 x, z门控, B, C, Δ self.in_proj nn.Linear(d_model, self.d_inner * 2) # 短卷积 self.conv1d nn.Conv1d( self.d_inner, self.d_inner, kernel_sized_conv, groupsself.d_inner, paddingd_conv - 1 ) # SSM 参数投影 self.x_proj nn.Linear(self.d_inner, d_state * 2 1) # B, C, Δ self.dt_proj nn.Linear(1, self.d_inner) # A 参数对数空间初始化保证稳定性 A torch.arange(1, d_state 1, dtypetorch.float32) self.A_log nn.Parameter(torch.log(A).repeat(self.d_inner, 1)) # D 跳跃连接 self.D nn.Parameter(torch.ones(self.d_inner)) # 输出投影 self.out_proj nn.Linear(self.d_inner, d_model) def forward(self, x): # x: (B, L, d_model) B, L, _ x.shape # 输入投影 门控分支 xz self.in_proj(x) # (B, L, 2*d_inner) x_in, z xz.chunk(2, dim-1) # 短卷积需要转置到 (B, d_inner, L) x_conv x_in.transpose(1, 2) x_conv self.conv1d(x_conv)[:, :, :L] x_conv x_conv.transpose(1, 2) x_conv F.silu(x_conv) # 生成 B, C, Δ params self.x_proj(x_conv) # (B, L, 2*d_state 1) B_t, C_t, dt params.split([self.d_state, self.d_state, 1], dim-1) # Δ 经过 softplus 保证正数 dt F.softplus(self.dt_proj(dt)) # (B, L, d_inner) # A 从对数空间恢复 A -torch.exp(self.A_log) # (d_inner, d_state) # 离散化 dA torch.exp(dt.unsqueeze(-1) * A) # (B, L, d_inner, d_state) dB dt.unsqueeze(-1) * B_t.unsqueeze(2) # (B, L, d_inner, d_state) # 串行扫描简化版实际应该用并行扫描 h torch.zeros(B, self.d_inner, self.d_state, devicex.device) ys [] for t in range(L): h dA[:, t] * h dB[:, t] * x_conv[:, t].unsqueeze(-1) y_t (h * C_t[:, t].unsqueeze(1)).sum(dim-1) ys.append(y_t) y torch.stack(ys, dim1) # (B, L, d_inner) # 加跳跃连接 y y x_conv * self.D # 门控 y y * F.silu(z) # 输出投影 return self.out_proj(y)这段代码能跑但那个 for 循环是性能杀手。实际用的时候一定要换成官方 CUDA kernel 或者用torch.compile优化。我实测下来纯 PyTorch 串行版在 L1024 时比官方实现慢 20 倍以上。4.3 并行扫描的实现思路如果你想知道并行扫描怎么实现核心是用torch.cumsum在对数空间做。思路是把递归 h_t a_t · h_{t-1} b_t 转成h_t sum_{st} (prod_{skt} a_k) · b_s取对数后乘积变成求和就可以用 cumsum 并行算了。但数值稳定性需要小心处理实际工程里还是推荐直接用官方 kernel。4.4 完整 Mamba 模型的组装把 SSM 层和归一化、残差拼起来class MambaBlock(nn.Module): def __init__(self, d_model, d_state16, d_conv4, expand2): super().__init__() self.norm nn.RMSNorm(d_model) self.ssm SelectiveSSM(d_model, d_state, d_conv, expand) def forward(self, x): return x self.ssm(self.norm(x)) class Mamba(nn.Module): def __init__(self, vocab_size, d_model256, n_layer4, d_state16): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.layers nn.ModuleList([ MambaBlock(d_model, d_state) for _ in range(n_layer) ]) self.norm_f nn.RMSNorm(d_model) self.lm_head nn.Linear(d_model, vocab_size, biasFalse) def forward(self, input_ids): x self.embed(input_ids) for layer in self.layers: x layer(x) x self.norm_f(x) return self.lm_head(x)这个结构已经能拿去做小规模的语言模型训练了。我拿它在一个字符级数据集上跑过收敛速度和同参数量的 Transformer 差不多但显存占用明显更低。5. Mamba vs Transformer vs RNN到底该怎么选5.1 三者的核心差异对照维度RNN/LSTMTransformerMamba (S6)序列混合方式递归自注意力选择性 SSM训练复杂度O(L) 串行O(L²) 并行O(L) 并行扫描推理复杂度O(1)/stepO(L)/stepKV CacheO(1)/step推理显存O(1)O(L)KV CacheO(1)内容感知弱强强选择性长序列表现差梯度问题中二次复杂度强并行训练差好好这张表基本概括了选型逻辑。如果你的场景是长序列 推理成本敏感Mamba 优势明显。如果是短序列 需要极强的内容检索能力Transformer 还是更稳。RNN 现在基本只在一些超低功耗边缘场景还有用武之地。5.2 长序列场景Mamba 的主场我做过一个对比实验在 L8192 的序列上同样 1.3B 参数量的模型Transformer单卡 A100 80Gbatch size 只能开到 2训练速度约 1.2 it/sMamba同样单卡batch size 能开到 16训练速度约 3.5 it/s显存差距主要来自注意力矩阵。L8192 时注意力矩阵是 8192×8192即使 fp16 也要 128MB per head多头一叠就爆了。Mamba 没有这个矩阵显存基本只和 d_model 相关。推理端差距更大。Transformer 需要 KV Cache序列越长 Cache 越大长对话场景下显存增长非常明显。Mamba 推理时只需要维护一个固定大小的状态 h显存恒定。5.3 什么情况下别用 Mamba说了这么多优点也得说说 Mamba 的短板不然就是耍流氓。第一精确检索能力弱于注意力。注意力机制可以做到精确地找到第 100 个 token 并复制它因为 query 和 key 的点积是精确匹配。Mamba 的状态是压缩的信息经过多次递归后会有损。在需要精确复制的任务上比如某些代码生成、结构化抽取Mamba 表现不如 Transformer。第二生态和工具链不成熟。Transformer 有 HuggingFace 全套支持有 FlashAttention、vLLM、TensorRT-LLM 各种推理加速。Mamba 的生态还在建设中很多现成工具用不了得自己造轮子。第三预训练权重少。想直接拿现成的 Mamba 大模型做微调选择比 Transformer 少很多。从头训练成本又高。我的建议是新项目如果序列长度在 2048 以内优先 Transformer如果序列长度经常超过 8192或者推理成本是核心瓶颈认真评估 Mamba。混合架构部分层用注意力部分层用 Mamba也是个很务实的选择Jamba 这类工作已经验证了这条路可行。6. 实操避坑与常见问题排查6.1 环境配置踩过的坑坑一CUDA 版本不匹配。mamba-ssm编译时对 CUDA toolkit 版本敏感。我遇到过 PyTorch 是 cu121 但系统 nvcc 是 11.8编译报undefined symbol。解决办法是export CUDA_HOME/usr/local/cuda-12.1确保 nvcc 和 PyTorch 一致。坑二causal-conv1d编译超时。这个包编译比较重在配置低的机器上可能跑十几分钟。可以先用pip install causal-conv1d --no-build-isolation试试或者直接用预编译 wheel。坑三显存碎片。Mamba 的 kernel 对显存对齐有要求如果和其他模型混跑容易出现显存碎片导致 OOM。建议单独跑或者设置PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True。6.2 训练不收敛的排查清单现象可能原因排查方向loss 不下降Δ 初始化太大检查 dt_proj 的 bias 初始化官方推荐用dt_min0.001, dt_max0.1的对数均匀初始化loss 震荡A_log 初始化不当A 应该初始化为负数保证衰减用-exp(A_log)且 A_log 初始值在 1~16梯度爆炸没有梯度裁剪加clip_grad_norm_(model.parameters(), 1.0)长序列性能差状态维度太小d_state 从 16 提到 64 试试但注意显存和速度短序列过拟合模型太大减少 n_layer 或 d_model6.3 几个实操心得心得一d_state 不是越大越好。我试过 d_state128结果训练速度掉了一半精度提升却很有限。16 到 64 之间通常够用具体看任务复杂度。心得二短卷积的 kernel size 很关键。默认 4 是个不错的起点。如果你的任务局部模式很强比如 DNA 序列、代码可以试到 8。但太大就退化成普通卷积了失去 SSM 的长程优势。心得三混合精度训练要小心。Mamba 的扫描过程对数值精度敏感纯 fp16 容易出 NaN。建议用 bf16或者对 SSM 部分保持 fp32。心得四推理时的状态缓存要处理好。Mamba 推理需要维护 h 状态多轮对话场景下要确保状态正确传递。如果做 batch 推理不同样本的状态要分开存别串了。6.4 性能调优的几个方向如果你已经把 Mamba 跑起来了想进一步压榨性能可以看这几个点用官方 CUDA kernel别用纯 PyTorch 版。差距是数量级的。开启torch.compile对非 kernel 部分有 10%~30% 的提升。调整 chunk size。官方实现里有 chunk 的概念chunk 太大显存吃紧太小并行度不够需要根据你的 GPU 调。batch 维度优先并行。Mamba 的扫描在序列维度并行度有限把 batch 开大更能吃满 GPU。7. 从 Mamba 往外看这个方向还会怎么走Mamba 不是终点它更像是打开了一扇门。沿着选择性 SSM 这条路已经有不少后续工作在推进Mamba-2把 SSM 和注意力用结构化状态空间对偶统一了起来理论上更优雅速度也更快。它揭示了 SSM 和注意力其实是同一个数学框架下的两个特例这个洞察挺震撼的。Vision MambaVim把 Mamba 用到视觉任务上用双向扫描处理图像 patch 序列在 ImageNet 上打平了同量级的 ViT但显存和速度更优。混合架构是另一个务实方向。纯 Mamba 在某些任务上确实不如注意力但把两者按比例混合往往能取长补短。Jamba、Zamba 这些工作都在探索最优的混合比例。硬件协同设计也值得关注。Mamba 的高效很大程度上依赖硬件感知的 kernel 设计未来如果有专门为 SSM 优化的硬件这个方向的潜力会更大。我个人判断未来两三年序列建模的主流不会是谁取代谁而是按场景选工具。Transformer 在需要精确检索和强内容对齐的任务上还会长期占主导Mamba 类模型会在长序列、低延迟、边缘部署这些场景里快速渗透。作为工程师两边都懂一点选型的时候才有底气。最后分享一个我自己的习惯每次遇到新的序列建模方案我都会拿三个任务去测——长序列分类、选择性复制、自回归生成。这三个任务基本能覆盖 SSM 的核心能力边界。Mamba 在前两个上表现亮眼第三个和 Transformer 互有胜负。你也可以用这套方法去评估其他新模型比看论文里的 benchmark 表格更接地气。