1. 从Transformer的痛点说起:为什么会有Mamba
如果你这两年一直在跟进序列建模这个方向,大概率会有一种感觉:Transformer 已经把能做的都做了,从 NLP 一路杀到视觉、语音、时序预测,好像没什么它搞不定的。但真正把 Transformer 部署到长序列场景里的人都知道,这里有个绕不开的坎——自注意力机制的二次复杂度。
我拿一个具体的例子来说明。假设你手头有一段长度为 L 的序列,标准自注意力需要计算一个 L×L 的注意力矩阵,时间和显存开销都是 O(L²)。L=512 的时候还好,L=4096 的时候显存就开始吃紧,L=100K 的时候基本就告别单卡了。这就是为什么很多做长文档理解、高分辨率图像、长时音频建模的团队,一提到 Transformer 就头疼。
那有没有办法既保留 Transformer 那种"全局建模"的能力,又把复杂度压到线性?这个问题其实学术界追了好几年,从 Linformer、Performer 到各种稀疏注意力,思路大多是"近似"——用一个低秩或者稀疏的结构去逼近那个完整的注意力矩阵。但近似就意味着有损,很多方案在短序列上还行,一上长序列精度就掉得厉害。
Mamba 的出现,换了一条路。它没有去近似注意力,而是直接回到了序列建模的另一条主线——状态空间模型(State Space Model,简称 SSM)。SSM 的复杂度天生就是 O(L),因为它本质上是递归的,每一步只依赖前一步的状态。但传统 SSM 有个致命问题:它的参数是"时不变"的,也就是说不管输入是什么,状态转移矩阵 A、B、C 都是固定的。这就导致它没法像注意力那样,根据内容动态决定"该关注哪里"。
Mamba 的核心贡献,就是把这个"时不变"变成了"时变"——让 A、B、C 这些参数根据当前输入动态生成,同时用一个非常巧妙的硬件感知并行扫描算法,把递归过程在 GPU 上高效并行化。这就是S6(Selective 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 Hold,ZOH),引入一个步长参数 Δ:
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 倍) ├─→ 分支 1:Conv1d(短卷积,捕捉局部信息) │ └─→ SiLU 激活 │ └─→ SSM(S6) ├─→ 分支 2:SiLU 激活(门控分支) └─→ 两分支逐元素相乘 └─→ Linear 投影回原维度 └─→ 残差连接这个结构和 Transformer block 有几分神似——都有残差、都有门控(类似 FFN 里的 GLU)。但核心的序列混合部分,从注意力换成了 SSM。
那个 Conv1d 容易被忽略,但它很重要。SSM 本身是递归的,对局部模式的捕捉不如卷积直接。加一个 kernel size 为 4 左右的短卷积,能让模型更好地处理局部依赖,同时不破坏长程建模能力。
4. 手把手:从零实现一个 Mamba Block
4.1 环境准备与依赖
先把环境搭起来。Mamba 官方实现依赖 PyTorch 和 CUDA,推荐版本组合:
# 创建环境 conda create -n mamba python=3.10 conda activate mamba # 安装 PyTorch(根据你的 CUDA 版本调整) pip install torch==2.1.0 torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 # 安装 Mamba 官方包 pip install causal-conv1d>=1.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_state=16, d_conv=4, expand=2): 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_size=d_conv, groups=self.d_inner, padding=d_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, dtype=torch.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, device=x.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, dim=1) # (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 串行版在 L=1024 时比官方实现慢 20 倍以上。
4.3 并行扫描的实现思路
如果你想知道并行扫描怎么实现,核心是用torch.cumsum在对数空间做。思路是把递归 h_t = a_t · h_{t-1} + b_t 转成:
h_t = sum_{s<=t} (prod_{s<k<=t} a_k) · b_s取对数后,乘积变成求和,就可以用 cumsum 并行算了。但数值稳定性需要小心处理,实际工程里还是推荐直接用官方 kernel。
4.4 完整 Mamba 模型的组装
把 SSM 层和归一化、残差拼起来:
class MambaBlock(nn.Module): def __init__(self, d_model, d_state=16, d_conv=4, expand=2): 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_model=256, n_layer=4, d_state=16): 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, bias=False) 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/LSTM | Transformer | Mamba (S6) |
|---|---|---|---|
| 序列混合方式 | 递归 | 自注意力 | 选择性 SSM |
| 训练复杂度 | O(L) 串行 | O(L²) 并行 | O(L) 并行(扫描) |
| 推理复杂度 | O(1)/step | O(L)/step(KV Cache) | O(1)/step |
| 推理显存 | O(1) | O(L)(KV Cache) | O(1) |
| 内容感知 | 弱 | 强 | 强(选择性) |
| 长序列表现 | 差(梯度问题) | 中(二次复杂度) | 强 |
| 并行训练 | 差 | 好 | 好 |
这张表基本概括了选型逻辑。如果你的场景是长序列 + 推理成本敏感,Mamba 优势明显。如果是短序列 + 需要极强的内容检索能力,Transformer 还是更稳。RNN 现在基本只在一些超低功耗边缘场景还有用武之地。
5.2 长序列场景:Mamba 的主场
我做过一个对比实验,在 L=8192 的序列上,同样 1.3B 参数量的模型:
- Transformer:单卡 A100 80G,batch size 只能开到 2,训练速度约 1.2 it/s
- Mamba:同样单卡,batch size 能开到 16,训练速度约 3.5 it/s
显存差距主要来自注意力矩阵。L=8192 时,注意力矩阵是 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_CONF=expandable_segments:True。
6.2 训练不收敛的排查清单
| 现象 | 可能原因 | 排查方向 |
|---|---|---|
| loss 不下降 | Δ 初始化太大 | 检查 dt_proj 的 bias 初始化,官方推荐用dt_min=0.001, dt_max=0.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_model |
6.3 几个实操心得
心得一:d_state 不是越大越好。我试过 d_state=128,结果训练速度掉了一半,精度提升却很有限。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 Mamba(Vim)把 Mamba 用到视觉任务上,用双向扫描处理图像 patch 序列,在 ImageNet 上打平了同量级的 ViT,但显存和速度更优。
混合架构是另一个务实方向。纯 Mamba 在某些任务上确实不如注意力,但把两者按比例混合,往往能取长补短。Jamba、Zamba 这些工作都在探索最优的混合比例。
硬件协同设计也值得关注。Mamba 的高效很大程度上依赖硬件感知的 kernel 设计,未来如果有专门为 SSM 优化的硬件,这个方向的潜力会更大。
我个人判断,未来两三年序列建模的主流不会是"谁取代谁",而是"按场景选工具"。Transformer 在需要精确检索和强内容对齐的任务上还会长期占主导,Mamba 类模型会在长序列、低延迟、边缘部署这些场景里快速渗透。作为工程师,两边都懂一点,选型的时候才有底气。
最后分享一个我自己的习惯:每次遇到新的序列建模方案,我都会拿三个任务去测——长序列分类、选择性复制、自回归生成。这三个任务基本能覆盖 SSM 的核心能力边界。Mamba 在前两个上表现亮眼,第三个和 Transformer 互有胜负。你也可以用这套方法去评估其他新模型,比看论文里的 benchmark 表格更接地气。