目标: 从零训练一个 Mamba/SSM 语言模型,或基于现有 SSM 模型微调,掌握状态空间模型的完整训练流程。
架构选择: Mamba-2(SSD 层)或 Mamba-3(最新),混合架构可选 Jamba/Zamba 风格。
硬件需求: 最小 1×A100 (80GB) 训练 130M 参数;8×A100 训练 1.3B-2.8B 参数。
资源清单
核心仓库
| 仓库 | 说明 | Stars |
|---|---|---|
| state-spaces/mamba | 官方 Mamba-1/2/3 实现 + 训练脚本 + CUDA kernel | ⭐18.2k |
| johnma2006/mamba-minimal | 单文件 PyTorch 极简 Mamba,适合学习 | ⭐2.9k |
| alxndrTL/mamba.py | 纯 PyTorch + MLX 实现,简洁高效 | ⭐1.5k |
| Zyphra/Zamba2 | 混合 SSM+Attention 训练代码(含完整训练 pipeline) | ⭐194 |
| microsoft/Samba | Microsoft 混合 Mamba+Attention 架构 | ⭐958 |
| HazyResearch/H3 | 最早 SSM+Attention 混合实验 | ⭐523 |
预训练模型(HuggingFace)
| 模型 | 参数量 | 下载 |
|---|---|---|
state-spaces/mamba-130m-hf | 130M | HF ↓389k |
state-spaces/mamba-370m-hf | 370M | HF |
state-spaces/mamba-790m-hf | 790M | HF |
state-spaces/mamba-1.4b-hf | 1.4B | HF |
state-spaces/mamba-2.8b-hf | 2.8B | HF |
state-spaces/mamba2-130m | 130M (Mamba-2) | HF |
state-spaces/mamba2-370m | 370M (Mamba-2) | HF |
state-spaces/mamba2-1.3b | 1.3B (Mamba-2) | HF |
state-spaces/mamba2-2.7b | 2.7B (Mamba-2) | HF |
mistralai/Mamba-Codestral-7B-v0.1 | 7B | HF |
tiiuae/falcon-mamba-7b | 7B | HF |
必读论文
| 论文 | 链接 | 关键词 |
|---|---|---|
| Mamba-1 | https://arxiv.org/abs/2312.00752 | 选择性 SSM,输入依赖参数,硬件感知算法 |
| Mamba-2 (SSD) | https://arxiv.org/abs/2405.21060 | 结构化状态空间对偶性,SSM↔Attention 数学统一 |
| Mamba-3 | https://arxiv.org/abs/2603.15569 | 推理优先设计,改进序列建模 |
| S4 | https://arxiv.org/abs/2111.00396 | 结构化状态空间奠基论文 |
| H3 | https://arxiv.org/abs/2212.14052 | SSM+Attention 混合先驱 |
| Jamba | https://arxiv.org/abs/2403.19887 | 生产级混合架构设计 |
| Zamba | https://arxiv.org/abs/2405.16712 | 小模型混合 SSM 方案 |
训练数据集
| 数据集 | 规模 | 用途 |
|---|---|---|
| SlimPajama | 627B tokens | 大规模预训练(RedPajama 去重版) |
| The Pile | 825GB | 经典预训练语料 |
| FineWeb-Edu | 1.3T tokens | 高质量教育内容 |
| C4 | 750GB | Colossal Clean Crawled Corpus |
| OpenWebText2 | — | 替代 WebText |
路径 A:从零预训练 Mamba
阶段 0:环境搭建
依赖安装:
# 安装 mamba-ssm(包含 CUDA kernel)pip install causal-conv1d>=1.4.0 --no-build-isolationpip install mamba-ssm --no-build-isolation
# 如需 Mamba-3(需要从源码安装)MAMBA_FORCE_BUILD=TRUE pip install --no-cache-dir --force-reinstall \ git+https://github.com/state-spaces/mamba.git --no-build-isolation
# 训练依赖pip install torch transformers datasets wandb einops阶段 1:最小原型(130M,单卡 A100,1 天)
参照 mamba-minimal 的结构,用官方 mamba-ssm 包搭建训练脚本:
任务 1.1:加载 tokenizer(GPT-2 tokenizer / 自定义 BPE)任务 1.2:构建 Mamba 模型配置(d_model, d_state, n_layer)任务 1.3:准备数据加载器(SlimPajama 子集,1B tokens)任务 1.4:训练循环(AdamW, cosine schedule, grad clip)任务 1.5:评估 perplexity,对比 Transformer 同参数量训练超参参考(Mamba-130M):
d_model=768,d_state=16,n_layer=24,expand=2lr=3e-3,batch_size=0.5M tokens,warmup=2K steps- 训练 10B-20B tokens(约 12h on A100)
阶段 2:扩量训练(370M-1.4B,多卡,3-7 天)
参照官方 state-spaces/mamba 的 benchmark 脚本:
任务 2.1:分布式训练配置(FSDP / DDP,参照 Zamba2 训练脚本)任务 2.2:数据 pipeline(SlimPajama 全量,预处理为 MosaicML 格式或 HF datasets 流式加载)任务 2.3:混合精度训练(bf16 + 梯度缩放)任务 2.4:使用 Mamba-2 SSD 层(吞吐 2-8x 提升)任务 2.5:评估 benchmark(LM eval harness:HellaSwag, PIQA, ARC, WinoGrande)关键超参(Mamba-1.4B):
d_model=2048,d_state=16,n_layer=48,expand=2lr=1.5e-3,batch_size=0.5M tokens,warmup=2K steps- 训练 100B+ tokens
- 使用 Mamba-2 的 SSD 层可在同等硬件上大幅加速
阶段 3:混合架构(可选,2.8B+,>8×A100)
参照 Zamba2 的 hybrid block 设计:
任务 3.1:实现 Jamba 风格混合层(每 N 层 SSM + 1 层 Attention)任务 3.2:Attention + SSM 比例调优(建议 attn_every=4 或 8)任务 3.3:长上下文训练(>128K tokens),利用 SSM 恒定显存优势路径 B:微调现成 Mamba 模型
如果不想从零训练,可以基于 HuggingFace 上已有的 Mamba/Mamba-2 模型微调。
框架选择
⚠️ Axolotl 和 Unsloth 目前主要针对 Transformer(Llama/Mistral/Qwen),对 Mamba 系列支持有限。 建议直接使用标准的
transformers+trl或自定义训练循环。
微调方案
| 方案 | 工具 | 适用场景 |
|---|---|---|
| 全量微调 | transformers Trainer + mamba-ssm | 领域适配,数据量较大 |
| LoRA | peft(已支持 Mamba,需测试) | 小数据,单卡即可 |
| SFT | trl.SFTTrainer(部分支持 Mamba) | 指令微调 |
# LoRA 微调示例框架from transformers import AutoModelForCausalLM, AutoTokenizerfrom peft import LoraConfig, get_peft_model
model = AutoModelForCausalLM.from_pretrained( "state-spaces/mamba-2.8b-hf", trust_remote_code=True, torch_dtype=torch.bfloat16)
lora_config = LoraConfig( r=16, lora_alpha=32, target_modules=["x_proj", "in_proj", "out_proj"], # Mamba 特有模块 lora_dropout=0.05, bias="none")model = get_peft_model(model, lora_config)注意事项
-
CUDA kernel 依赖:
mamba-ssm需要causal-conv1dCUDA 扩展,必须--no-build-isolation用已有 PyTorch CUDA 编译。 -
Mamba vs Mamba-2 选型:
- 从零训练优选 Mamba-2(SSD),速度 2-8x 提升
- 混合架构优选 Mamba-1(Jamba 发现 Mamba-1+Attention > Mamba-2+Attention)
- 如需试验新架构,可选 Mamba-3(需源码安装)
-
长上下文训练:SSM 的 O(N) 优势在 >57K tokens 时反超 Transformer,短序列反而慢,所以训练时 seq_len 至少 2K-8K,评估测 32K-128K。
-
微调限制:标准 fine-tuning 框架(Axolotl/Unsloth)对 Mamba 支持不完整,建议走自定义训练循环或直接用
transformersTrainer。 -
内存预算:
- Mamba-130M:< 10GB VRAM 训练
- Mamba-2.8B:~ 40GB VRAM 训练
- 长上下文 128K:SSM 内存恒定增长,远优于 Transformer
-
无 KV Cache 意味着:推理时不需要 KV cache,但 Mamba 推理是 RNN 式逐 token 串行,长生成速度可能低于 Transformer(后者有 KV cache 批量推理)。推理阶段需实测吞吐量与延迟。
推荐执行顺序
以下为计划目标。训练耗时和显存需求取决于数据量、序列长度、批量大小及实现,应先用小规模测试估算。
1 天:阶段 0(环境)+ 阶段 1(130M 原型)→ 验证 pipeline 可用1 周:阶段 2(1.4B 全量训练)→ 产出可用模型可选:阶段 3(混合架构 2.8B+)→ 寻找 SSM 与 Attention 的合适比例扩展阅读
- Awesome Efficient Architectures — 高效架构综述
- mamba-minimal blog — 极简 Mamba 解读
- The Annotated S4 — S4 逐行注释
- state-spaces/s4 — S4 官方仓库
- NVIDIA “Attention was never enough” blog — 混合架构趋势