Mamba 状态空间模型原理与 Transformer 对比

本文解读 Mamba 选择性状态空间模型,针对 Transformer 自注意力二次复杂度与固定上下文窗口的瓶颈,提出让 SSM 参数随输入动态变化的选择机制,并设计硬件感知并行 scan 算法,在保持线性序列复杂度的同时达到 Transformer 级别的建模能力。

A
AGISeed Team
AGISeed 作者

Mamba 状态空间模型原理与 Transformer 对比

Figure 3:(Architecture.) Our simplified block design combine Figure 3:(Architecture.) Our simplified block design combines the H3 block, which is the basis of most SSM architectures, with the ubiquitous MLP block of modern neural networks. Instead of interleaving these two blocks, we simply repeat the Mamba block homogenously. Compared to the H3 block, Mamba replaces the first multiplicative gate with an activation function. Compared to the MLP block, Mamba adds an SSM to the main branch. Forσ\sigmawe use the SiLU / Swish activation[hendrycks2016gaussian,ramachandran2017swish].

文章分类:大模型与架构
文章子分类:模型架构


1. 引言与动机

当前,基础模型(Foundation Models, FMs)几乎由 Transformer 主导。其核心 self-attention 机制能够在上下文窗口内密集地传递信息,从而建模复杂数据。然而,这一特性也带来了两个根本缺陷:

  1. 二次复杂度:attention 的计算量和内存随窗口长度呈二次增长;
  2. 固定上下文窗口:无法直接建模窗口外的信息。

为克服上述问题,研究者提出了大量次二次时间复杂度的架构,包括 linear attention、gated convolution、RNN 变体以及结构化状态空间模型(Structured State Space Models, SSMs,如 S4)。这些模型可视为 RNN 与 CNN 的结合,能够以线性或近线性复杂度进行序列建模,并在音频、视觉等连续信号模态上表现优异。

然而,在文本、DNA 等离散、信息密集的模态上,这些 SSMs 的表现仍不及 Transformer。作者指出,其核心瓶颈在于缺乏基于内容的选择能力(content-based selection):传统 SSM 的参数对时间/输入保持不变,无法根据当前 token 选择性地保留或遗忘信息。

针对这一问题,Mamba 提出了选择性状态空间模型(Selective State Space Models)。通过让 SSM 参数成为输入的函数,模型获得类似 attention 的内容选择能力;同时,借助硬件感知的并行 scan 算法,Mamba 在保持线性序列复杂度的同时,达到了 Transformer 级别的建模能力。


2. 状态空间模型基础

2.1 连续与离散状态空间

结构化状态空间模型 S4 受经典状态空间模型启发,通过一个隐式潜状态 $h(t)$ 将输入 $x(t)$ 映射到输出 $y(t)$。连续系统定义为:

$$ h’(t) = \mathbf{A}h(t) + \mathbf{B}x(t) \tag{1a} $$

$$ y(t) = \mathbf{C}h(t) \tag{1b} $$

为了在离散序列上计算,需要通过离散化规则(如 zero-order hold, ZOH)将连续参数 $(\Delta, \mathbf{A}, \mathbf{B})$ 映射为离散参数 $(\overline{\mathbf{A}}, \overline{\mathbf{B}})$:

$$ \overline{\mathbf{A}} = \exp(\Delta \mathbf{A}), \quad \overline{\mathbf{B}} = (\Delta \mathbf{A})^{-1}(\exp(\Delta \mathbf{A}) - \mathbf{I}) \cdot \Delta \mathbf{B} \tag{2} $$

离散化后,模型可以按两种方式计算:

  • Recurrent 形式

$$ h_t = \overline{\mathbf{A}} h_{t-1} + \overline{\mathbf{B}} x_t \tag{3a} $$

$$ y_t = \mathbf{C} h_t \tag{3b} $$

  • Convolution 形式

$$ \overline{\mathbf{K}} = (\mathbf{C}\overline{\mathbf{B}}, \mathbf{C}\overline{\mathbf{A}}\overline{\mathbf{B}}, \dots, \mathbf{C}\overline{\mathbf{A}}^k\overline{\mathbf{B}}, \dots) \tag{4a} $$

$$ y = x * \overline{\mathbf{K}} \tag{4b} $$

2.2 线性时不变性(LTI)的局限

传统 SSM 要求参数 $(\Delta, \mathbf{A}, \mathbf{B}, \mathbf{C})$ 在时间/输入上保持不变,这一性质称为线性时不变性(Linear Time Invariance, LTI)。LTI 使得模型可以高效地通过卷积并行计算,但也意味着模型无法根据输入内容动态调整状态转移,从而在文本等离散模态上表现受限。


3. 选择性状态空间机制

3.1 核心思想:选择即压缩

序列建模的一个根本问题是如何将上下文压缩到有限状态中。Attention 不压缩上下文,因此效果好但效率低;RNN/SSM 压缩状态,因此效率高但效果受限于压缩质量。

Mamba 从两个合成任务出发说明问题:

  • Selective Copying:在 Copying 任务基础上随机化待记忆 token 的位置,要求模型根据内容选择性地记住相关 token、忽略噪声;
  • Induction Heads:要求模型根据上下文进行关联回忆。

LTI 模型在这两个任务上表现不佳,因为其动态固定,无法根据输入选择性地传播或过滤信息。

3.2 让 SSM 参数依赖于输入

Mamba 的选择机制核心在于:让影响序列交互的参数成为输入的函数。具体地,参数 $\Delta, \mathbf{B}, \mathbf{C}$ 被设为输入 $x$ 的函数:

$$ s_B(x) = \text{Linear}_N(x), \quad s_C(x) = \text{Linear}N(x), \quad s\Delta(x) = \text{Broadcast}_D(\text{Linear}_1(x)) $$

并通过 softplus 等变换得到最终的 $\Delta$。这一改动使模型从时间不变变为时间/输入可变,赋予其内容选择能力,但也打破了与卷积形式的等价性,对计算效率提出挑战。

3.3 选择机制的解释

  • $\Delta$ 的作用:控制对当前输入的关注或忽略程度,可视为 RNN 门控机制的推广;
  • $\mathbf{B}$ 与 $\mathbf{C}$ 的作用:分别控制输入是否进入状态,以及状态是否参与输出,实现基于内容和上下文的状态调制;
  • 边界重置与上下文过滤:选择性使模型能够在任意位置重置状态、过滤无关历史,因此随着上下文增长性能可持续提升。

4. 硬件感知的并行算法

由于参数随输入变化,Mamba 无法使用高效的卷积计算,必须回到 recurrent scan。然而,直接物化形状为 $(B, L, D, N)$ 的潜状态会占用大量内存并造成 IO 瓶颈。

4.1 核心优化策略

为应对上述挑战,Mamba 采用三种经典技术:

  1. Kernel fusion:将离散化、recurrence 等操作融合为一个 CUDA kernel,减少 HBM 与 SRAM 之间的数据搬运;
  2. Parallel scan(scan):利用 work-efficient 的并行 scan 算法打破 recurrence 的串行性;
  3. Recomputation:前向传播时不保存中间状态,反向传播时从 HBM 重新加载输入并在 SRAM 中重新计算。

具体流程为:将 SSM 参数 $(\Delta, \mathbf{A}, \mathbf{B}, \mathbf{C})$ 从慢速 HBM 加载到快速 SRAM,在 SRAM 中完成离散化和 recurrence,最终只将输出 $(B, L, D)$ 写回 HBM。

4.2 复杂度与速度

  • 理论上,selective scan 的 FLOPs 为 $O(BLDN)$,序列长度严格线性扩展;
  • 在 A100 等硬件上,优化后的 scan 比标准 PyTorch 实现快约 20–40 倍,比卷积式 SSM 快约 3 倍;
  • 端到端推理吞吐比同规模 Transformer 高约 5 倍,因为无需维护历史 KV cache。

5. Mamba 整体架构

Figure 9:(Audio Pretraining.) Mamba improves performance ove Figure 9:(Audio Pretraining.) Mamba improves performance over prior state-of-the-art (Sashimi) in autoregressive audio modeling, while improving up to minute-long context or million-length sequences (controlling for computation).

Figure 1:(Overview.)
Structured SSMs independently map each Figure 1:(Overview.) Structured SSMs independently map each channel (e.g.D=5D=5) of an inputxxto outputyythrough a higher dimensional latent statehh(e.g.N=4N=4). Prior SSMs avoid materializing this large effective state (D​NDN, times batch sizeBBand sequence lengthLL) through clever alternate computation paths requiring time-invariance: the(Δ,𝑨,𝑩,𝑪)(\Delta,\bm{A},\bm{B},\bm{C})parameters are constant across time. Our selection mechanism adds back input-dependent dynamics, which also requires a careful hardware-aware algorithm to only materialize the expanded states in more efficient levels of the GPU memory hierarchy.

Figure 2:(Left) The standard version of the Copying task inv Figure 2:(Left) The standard version of the Copying task involves constant spacing between input and output elements and is easily solved by time-invariant models such as linear recurrences and global convolutions. (Right Top) The Selective Copying task has random spacing in between inputs and requires time-varying models that canselectivelyremember or ignore inputs depending on their content. (Right Bottom) The Induction Heads task is an example of associative recall that requires retrieving an answer based on context, a key ability for LLMs.

Figure 5:(Induction Heads.)
Models are trained on sequence l Figure 5:(Induction Heads.) Models are trained on sequence length28=2562^{8}=256, and tested on increasing sequence lengths of26=642^{6}=64up to220=10485762^{20}=1048576. Full numbers inLABEL:tab:induction.

Mamba 将 SSM 模块与 Transformer 的 MLP 块融合为单一、同质的 Mamba 块,去除了独立的 attention 或 MLP 块。每个 Mamba 块通过可控的扩展因子 $E$ 扩展模型维度 $D$,其中大部分参数集中在线性投影,内部 SSM 参数占比较小。

整体架构特点如下:

  • 端到端、纯 recurrent 序列模型
  • 训练时计算和内存随序列长度线性增长;
  • 自回归推理时每步仅需常数时间,无需 KV cache;
  • 支持百万级上下文长度。

6. 实验验证与性能

6.1 合成任务

在 Selective Copying 和 Induction Heads 任务上,Mamba 不仅能轻松解决,还能外推到超过 100 万 token 的序列长度,而 LTI 模型和 attention 在更长序列上表现受限。

Selective Copying 结果

ModelArch.LayerAcc.
S4No gateS418.3
-No gateS697.0
H3H3S457.0
HyenaH3Hyena30.1
-H3S699.7
-MambaS456.4
-MambaHyena28.4
MambaMambaS699.8

6.2 语言建模

Mamba-3B 在预训练困惑度和下游任务上匹敌两倍参数 Transformer,生成吞吐提升约 5 倍。

下游零样本评估(部分结果)

ModelToken.Pile ppl↓LAMBADA ppl↓LAMBADA acc↑HellaSwagPIQAArc-EArc-CWinoGrandeAverage
Hybrid H3-130MGPT289.4825.7731.764.244.424.250.640.1
Pythia-160MNeoX29.6438.1033.030.261.443.224.151.940.6
Mamba-130MNeoX10.5616.0744.335.364.548.024.351.944.7
Hybrid H3-360MGPT212.5848.041.568.151.424.754.148.0
Pythia-410MNeoX9.9510.8451.440.666.952.124.653.848.2
Mamba-370MNeoX8.288.1455.646.569.555.128.055.350.0
Pythia-1BNeoX7.827.9256.147.270.757.027.153.551.9
Mamba-790MNeoX7.336.0262.755.172.161.229.556.157.1
GPT-Neo 1.3BGPT27.5057.248.971.156.225.954.952.4
Hybrid H3-1.3BGPT211.2549.652.671.359.228.156.953.0
OPT-1.3BOPT6.6458.053.772.456.729.659.555.0
Pythia-1.4BNeoX7.516.0861.752.171.060.528.557.255.2
RWKV-1.5BNeoX7.707.0456.452.572.460.529.454.654.3
Mamba-1.4BNeoX6.805.0464.959.174.265.532.861.559.7
GPT-Neo 2.7BGPT25.6362.255.872.161.130.257.656.5
Hybrid H3-2.7BGPT27.9255.759.773.365.632.361.458.0
OPT-2.7BOPT5.1263.660.674.860.831.361.058.7
Pythia-2.8BNeoX6.735.0464.759.374.064.132.959.759.1
RWKV-3BNeoX7.005.2463.959.673.767.833.159.659.6
Mamba-2.8BNeoX6.224.2369.266.175.269.736.363.563.3
GPT-J-6BGPT24.1068.366.375.467.036.664.163.0
OPT-6.7BOPT4.2567.767.276.365.634.965.562.9
Pythia-6.9BNeoX6.514.4567.164.075.267.335.561.361.7
RWKV-7.4BNeoX6.314.3867.265.576.167.837.561.062.5

6.3 基因组学(DNA)

在 HG38 人类基因组数据上,Mamba 随模型规模和上下文长度增长均表现更好。当上下文长度从 1K 增加到 1M 时,Mamba 的预训练困惑度持续改善,而 HyenaDNA 等 LTI 模型性能反而下降。在区分大猩猩、黑猩猩等近缘物种的分类任务上,Mamba 也能利用长达百万的上下文。

6.4 音频建模与生成

在 YouTubeMix 钢琴音乐和 SC09 语音生成任务上,Mamba 超越了 SaShiMi、Hyena 和 Transformer。Mamba-UNet 在 SC09 上的 FID 显著低于更大规模的 GAN/diffusion 模型。

ModelParams (M)NLL↓FID↓IS↑mIS↑AM↓
SampleRNN35.02.0428.961.713.021.76
WaveNet4.21.9255.082.275.801.47
SaShiMi5.81.8731.995.1342.570.74
WaveGAN19.1-2.034.9036.100.80
DiffWave24.1-1.925.2651.210.68
+ SaShiMi23.0-1.425.9469.170.59
Mamba6.11.8520.946.2688.540.52
Mamba24.31.8600.677.33144.90.36
OuterCenterNLL↓FID↓IS↑mIS↑AM↓
S4+MLPMHA+MLP1.8591.455.0647.030.70
S4+MLPS4+MLP1.8671.435.4253.540.65
S4+MLPMamba1.8591.425.7156.510.64
MambaMHA+MLP1.8501.375.6358.230.62
MambaS4+MLP1.8531.076.0573.340.55
MambaMamba1.8520.946.2688.540.52

6.5 消融实验

架构与 SSM 层消融

ModelArch.SSM LayerPerplexity
HyenaH3Hyena10.24
H3H3S4 (complex)10.30
-H3S4 (real)10.34
-H3S68.95
-MambaHyena10.75
-MambaS4 (complex)10.54
-MambaS4 (real)10.56
MambaMambaS68.69

选择性参数消融

Selective $\Delta$Selective $\mathbf{B}$Selective $\mathbf{C}$Perplexity
10.93
10.15
9.98
9.81
8.71

$\mathbf{A}$ 初始化消融

$\mathbf{A}_n$ InitializationFieldPerplexity
$-1/2 + ni$Complex9.16
$-1/2$Real8.85
$-(n+1)$Real8.71
$\exp(\mathcal{N}(0,1))$Real8.71

$\Delta$ 投影维度消融

Size of $\Delta$ proj.Params (M)Perplexity
-358.99.12
1359.18.97
2359.38.97
4359.78.91
8360.58.83
16362.18.84
32365.28.80
64371.58.71

状态维度 $N$ 消融(注:原文表格未明确标注上下两组的完整条件,此处按原文列出):

State dimension $N$Params (M)Perplexity
1367.19.88
2367.49.86
4368.09.82
8369.19.82
16371.59.81
1367.19.73
2367.49.40
4368.09.09
8369.18.84
16371.58.71

6.6 总结

Mamba 通过选择性状态空间机制,使 SSM 获得内容选择能力,同时借助硬件感知 scan 算法保持线性复杂度。实验表明,Mamba 在语言、音频、基因组学等多个模态上达到或超越 Transformer 水平,尤其在长序列场景下优势显著,并具备更高的推理吞吐。


原文来源

Gu, A., & Dao, T. (2023). Mamba: Linear-Time Sequence Modeling with Selective State Spaces. arXiv:2312.00752.
论文链接:https://arxiv.org/abs/2312.00752
代码与预训练模型:https://github.com/state-spaces/mamba

原文链接

https://arxiv.org/html/2312.00752

相关文章

大模型与架构

大语言模型是怎么工作的?

大语言模型工作原理深度解析:从Transformer架构到预训练、微调、推理的完整技术链路。

阅读更多
大模型与架构

Prompt Engineering 技巧大全

Prompt Engineering技巧大全:链式思考、少样本学习、角色设定等高级提示工程方法论。

阅读更多
大模型与架构

Transformer 架构详解:从 Attention 机制到行业范式

本文全面解析Transformer:首个完全基于注意力机制的序列转换模型。文章先对比RNN/CNN的不足,再介绍编码器-解码器架构、多头自注意力、Scaled Dot-Product Attention、位置编码与训练细节,最后给出WMT机器翻译实验结果,说明其训练效率与翻译质量优势。

阅读更多