Mamba 选择性状态空间模型:亚二次复杂度大语言模型推理工程实践
从 S4 到 Mamba 的架构演进,选择性状态空间如何实现 Transformer 级别的建模能力与 RNN 级别的推理效率?本文深入剖析 SSM 的数学本质、选择性机制的硬件感知实现,以及 Mamba 在百亿参数规模的工程部署实践。
一、Transformer 推理的根本瓶颈
自 2017 年 Attention is All You Need 发表以来,Transformer 凭借其卓越的并行训练能力和全局依赖建模,成为大语言模型的事实标准。然而,Transformer 的自注意力机制存在两个核心效率瓶颈:
训练阶段:注意力矩阵的计算复杂度为 O(N²),其中 N 为序列长度。虽然 Flash Attention 在 GPU HBM 与 SRAM 之间做了精巧的 IO 感知分块,但计算量本身并未减少。
推理阶段:自回归生成时,模型需要维护一个 Key-Value Cache(KV Cache),其内存占用与序列长度呈线性增长。对于 70B 模型、8K 上下文长度,KV Cache 可达数 GB,严重制约了推理吞吐和批处理能力。
业界已经提出了诸多方案试图缓解这一问题:线性注意力(Linear Attention)牺牲了精确记忆能力;稀疏注意力(如 BigBird、Longformer)仅适用于特定任务;RWKV 尝试用线性时间递推的 RNN 结构替代注意力。而 2023 年 12 月提出的 Mamba(Gu & Dao, 2023)则走出了一条全新路线——选择性状态空间模型(Selective State Space Model, SSM),在保持 Transformer 级建模能力的同时,实现了推理阶段的线性复杂度。
二、状态空间模型:从控制论到深度学习
2.1 SSM 的连续时间形式
状态空间模型起源于控制论,描述一个连续时间动态系统:
h'(t) = A · h(t) + B · x(t)
y(t) = C · h(t) + D · x(t)
其中 x(t) 是输入信号,h(t) 是隐藏状态(latent state),y(t) 是输出,矩阵 A 控制状态演化,B 控制输入如何注入状态,C 控制状态如何映射到输出,D 是直连跳跃项(Connection)。
这个 ODE 的含义很直观:h(t) 持续记忆历史输入信息,其演化方式由 A 的特征值决定——适当的 A 设计能使状态在长期依赖中既不死机(梯度消失)也不爆炸(梯度发散)。
2.2 离散化:从连续到神经网络层
在实际深度学习中,我们处理的是离散序列 x₀, x₁, ..., xₜ。通过零阶保持(Zero-Order Hold, ZOH)离散化,步长为 Δ(采样间隔),SSM 变为:
Ā = exp(ΔA)
B̄ = (ΔA)⁻¹ · (exp(ΔA) - I) · B
h_t = Ā · h_{t-1} + B̄ · x_t
y_t = C · h_t + D · x_t
这里有一个关键观察:离散化后的 SSM 在形式上就是一个 RNN。输入 x_t 被累加到隐藏状态 h_t,然后 h_t 作为下一时间步的记忆。与标准 RNN 不同的是,SSM 的 A 矩阵来源于连续时间系统的解析离散化,而非直接学习,这使得状态演化具有更好的数学稳定性和表达能力。
2.3 并行计算的卷积视角
将上述递推展开,可以得到 y_t 关于所有历史输入的显式表达:
y_t = Σ(k=0..t) C · Āᵏ · B̄ · x_{t-k} + D · x_t
这正是一个卷积运算,滤波器核为:
K = (C·B̄, C·Ā·B̄, C·Ā²·B̄, ..., C·Ā^{N-1}·B̄)
因此,SSM 的并行训练可以通过高效卷积实现,使用 FFT 将复杂度从 O(N²) 降至 O(N log N)。这就是 S4(Structured State Space Sequence Model)高效训练的数学基础。
三、Mamba 的核心创新:选择性机制
3.1 S4 的局限性
S4(Gu et al., 2021)通过以下设计取得了突破:
- A 矩阵采用 HiPPO 初始化的高阶正交结构,最大化长期记忆保持能力
- NPLR(Normal Plus Low-Rank)参数化确保矩阵幂运算的数值稳定性
- 对角化结构高效计算矩阵幂
但 S4 有一个致命缺点:模型参数 A, B, C 在所有时间步共享,与输入无关。这意味着 SSM 无法对不同的输入 token 采取不同的策略——它无法判断"这个 token 很重要,我需要更努力地记住它"或"这个 token 无关紧要,我可以快速遗忘它"。
类比注意力机制:Transformer 通过 softmax 计算权重,使得重要 token 获得更多关注。而 S4/SSM 对所有 token 一视同仁地应用固定变换,无法动态地选择记忆内容。
3.2 选择性状态空间
Mamba 的核心创新是让 B, C, Δ 参数变成输入依赖的动态参数:
B_t = Linear_B(x_t)
C_t = Linear_C(x_t)
Δ_t = softplus(Linear_Δ(x_t) + broadcast_Δ)
A = 全局学习参数(不依赖于输入)
Δ 的作用尤其关键——它控制状态更新的"速度":
- 当 Δ 很大时,Ā = exp(-exp(Δ) · A) 快速衰减,历史记忆被迅速遗忘,模型只关注最近输入
- 当 Δ 很小时,Ā 接近恒等矩阵,状态缓慢衰减,历史记忆被长期保留
这种选择性赋予 Mamba 两个能力:
- 有选择地记住或忽略输入:对于"虽然"、"但是"等转折词后的关键信息,增大 Δ 让模型更快更新状态;对于标点符号、填充词等无关信息,减小 Δ 让状态不发生变化
- 信息滤波:在推理时,无关信息被直接从记忆中过滤掉,不会污染上下文
- 序列的每个位置对应一个时间步的递推
- 每个 token 对应一个隐藏状态 h_t,维度较小(典型值 d_state=128,远小于 KV Cache 的维度)
- 在 kernel 融合中,系统矩阵 (A, B_t, C_t, Δ_t) 和隐藏状态 h_t 全部保留在 GPU SRAM
- 只在序列处理完成后,才将最终状态 h_N 写回 HBM(用于下一层的输入或下一次推理的初始状态)
- 显著优于同参数量的 Transformer 模型(如 Pythia-2.8B、OPT-2.7B)
- 在长上下文任务中(PG19、Books3),SSM 的困惑度随上下文长度增加而持续下降,而 Transformer 的困惑度在较长上下文后不再改善
- 大幅超越同规模的 RWKV-3B(约 1.5-2.0 ppl 差距)
- 引入了类似 Attention 的双路径结构,增强信息交互
- 使用更大的状态维度(128→256),提高记忆容量
- 改进了并行扫描的原生并行度
- 前缀状态缓存(Prefix State Caching):在处理相同前缀的多个请求时直接复用已计算好的层状态
- 连续批处理优化:对于新到达的 token,只需增量递推一步,无需从头计算
- 按需前向计算(Continuous Batching):仅在新生 token 上做前向,但同一批内不同请求需独立执行扫描
- 定制的 CUDA Graph:针对不同序列长度预编译多个 Graph,按实际长度选择执行
- Warp 级并行化:将批内请求的状态排列在连续线程中,利用 GPU SIMT 并行执行
- Mamba 解决序列长度维度的效率瓶颈(注意力复杂度)
- MoE 解决参数总量维度的效率瓶颈(各 token 仅激活部分专家)
- 混合架构:Jamba(AI21 Labs)、Zamba 等模型在部分层使用全局注意力,其他层使用 Mamba
- 外部记忆:将精确信息存入外部数据库,由 SSM 控制检索时机
- 状态维度扩展:增大 d_state(代价是状态不再适合完整驻留 SRAM)
- Attention 架构借鉴 RNN:RetNet、RWKV-6 引入"循环等价"推理模式
- SSM 架构借鉴 Attention:Mamba 2 的双路径设计本质上类似"内容寻址的注意力"
- 混合架构涌现:Zamba、Jamba、Griffin 等模型混合注意力与循环结构
- 硬件层面趋同:各架构都在向"状态驻留 SRAM、最小化 HBM 读写"方向优化
- 模型转换:现有 Attention 模型无法直接转为 Mamba,需从头训练或蒸馏
- 推理框架:vLLM 的 Mamba 支持仍在早期阶段,自研内核需一定工程投入
- 监控指标:需要新增"状态维度利用率"、"SRAM 与 HBM 读写比例"等性能指标
3.3 选择性机制的直觉解释
以一个具体例子说明选择性如何工作:
输入序列:"赵明是公司的技术总监,但他昨天提交了辞职报告。"
在处理到"但"时,模型需要抛弃先前的部分记忆来重点关注转折后的内容。Mamba 通过增大 Δt 快速衰减旧状态来实现这一点。而 S4 的固定参数无法感知"这里发生了转折",会无差别地将先前的状态混入后续计算。
另一个场景是精确记忆:在处理需要长距离回忆的任务时,模型需要精确记住某个实体名称。Mamba 可以通过极小的 Δ 让这个信息几乎不受干扰地保存在状态中——类似 Attention 中的硬 select,但实现方式完全不同。
四、硬件感知的并行算法
4.1 并行扫描(Parallel Scan)
选择性机制引入输入依赖的 B_t, C_t, Δ_t 后,并行卷积方案不再适用。Mamba 需要回归递推形式:
h_t = Ā_t · h_{t-1} + B̄_t · x_t
但这个递推可以通过并行前缀和(Parallel Scan / Blelloch 算法)高效计算。对于结合算子 ⊗ 的定义:
(a₁, b₁) ⊗ (a₂, b₂) = (a₂ · a₁, a₂ · b₁ + b₂)
递推的每一步可以表示为 (h_t, x_t) → (Ā_t · h_{t-1} + B̄_t · x_t, x_t),而并行扫描用分治法在 O(log N) 步内完成整个序列的计算。在 GPU 上这被优化为高效的并行归约操作。
4.2 内存层级优化
Mamba 最重要的工程贡献:状态驻留在 SRAM,永不写入 HBM。
在 Transformer 推理中,KV Cache 必须常驻 HBM(几十 GB 显存),每次前向计算都需要从 HBM 读入 KV。对于长序列,HBM 带宽成为吞吐瓶颈。
Mamba 的策略:
这种设计与 Flash Attention 的思路一脉相承——都是将中间计算结果保留在 SRAM 中,最小化 HBM 读写。但 Mamba 的优势更加彻底:跨 token 的隐藏状态更新完全不需要任何 HBM 访问,而 Flash Attention 仍需读写分块的注意力矩阵。
4.3 内核融合实践
Mamba 通过 Triton GPU 编程框架实现了高度 fuse 的内核:
# Mamba 的简化伪代码(选择性扫描内核)
import triton
import triton.language as tl
@triton.jit
def selective_scan_kernel(
A_ptr, B_ptr, C_ptr, delta_ptr,
hidden_ptr, x_ptr, y_ptr,
seq_len, state_dim, block_size: tl.constexpr
):
pid = tl.program_id(0)
# 从 HBM 加载当前时间步的输入和动态参数
x = tl.load(x_ptr + pid * state_dim + tl.arange(0, state_dim))
b = tl.load(B_ptr + pid * state_dim + tl.arange(0, state_dim))
c = tl.load(C_ptr + pid * state_dim + tl.arange(0, state_dim))
delta = tl.load(delta_ptr + pid * state_dim + tl.arange(0, state_dim))
# 离散化:计算当前步的 Ā 和 B̄
a_base = tl.load(A_ptr + tl.arange(0, state_dim)) # 全局 A 矩阵
dA = tl.exp(delta * a_base)
dB = delta * b
# 递推更新(保持在 SRAM)
h = tl.load(hidden_ptr + tl.arange(0, state_dim))
h = dA * h + dB * x
y = c * h
tl.store(y_ptr + pid * state_dim + tl.arange(0, state_dim), y)
Mamba v2(2024 年 4 月发布)进一步优化,将状态维度增大、并行扫描内核的并行度提升,在 NVIDIA A100/H100 上实现了比 Flash Attention 2 更优的吞吐。
五、Mamba 在大规模语言模型中的表现
5.1 语言建模基准
Mamba-2.8B(2.8B 参数)在 Pile 语言建模基准上,其困惑度(perplexity):
Mamba-1.4B 在零样本推理任务(HellaSwag、PIQA、Winogrande 等)中普遍超越 Pythia-1.4B 和 RWKV-1.5B,在部分任务上(如 ARC-Challenge、Lambada)甚至接近 Transformer-6.9B。
5.2 推理效率
Mamba 在推理速度上与 Transformer 的对比(基于官方数据,NVIDIA A100 80GB):
| 指标 | Transformer-1.3B | Mamba-1.4B |
|---|---|---|
| 训练吞吐量(tokens/s) | ~360K | ~380K |
| 推理延迟(8K上下文) | ~45ms | ~12ms |
| 推理显存(8K上下文) | ~16GB | ~6GB |
| 生成吞吐(tokens/s) | ~22 | ~83 |
Mamba 的推理延迟与序列长度无关(固定状态维度),而 Transformer 每个生成都需要全量 Attention 计算。这使得 Mamba 在长上下文场景中优势极为突出——8K 上下文时延迟差 3-4 倍,64K 上下文时差距超过 10 倍。
5.3 Mamba 2 的架构演进
原始的 Mamba 仍采用 Transformer 风格的 MLP+Mamba 混合架构。Mamba 2(2024)进一步修改了 SSM 结构:
Mamba 2-2.7B 在几乎所有语言建模基准上超越了 Llama-3-8B(8B 参数),参数量比仅为 1:3,这是 Mamba 系列最惊人的数据之一。
六、工程实践:部署与优化
6.1 快速体验
目前最成熟的 Mamba 推理实现是 mamba-ssm 库(基于 Triton)和 transformers 库。
import torch
from transformers import Mamba2ForCausalLM, AutoTokenizer
# 加载模型与分词器
model = Mamba2ForCausalLM.from_pretrained(
"state-spaces/mamba2-2.7b",
dtype=torch.bfloat16,
device="cuda"
)
tokenizer = AutoTokenizer.from_pretrained("state-spaces/mamba2-2.7b")
# 自回归生成
def generate(prompt, max_new_tokens=256, temperature=0.7):
input_ids = tokenizer.encode(prompt, return_tensors="pt").cuda()
for _ in range(max_new_tokens):
with torch.no_grad():
output = model(input_ids, cache_params=None, stateful=True)
logits = output.logits[:, -1, :]
if temperature == 0:
next_token = logits.argmax(dim=-1, keepdim=True)
else:
probs = torch.softmax(logits / temperature, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
input_ids = torch.cat([input_ids, next_token], dim=1)
if next_token.item() == tokenizer.eos_token_id:
break
return tokenizer.decode(input_ids[0], skip_special_tokens=True)
# 调用
response = generate("请用简单的语言解释状态空间模型的工作原理。")
print(response)
6.2 状态缓存与推理服务
在实际部署场景中,Mamba 的独特状态管理给服务框架带来了新的挑战和机会。
状态压缩与池化(State Pooling)
Mamba 的状态是一个维度固定的向量(如 d_state=256),不需要按 token 分页。可以采用:
# 概念性的状态缓存复用示例
class MambaStateCache:
"""缓存不同前缀对应的 Mamba 层状态"""
def __init__(self, model_config, max_cache_size=1024):
self.cache = {} # prefix_hash -> [layer_states]
self.max_size = max_cache_size
def get(self, input_ids: torch.Tensor):
"""获取已缓存的状态"""
key = hash(input_ids.tolist())
if key in self.cache:
return self.cache[key]
return None
def put(self, input_ids: torch.Tensor, states):
if len(self.cache) >= self.max_size:
# LRU 淘汰
self.cache.pop(next(iter(self.cache)))
key = hash(input_ids.tolist())
self.cache[key] = states
6.3 批处理的加速挑战
Mamba 推理的一大挑战是动态形状批处理。不同请求的长度和已生成位置不同,导致递推步数不一致。
目前工程上的解决方案:
6.4 与 MoE 架构的互补
Mamba 与 Mixture-of-Experts(MoE)形成互补:
MambaMoE(2024 年提出)将 Mamba 层替换 Attention,保留 MoE 前馈网络。在部分基准上,Mamba-MoE-3B(激活 1.3B)表现接近 Transformer-MoE-8B(激活 4B),展示了极具效率潜力的新架构方向。
七、Mamba 的局限与当前挑战
7.1 精确召回能力
Mamba 通过固定维度的连续状态向量记忆信息,其记忆容量存在上限。对于需要精确区分某个信息是否"在上下文中出现过"的任务(如 key-value 精确检索),Mamba 弱于 Transformer。
解决方案方向:
7.2 超长序列续推
Mamba 在推理开始时通常从全零状态初始化。扩展到 100K+ 上下文时,需要状态续推(State Continuation)机制——在分块处理超长序列时保持跨块状态的连续性,目前工程实践中仍缺乏标准化方案。
7.3 硬件生态
Mamba 依赖定制的 Triton 内核,目前 ONLY 对 NVIDIA GPU(Hopper/Ada/Ampere 架构)有最优支持。AMD GPU(ROCm)和 Apple Silicon(MPS)的支持仍在完善。在 Intel Gaudi、AWS Trainium/Inferentia 等平台上,Mamba 的性能优势不明显。
八、未来方向与选型建议
8.1 SSM 与 Transformer 的趋同
2024 年下半年,SSM 与 Transformer 呈现出有趣的趋同趋势:
未来一年,混合架构可能取代纯 Transformer 成为长上下文部署的新标准。
8.2 选型框架
| 场景 | 推荐架构 | 原因 |
|---|---|---|
| 超长上下文推理(32K+) | Mamba / 混合架构 | 延迟与上下文长度无关,显存恒定 |
| 需要精确信息检索 | Transformer / 混合架构 | 状态压缩无法保证精确召回 |
| 边缘设备推理(低显存) | Mamba | 无 KV Cache 显存压力 |
| 训练新基础模型 | Transformer | 生态成熟、训练效率稳定 |
| 长文档摘要/代码生成 | Mamba | 依赖全局模式而非精确 token 召回 |
| 实时流式生成(聊天) | Mamba | 恒定低延迟 |
8.3 基础设施准备
对于已有 Transformer 部署的团队,引入 Mamba 需要考量:
结语
Mamba 不是 Transformer 的完美替代者,而是一个全新的模型架构范式——选择性状态空间。它在序列建模中引入了"内容感知"的稀疏记忆能力,同时在计算层面实现了硬件感知的并行扫描和状态压缩。
从 S4(结构化状态空间)到 Mamba(选择性状态空间)再到混合架构(如 Jamba、Zamba),模型设计从"结构归一化"向"能力组合化"演进的趋势日益明显。在推理成本和长上下文需求成为核心瓶颈的 2025 年,SSM 家族在工程部署中的价值将愈发凸显。
对于工程师而言,理解 Mamba 的选择性原理和硬件感知实现,不仅有助于在合适的场景中做出更好的架构选择,也提供了看待"记忆与计算"这一基础命题的新视角——连续时间动力学系统如何高效地离散化为神经网络的计算图,可能是接下来两年深度学习基础设施进化的核心方向之一。

发表评论 取消回复