
Mamba: 使用选择性状态空间的线性时间序列建模\
Albert Gu, Tri Dao\
论文: https://arxiv.org/abs/2312.00752

Transformer 即 SSM:通过结构化状态空间对偶的广义模型与高效算法\
Tri Dao, Albert Gu\
论文: https://arxiv.org/abs/2405.21060

Mamba-3:利用状态空间原理改进序列建模\
Aakash Lahoti, Kevin Y. Li, Berlin Chen, Caitlin Wang, Aviv Bick, J. Zico Kolter, Tri Dao†, Albert Gu†\
论文: https://arxiv.org/abs/2603.15569
Mamba 是一种新的状态空间模型架构,在语言建模等信息密集型数据上表现出良好的性能,而之前的次二次复杂度模型在此类任务上不如 Transformer。
它基于结构化状态空间模型的研究进展,并借鉴了 FlashAttention 的设计理念,采用了高效、硬件感知的实现方式。
首先安装 PyTorch。默认情况下,mamba-ssm 将安装核心包,而不编译 selective_scan_cuda 扩展,也不会下载缓存的 CUDA 轮子。
| 安装模式 | 命令 | 行为 |
|---|---|---|
| 核心包 | pip install mamba-ssm --no-build-isolation |
安装不包含 selective_scan_cuda 的 mamba-ssm,不编译 CUDA 扩展。 |
| 核心包加因果卷积 | pip install "mamba-ssm[causal-conv1d]" --no-build-isolation |
安装核心包及 causal-conv1d 额外依赖,仍然不包含 selective_scan_cuda。 |
| 强制本地构建核心包 | MAMBA_FORCE_BUILD=TRUE pip install mamba-ssm --no-build-isolation |
在本地构建默认纯 Python 轮子,仍然不包含 selective_scan_cuda。 |
| 启用 CUDA 选择性扫描 | MAMBA_KEEP_CUDA_BUILD=TRUE pip install mamba-ssm --no-build-isolation |
安装 selective_scan_cuda;pip 会先尝试匹配预先构建的 CUDA/HIP 轮子,若没有则本地编译。 |
| 强制本地构建 CUDA 选择性扫描 | MAMBA_FORCE_BUILD=TRUE MAMBA_KEEP_CUDA_BUILD=TRUE pip install mamba-ssm --no-build-isolation |
跳过缓存的轮子,在本地编译 selective_scan_cuda。 |
注意:CUDA 构建需要 --no-build-isolation,这样 pip 会使用你现有的支持 CUDA 的 PyTorch,而不是在隔离的构建环境中安装 torch-cpu。
从源码安装使用相同的默认值和可选项标志:
| 安装模式 | 命令 |
|---|---|
源码默认安装,不包含 selective_scan_cuda |
pip install . --no-build-isolation |
从 GitHub 源码默认安装,不包含 selective_scan_cuda |
pip install git+https://github.com/state-spaces/mamba.git --no-build-isolation |
源码强制本地构建,不包含 selective_scan_cuda |
MAMBA_FORCE_BUILD=TRUE pip install . --no-build-isolation |
| 源码安装启用 CUDA 选择性扫描 | MAMBA_KEEP_CUDA_BUILD=TRUE pip install . --no-build-isolation |
| 源码强制本地构建 CUDA 选择性扫描 | MAMBA_FORCE_BUILD=TRUE MAMBA_KEEP_CUDA_BUILD=TRUE pip install . --no-build-isolation |
注意:如需使用最新的 Mamba-3,请从源码安装。默认源码安装(不包含 selective_scan_cuda)使用 pip install git+https://github.com/state-spaces/mamba.git --no-build-isolation。要包含 CUDA 选择性扫描扩展,请添加 MAMBA_KEEP_CUDA_BUILD=TRUE;若需强制本地 CUDA 编译,再添加 MAMBA_FORCE_BUILD=TRUE。
核心要求:
- Linux
- Python 3.10+
- PyTorch 1.12+
CUDA selective_scan_cuda 构建和 GPU 执行的额外要求:
- NVIDIA GPU
- CUDA 11.6+
对于 AMD 显卡,请参阅下面的额外先决条件。
我们提供了多种级别的 Mamba 模型接口。
Mamba 基于选择性 SSM 层,这是论文的重点(第 3 节;算法 2)。
源代码: ops/selective_scan_interface.py.
本仓库的主要模块是封装了选择性 SSM 的 Mamba 架构块。
源代码: modules/mamba_simple.py.
用法:
import torch
from mamba_ssm import Mamba
batch, length, dim = 2, 64, 16
x = torch.randn(batch, length, dim).to("cuda")
model = Mamba(
# 此模块大约使用 3 * expand * d_model^2 个参数
d_model=dim, # 模型维度 d_model
d_state=16, # SSM 状态扩展因子
d_conv=4, # 局部卷积宽度
expand=2, # 块扩展因子
).to("cuda")
y = model(x)
assert y.shape == x.shape
Mamba-2 块实现位于 modules/mamba2.py。
更简化的版本位于 modules/mamba2_simple.py
用法与 Mamba(-1) 类似:
from mamba_ssm import Mamba2
model = Mamba2(
# 此模块大约使用 3 * expand * d_model^2 个参数
d_model=dim, # 模型维度 d_model
d_state=64, # SSM 状态扩展因子,通常为 64 或 128
d_conv=4, # 局部卷积宽度
expand=2, # 块扩展因子
).to("cuda")
y = model(x)
assert y.shape == x.shape
内部 SSD 模块的最小版本(来自 Mamba-2 论文的列表 1)以及“离散”和“连续”SSM 版本之间的转换,位于 modules/ssd_minimal.py。
Mamba-3 块实现位于 modules/mamba3.py。
用法如下:
from mamba_ssm import Mamba3
batch, length, dim = 2, 2048, 768
x = torch.randn(batch, length, dim).to(torch.bfloat16).to("cuda")
model = Mamba3(
# 此模块大约使用 6 * d_model^2 个参数
d_model=dim, # 模型维度 d_model
d_state=128, # SSM 状态大小
headdim=64, # SSM 头维度
is_mimo=True, # 使用 MIMO 模式
mimo_rank=4, # 当 is_mimo=True 时的 MIMO 秩
chunk_size=16, # 当 x 为 bf16 时为 64/mimo_rank,否则为 32/mimo_rank
is_outproj_norm=False, # 设置 SSM 后额外归一化
dtype=torch.bfloat16,
).to("cuda")
y = model(x)
assert y.shape == x.shape
最后,我们提供了一个完整语言模型的示例:一个深层序列模型主干(带有重复的 Mamba 块)+ 语言模型头。
源代码: models/mixer_seq_simple.py.
这是一个如何将 Mamba 集成到端到端神经网络中的示例。
下面的生成脚本中使用了这个示例。
预训练模型已上传到
Hugging Face: mamba-130m, mamba-370m,
mamba-790m, mamba-1.4b, mamba-2.8b, mamba2-130m, mamba2-370m,
mamba2-780m, mamba2-1.3b, mamba2-2.7b, transformerpp-2.7b, mamba2attn-2.7b,在 Pile 数据集上训练了 300B 个标记,还有 mamba-2.8b-slimpj(在 SlimPajama 数据集上训练了 600B 个标记)。
模型将由下面的生成脚本自动下载。
这些模型在 Pile 上训练,遵循 GPT-3 描述并被许多开源模型采用的标准模型维度:
| 参数 | 层数 | 模型维度 |
|---|---|---|
| 130M | 24 | 768 |
| 370M | 48 | 1024 |
| 790M | 48 | 1536 |
| 1.4B | 48 | 2048 |
| 2.8B | 64 | 2560 |
(Mamba 的层数是相似大小 Transformer 的两倍,因为 Transformer 的每个“层”(MHA 块 + MLP 块)需要两个 Mamba 块。)
注意:这些基础模型仅训练了 300B 个标记,没有经过任何下游修改(如指令微调等)。
性能预计与在类似数据上训练的其他架构相当或更好,但不会匹配更大或微调过的模型。
要运行模型的零样本评估(对应论文的表 3),
我们使用 lm-evaluation-harness 库。
pip install lm-eval==0.4.2 安装 lm-evaluation-harness。lm_eval --model mamba_ssm --model_args pretrained=state-spaces/mamba-130m --tasks lambada_openai,hellaswag,piqa,arc_easy,arc_challenge,winogrande,openbookqa --device cuda --batch_size 256
python evals/lm_harness_eval.py --model hf --model_args pretrained=EleutherAI/pythia-160m --tasks lambada_openai,hellaswag,piqa,arc_easy,arc_challenge,winogrande --device cuda --batch_size 64
要复现博客文章中关于 mamba-2.8b-slimpj 模型的结果:
lm_eval --model mamba_ssm --model_args pretrained=state-spaces/mamba-2.8b-slimpj --tasks boolq,piqa,hellaswag,winogrande,arc_easy,arc_challenge,openbookqa,race,truthfulqa_mc2 --device cuda --batch_size 256
lm_eval --model mamba_ssm --model_args pretrained=state-spaces/mamba-2.8b-slimpj --tasks mmlu --num_fewshot 5 --device cuda --batch_size 256
要运行 Mamba-2 模型的评估,只需替换模型名称:
lm_eval --model mamba_ssm --model_args pretrained=state-spaces/mamba2-2.7b --tasks lambada_openai,hellaswag,piqa,arc_easy,arc_challenge,winogrande,openbookqa --device cuda --batch_size 256
lm_eval --model mamba_ssm --model_args pretrained=state-spaces/transformerpp-2.7b --tasks lambada_openai,hellaswag,piqa,arc_easy,arc_challenge,winogrande,openbookqa --device cuda --batch_size 256
lm_eval --model mamba_ssm --model_args pretrained=state-spaces/mamba2attn-2.7b --tasks lambada_openai,hellaswag,piqa,arc_easy,arc_challenge,winogrande,openbookqa --device cuda --batch_size 256
请注意,由于评估过程中的噪声,每个任务的结果可能与报告的值有 0.1-0.3 的差异。
脚本 benchmarks/benchmark_generation_mamba_simple.py 能够:
1. 从 Hugging Face Hub 自动加载模型,
2. 生成用户指定提示的补全,
3. 基准测试此生成的推断速度。
其他可配置选项包括 top-p(核采样)概率和 softmax 温度。
要测试不同采样策略下的生成延迟(例如批大小 = 1):
python benchmarks/benchmark_generation_mamba_simple.py --model-name "state-spaces/mamba-2.8b" --prompt "My cat wrote all this CUDA code for a new language model and" --topp 0.9 --temperature 0.7 --repetition-penalty 1.2
python benchmarks/benchmark_generation_mamba_simple.py --model-name "EleutherAI/pythia-2.8b" --prompt "My cat wrote all this CUDA code for a new language model and" --topp 0.9 --temperature 0.7 --repetition-penalty 1.2
python benchmarks/benchmark_generation_mamba_simple.py --model-name "state-spaces/mamba-2.8b" --prompt "My cat wrote all this CUDA code for a new language model and" --minp 0.05 --topk 0 --temperature 0.7 --repetition-penalty 1.2
要测试随机提示下的生成吞吐量(例如大批大小):
python benchmarks/benchmark_generation_mamba_simple.py --model-name "state-spaces/mamba-2.8b" --batch 64
python benchmarks/benchmark_generation_mamba_simple.py --model-name "EleutherAI/pythia-2.8b" --batch 64
使用 Mamba-2,只需更改模型名称:
python benchmarks/benchmark_generation_mamba_simple.py --model-name "state-spaces/mamba2-2.7b" --prompt "My cat wrote all this CUDA code for a new language model and" --topp 0.9 --temperature 0.7 --repetition-penalty 1.2
我们的模型使用 PyTorch AMP 进行混合精度训练。AMP 将模型参数保存在 float32 中,并在必要时转换为半精度。
另一方面,其他框架如 DeepSpeed 将参数存储在 float16 中,并在必要时向上转换(例如用于优化器累积)。
我们观察到,主要模型参数可能需要更高精度,因为 SSM 对其循环动态很敏感。如果您遇到不稳定性,
首先请尝试使用将参数存储在 fp32 中的框架(如 AMP)。
模型的某些部分继承了先前 S4 模型工作的初始化。
例如,此处,通过初始化其线性投影的偏置,$\Delta$ 参数具有目标范围。
但是,某些框架可能有初始化后钩子(例如将所有 nn.Linear 模块的偏置项设置为零)。
如果是这种情况,您可能需要添加自定义逻辑(例如,这行 在我们的训练器中关闭了重新初始化,但在任何其他框架中将是空操作),
具体取决于训练框架。
如果您使用的是 ROCm 6.0,请执行以下步骤以避免编译期间的错误。ROCm 6.1 及更高版本不需要此操作。
找到您的 ROCm 安装目录。这通常位于 /opt/rocm/,但可能因安装方式而异。
应用补丁。如果遇到权限问题,请使用 sudo 运行。
bash
patch /opt/rocm/include/hip/amd_detail/amd_hip_bf16.h < rocm_patch/rocm6_0.patch
如果您使用此代码库,或以其他方式认为我们的工作有价值,请引用 Mamba:
@article{mamba,
title={Mamba: Linear-Time Sequence Modeling with Selective State Spaces},
author={Gu, Albert and Dao, Tri},
journal={arXiv preprint arXiv:2312.00752},
year={2023}
}
@inproceedings{mamba2,
title={Transformers are {SSM}s: Generalized Models and Efficient Algorithms Through Structured State Space Duality},
author={Dao, Tri and Gu, Albert},
booktitle={International Conference on Machine Learning (ICML)},
year={2024}
}
@misc{lahoti2026mamba3improvedsequencemodeling,
title={Mamba-3: Improved Sequence Modeling using State Space Principles},
author={Aakash Lahoti and Kevin Y. Li and Berlin Chen and Caitlin Wang and Aviv Bick and J. Zico Kolter and Tri Dao and Albert Gu},
year={2026},
eprint={2603.15569},
archivePrefix={arXiv},
primaryClass={cs.LG},
url={https://arxiv.org/abs/2603.15569},
}