深入理解 Mamba:选择性状态空间模型的架构与实现

5571 字
28 分钟
深入理解 Mamba:选择性状态空间模型的架构与实现

1. Mamba 诞生的背景#

1.1 Transformer 的局限#

虽然 Transformer 主导了深度学习领域,但它有两大根本性缺陷

  1. 二次方复杂度:注意力机制需要计算所有 token 对之间的关系,复杂度为 O(L2)O(L^2)
  2. 无限上下文无法扩展:KV 缓存随序列长度线性增长,无法处理真正长的序列

核心问题

Transformer 的”全连接注意力”是一种”过度密集”的机制——所有 token 都必须相互通信,无论它们之间是否真的相关。

1.2 RNN 的复兴希望#

RNN 具有线性复杂度 O(L)O(L),但传统 RNN 有两大问题:

  • 训练困难:反向传播需要展开整个时间序列,梯度消失/爆炸
  • 并行性差:顺序依赖使得训练必须逐 token 进行

理想状态

像 RNN 一样高效(线性),像 Transformer 一样强大(长距离依赖),像 CNN 一样可并行训练。

这正是 Mamba 的设计目标!

1.3 SSM 的前世今生#

SSM 的发展路径:

1960s: 卡尔曼滤波
│ 现代 SSM 的理论基础
2019: HiPPO (Dao et al.)
│ "用最优投影递归记忆历史"
│ 解决了"如何压缩历史"的问题
2021: S4 (Gu, Goel, Re)
│ "结构化状态空间序列模型"
│ 将 SSM 引入深度学习,验证可行性
2023: Mamba (Gu & Dao)
│ "选择性状态空间模型"
│ 解决了 SSM "无法选择性遗忘"的关键问题
2024: Mamba-2 (Dao & Gu)
│ "状态空间对偶性 (SSD)"
│ 将 SSM 与注意力机制统一起来
2024-2025: 混合架构爆发
│ Jamba, StripedHyena, Falcon-Mamba, ...
│ SSM + Transformer 成为新趋势

1.4 Mamba 的革命性意义#

Mamba 是深度学习序列建模的范式转变

维度TransformerMamba
计算复杂度O(L2)O(L^2)O(L)O(L)
状态存储O(L)O(L) (KV cache)O(N)O(N) (固定)
训练并行性完全并行高度并行 (并行扫描)
推理速度慢 (随长度增长)快 (常数时间增量)
长序列能力有限极强
检索能力中等
选择性选择性 (选择性扫描)

2. Mamba 的核心思想:选择性状态空间#

2.1 SSM 的根本缺陷#

让我们回顾标准 SSM 的方程:

x(t)=Ax(t)+Bu(t)x'(t) = \mathbf{A} x(t) + \mathbf{B} u(t)y(t)=Cx(t)y(t) = \mathbf{C} x(t)

问题A\mathbf{A}B\mathbf{B}C\mathbf{C}与输入无关的常数矩阵

这意味着:

  • 所有输入都被同等对待(无差别压缩)
  • 无法根据上下文调整行为(缺乏选择性)
  • 不擅长检索任务(无法精准访问特定历史)

2.2 选择性的直觉#

想象你在读一本侦探小说:

标准 SSM(无选择性)

  • 像一个”匀速抹除器”,每读一页就按固定比例覆盖前面的记忆
  • 不管是重要线索还是废话,一视同仁

Mamba(有选择性)

  • 像一个”智能读者”,会主动记笔记
  • 遇到重要线索时:主动加强记忆(增大 B\mathbf{B} 的更新幅度)
  • 遇到无关内容时:主动遗忘(降低更新幅度)

2.3 选择性的数学表达#

Mamba 的核心创新:让 B\mathbf{B}C\mathbf{C}Δ\mathbf{\Delta} 成为输入的函数

标准 SSM(线性时不变,LTI):

xk+1=Aˉxk+Bˉukx_{k+1} = \mathbf{\bar{A}} x_k + \mathbf{\bar{B}} u_k

Mamba SSM(线性时变,LTV):

xk+1=Aˉkxk+Bˉkukx_{k+1} = \mathbf{\bar{A}}_k x_k + \mathbf{\bar{B}}_k u_k

其中:

  • Bˉk=LinearB(xk)\mathbf{\bar{B}}_k = \text{Linear}_B(x_k)
  • Cˉk=LinearC(xk)\mathbf{\bar{C}}_k = \text{Linear}_C(x_k)
  • Aˉk=exp(ΔkA)\mathbf{\bar{A}}_k = \exp(\Delta_k \cdot \mathbf{A})Δk=softplus(LinearΔ(xk))\Delta_k = \text{softplus}(\text{Linear}_\Delta(x_k))
选择性的关键意义
  1. 使 SSM 变为非线性:依赖输入的参数打破了线性时不变 (LTI) 假设
  2. 使 SSM 具有因果选择性:可以”遗忘”和”记忆”
  3. 理论上变强大:可以模拟任何依赖历史的算法(包括有限状态机)

2.4 Mamba vs 线性注意力#

线性注意力也可以重写为递归形式:

线性注意力 (RetNet):
s_k = γ s_{k-1} + k_k^T v_k
y_k = q_k s_k
Mamba:
x_k = A_k x_{k-1} + B_k u_k
y_k = C_k x_k

两者形式类似,但关键区别

  • 线性注意力的 γ\gamma 通常是固定常数
  • Mamba 的 Ak\mathbf{A}_k输入相关的

这种输入依赖性是 Mamba 性能优异的关键!

3. Mamba 架构详解#

3.1 整体架构#

Mamba 模型由多个 Mamba 块 堆叠而成,结构类似 Transformer 但去掉了 attention 层:

┌─────────────────────────────────────────────────────────────────┐
│ 完整 Mamba 架构 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ Token Embedding │
│ │ │
│ ▼ │
│ ┌─────────────────────────────────────┐ │
│ │ Mamba Block (×N 层) │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ 输入归一化 │ │ │
│ │ └──────────────┬────────────────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ 输入投影 → Split │ │ │
│ │ │ (x_inner, gate) │ │ │
│ │ └──────────────┬────────────────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ 1D 深度可分离卷积 │ │ │
│ │ │ (捕获局部上下文) │ │ │
│ │ └──────────────┬────────────────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ SiLU 激活 │ │ │
│ │ └──────────────┬────────────────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ 选择性 SSM │ │ │
│ │ │ (选择性扫描 + 状态更新) │ │ │
│ │ └──────────────┬────────────────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ 门控乘法 (× SiLU(gate)) │ │ │
│ │ └──────────────┬────────────────┘ │ │
│ │ │ │ │
│ │ ▼ │ │
│ │ ┌───────────────────────────────┐ │ │
│ │ │ 输出投影 │ │ │
│ │ └───────────────────────────────┘ │ │
│ └─────────────────┬───────────────────┘ │
│ │ │
│ ▼ │
│ 输出归一化 │
│ │
└─────────────────────────────────────────────────────────────────┘

3.2 Mamba 块的内部结构#

每个 Mamba 块包含以下组件:

class MambaBlock(nn.Module):
"""
Mamba 选择性状态空间模型块
核心组件:
1. 输入投影 (In-Projection): 将输入映射到高维
2. 1D 卷积: 捕获局部上下文
3. 选择性 SSM: 长程依赖建模
4. 门控机制: 动态特征选择
5. 输出投影 (Out-Projection): 映射回原维度
"""
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.d_conv = d_conv
self.d_inner = int(expand * d_model)
# 1. 输入投影
# 将 d_model 投影到 d_inner,并分成 x 和 gate 两部分
self.in_proj = nn.Linear(d_model, self.d_inner * 2, bias=False)
# 2. 1D 深度可分离卷积
self.conv1d = nn.Conv1d(
in_channels=self.d_inner,
out_channels=self.d_inner,
kernel_size=d_conv,
padding=d_conv - 1,
groups=self.d_inner, # 深度可分离
bias=True
)
# 3. SSM 参数投影(输入相关)
# 输入依赖的 B, C, Δ 参数
self.x_proj = nn.Linear(
self.d_inner, d_state * 2 + 1, bias=False
) # 输出: B (N), C (N), Δ (1)
# 4. Δ 的投影层
self.dt_proj = nn.Linear(1, self.d_inner, bias=True)
# 5. A 矩阵(log 形式以保证稳定性)
# 初始化为 HiPPO 结构
A_init = self._init_hippo_matrix(d_state, self.d_inner)
self.A_log = nn.Parameter(torch.log(A_init))
self.A_log._no_weight_decay = True
# 6. D 矩阵(直接馈通 / skip connection)
self.D = nn.Parameter(torch.ones(self.d_inner))
# 7. 输出投影
self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)

3.3 HiPPO 初始化#

Mamba 使用 HiPPO 矩阵的变体来初始化 A\mathbf{A}

def _init_hippo_matrix(self, N, D):
"""
HiPPO-LegS 矩阵初始化
N: 状态维度
D: 并行通道数(d_inner)
"""
A = torch.zeros(D, N)
# 基于 HiPPO-LegS 的简化形式
for d in range(D):
for n in range(N):
if n == 0:
# 第一行特殊处理
A[d, n] = 1.0
else:
# 负的对角线,捕获长期依赖
A[d, n] = -(2 * n + 1) ** 0.5
return A
为什么 HiPPO 矩阵如此重要?

HiPPO 矩阵初始化让 SSM 能够最优地近似任意连续函数。如果随机初始化,SSM 几乎无法训练。

4. 选择性扫描 (Selective Scan)#

4.1 为什么需要选择性扫描?#

选择性 SSM 的计算需要按顺序进行

x1=Aˉ1x0+Bˉ1u1x_1 = \mathbf{\bar{A}}_1 x_0 + \mathbf{\bar{B}}_1 u_1x2=Aˉ2x1+Bˉ2u2x_2 = \mathbf{\bar{A}}_2 x_1 + \mathbf{\bar{B}}_2 u_2x3=Aˉ3x2+Bˉ3u3x_3 = \mathbf{\bar{A}}_3 x_2 + \mathbf{\bar{B}}_3 u_3

这看似无法并行!但实际上,Mamba 使用了并行扫描算法 (Parallel Scan) 来加速。

4.2 并行扫描原理#

并行扫描通过树形规约实现 O(logN)O(\log N) 的并行步骤:

序列扫描 (顺序): O(N) 步
x_1 → x_2 → x_3 → ... → x_N
↓ ↓ ↓ ↓
T1 T1 T1 T1
并行扫描 (树形): O(log N) 步
┌─────┬─────┐
│ │ │
┌───┴─┐ ┌─┴───┐ │
│ │ │ │ │
┌─┴─┐ ┌─┴─┐ ┌─┴─┐ ┌─┴─┐
T1 T2 T3 T4 T5 T6 T7 T8

4.3 并行扫描的实现#

def parallel_scan(a, b, x):
"""
并行扫描算法
计算 y_i = sum_{j=0}^{i} (prod_{k=j+1}^{i} a_k) * b_j * x_j
a: (T, N, N) - 状态转移矩阵
b: (T, N) - 输入矩阵
x: (T, N) - 输入
返回: y (T, N) - 输出
"""
T, N = x.shape
# 初始化累加器
acc = torch.zeros(T, N, N, dtype=torch.float64)
y = torch.zeros_like(x)
# 树形规约(work-efficient 扫描)
chunk_size = 1
while chunk_size < T:
for i in range(0, T - chunk_size, 2 * chunk_size):
# 组合相邻块
a_left = a[i + chunk_size - 1]
a_right = a[i + 2 * chunk_size - 1] if i + 2 * chunk_size - 1 < T else torch.eye(N)
# 累积矩阵乘积
acc[i + 2 * chunk_size - 1] = a_right @ acc[i + chunk_size - 1]
if i + 2 * chunk_size - 1 < T:
acc[i + 2 * chunk_size - 1] += a[i + 2 * chunk_size - 1] @ a[i + chunk_size - 1]
chunk_size *= 2
return y

4.4 完整的选择性扫描#

def selective_scan(self, u, delta, A, B, C, D):
"""
选择性扫描算法 - Mamba 的核心
u: (batch, L, d_inner) - 输入序列
delta: (batch, L, d_inner) - 输入依赖的步长
A: (d_inner, d_state) - 状态转移矩阵
B: (batch, L, d_state) - 输入依赖的输入矩阵
C: (batch, L, d_state) - 输入依赖的输出矩阵
D: (d_inner,) - 直接馈通
返回: (batch, L, d_inner) - 输出序列
"""
batch, L, d_inner = u.shape
N = A.shape[-1]
# 1. 离散化 A
# delta: (batch, L, d_inner) → (batch, L, d_inner, 1)
# A: (d_inner, N) → (1, 1, d_inner, N)
delta_A = torch.exp(delta.unsqueeze(-1) * A.unsqueeze(0).unsqueeze(0))
# 形状: (batch, L, d_inner, N)
# 2. 离散化 B
# B: (batch, L, N) → (batch, L, 1, N)
# u: (batch, L, d_inner) → (batch, L, d_inner, 1)
delta_B_u = delta.unsqueeze(-1) * B.unsqueeze(2) * u.unsqueeze(-1)
# 形状: (batch, L, d_inner, N)
# 3. 顺序扫描(训练)或 并行扫描(推理优化)
h = torch.zeros(batch, d_inner, N, device=u.device, dtype=u.dtype)
hs = []
for i in range(L):
h = delta_A[:, i] * h + delta_B_u[:, i]
hs.append(h)
# 4. 收集所有隐藏状态
hs = torch.stack(hs, dim=1) # (batch, L, d_inner, N)
# 5. 通过 C 计算输出
# C: (batch, L, N) → (batch, L, 1, N)
y = (hs * C.unsqueeze(2).unsqueeze(-1)).sum(-1)
# 形状: (batch, L, d_inner)
# 6. 添加直接馈通 D
y = y + u * D
return y

4.5 选择性扫描的并行优化#

实际部署中,Mamba 使用自定义 CUDA 内核实现真正的高效并行扫描:

优化技术效果
内存访问合并减少 IO 瓶颈
寄存器重用减少内存访问
Warp 级别原语利用 GPU 硬件特性
分块计算平衡并行度和效率

5. Mamba 完整实现#

5.1 RMSNorm 实现#

class RMSNorm(nn.Module):
"""
Root Mean Square Layer Normalization
比 LayerNorm 更简单高效
"""
def __init__(self, d_model, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(d_model))
def forward(self, x):
# 计算 RMS
norm = x.norm(dim=-1, keepdim=True) * (x.shape[-1] ** -0.5)
return self.weight * x / (norm + self.eps)

5.2 完整 Mamba 块#

class MambaBlock(nn.Module):
"""
完整的 Mamba 块实现
"""
def __init__(self, d_model, d_state=16, d_conv=4, expand=2, dt_rank="auto"):
super().__init__()
self.d_model = d_model
self.d_state = d_state
self.d_conv = d_conv
self.d_inner = int(expand * d_model)
self.dt_rank = d_model // 16 if dt_rank == "auto" else dt_rank
# 输入投影
self.in_proj = nn.Linear(d_model, self.d_inner * 2, bias=False)
# 1D 深度卷积
self.conv1d = nn.Conv1d(
self.d_inner, self.d_inner,
kernel_size=d_conv,
padding=d_conv - 1,
groups=self.d_inner,
bias=True
)
# SSM 参数
self.x_proj_weight = nn.Parameter(torch.empty(
self.dt_rank + d_state * 2, # dt, B, C
self.d_inner
))
# 初始化
nn.init.kaiming_uniform_(self.x_proj_weight, a=math.sqrt(5))
self.dt_proj_weight = nn.Parameter(torch.empty(self.d_inner, self.dt_rank))
nn.init.kaiming_uniform_(self.dt_proj_weight, a=math.sqrt(5))
# A 矩阵(log)
A = torch.arange(1, d_state + 1, dtype=torch.float32).repeat(self.d_inner, 1)
A = -torch.log(A) # 负值,HiPPO 风格
self.A_log = nn.Parameter(torch.log(A))
self.A_log._no_weight_decay = True
# D 参数
self.D = nn.Parameter(torch.ones(self.d_inner))
# 输出投影
self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)
def forward(self, x):
"""
x: (batch, L, d_model)
返回: (batch, L, d_model)
"""
batch, L, d = x.shape
# 1. 输入投影 + 分裂
xz = self.in_proj(x) # (batch, L, 2*d_inner)
x_inner, z = xz.chunk(2, dim=-1) # 各 (batch, L, d_inner)
# 2. 1D 卷积 + 激活
x_conv = self.conv1d(x_inner.transpose(1, 2))[:, :, :L].transpose(1, 2)
x_conv = F.silu(x_conv)
# 3. 选择性 SSM
# 3a. 投影得到 Δ, B, C
x_dbl = F.linear(x_conv, self.x_proj_weight) # (batch, L, dt_rank + 2*d_state)
dt, B, C = x_dbl.split([self.dt_rank, self.d_state, self.d_state], dim=-1)
# 3b. dt 投影
dt = F.linear(dt, self.dt_proj_weight) # (batch, L, d_inner)
# 3c. 选择性扫描
y = self._selective_scan(x_conv, dt, self.A_log, B, C, self.D)
# 4. 门控
y = y * F.silu(z)
# 5. 输出投影
output = self.out_proj(y)
return output
def _selective_scan(self, u, delta, A_log, B, C, D):
"""完整的选择性扫描"""
# A_log → A
A = -torch.exp(A_log.float()) # (d_inner, d_state)
batch, L, d_inner = u.shape
N = A.shape[1]
# 离散化
dA = torch.exp(delta.unsqueeze(-1) * A) # (batch, L, d_inner, N)
dB = delta.unsqueeze(-1) * B.unsqueeze(2) # (batch, L, d_inner, N)
dB_u = dB * u.unsqueeze(-1) # (batch, L, d_inner, N)
# 顺序扫描
h = torch.zeros(batch, d_inner, N, device=u.device, dtype=torch.float32)
hs = []
for i in range(L):
h = dA[:, i] * h + dB_u[:, i]
hs.append(h)
hs = torch.stack(hs, dim=1) # (batch, L, d_inner, N)
# 输出
y = (hs * C.unsqueeze(2).unsqueeze(-1)).sum(-1).to(u.dtype)
y = y + u * D
return y

5.3 完整 Mamba 模型#

class MambaLanguageModel(nn.Module):
"""
完整的 Mamba 语言模型
"""
def __init__(self, vocab_size, d_model=768, n_layers=24,
d_state=16, d_conv=4, expand=2, dropout=0.0):
super().__init__()
self.vocab_size = vocab_size
self.d_model = d_model
# Token 嵌入
self.embedding = nn.Embedding(vocab_size, d_model)
# Mamba 块堆叠
self.layers = nn.ModuleList([
MambaBlock(d_model, d_state, d_conv, expand)
for _ in range(n_layers)
])
# 归一化
self.norm_f = RMSNorm(d_model)
# 输出头
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
# Dropout
self.dropout = nn.Dropout(dropout)
# 权重绑定
self.lm_head.weight = self.embedding.weight
def forward(self, input_ids, labels=None):
"""
input_ids: (batch, seq_len)
"""
# 嵌入
x = self.embedding(input_ids) # (batch, L, d_model)
x = self.dropout(x)
# Mamba 层
for layer in self.layers:
x = x + layer(x) # 残差连接
x = self.dropout(x)
# 最终归一化
x = self.norm_f(x)
# LM 头
logits = self.lm_head(x) # (batch, L, vocab_size)
# 损失计算
loss = None
if labels is not None:
loss = F.cross_entropy(
logits.view(-1, self.vocab_size),
labels.view(-1)
)
return logits, loss

6. Mamba-2:状态空间对偶性 (SSD)#

6.1 SSD 理论#

Mamba-2 的核心理论突破:状态空间模型与注意力的对偶性

关键洞察

当 SSM 矩阵 A\mathbf{A} 是对角矩阵时,SSM 可以被重写为一种特殊的注意力形式——这就是状态空间对偶性 (SSD)。

SSD 的数学形式:

SSM 形式

xk+1=Axk+Bukx_{k+1} = \mathbf{A} x_k + \mathbf{B} u_kyk=Cxky_k = \mathbf{C} x_k

SSD 形式(对偶注意力)

yk=j=1kCkAk:jBjujy_k = \sum_{j=1}^{k} \mathbf{C}_k \mathbf{A}_{k:j} \mathbf{B}_j u_j

其中 Ak:j=AkAk1Aj+1\mathbf{A}_{k:j} = \mathbf{A}_{k} \cdot \mathbf{A}_{k-1} \cdots \mathbf{A}_{j+1}

这与注意力机制的形式非常相似:

yk=j=1kattn(k,j)vjy_k = \sum_{j=1}^{k} \text{attn}(k, j) \cdot v_j

6.2 Mamba-2 架构改进#

特性MambaMamba-2
A 矩阵一般矩阵(HiPPO 初始化)强制对角
计算方式顺序扫描并行矩阵乘积
算法自定义 CUDA利用矩阵乘法原语
性能5× 注意力8× 注意力
与注意力关系隐式显式 (SSD)

6.3 Mamba-2 的实际收益#

Mamba-2 相比 Mamba 的核心改进:

  1. 2× 训练速度提升:利用硬件优化过的矩阵乘法
  2. 理论清晰:与注意力机制的形式统一
  3. 更好的扩展性:在更大规模上表现更稳定

7. 硬件感知算法 (Hardware-Aware Algorithm)#

7.1 为什么 Mamba 特别快?#

Mamba 论文中提出的硬件感知算法是其速度优势的关键:

传统顺序计算:
for i in range(L):
h = A * h + B * x
y[i] = C * h
每一步:内存 IO 密集,无法利用 GPU 并行性
Mamba 硬件感知扫描:
- 在 GPU HBM (高带宽内存) 和 SRAM (片上缓存) 之间高效数据搬运
- 融合多个操作为单个 CUDA 内核
- 最大化计算密度

7.2 GPU 内存层次#

GPU 内存层次(从快到慢):
┌──────────────────┐
│ Register │ < 1 cycle ← 最快
├──────────────────┤
│ L1 Cache │ ~30 cycles
├──────────────────┤
│ Shared Memory │ ~30 cycles ← 片上 (SRAM)
├──────────────────┤
│ L2 Cache │ ~200 cycles
├──────────────────┤
│ HBM (VRAM) │ ~500 cycles ← 高带宽内存
└──────────────────┘

关键:HBM 与 SRAM 之间的 IO 速度差是数量级的。Mamba 通过减少 IO 实现加速。

7.3 硬件感知扫描的具体优化#

# 伪代码:Mamba 的硬件感知扫描
def mamba_hardware_aware_scan(A, B, C, x, chunk_size=64):
"""
分块 + 重计算优化
核心思想:
1. 将序列分成固定大小的块
2. 每个块内的状态保存到 HBM
3. 块内重计算隐藏状态(用计算换 IO)
"""
L = x.shape[0]
num_chunks = L // chunk_size
# 阶段 1: 块间顺序处理(IO 密集)
states = [] # 保存到 HBM
h = zero_state
for i in range(num_chunks):
# 块内计算
h = process_chunk(A, B, C, x[i*chunk_size:(i+1)*chunk_size], h)
states.append(h)
# 阶段 2: 块内并行处理(计算密集)
for i in range(num_chunks):
# 重计算隐藏状态(不访问 HBM)
local_h = recompute_states(A, B, C, x[i*chunk_size:(i+1)*chunk_size], states[i])
# 并行计算输出
y[i*chunk_size:(i+1)*chunk_size] = compute_outputs(local_h, C)

7.4 性能对比#

在不同序列长度下的推理速度(Mamba vs Transformer):

序列长度 Mamba Transformer 加速比
─────────────────────────────────────────
512 1.0 ms 1.0 ms 1×
2K 4.2 ms 8.5 ms 2×
8K 16 ms 130 ms 8×
32K 65 ms 2100 ms 32×
128K 260 ms 33000 ms 127×
Mamba 的核心优势
  • 线性复杂度:序列长度增长只带来线性增长
  • 常数增量推理:每生成一个新 token 的时间是常数
  • 长上下文友好:处理 100K+ tokens 不卡顿

8. Mamba 与 Transformer 的实验对比#

8.1 语言建模基准#

基准Mamba (130M)Pythia (130M)差异
WikiText103 (PPL)18.521.4Mamba 优 13%
LAMBADA (Acc)51.5%46.7%Mamba 优 4.8%
PIQA (Acc)65.4%64.0%Mamba 优 1.4%
HellaSwag (Acc)53.5%49.1%Mamba 优 4.4%
ARC-e (Acc)60.8%59.1%Mamba 优 1.7%

8.2 长上下文能力测试#

有效上下文长度测试

有效上下文长度 (Effective Context Length)
Pythia-160M: ████████░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 2K
Mamba-130M: ████████████████████░░░░░░░░░░░░░░░░░░ 8K
Mamba-370M: ██████████████████████████░░░░░░░░░░░ 16K
Mamba-780M: ███████████████████████████████░░░░░░ 32K

8.3 检索任务#

Mamba 在”检索”任务上仍有差距,这是 Transformer 的传统强项:

任务 Transformer Mamba 差距
─────────────────────────────────────────────────────
查找特定字符串出现次数 95% 45% 50%
复制粘贴模式检索 90% 52% 38%
精确长距离依赖检索 88% 41% 47%
检索能力差距

Mamba 的有限状态空间使其难以做”精确检索”。但通过混合架构(Mamba + Transformer)可以弥补。

9. Mamba 的训练技巧#

9.1 学习率策略#

# 推荐的学习率调度
optimizer = torch.optim.AdamW(
model.parameters(),
lr=6e-4, # 基础学习率
betas=(0.9, 0.95), # Mamba 推荐 β2
weight_decay=0.1
)
# Cosine 学习率调度
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer,
T_max=total_steps,
eta_min=6e-5
)
# Warmup
warmup_steps = 2000

9.2 训练稳定性技巧#

# 1. A_log 不应用权重衰减
no_decay_params = ['A_log', 'D']
param_groups = [
{'params': [p for n, p in model.named_parameters()
if not any(nd in n for nd in no_decay_params)],
'weight_decay': 0.1},
{'params': [p for n, p in model.named_parameters()
if any(nd in n for nd in no_decay_params)],
'weight_decay': 0.0}
]
# 2. 梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# 3. 混合精度训练
with torch.cuda.amp.autocast():
logits, loss = model(input_ids, labels=labels)

9.3 数据处理建议#

# 推荐的数据配比
data_mix = {
'web_text': 0.40, # 高质量网页文本
'books': 0.15, # 长篇书籍
'code': 0.20, # 编程代码
'academic': 0.10, # 学术论文
'dialog': 0.10, # 对话数据
'multilingual': 0.05, # 多语言
}
# 推荐序列打包
def pack_sequences(sequences, max_length=4096):
"""
打包短序列到固定长度
提高训练效率
"""
packed = []
current = []
for seq in sequences:
if sum(len(s) for s in current) + len(seq) > max_length:
packed.append(current)
current = [seq]
else:
current.append(seq)
if current:
packed.append(current)
return packed

10. Mamba 的应用生态#

10.1 开源 Mamba 模型#

模型参数量开发者特点
Mamba (130M-3B)130M-3BAlbert Gu原始论文
Mamba-2.8B2.8BAlbert GuSSD 架构
Mistral-Mamba12BMistral AI生产级
Falcon-Mamba7BTII阿联酋团队
Codestral-Mamba7BMistral AI代码专用
Jamba-Mini12BAI21 LabsMamba+Transformer
Zamba1.4B-7BZyphra混合架构
RWKV-Mamba7BBo PengRWKV 系列

10.2 垂直应用#

# Mamba 在不同领域的应用示例
class GenomicMamba(nn.Module):
"""基因组学专用 Mamba"""
def __init__(self, vocab_size=4, d_model=512, n_layers=24):
super().__init__()
self.embed = nn.Embedding(vocab_size, d_model)
self.mamba_layers = nn.ModuleList([
MambaBlock(d_model, d_state=16)
for _ in range(n_layers)
])
self.head = nn.Linear(d_model, 2) # 分类
class SpeechMamba(nn.Module):
"""语音识别 Mamba"""
def __init__(self, n_mels=80, d_model=512):
super().__init__()
self.feature_extractor = nn.Conv1d(n_mels, d_model, kernel_size=3, padding=1)
self.mamba = nn.ModuleList([
MambaBlock(d_model) for _ in range(20)
])
self.ctc_head = nn.Linear(d_model, vocab_size)
class TimeSeriesMamba(nn.Module):
"""时间序列预测 Mamba"""
def __init__(self, input_dim, d_model=256):
super().__init__()
self.proj = nn.Linear(input_dim, d_model)
self.mamba = MambaBlock(d_model, d_state=32)
self.decoder = nn.Linear(d_model, input_dim)

10.3 视觉领域的 Mamba#

架构应用发布
Vision Mamba (Vim)图像分类2024
VMamba分层视觉2024
PlainMamba简洁视觉2024
Mamba-ND多维信号2024

11. Mamba 的局限性与未来#

11.1 已知局限#

局限描述影响
检索能力弱难以”查找”远距离特定信息复制粘贴等任务表现差
上下文外推训练后的上下文长度有限难以扩展到训练外的长度
多模态融合与视觉/音频融合仍需探索通用模型能力受限
解释性内部状态难以解释黑盒程度高

11.2 改进方向#

11.2.1 混合架构:Mamba + Transformer#

class HybridMambaTransformerBlock(nn.Module):
"""
混合架构层:Mamba + Transformer
比例可调(如 3:1)
"""
def __init__(self, d_model, num_heads, ratio_mamba=3):
super().__init__()
# 每 3 个 Mamba 层 + 1 个 Attention 层
self.mamba_blocks = nn.ModuleList([
MambaBlock(d_model) for _ in range(ratio_mamba)
])
self.attn = nn.MultiheadAttention(d_model, num_heads, batch_first=True)
self.norm = RMSNorm(d_model)
def forward(self, x):
for mb in self.mamba_blocks:
x = x + mb(x)
attn_out, _ = self.attn(x, x, x)
x = self.norm(x + attn_out)
return x

11.2.2 层次化 SSM#

class HierarchicalMamba(nn.Module):
"""
层次化 Mamba: 多层 SSM 处理不同时间尺度
"""
def __init__(self, d_model):
super().__init__()
self.low_freq_mamba = MambaBlock(d_model) # 处理短程依赖
self.mid_freq_mamba = MambaBlock(d_model) # 处理中程依赖
self.high_freq_mamba = MambaBlock(d_model) # 处理长程依赖
def forward(self, x):
# 不同时间尺度建模
x_low = self.low_freq_mamba(x)
x_mid = self.mid_freq_mamba(x)
x_high = self.high_freq_mamba(x)
# 融合
return x_low + x_mid + x_high

11.3 未来研究方向#

  1. 更强的选择性:让 SSM 更智能地选择”记住”或”遗忘”
  2. 真正的无限上下文:消除上下文长度限制
  3. 统一多模态架构:文本/图像/音频共享 SSM 骨干
  4. 硬件协同设计:针对 SSM 特性的专用芯片
  5. SSM 与 LLM 结合:SSM 作为大模型的高效替代层

12. 核心公式总结#

  1. SSM 状态方程(连续)
dxdt=Ax+Bu\frac{dx}{dt} = \mathbf{A}x + \mathbf{B}u
  1. SSM 离散化
Aˉ=eAΔ,Bˉ=(A1(eAΔI))B\mathbf{\bar{A}} = e^{\mathbf{A}\Delta}, \quad \mathbf{\bar{B}} = (\mathbf{A}^{-1}(e^{\mathbf{A}\Delta} - I))\mathbf{B}
  1. Mamba 选择性状态方程
xk+1=Aˉkxk+Bˉkukx_{k+1} = \mathbf{\bar{A}}_k x_k + \mathbf{\bar{B}}_k u_k
  1. Mamba 输出
yk=Ckxk+Duky_k = \mathbf{C}_k x_k + \mathbf{D} u_k
  1. SSD 注意力对偶
yk=j=1kCk(l=j+1kAl)Bjujy_k = \sum_{j=1}^{k} \mathbf{C}_k \left(\prod_{l=j+1}^{k} \mathbf{A}_l\right) \mathbf{B}_j u_j
  1. A 矩阵的 HiPPO 初始化
An=2n+1,n=1,2,,N\mathbf{A}_n = -\sqrt{2n+1}, \quad n = 1, 2, \ldots, N

13. 总结#

Mamba 的核心创新#

  1. 选择性扫描:让 SSM 从线性时不变 (LTI) 变为线性时变 (LTV),具备输入依赖性
  2. 硬件感知算法:通过分块和重计算实现 GPU 内存层次的最大化利用
  3. SSD 对偶性:揭示 SSM 与注意力的深层联系,统一两种架构
  4. 线性时间复杂度O(L)O(L) 的序列处理,颠覆 Transformer 的 O(L2)O(L^2)

Mamba 的定位#

Mamba 不是 Transformer 的替代品,而是序列建模领域的”第二选择”。

场景推荐架构
短序列(< 2K tokens)+ 高精度需求Transformer
长序列(> 8K tokens)+ 效率需求Mamba
需要精确检索的任务Transformer
通用长上下文理解Mamba + Transformer 混合
流式实时推理Mamba

学习建议#

  1. 理论学习:从 HiPPO → S4 → Mamba 的脉络循序渐进
  2. 代码实践:从最简单的 SSM 实现开始,逐步添加选择性扫描
  3. 对比实验:在相同任务上对比 Transformer 和 Mamba 的性能
  4. 关注前沿:Mamba 系列仍在快速演进,关注最新论文

Mamba 代表了深度学习序列建模的一个重要突破,它向我们展示了状态空间模型作为 Transformer 替代方案的巨大潜力。随着研究的深入和硬件的优化,Mamba 有望在更多场景中发挥重要作用,成为下一代 AI 基础设施的核心组件之一。

参考资料#

  1. Gu, A., & Dao, T. (2023). “Mamba: Linear-Time Sequence Modeling with Selective State Spaces.” arXiv.
  2. Gu, A., Goel, K., & Re, C. (2022). “Efficiently Modeling Long Sequences with Structured State Spaces.” ICLR.
  3. Dao, T., & Gu, A. (2024). “Mamba-2: State Space Models are a New Kind of Attention.” arXiv.
  4. Lieber, O., et al. (2024). “Jamba: A Hybrid Transformer-Mamba Language Model.” arXiv.
  5. Poli, M., et al. (2023). “Hyena Hierarchy: Towards Larger Convolutional Language Models.” ICML.
  6. Mehta, H., et al. (2023). “Simple Hardware-Efficient Long Convolutions for Sequence Modeling.” ICML.
  7. De, S., et al. (2024). “Griffin: Mixing Gated Linear Recurrences with Local Attention for Efficient Language Models.” arXiv.
  8. Zhu, L., et al. (2024). “Vision Mamba: Efficient Visual Representation Learning with Bidirectional State Space Model.” ICML.
  9. Behrouz, A., et al. (2024). “Titans: Learning to Memorize at Test Time.” arXiv.
  10. Park, J., et al. (2024). “Falcon Mamba: The First Competitive Attention-free 7B Language Model.” arXiv.

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

深入理解 Mamba:选择性状态空间模型的架构与实现
https://aiattnstudio.link/posts/mamba/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签