2023 年,Albert Gu 和 Tri Dao 发表的论文《Mamba: Linear-Time Sequence Modeling with Selective State Spaces》在深度学习架构领域投下了一颗重磅炸弹。当整个行业还在为 Transformer 的二次复杂度焦头烂额时,Mamba 用 O(n) 的推理复杂度和接近 Transformer 的性能,证明了"替代方案"并非空中楼阁。
一、Transformer 的阿喀琉斯之踵
Transformer 的核心是 Self-Attention 机制。给定长度为 n 的序列,Attention 需要计算所有 token 对之间的相似度,时间复杂度为 O(n²)。这意味着:
- 序列长度翻倍,计算量翻四倍
- 推理时 KV Cache 的内存随序列线性增长
- 长文档处理成本急剧膨胀
虽然 FlashAttention、Ring Attention、MQA/GQA 等技术在工程层面做了大量优化,但 O(n²) 的复杂度本质没有改变。当我们尝试处理 100K、1M 甚至更长序列时,纯粹的工程技巧已触及天花板。
问题的根源不在工程,而在数学。
二、状态空间模型的数学直觉
状态空间模型(State Space Model, SSM)是控制论中的经典工具。其核心思想可以用一个简洁的微分方程描述:
h'(t) = A·h(t) + B·x(t)
y(t) = C·h(t)
其中 x(t) 是输入信号,h(t) 是隐藏状态,y(t) 是输出。离散化后得到递推形式:
h_k = Ā · h_{k-1} + B̄ · x_k
y_k = C · h_k
这个递推结构天然具备两个关键优势:
- 线性时间复杂度:每个时间步只需一次矩阵-向量乘法,总体 O(n)
- 有界内存:隐状态 h 的维度固定,不随序列长度增长
- Transformer 的完整 KV Cache 可以看作是无损存储——完整保留了所有历史信息
- Mamba 的隐状态是有损压缩——只保留对后续预测"有用"的信息
- 实时对话系统:生成延迟稳定,无随对话长度增长而变慢的问题
- 流式处理:可以无限处理输入,内存占用恒定
- 批处理推理:吞吐量不会因为最长序列的存在而受到严重影响
- SSM 适合处理"连续流"——正文、对话等自然序列
- Attention 适合处理"结构跳变"——段落开头、列表项等边界
- 比例比类型更重要——适量 Attention + 大量 SSM 是当前最佳实践
- 纯 Transformer:KV Cache ~ 128GB,单步延迟 ~ 45ms
- 纯 Mamba:状态 ~ 16GB(固定),单步延迟 ~ 8ms
- 混合架构(20:1):状态 ~ 20GB + KV Cache ~ 6.4GB,单步延迟 ~ 12ms
- Gu, A., & Dao, T. (2023). Mamba: Linear-Time Sequence Modeling with Selective State Spaces. arXiv:2312.00752
- Gu, A., et al. (2022). Efficiently Modeling Long Sequences with Structured State Spaces (S4). ICLR 2022
- Dao, T., & Gu, A. (2024). Mamba2: Subquadratic-State Space Models Are As Efficient As Transformers. arXiv:2405.21060
- Lieber, O., et al. (2024). Jamba: A Hybrid Transformer-Mamba Language Model. arXiv:2408.15668
但传统 SSM(如 S4)有一个致命缺陷:模型参数 A、B、C 是固定的,对所有输入一视同仁。这意味着它无法像 Attention 那样选择性关注某些 token。
三、Mamba 的核心突破:选择性状态空间
Mamba 的第一个关键创新是选择性(Selectivity)。传统 S4 的 Δ、B、C 是常数,Mamba 让它们依赖输入:
Δ_k = Δ(x_k), B_k = B(x_k), C_k = C(x_k)
这个看似简单的改动带来了质的飞跃:
# S4: 固定参数,无差别处理
h_k = A_bar * h_{k-1} + B_bar * x_k # 对所有 token 一样
# Mamba: 输入依赖,选择性遗忘与记忆
delta = softplus(W_delta @ x_k) # 依赖输入的离散化步长
B = W_B @ x_k # 依赖输入的投影
C = W_C @ x_k # 依赖输入的输出投影
h_k = (1 - delta) * h_{k-1} + delta * (A @ h_{k-1} + B * x_k)
当某个 token 不重要时,模型可以让 Δ→0,此时 h_k ≈ h_{k-1},相当于跳过了该 token;当 token 重要时,Δ 变大,模型更新状态。这种机制在功能上等价于 Attention 的"选择性关注",但复杂度是线性的。
3.1 硬件感知的并行算法
仅靠选择性不够。递推结构虽然理论上线性,但不像矩阵乘法那样天然适合 GPU 并行。Mamba 作者设计了一个硬件感知(Hardware-Aware)的并行算法:
# 将递推转化为并行前缀和(parallel scan)
# 利用 GPU 的 Tensor Core 加速
def selective_scan(u, delta, A, B, C, D):
"""
u: 输入序列 [B, L, d]
delta: 离散化步长 [B, L, d]
A: 状态矩阵 [d, N]
B, C: 投影 [B, L, N]
D: 跳跃连接 [d]
"""
# 预计算每个时间步的 effective A
A_bar = torch.exp(delta.unsqueeze(-1) * A) # [B, L, d, N]
B_bar = delta.unsqueeze(-1) * B.unsqueeze(2) # [B, L, d, N]
# 并行前缀和: 将递推转化为并行归约
# 使用 Blelloch 扫描算法,O(log n) 的并行步骤
h = parallel_scan(A_bar, B_bar * u.unsqueeze(-1))
# 输出投影
y = (h @ C.unsqueeze(-1)).squeeze(-1) + D * u
return y
这个算法的精妙之处在于:它牺牲了一点数学上的等价性(用并行前缀和替代严格递推),换取了 GPU 上的高效并行。现代 GPU 的 Tensor Core 可以极大加速这个计算流程。
四、架构对比:为什么 Mamba 能行
4.1 与 Transformer 的核心差异
| 维度 | Transformer | Mamba |
|---|---|---|
4.2 理论分析:选择性信息的压缩视角
从信息论角度看,Mamba 做的事情是:将整个序列选择性地压缩进一个固定大小的隐状态中。这类似于无损压缩 vs 有损压缩的权衡:
关键问题是:这种有损压缩的质量如何?Mamba 论文的实验表明,在 选择性任务(selective tasks)——如选择性复制(Selective Copying)、诱导头(Induction Heads)等——上,Mamba 能完美完成,因为它学会了丢弃无关信息、保留关键信息。但在需要精确 recall 的任务上,Mamba 仍有差距。
五、代码实战:从零实现选择性扫描
让我们从底层实现一个选择性扫描(Selective Scan)模块:
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
class SelectiveScan(nn.Module):
"""简化的选择性扫描实现"""
def __init__(self, d_model, d_state=16):
super().__init__()
self.d_model = d_model
self.d_state = d_state
# A 矩阵:初始化为 S4 论文建议的特征值分布
# 这里简化为可学习的下三角矩阵
self.A_log = nn.Parameter(torch.log(
torch.arange(1, d_state + 1, dtype=torch.float32).unsqueeze(0).repeat(d_model, 1)
))
# 投影层
self.x_proj = nn.Linear(d_model, d_state * 3, bias=False) # delta, B, C
self.dt_proj = nn.Linear(d_state, d_model, bias=True)
self.out_proj = nn.Linear(d_model, d_model, bias=False)
# 初始化 dt_proj 的 bias,使 delta 初始值较小
nn.init.uniform_(self.dt_proj.bias, -0.5, 0.5)
self.D = nn.Parameter(torch.randn(d_model))
def forward(self, u):
"""
u: [batch, seq_len, d_model]
returns: [batch, seq_len, d_model]
"""
batch, seq_len, d_model = u.shape
# 计算 A(保证稳定性)
A = -torch.exp(self.A_log) # [d_model, d_state]
# 投影得到 delta, B, C
x_proj_out = self.x_proj(u) # [batch, seq_len, d_state * 3]
delta, B, C = x_proj_out.chunk(3, dim=-1)
# delta 通过 softplus 确保为正
delta = F.softplus(self.dt_proj(delta)) # [batch, seq_len, d_model]
# 执行选择性扫描(简化版递推实现)
# 实际生产中应使用 CUDA 优化的 parallel scan
h = torch.zeros(batch, d_model, self.d_state, device=u.device)
ys = []
for t in range(seq_len):
# 更新规则:h_t = exp(-delta * A) * h_{t-1} + delta * B * x
dA = torch.exp(-delta[:, t].unsqueeze(-1) * A.unsqueeze(0)) # [batch, d_model, d_state]
dB = delta[:, t].unsqueeze(-1) * B[:, t].unsqueeze(1) # [batch, d_model, d_state]
h = dA * h + dB * u[:, t].unsqueeze(-1)
# 输出:y = C @ h
y_t = (h * C[:, t].unsqueeze(1)).sum(dim=-1) # [batch, d_model]
ys.append(y_t)
y = torch.stack(ys, dim=1) # [batch, seq_len, d_model]
# 跳跃连接
y = y + self.D.unsqueeze(0).unsqueeze(0) * u
return self.out_proj(y)
class MambaBlock(nn.Module):
"""完整的 Mamba 块:SSM + MLP"""
def __init__(self, d_model, d_state=16, d_conv=4, expand=2):
super().__init__()
self.d_inner = d_model * expand
# 输入投影
self.in_proj = nn.Linear(d_model, self.d_inner * 2, bias=False)
# 卷积层(提取局部特征)
self.conv1d = nn.Conv1d(
self.d_inner, self.d_inner,
kernel_size=d_conv,
padding=d_conv-1,
groups=self.d_inner
)
# 选择性扫描
self.ssm = SelectiveScan(self.d_inner, d_state)
# 门控
self.act = nn.SiLU()
self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)
def forward(self, x):
"""
x: [batch, seq_len, d_model]
"""
# 残差分支
x_res = x
xz = self.in_proj(x)
x, z = xz.chunk(2, dim=-1)
# 卷积提取局部特征
x = rearrange(x, 'b l d -> b d l')
x = self.conv1d(x)[..., :x.shape[-1]] # 截断 padding
x = rearrange(x, 'b d l -> b l d')
x = self.act(x)
# 选择性扫描
x = self.ssm(x)
# 门控
z = self.act(z)
x = x * z
return self.out_proj(x) + x_res
class SimpleMamba(nn.Module):
"""一个极简但完整的 Mamba 语言模型"""
def __init__(self, vocab_size, d_model=256, n_layers=6, d_state=16, max_len=2048):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_emb = nn.Embedding(max_len, d_model)
self.layers = nn.ModuleList([
MambaBlock(d_model, d_state) for _ in range(n_layers)
])
self.norm = nn.RMSNorm(d_model)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
def forward(self, input_ids):
B, L = input_ids.shape
pos = torch.arange(L, device=input_ids.device).unsqueeze(0).expand(B, -1)
x = self.embedding(input_ids) + self.pos_emb(pos)
for layer in self.layers:
x = layer(x)
x = self.norm(x)
logits = self.lm_head(x)
return logits
六、工程实践中的关键挑战
6.1 训练不稳定性
Mamba 训练的稳定性比 Transformer 更具挑战性。选择性步长 Δ 如果训练不当,会导致状态爆炸或消失。实践中需要注意:
# 训练技巧 1:梯度裁剪必须更激进
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.1)
# 训练技巧 2:Δ 的初始化非常关键
# 初始时应让 delta 较小(接近 0.1),避免训练初期状态剧烈变化
# 论文建议 dt_proj 的 bias 初始化在 [-0.5, 0.5]
# 训练技巧 3:A 矩阵的特征值分布
# 建议使用 HiPPO 初始化或类似的连续时间先验
# 确保特征值的负实部足够大,保证系统稳定性
6.2 长序列中的遗忘问题
尽管理论上 SSM 可以建模无限长序列的依赖,但在实践中,当序列超过一定长度后,早期信息会被逐渐稀释。这是因为选择性步长 Δ 的累积效应——如果后面不断有新信息输入,早期状态的对数概率贡献会指数衰减。
实验表明,Mamba 在 DNA 序列建模(需要 64K+ 上下文)上表现出色,但在需要精确回忆长距离 token 的任务上,仍不如 Transformer。
6.3 推理引擎的优化
Mamba 的推理有一个独特优势:每一步的计算量是固定的(不像 Transformer 需要访问整个 KV Cache)。这使得它在以下场景有显著优势:
# Mamba 推理伪代码:每步固定计算量
def mamba_generate_step(model, token, state):
"""每次调用只执行固定计算,不访问历史"""
x = model.embedding(token) # [B, 1, d]
for layer in model.layers:
# 状态更新:h_new = f(h_old, x)
state[layer.idx], x = layer.sigmoid_ssm(x, state[layer.idx])
logits = model.lm_head(x)
return logits, state # 返回新状态,下一步使用
七、混合架构:SSM 与 Attention 的联姻
2024 年下半年以来,越来越多的模型开始探索混合架构——将 SSM 的高效和 Attention 的精确结合:
7.1 Jamba 架构(AI21 Labs)
Jamba 采用 8:1 的 Mamba-to-Attention 层比例,在保持接近纯 Transformer 性能的同时,将内存占用降低了 4-8 倍:
class JambaLayer(nn.Module):
def __init__(self, config, use_attention=False):
super().__init__()
if use_attention:
self.layer = AttentionBlock(config)
else:
self.layer = MambaBlock(config)
def forward(self, x):
return self.layer(x)
# 配置:每 8 层 Mamba 加 1 层 Attention
jamba_config = {
'total_layers': 52,
'attention_every_k': 8, # 第 8, 16, 24... 层使用 Attention
'experts_per_layer': 4 # MoE 化的中间层
}
7.2 Zamba(Zyphra)
Zamba 更进一步,仅用 1 层 Attention 就恢复了完整的 Transformer 性能。其核心研究发现:Attention 的价值不在于数量,而在于位置——在信息流的"瓶颈"处(如跨段落边界)放置 Attention 层,效果远胜于均匀分布。
7.3 工程启示
这对架构设计的启示是:
八、生产部署考量
8.1 推理框架支持
截至 2026 年,Mamba 的推理生态已日趋成熟:
| 框架 | SSM 支持 | 特性 |
|---|---|---|
8.2 内存与吞吐对比
以 Llama-3 70B 大小的模型为例,处理 32K 上下文时:
在长上下文场景下,Mamba 的优势不仅是成本,更是可行性——有些场景对 Transformer 而言根本不可行(如整本书籍的实时分析),Mamba 让这些场景成为可能。
九、未来方向
SSM 领域仍在快速演进。以下几个方向值得关注:
9.1 Mamba-2:对角化与tensor并行
Mamba-2 将 SSM 参数化为对角矩阵(Diagonal SSM),使得矩阵幂运算从 O(N²) 降至 O(N)。更关键的是,对角化让张量并行变得可以直接实现:
# Mamba-2 的对角 SSM:可以按状态维度切分到不同 GPU
# 每个 GPU 负责 d_state / world_size 个状态维度
# 通信开销远小于 Attention 的 KV Cache 同步
class DiagonalSSM(nn.Module):
def __init__(self, d_model, d_state):
# 只需要存储对角元素,参数减少 N 倍
self.Lambda = nn.Parameter(torch.randn(d_model, d_state)) # 对角矩阵
def forward(self, u):
# Lambda^k 只需逐元素幂运算,无需矩阵乘法
Lambda_power = self.Lambda.unsqueeze(0).unsqueeze(0) ** torch.arange(seq_len)
# ... 向量化实现
9.2 结构化稀疏 SSM
将 SSM 的 A 矩阵从稠密改为结构化稀疏(块对角、带状),可以在保持长距离建模能力的同时,进一步降低计算开销。
9.3 硬件亲和设计
Mamba 的固定计算图特性天然适合专用 AI 加速器。与 Attention 需要不同形状的矩阵乘加不同,SMA 每一步的计算模式完全相同,这为 FPGA/ASIC 实现提供了天然优势。
结语
Mamba 和状态空间模型代表了一个重要趋势:当我们从第一性原理审视序列建模问题时,会发现 Transformer 并非唯一的答案,甚至不是最优答案。
选择性状态空间的核心贡献不在于它"替代"了 Transformer,而在于它证明了另一条路径的存在——一条路径可以用线性复杂度处理序列,用固定内存承载长距离依赖,同时保持与 Transformer 相媲美的性能。
对于工程师而言,2026 年的技术版图已经清晰:Transformer 不会消失,但它不再是唯一的选择。理解 SSM、理解选择性扫描、理解混合架构的设计权衡,已经成为现代 AI 基础设施工程师的核心素养之一。
参考论文:

发表评论 取消回复