Skip to content
0

文章发布较早,内容可能过时,阅读注意甄别。

Mamba模型笔记

1、为什么会有 Mamba

在序列建模里,常见路线大致经历了 RNN -> Transformer -> SSM/Mamba 的演化。

1.1、传统 RNN 的优缺点

传统 RNN 的隐藏状态递推可以写成:

ht=tanh(Whht1+Wxxt+b)

它的优点是天然适合处理序列,推理时只需要维护当前状态;但问题也很明显:

  1. 长序列下容易出现梯度消失或遗忘较远信息的问题。
  2. 时间步之间强依赖,训练阶段难以像 Transformer 那样高效并行。
  3. 面对超长上下文时,表达能力和训练效率都容易受限。

1.2、Transformer 为什么强,又为什么贵

Transformer 通过自注意力机制直接建立任意两个 token 之间的关系,因此在内容建模、长距离依赖和并行训练方面很强,这也是它成为大模型主流架构的原因。

但它也有代价:

  1. 标准自注意力的时间复杂度和显存复杂度通常随序列长度呈二次增长。
  2. 当上下文特别长时,训练和推理成本会明显上升。
  3. 推理阶段的 KV Cache 也会带来额外显存开销。

1.3、Mamba 要解决什么问题

Mamba 的目标可以概括为一句话:

在尽量保留强序列建模能力的前提下,获得接近线性复杂度的长序列建模能力与更高的硬件效率。

因此,Mamba 的出发点不是简单替代 Transformer 的全部优势,而是试图在以下几件事之间找到更好的平衡:

  • 长上下文建模能力
  • 训练与推理效率
  • 显存占用
  • 对序列内容的选择性建模能力

2、状态空间模型(SSM)基础

Mamba 的理论基础来自状态空间模型(State Space Model, SSM)。

2.1、连续时间形式

经典 SSM 可以写成两部分:

状态方程:

h(t)=Ah(t)+Bx(t)

观测方程:

y(t)=Ch(t)+Dx(t)

其中:

  • x(t) 表示输入
  • h(t) 表示隐状态,也可以理解为系统记忆
  • y(t) 表示输出
  • A 控制状态如何随时间演化
  • B 控制输入如何写入状态
  • C 控制如何从状态中读出输出
  • D 是输入到输出的直接通路

2.2、离散时间形式

在深度学习里我们处理的是离散 token 序列,因此更常见的是离散形式:

ht=A¯ht1+B¯xtyt=Cht+Dxt

它的直观含义是:

  1. 当前状态由“上一时刻记忆”与“当前输入”共同决定。
  2. 输出不是直接只看当前输入,而是先经过状态累积,再从状态中读出。
  3. 如果状态设计得好,模型就能用较小代价保留很长的历史信息。

2.3、为什么 SSM 对长序列有吸引力

相较于显式计算两两 token 关系的注意力,SSM 更像是在维护一个不断更新的“压缩记忆”。

它的优势在于:

  • 推理时只需维护状态,不必保留完整历史序列。
  • 理论上更适合超长序列。
  • 如果实现得当,复杂度可以做到接近线性。

但早期 SSM 也有明显短板:

  • 对离散符号内容的“按内容选择”能力较弱。
  • 虽然高效,但在语言建模等任务上不一定比注意力更强。

3、从结构化 SSM 到 Mamba

在 Mamba 之前,已经有一条重要的发展脉络:通过结构化状态空间模型提升 SSM 的表达能力与可计算性。

3.1、离散化

连续时间 SSM 需要先离散化,才能作用在 token 序列上。离散化之后,模型就能像递推系统一样逐步处理序列。

3.2、卷积化与并行化

对于某些特殊结构的 SSM,可以把递推过程改写成卷积形式,从而在训练阶段利用并行计算能力。这是早期结构化 SSM 非常关键的一步。

3.3、长距离依赖增强

为了更稳定地记住远距离信息,结构化 SSM 往往会对状态转移矩阵 A 做特殊设计。常见思路之一是借助 HiPPO(High-order Polynomial Projection Operator,高阶多项式投影算子)一类方法,让系统更擅长压缩和保留历史信息。

可以把这条路线粗略理解为:

先让 SSM 变得“能记、能算、能并行”,再进一步让它“知道该记什么、该忘什么”,于是就走到了 Mamba 的选择性状态空间模型。


4、Mamba 的核心思想

Mamba 的关键词是:Selective State Space Model,即选择性状态空间模型。

4.1、核心改进:让参数依赖输入

普通 SSM 的一个问题是参数相对固定,对不同输入内容不够敏感。Mamba 的关键改进是:

  • BCΔ 随输入动态变化
  • 让模型根据当前 token 内容,决定“写入什么、保留什么、读出什么”

其中 Δ 可以理解为和“时间步长”或“遗忘强度”相关的量。直观上,它有点类似门控机制中的遗忘门:

  • Δ 较大时,更偏向快速更新状态
  • Δ 较小时,更偏向保留已有记忆

这正是你原始笔记里提到的“选择性重视/遗忘输入”的核心含义。

4.2、为什么这很重要

Mamba 论文强调:很多早期高效序列模型不如 Transformer 的一个重要原因,是它们缺乏足够强的内容选择能力,也就是不擅长做 content-based reasoning。

而 Mamba 的选择机制本质上是在说:

不是所有历史信息都同等重要,模型需要根据当前输入内容来决定该保留哪些信息。

这使它比传统 SSM 更适合文本这类离散、语义敏感的序列数据。

4.3、Mamba Block 可以怎样理解

从实现视角看,可以把一个 Mamba Block 粗略理解成下面这条链路:

  1. 输入先经过线性投影。
  2. 一支分支用于产生门控信号。
  3. 另一支分支经过局部卷积,增强短程局部建模能力。
  4. 然后进入选择性 SSM,对序列进行状态更新与信息传递。
  5. 最后经过输出投影得到结果。

这样设计的直觉是:

  • 卷积负责局部模式
  • 选择性 SSM 负责长距离记忆
  • 门控机制负责信息筛选

4.4、每个通道一个 SSM

一个常见的理解方式是:Mamba 可以看作对不同通道并行地应用状态空间更新。这样既保留了状态模型的递推特性,也更利于高维特征建模。

4.5、为什么还能高效

一旦 BCΔ 变成输入相关参数,很多早期 SSM 可用的“直接卷积化”技巧就不再能原封不动使用了。

Mamba 为此引入了硬件友好的并行扫描(parallel scan / selective scan)思路:

  • 保留递推形式的表达能力
  • 同时尽量利用 GPU 的并行能力
  • 减少中间状态显式展开带来的显存负担

所以,Mamba 的关键不是“只用了 SSM”,而是:

它把“选择性”与“高效扫描算法”配套设计到了一起。


5、Mamba2 的核心改进

Mamba2 可以看作 Mamba 的进一步发展版本,核心关键词是:SSD(Structured State Space Duality,结构化状态空间对偶)

5.1、SSD 是什么

SSD 的核心思想是建立 SSM 与注意力机制之间更清晰的统一视角,把二者都放到结构化矩阵的框架下理解。

直观理解:

  • Transformer 擅长矩阵并行计算
  • SSM 擅长线性复杂度递推
  • SSD 试图把这两种优势联系起来

这样一来,Mamba2 不只是“更快的 Mamba”,它还提供了一个更强的理论桥梁,说明 SSM 与某些注意力形式并不是完全割裂的两类方法。

5.2、Mamba2 带来了什么变化

从学习和工程角度,可以抓住以下几点:

  1. 更适合现代硬件加速。
  2. 更容易利用矩阵乘法等高吞吐计算。
  3. block 结构有更新,整体训练和推理效率更好。
  4. 论文中强调在不少设置下可比 Mamba 获得更高速度。

5.3、如何理解 Mamba 与 Mamba2 的关系

可以把二者理解成:

  • Mamba:把选择性状态空间模型真正做成了可用的强序列模型。
  • Mamba2:在理论统一与硬件效率上更进一步,把这条路线推得更成熟。

6、Mamba、Transformer、RNN 的对比理解

模型主要记忆方式长距离建模训练并行性长序列复杂度推理缓存特点
RNN隐状态递推一般近似线性只维护状态
Transformer自注意力通常二次需要 KV Cache
Mamba选择性状态递推较强近似线性主要维护状态

可以这样记:

  • RNN 是“有状态,但不够强”。
  • Transformer 是“表达力强,但长序列代价高”。
  • Mamba 是“想保留状态模型的高效,同时增强内容选择能力”。

7、Mamba 的安装与环境配置

7.1、安装前先明确几个原则

Mamba 安装最容易踩坑的地方不在 Python 代码,而在版本匹配:

  1. Python 版本要和 whl 对应。
  2. PyTorch 版本要和 mamba-ssmcausal-conv1d 对应。
  3. CUDA / nvcc / cudatoolkit 版本要和 PyTorch 编译版本匹配。
  4. 还要注意 cxx11abiTRUE/FALSE 是否一致。

IMPORTANT

下面整理的是学习与实操笔记中的可行经验组合,重点在“能装通、能跑通”。如果未来版本升级,请优先按对应 release 的 whl 命名规则核对版本。

7.2、基础安装流程

bash
# 1、创建 conda 虚拟环境
conda create -n your_env_name python=3.10
conda activate your_env_name

# 2、安装 cudatoolkit
# 安装前可先搜索可用版本
conda search cudatoolkit -c nvidia
conda install cudatoolkit==11.8 -c nvidia

# 3、安装 PyTorch、torchvision、torchaudio
# 官方 CUDA 11.8 版本
pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 --index-url https://download.pytorch.org/whl/cu118

# 国内版-阿里镜像
pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 -f https://mirrors.aliyun.com/pytorch-wheels/cu118

# 4、安装 nvcc
# 安装前进行search
conda search -c nvidia cuda-nvcc
# 安装对应的版本
conda install -c "nvidia/label/cuda-11.8.0" cuda-nvcc
conda install packaging

# 5、在终端运行以下 Python 命令,查看当前 PyTorch 的 ABI 状态:
# 如果输出 False,去 GitHub Releases 页面下载带有 cxx11abiFALSE 标签的 causal_conv1d 和 mamba_ssm 的 whl 文件。
# 如果输出 True,则下载 cxx11abiTRUE 的文件。(注意:causal_conv1d 必须先于 mamba_ssm 安装)

python -c "import torch; print(torch._C._GLIBCXX_USE_CXX11_ABI)"

# 6、安装 mamba 相关包
pip install causal-conv1d==1.4.0
pip install mamba-ssm==2.2.4

1、对于causal_conv1dmamba_ssm 可以手动下载

(1) causal_conv1d下载链接:

Dao-AILab/causal-conv1d: Causal depthwise conv1d in CUDA, with a PyTorch interface的release

(2)mamba_ssm 下载链接:

state-spaces/mamba: Mamba SSM architecture的release

请下载对应版本的包:

例如,causal_conv1d-1.4.0+cu118torch2.4cxx11abiFALSE-cp310-cp310-linux_x86_64.whl表示cuda版本为11.x,torch版本为2.4,cp表示python版本为3.10.x,并且请下载含有FALSE的包

下载完之后,进行pip install causal_conv1d-1.4.0xxx.whl 安装

注意:

causal_conv1d 一般先于 mamba_ssm 安装。

7.5、实操中记录过的可行组合

组合一

text
nvcc 11.8
python 3.10
torch 2.2.2
torchaudio 2.2.2
torchvision 0.17.2
causal-conv1d 1.1.3
mamba-ssm 1.1.3

组合二

text
nvcc 11.8
python 3.10
torch 2.4.0
torchaudio 2.4.0
torchvision 0.19.0
causal-conv1d 1.4.0
mamba-ssm 2.2.4

组合三

text
nvcc 12.4
python 3.10
torch 2.4.0
torchaudio 2.4.0
torchvision 0.19.0
causal-conv1d 1.4.0
mamba-ssm 2.2.4

7.6、一份更贴近实操的安装记录

当机器上已经存在 nvcc 12.4 时,可以参考下面这组安装思路:

bash
# 已存在 nvcc
nvcc -V

# 1、创建 conda 虚拟环境
conda create -n mamba python=3.10 -y
conda activate mamba

# 2、如果系统 CUDA 已经满足,cudatoolkit 可按需决定是否安装

# 3、安装 PyTorch CUDA 12.4 版本
pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 --index-url https://download.pytorch.org/whl/cu124

# 4、查看当前 PyTorch 的 ABI 状态
python -c "import torch; print(torch._C._GLIBCXX_USE_CXX11_ABI)"

# 5、安装匹配的 wheel
pip install causal_conv1d-1.4.0+cu122torch2.4cxx11abiFALSE-cp310-cp310-linux_x86_64.whl
pip install mamba_ssm-2.2.4+cu12torch2.4cxx11abiFALSE-cp310-cp310-linux_x86_64.whl

# 6、如果 transformers 版本不兼容,可降级
pip uninstall transformers -y
pip install transformers==4.36.2

7.7、常见踩坑总结

  1. torch 版本对不上,导致 mamba-ssm 编译或导入失败。
  2. cxx11abiTRUE/FALSE 选错,导致 wheel 安装后仍无法正常导入。
  3. causal_conv1d 没先装,后续 mamba_ssm 可能失败。
  4. CUDA 版本、PyTorch CUDA 版本、nvcc 版本混用,容易出现底层依赖错误。
  5. transformers 版本过新时,和某些 mamba-ssm 组合可能有兼容性问题。

8、Mamba 与 Mamba2 的代码测试

8.1、Mamba 测试

python
import torch
from mamba_ssm import Mamba

batch, length, dim = 2, 32, 16
x = torch.randn(batch, length, dim).to("cuda")
model = Mamba(
    # This module uses roughly 3 * expand * d_model^2 parameters
    d_model=dim,   # 模型维度,需要和输入最后一维一致
    d_state=64,    # SSM 状态扩展维度
    d_conv=4,      # 局部卷积宽度
    expand=2,      # Block 扩展因子
).to("cuda")
y = model(x)
print("参数量", sum(p.numel() for p in model.parameters()))
print(y.shape)
assert y.shape == x.shape
'''
参数量 7968
torch.Size([2, 32, 16])
'''

TIP

d_model 必须和输入张量最后一维 dim 对应,否则会维度不匹配。

8.2、Mamba2 测试

python

import torch
from mamba_ssm import Mamba

from mamba_ssm import Mamba2
batch, length, dim = 2, 64, 1024
x = torch.randn(batch, length, dim).to("cuda")

model = Mamba2(
    d_model=dim,   # 模型维度
    d_state=64,    # SSM state expansion factor,通常可设为 64 或 128
    d_conv=4,      # 局部卷积宽度
    expand=2,      # Block 扩展因子
).to("cuda")
y = model(x)
print("参数量", sum(p.numel() for p in model.parameters()))
print(y.shape)
assert y.shape == x.shape

输出结果:

text
参数量 6468320
torch.Size([2, 64, 1024])

8.3、建议增加一个最小化导入测试

在正式跑模型前,可以先做一个 smoke test:

python
import torch
import mamba_ssm
from mamba_ssm import Mamba, Mamba2

print(torch.__version__)
print(mamba_ssm.__file__)
print("import success")

如果这里都导入失败,优先回头检查:

  • PyTorch 版本
  • CUDA / nvcc 版本
  • causal_conv1d 是否已正确安装
  • ABI 是否匹配

9、学习 Mamba 时最值得抓住的几个点

如果只保留最核心的理解,我认为可以抓住下面 5 句:

  1. Mamba 不是普通 RNN,而是建立在状态空间模型上的高效序列模型。
  2. 它的关键创新是“选择性”,也就是让部分 SSM 参数随输入动态变化。
  3. Δ 可以直观理解为一种控制保留与遗忘强度的机制。
  4. Mamba 的强点在于长序列下接近线性的复杂度和较高推理效率。
  5. Mamba2 则进一步从理论和硬件实现上把这条路线推向成熟。

10、总结

Mamba 的价值不只是提出了一个新模型名字,而是在序列建模这件事上提供了一条非常有代表性的思路:

  • 不一定所有问题都要用显式注意力解决。
  • 记忆机制、状态更新、内容选择,也可以组合成强大的序列模型。
  • 当序列长度继续增长时,线性复杂度模型会越来越有吸引力。

因此,理解 Mamba 的最好方式不是把它当成一个孤立模型,而是把它放回整个发展脉络中看:

RNN 提供“状态记忆”的直觉,Transformer 提供“内容交互”的强大能力,而 Mamba 试图把“高效状态更新”和“内容选择”重新结合起来。

最近更新