Mamba模型笔记
1、为什么会有 Mamba
在序列建模里,常见路线大致经历了 RNN -> Transformer -> SSM/Mamba 的演化。
1.1、传统 RNN 的优缺点
传统 RNN 的隐藏状态递推可以写成:
它的优点是天然适合处理序列,推理时只需要维护当前状态;但问题也很明显:
- 长序列下容易出现梯度消失或遗忘较远信息的问题。
- 时间步之间强依赖,训练阶段难以像 Transformer 那样高效并行。
- 面对超长上下文时,表达能力和训练效率都容易受限。
1.2、Transformer 为什么强,又为什么贵
Transformer 通过自注意力机制直接建立任意两个 token 之间的关系,因此在内容建模、长距离依赖和并行训练方面很强,这也是它成为大模型主流架构的原因。
但它也有代价:
- 标准自注意力的时间复杂度和显存复杂度通常随序列长度呈二次增长。
- 当上下文特别长时,训练和推理成本会明显上升。
- 推理阶段的 KV Cache 也会带来额外显存开销。
1.3、Mamba 要解决什么问题
Mamba 的目标可以概括为一句话:
在尽量保留强序列建模能力的前提下,获得接近线性复杂度的长序列建模能力与更高的硬件效率。
因此,Mamba 的出发点不是简单替代 Transformer 的全部优势,而是试图在以下几件事之间找到更好的平衡:
- 长上下文建模能力
- 训练与推理效率
- 显存占用
- 对序列内容的选择性建模能力
2、状态空间模型(SSM)基础
Mamba 的理论基础来自状态空间模型(State Space Model, SSM)。
2.1、连续时间形式
经典 SSM 可以写成两部分:
状态方程:
观测方程:
其中:
x(t)表示输入h(t)表示隐状态,也可以理解为系统记忆y(t)表示输出A控制状态如何随时间演化B控制输入如何写入状态C控制如何从状态中读出输出D是输入到输出的直接通路
2.2、离散时间形式
在深度学习里我们处理的是离散 token 序列,因此更常见的是离散形式:
它的直观含义是:
- 当前状态由“上一时刻记忆”与“当前输入”共同决定。
- 输出不是直接只看当前输入,而是先经过状态累积,再从状态中读出。
- 如果状态设计得好,模型就能用较小代价保留很长的历史信息。
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 的关键改进是:
- 让
、 、 随输入动态变化 - 让模型根据当前 token 内容,决定“写入什么、保留什么、读出什么”
其中
较大时,更偏向快速更新状态 较小时,更偏向保留已有记忆
这正是你原始笔记里提到的“选择性重视/遗忘输入”的核心含义。
4.2、为什么这很重要
Mamba 论文强调:很多早期高效序列模型不如 Transformer 的一个重要原因,是它们缺乏足够强的内容选择能力,也就是不擅长做 content-based reasoning。
而 Mamba 的选择机制本质上是在说:
不是所有历史信息都同等重要,模型需要根据当前输入内容来决定该保留哪些信息。
这使它比传统 SSM 更适合文本这类离散、语义敏感的序列数据。
4.3、Mamba Block 可以怎样理解
从实现视角看,可以把一个 Mamba Block 粗略理解成下面这条链路:
- 输入先经过线性投影。
- 一支分支用于产生门控信号。
- 另一支分支经过局部卷积,增强短程局部建模能力。
- 然后进入选择性 SSM,对序列进行状态更新与信息传递。
- 最后经过输出投影得到结果。
这样设计的直觉是:
- 卷积负责局部模式
- 选择性 SSM 负责长距离记忆
- 门控机制负责信息筛选
4.4、每个通道一个 SSM
一个常见的理解方式是:Mamba 可以看作对不同通道并行地应用状态空间更新。这样既保留了状态模型的递推特性,也更利于高维特征建模。
4.5、为什么还能高效
一旦
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 带来了什么变化
从学习和工程角度,可以抓住以下几点:
- 更适合现代硬件加速。
- 更容易利用矩阵乘法等高吞吐计算。
- block 结构有更新,整体训练和推理效率更好。
- 论文中强调在不少设置下可比 Mamba 获得更高速度。
5.3、如何理解 Mamba 与 Mamba2 的关系
可以把二者理解成:
Mamba:把选择性状态空间模型真正做成了可用的强序列模型。Mamba2:在理论统一与硬件效率上更进一步,把这条路线推得更成熟。
6、Mamba、Transformer、RNN 的对比理解
| 模型 | 主要记忆方式 | 长距离建模 | 训练并行性 | 长序列复杂度 | 推理缓存特点 |
|---|---|---|---|---|---|
| RNN | 隐状态递推 | 一般 | 弱 | 近似线性 | 只维护状态 |
| Transformer | 自注意力 | 强 | 强 | 通常二次 | 需要 KV Cache |
| Mamba | 选择性状态递推 | 强 | 较强 | 近似线性 | 主要维护状态 |
可以这样记:
- RNN 是“有状态,但不够强”。
- Transformer 是“表达力强,但长序列代价高”。
- Mamba 是“想保留状态模型的高效,同时增强内容选择能力”。
7、Mamba 的安装与环境配置
IMPORTANT
参考博客:
1、Linux 下 Mamba 环境安装踩坑问题汇总(重置版)_linux安装mamba ssm1.1.3-CSDN博客
2、mambassm和causal-conv1d安装教程不同torch版本的mamba-ssm-CSDN博客
3、Anaconda虚拟环境中安装cudatoolkit和cudnn包并配置tensorflow-gpu_conda install cudatoolkit-CSDN博客
7.1、安装前先明确几个原则
Mamba 安装最容易踩坑的地方不在 Python 代码,而在版本匹配:
Python版本要和 whl 对应。PyTorch版本要和mamba-ssm、causal-conv1d对应。CUDA/nvcc/cudatoolkit版本要和 PyTorch 编译版本匹配。- 还要注意
cxx11abiTRUE/FALSE是否一致。
IMPORTANT
下面整理的是学习与实操笔记中的可行经验组合,重点在“能装通、能跑通”。如果未来版本升级,请优先按对应 release 的 whl 命名规则核对版本。
7.2、基础安装流程
# 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.41、对于
causal_conv1d和mamba_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、实操中记录过的可行组合
组合一
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组合二
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组合三
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.47.6、一份更贴近实操的安装记录
当机器上已经存在 nvcc 12.4 时,可以参考下面这组安装思路:
# 已存在 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.27.7、常见踩坑总结
torch版本对不上,导致mamba-ssm编译或导入失败。cxx11abiTRUE/FALSE选错,导致 wheel 安装后仍无法正常导入。causal_conv1d没先装,后续mamba_ssm可能失败。- CUDA 版本、PyTorch CUDA 版本、nvcc 版本混用,容易出现底层依赖错误。
transformers版本过新时,和某些mamba-ssm组合可能有兼容性问题。
8、Mamba 与 Mamba2 的代码测试
8.1、Mamba 测试
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 测试
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输出结果:
参数量 6468320
torch.Size([2, 64, 1024])8.3、建议增加一个最小化导入测试
在正式跑模型前,可以先做一个 smoke test:
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 句:
Mamba不是普通 RNN,而是建立在状态空间模型上的高效序列模型。- 它的关键创新是“选择性”,也就是让部分 SSM 参数随输入动态变化。
可以直观理解为一种控制保留与遗忘强度的机制。 Mamba的强点在于长序列下接近线性的复杂度和较高推理效率。Mamba2则进一步从理论和硬件实现上把这条路线推向成熟。
10、总结
Mamba 的价值不只是提出了一个新模型名字,而是在序列建模这件事上提供了一条非常有代表性的思路:
- 不一定所有问题都要用显式注意力解决。
- 记忆机制、状态更新、内容选择,也可以组合成强大的序列模型。
- 当序列长度继续增长时,线性复杂度模型会越来越有吸引力。
因此,理解 Mamba 的最好方式不是把它当成一个孤立模型,而是把它放回整个发展脉络中看:
RNN提供“状态记忆”的直觉,Transformer提供“内容交互”的强大能力,而Mamba试图把“高效状态更新”和“内容选择”重新结合起来。
