深入理解状态空间模型 (SSM):RNN 与 Transformer 的优雅融合

4868 字
24 分钟
深入理解状态空间模型 (SSM):RNN 与 Transformer 的优雅融合

1. 状态空间模型的引入#

1.1 为什么需要 SSM?#

在深度学习的序列建模领域,我们长期面临一个核心问题:

Transformer 的长序列建模能力很强,但计算复杂度随序列长度呈二次方增长。

RNN 的计算是线性的,但训练不稳定,难以捕捉长距离依赖。

状态空间模型 (State Space Model, SSM) 提供了一个优雅的解决方案:

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

1.2 SSM 的直观理解#

状态空间模型的核心思想可以用录像机来理解:

┌─────────────────────────────────────────────────────────────────┐
│ 录像机比喻 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 输入序列 ──▶ ┌──────────────┐ ──▶ 输出序列 │
│ │ 状态空间 │ │
│ │ │ │
│ │ ● 磁带位置 │ ← 状态:记录当前位置 │
│ │ ● 播放速度 │ ← 状态:控制快进/快退速度 │
│ │ ● 时间戳 │ ← 状态:当前播放到的时间 │
│ │ │ │
│ └──────────────┘ │
│ │
│ 输入 = 快进/快退/暂停命令 │
│ 输出 = 录像画面 │
│ 状态 = 录像机的内部状态(磁带位置、播放模式等) │
└─────────────────────────────────────────────────────────────────┘

关键洞察

  • 录像机不需要记住”看过的每一帧画面”,只需要记住”当前磁带位置”
  • 这就是状态压缩的思想——用有限的状态表示无限的历史信息

1.3 SSM 发展简史#

1980s: 卡尔曼滤波
│ 线性系统状态估计的理论基础
1990s: 隐马尔可夫模型 (HMM)
│ 离散状态的序列建模
2013-2020: RNN/LSTM 时代
│ 深度学习序列建模主流方法
2021: S4 (Structured State Space Sequence Model)
│ Albert Gu & Tri Dao, UC Berkeley
│ 将 SSM 引入深度学习,实现超长序列建模
2023: Mamba
│ 选择性状态空间模型 (Selective State Spaces)
│ 线性时间序列建模的新范式
2024-2025: SSM 大爆发
│ Jamba, Mistral Mamba, Griffin, Hawk, Mamba-2, ...
│ SSM 与 Transformer 的融合成为主流

2. SSM 的数学基础#

2.1 连续时间状态空间方程#

经典的连续时间状态空间模型由两个微分方程定义:

状态方程(描述状态如何随时间演变):

dx(t)dt=Ax(t)+Bu(t)\frac{dx(t)}{dt} = \mathbf{A} x(t) + \mathbf{B} u(t)

输出方程(描述如何从状态产生输出):

y(t)=Cx(t)+Du(t)y(t) = \mathbf{C} x(t) + \mathbf{D} u(t)

其中:

  • x(t)x(t)状态向量 (state),维度为 NN
  • u(t)u(t)输入向量 (input)
  • y(t)y(t)输出向量 (output)
  • A\mathbf{A}状态转移矩阵 (state transition matrix)
  • B\mathbf{B}输入矩阵 (input matrix)
  • C\mathbf{C}输出矩阵 (output matrix)
  • D\mathbf{D}直接馈通矩阵 (feedthrough matrix)

2.2 离散化#

在计算机中,我们需要将连续时间模型离散化。假设采样间隔为 Δ\Delta

xk+1=Adxk+Bdukx_{k+1} = \mathbf{A}_d x_k + \mathbf{B}_d u_kyk=Cdxk+Dduky_k = \mathbf{C}_d x_k + \mathbf{D}_d u_k

零阶保持离散化

Ad=eAΔ\mathbf{A}_d = e^{\mathbf{A} \Delta}Bd=(A1(eAΔI))B\mathbf{B}_d = (\mathbf{A}^{-1}(e^{\mathbf{A}\Delta} - I)) \mathbf{B}

在实际实现中,我们使用更简洁的近似:

AdAΔ+I\mathbf{A}_d \approx \mathbf{A} \Delta + IBdBΔ\mathbf{B}_d \approx \mathbf{B} \Delta

2.3 状态空间模型的参数#

SSM 的参数是 (A,B,C,D)(\mathbf{A}, \mathbf{B}, \mathbf{C}, \mathbf{D}) 四个矩阵:

矩阵形状作用
A\mathbf{A}(N,N)(N, N)状态转移:控制历史信息如何传递到未来
B\mathbf{B}(N,D)(N, D)输入投影:将输入映射到状态空间
C\mathbf{C}(H,N)(H, N)状态输出:将状态映射到输出空间
D\mathbf{D}(H,D)(H, D)直接馈通:输入直接传递到输出(跳步连接)

其中 NN状态维度(SSM 的隐藏状态大小),DD输入维度HH输出维度

2.4 状态维度的物理意义#

状态维度 NN 是 SSM 的核心超参数,它决定了模型的”记忆容量”:

N 值特点适用场景
小 (4-16)内存紧凑,计算快简单模式
中 (32-64)平衡选择一般序列任务
大 (128-256)长程依赖强长序列、复杂任务
状态维度 vs 注意力头
  • 状态维度 NN 类比于 Transformer 的”隐状态维度”
  • 但 SSM 的状态是压缩表示,而 Transformer 存储完整的 KV 缓存
  • Mamba 使用 N=16N=16,在效率和性能间取得良好平衡

3. S4:结构化状态空间序列模型#

3.1 S4 的核心创新#

S4 (Structured State Space Sequence Model) 由 Albert Gu 和 Tri Dao 在 2021 年提出,是现代 SSM 的奠基之作。

S4 的三大创新

  1. HiPPO 矩阵初始化:用高阶多项色 Hirschberger (HiPPO) 框架初始化 A\mathbf{A} 矩阵,捕捉长程依赖
  2. 对角+低秩结构:将 A\mathbf{A} 矩阵参数化为对角+低秩形式,支持高效计算
  3. 卷积计算:将 SSM 重写为循环卷积,支持并行训练

3.2 HiPPO:高效的矩阵初始化#

S4 的关键洞察是:A\mathbf{A} 矩阵的初始化对 SSM 的性能至关重要

HiPPO 框架提出用特定的多项式基函数来初始化 A\mathbf{A}

对于Legendre 多项式基:

An,k={2n+1if k=n+1nif k=n0otherwise\mathbf{A}_{n,k} = \begin{cases} \sqrt{2n+1} & \text{if } k = n+1 \\ -\sqrt{n} & \text{if } k = n \\ 0 & \text{otherwise} \end{cases}
import torch
import torch.nn as nn
import math
def hippo_legendre(N):
"""
生成 HiPPO Legendre 矩阵 A
N: 状态维度
"""
A = torch.zeros(N, N)
for n in range(N):
for k in range(N):
if k == n + 1:
A[n, k] = math.sqrt(2 * n + 1)
elif k == n:
A[n, k] = -math.sqrt(n)
return A

这个初始化确保 SSM 能够近似表示任意连续函数,从而高效地建模长序列。

3.3 SSM 的卷积形式#

SSM 可以重写为离散的线性卷积形式:

定义卷积核 KRL\mathbf{K} \in \mathbb{R}^L

yk=CAkBu0+CAk1Bu1++CBuky_k = \mathbf{C} \mathbf{A}^k \mathbf{B} u_0 + \mathbf{C} \mathbf{A}^{k-1} \mathbf{B} u_1 + \cdots + \mathbf{C} \mathbf{B} u_k

这相当于:

y=uKy = u * \mathbf{K}

其中卷积核 K=(CB,CAB,CA2B,)\mathbf{K} = (\mathbf{C}\mathbf{B}, \mathbf{C}\mathbf{A}\mathbf{B}, \mathbf{C}\mathbf{A}^2\mathbf{B}, \ldots)

def ssm_to_conv_kernel(A, B, C, L):
"""
将 SSM 参数转换为卷积核
A: (N, N) 状态转移矩阵
B: (N, D) 输入矩阵
C: (H, N) 输出矩阵
L: 序列长度
"""
N = A.shape[0]
D = B.shape[1]
# 计算卷积核
K = []
for i in range(L):
# C @ A^i @ B
power = torch.matrix_power(A, i)
k_i = (C @ power @ B).squeeze()
K.append(k_i)
return torch.stack(K, dim=0) # (L, D, H) -> (L,) for D=H=1

3.4 S4 的完整实现#

import torch
import torch.nn as nn
import torch.nn.functional as F
from scipy.signal import cont2discrete
class S4Block(nn.Module):
"""
S4 (Structured State Space Sequence) 块
"""
def __init__(self, d_model, d_state=16, lr=1.0,
discretization='zoh', mode='nplr'):
super().__init__()
self.d_model = d_model
self.d_state = d_state
# 状态维度
N = d_state
# HiPPO 初始化
A, B = self._init_HiPPO(N, d_model)
# 缩放
A = A * lr
B = B * lr
# 离散化
if discretization == 'zoh':
Ad, Bd, _, _, _ = cont2discrete(
(A.cpu().numpy(), B.cpu().numpy(), torch.zeros(d_model, d_model).numpy(), 0),
dt=1.0, method='zoh'
)
self.A = torch.tensor(Ad, dtype=torch.float)
self.B = torch.tensor(Bd, dtype=torch.float).unsqueeze(0) # (1, N, D)
else:
self.A = A + torch.eye(N) # 简单欧拉
self.B = B
# 可学习参数
self.C = nn.Parameter(torch.randn(d_model, N))
self.D = nn.Parameter(torch.randn(d_model)) # 直接馈通
# 初始化
nn.init.normal_(self.C, std=0.5)
nn.init.normal_(self.D, std=1.0)
def _init_HiPPO(self, N, D):
"""HiPPO Legendre 初始化"""
A = torch.zeros(N, N)
B = torch.zeros(N, D)
for n in range(N):
for k in range(N):
if k == n + 1:
A[n, k] = (2 * n + 1) ** 0.5
elif k == n:
A[n, k] = -(n + 1) ** 0.5
# Legendre 初始 B
for n in range(N):
B[n] = ((2 * n + 1) ** 0.5) * torch.ones(D)
return A, B
def forward(self, u):
"""
u: (batch, seq_len, d_model)
返回: (batch, seq_len, d_model)
"""
batch, L, d = u.shape
# 计算卷积核
K = self._compute_kernel(L) # (L, d, d)
# 线性卷积(使用 FFT 加速)
y = F.conv1d(
u.view(batch, 1, -1),
K.unsqueeze(0).expand(batch, -1, -1),
padding=L-1,
groups=batch
)
y = y[:, :, :L] # 截断到原始长度
y = y + u * self.D # 加上直接馈通
return y
def _compute_kernel(self, L):
"""计算 SSM 卷积核"""
N = self.d_state
# 展平计算 C @ A^i @ B
K = torch.zeros(L, self.d_model, self.d_model)
# 快速计算(截断 HiPPO 的指数衰减)
A_powers = torch.matrix_power(self.A, 0)
AB = self.A @ self.B.squeeze(0) # (N, D)
for i in range(min(L, 2 * N)): # HiPPO 矩阵的谱半径 < 1,快速收敛
K[i] = self.C @ A_powers @ self.B.squeeze(0)
A_powers = A_powers @ self.A
return K

4. Mamba:选择性状态空间模型#

4.1 Mamba 的核心创新#

Mamba 由 Carnegie Mellon 大学的 Albert Gu 和 Tri Dao 于 2023 年底提出,是对 S4 的重大改进。

Mamba 的关键洞察

不是所有输入都应该以相同的方式影响状态!

在标准 SSM 中,A\mathbf{A}, B\mathbf{B}, C\mathbf{C} 矩阵是与输入无关的常数。这限制了 SSM 根据输入内容动态调整的能力。

Mamba 通过选择性扫描机制 (Selective Scan) 解决了这个问题。

4.2 选择性 SSM#

Mamba 的核心改变是让 B\mathbf{B}, C\mathbf{C}(以及离散化的 Aˉ\mathbf{\bar{A}}, Bˉ\mathbf{\bar{B}})变成输入相关的函数

标准 SSM(输入无关)

x=Ax+Bux' = \mathbf{A} x + \mathbf{B} u

Mamba SSM(输入相关)

x=A(u)x+B(u)ux' = \mathbf{A}(u) x + \mathbf{B}(u) u

其中 B(u)=LinearB(u)\mathbf{B}(u) = \text{Linear}_B(u)A(u)=ASoftplus(LinearA(u))\mathbf{A}(u) = \mathbf{A} \cdot \text{Softplus}(\text{Linear}_A(u))

class MambaBlock(nn.Module):
"""
Mamba 选择性状态空间模型块
"""
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)
# 输入投影
self.in_proj = nn.Linear(d_model, self.d_inner * 2, bias=False)
# 卷积层(局部上下文)
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
)
# SSM 参数投影(输入相关)
self.x_proj = nn.Linear(self.d_inner, d_state * 2 + 1, bias=False) # B, C, Δ
# Δ 的参数
self.dt_proj = nn.Linear(d_state, self.d_inner, bias=True)
# A 矩阵(初始化为 HiPPO)
A = self._init_A(d_state, self.d_inner)
self.A_log = nn.Parameter(torch.zeros(self.d_inner, d_state))
self.A_log.copy_(torch.log(A))
# D 矩阵(直接馈通)
self.D = nn.Parameter(torch.ones(self.d_inner))
# 输出投影
self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)
def _init_A(self, N, D):
"""HiPPO 矩阵初始化"""
A = torch.zeros(D, N)
for d in range(D):
for n in range(N):
if n == 0:
A[d, n] = (n + 1) ** 0.5
else:
A[d, n] = (2 * n + 1) ** 0.5
if n < N - 1:
A[d, n + 1] = -(n + 1) ** 0.5
return A
def forward(self, x):
"""
x: (batch, seq_len, d_model)
"""
batch, L, d = x.shape
# 输入投影并分割
xz = self.in_proj(x) # (batch, L, 2 * d_inner)
x_inner, z = xz.chunk(2, dim=-1) # 各 (batch, L, d_inner)
# 局部卷积
x_conv = self.conv1d(x_inner.transpose(1, 2))[:, :, :L].transpose(1, 2)
x_conv = F.silu(x_conv)
# SSM 参数(选择性:依赖输入)
x_ssm = x_conv # (batch, L, d_inner)
x_proj_out = self.x_proj(x_ssm) # (batch, L, d_state * 2 + 1)
B, C, delta = x_proj_out.split([self.d_state, self.d_state, 1], dim=-1)
delta = F.softplus(self.dt_proj(delta)) # (batch, L, d_inner)
# 选择性扫描(核心)
y = self.selective_scan(
x_conv, delta, self.A_log.exp(), B, C, self.D, z
)
# 门控
y = y * F.silu(z)
# 输出投影
output = self.out_proj(y)
return output
def selective_scan(self, u, delta, A, B, C, D, z):
"""
选择性扫描算法
这是 Mamba 的核心:输入决定如何扫描序列
"""
batch, L, d_inner = u.shape
N = A.shape[-1]
# 离散化
deltaA = torch.exp(delta.unsqueeze(-1) * A) # (batch, L, d_inner, N)
deltaB_u = delta.unsqueeze(-1) * B.unsqueeze(2) * u.unsqueeze(-1) # (batch, L, d_inner, N)
# 扫描
y = torch.zeros(batch, L, d_inner, N, device=u.device, dtype=u.dtype)
for i in range(L):
y[:, i] = deltaA[:, i] * y[:, i-1] + deltaB_u[:, i] if i > 0 else deltaB_u[:, i]
# 输出
y = (y * C.unsqueeze(1).unsqueeze(-1)).sum(-1) # (batch, L, d_inner)
return y + u * D

4.3 选择性机制的可视化#

标准 SSM(所有输入同等对待):
输入: "The cat sat on the mat"
↓ ↓ ↓ ↓ ↓
A,B,C: 常数 常数 常数 常数 常数
↓ ↓ ↓ ↓ ↓
状态: 线性更新 线性更新 线性更新 线性更新 线性更新
Mamba SSM(选择性处理):
输入: "The cat sat on the mat"
↓ ↓ ↓ ↓ ↓
A,B,C: 动态 动态 动态 动态 动态
↓ ↓ ↓ ↓ ↓
状态: 选择性更新,选择性遗忘,选择性记忆

4.4 Mamba vs S4 关键差异#

特性S4Mamba
A, B, C 矩阵固定,与输入无关输入相关,可学习
扫描方式并行卷积顺序扫描(可并行优化)
选择性
长序列建模更强
因果建模隐式显式选择
计算效率高(卷积)高(扫描 + 并行)

5. SSM 与其他架构的对比#

5.1 计算复杂度对比#

架构前向计算内存长序列适应性
Transformer (Full)O(L2)O(L^2)O(L2)O(L^2)需位置编码外推
Transformer (Flash Attention)O(L2)O(L^2)O(L)O(L)有限
RNN/LSTMO(L)O(L)O(1)O(1)梯度消失
SSM (S4)O(LN)O(L \cdot N)O(LN)O(L \cdot N)
SSM (Mamba)O(LN)O(L \cdot N)O(N)O(N)很强

其中 LL 是序列长度,NN 是状态维度(通常 NLN \ll L)。

5.2 并行化能力对比#

Transformer: 完全并行(注意力矩阵可并行计算)
│███████████████████████████████████████
RNN: 顺序计算(无法并行)
│ →
│ →
│ →
│ →
│ →
│ →
SSM (S4): 可并行卷积
│███████████████ (FFT + 卷积)
SSM (Mamba): 并行扫描(并行扫描算法)
│▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓ (并行扫描)

5.3 线性注意力:SSM 的等价形式#

有趣的是,SSM 可以被重写为线性注意力的形式

标准线性注意力:

yi=j=1iexp(qikj)l=1iexp(qikl)vjy_i = \sum_{j=1}^{i} \frac{\exp(q_i \cdot k_j)}{\sum_{l=1}^{i} \exp(q_i \cdot k_l)} v_j

exp(qikj)=ϕ(sj1)ψ(xj)\exp(q_i \cdot k_j) = \phi(s_{j-1}) \cdot \psi(x_j) 时,上式可以递归化为:

si=ϕ(si1)ψ(xi)+si1s_i = \phi(s_{i-1}) \cdot \psi(x_i) + s_{i-1}yi=θ(si)y_i = \theta(s_i)

这正是 SSM 的递归形式!其中 ϕ\phi 对应 A\mathbf{A}ψ\psi 对应 B\mathbf{B}θ\theta 对应 C\mathbf{C}

5.4 SSM 与 Transformer 的融合#

现代架构经常将 SSM 与 Transformer 结合:

模型融合方式
MambaFormer交替使用 Mamba 层和 Transformer 层
JambaTransformer 为主,插入 Mamba 层
Mistral Mamba纯 Mamba,替换 Transformer
StripedHyena混合 SSM、Attention、卷积
class HybridMambaTransformer(nn.Module):
"""
Mamba + Transformer 混合架构
"""
def __init__(self, d_model, num_heads, d_state,
transformer_ratio=0.7, num_layers=12):
super().__init__()
self.num_layers = num_layers
self.transformer_layers = int(num_layers * transformer_ratio)
self.layers = nn.ModuleList()
for i in range(num_layers):
if i < self.transformer_layers:
# Transformer 层
self.layers.append(
TransformerLayer(d_model, num_heads)
)
else:
# Mamba 层
self.layers.append(
MambaBlock(d_model, d_state)
)
def forward(self, x):
for layer in self.layers:
x = layer(x)
return x

6. SSM 的硬件优化#

6.1 内存高效的状态存储#

SSM 的一个关键优势是状态压缩

# Transformer 的 KV 缓存
kv_cache = {
'k': torch.zeros(batch, num_heads, seq_len, head_dim),
'v': torch.zeros(batch, num_heads, seq_len, head_dim)
} # O(L) 内存增长
# SSM 的状态
h_state = torch.zeros(batch, d_model, d_state) # O(1) 恒定内存

6.2 并行扫描算法#

Mamba 使用并行扫描来处理递归依赖:

def parallel_scan(log_a, log_b, x):
"""
并行扫描算法
计算 y = sum_{i=0}^{N-1} (prod_{j=i+1}^{N-1} a_j) * b_i * x_i
这是通过树形结构实现的,复杂度 O(log N) 而非 O(N)
"""
N = x.shape[0]
# 树形扫描
# 级别 1: 相邻元素组合
# 级别 2: 相邻块组合
# ... 直到根节点
# 简化的并行扫描实现
if N == 1:
return log_a.exp() * log_b.exp() * x
# 递归实现
mid = N // 2
left = parallel_scan(log_a[:mid], log_b[:mid], x[:mid])
right = log_a[mid:].exp() * left + log_b[mid:].exp() * x[mid:]
return torch.cat([left, right], dim=0)

6.3 融合内核#

现代 SSM 实现使用 CUDA 融合内核来最大化效率:

优化技术描述效果
融合扫描将扫描的多个操作融合为单一内核减少内存访问
梯度检查点用计算换内存减少显存占用
量化状态INT8/FP16 状态减少状态内存
异步调度计算与通信重叠提高吞吐量

7. 主流 SSM 模型#

7.1 S4 系列#

模型开发者特点
S4UC Berkeley奠基之作
S4-NDUC Berkeley支持任意维度
S4-DUC Berkeley对角结构优化
BlackMambaMuse & Carper.ai量化友好

7.2 Mamba 系列#

模型开发者特点
MambaCMU & Tri Dao选择性 SSM
Mamba-2Tri Dao改进并行性
Mistral MambaMistral AI生产级
Mamba-2-78MMLC轻量高效

7.3 混合架构#

模型架构开发者
JambaTransformer + MambaAI21 Labs
MambaFormer交替研究
StripedHyena多专家混合Together AI
RavenRWKV + SSM研究

7.4 开源模型对比#

模型参数量上下文类型
Mamba-2.8B2.7B256K纯 Mamba
Mistral-Nemo-Mamba12B128K纯 Mamba
Jamba-Mini12B256K混合
Falcon-Mamba7B2048纯 Mamba
SeqGPT1B8K纯 Mamba

8. SSM 的实际应用#

8.1 时间序列预测#

SSM 在时间序列预测中表现优异:

class SSMTimeSeriesForecast(nn.Module):
"""
基于 SSM 的时间序列预测模型
"""
def __init__(self, d_model, d_state, d_output):
super().__init__()
self.encoder = nn.Linear(1, d_model)
self.ssm = nn.Sequential(
MambaBlock(d_model, d_state),
MambaBlock(d_model, d_state),
MambaBlock(d_model, d_state),
)
self.decoder = nn.Linear(d_model, d_output)
def forward(self, x):
"""
x: (batch, seq_len, 1) - 单变量时间序列
返回: (batch, pred_len, 1) - 预测的未来值
"""
x = self.encoder(x)
x = self.ssm(x)
return self.decoder(x)

应用场景

  • 股票价格预测
  • 天气预报
  • 能源消耗预测
  • 传感器数据分析

8.2 语音处理#

SSM 在语音任务中表现出色:

任务SSM 优势代表模型
语音识别长音频建模S4-CTC
语音合成高效并行MambaTTS
语音增强因果建模SSM-U-Net
音乐生成序列生成MusicSSM

8.3 基因组学#

DNA/RNA 序列分析是 SSM 的强项:

class GenomicSSM(nn.Module):
"""
DNA 序列分析模型
DNA token: A, C, G, T -> 4 类
"""
def __init__(self, vocab_size=4, d_model=256, d_state=16):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.ssm_layers = nn.ModuleList([
MambaBlock(d_model, d_state)
for _ in range(24)
])
self.classifier = nn.Linear(d_model, 2) # 二分类:基因/非基因
def forward(self, dna_sequence):
x = self.embedding(dna_sequence)
for layer in self.ssm_layers:
x = layer(x)
return self.classifier(x.mean(dim=1)) # 全局池化

为什么 SSM 适合基因组学

  • DNA 序列可能长达数百万碱基对
  • SSM 的线性复杂度完美匹配
  • 长程依赖对基因调控至关重要

8.4 视觉任务#

SSM 正在进入计算机视觉领域:

架构方法任务
Vision MambaVim图像分类
SSM-UNetSSM + U-Net分割
DiSDiffusion + SSM生成
class VisionMambaBlock(nn.Module):
"""
用于视觉的 Mamba 块
处理 2D 图像补丁
"""
def __init__(self, d_model, d_state, patch_size=16, img_size=224):
super().__init__()
self.patch_size = patch_size
self.num_patches = (img_size // patch_size) ** 2
# 将图像划分为补丁
self.patch_embed = nn.Conv2d(
3, d_model, kernel_size=patch_size, stride=patch_size
)
# 位置嵌入
self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, d_model))
# SSM 层
self.ssm = MambaBlock(d_model, d_state)
def forward(self, x):
# x: (B, C, H, W)
B = x.shape[0]
# 划分补丁
x = self.patch_embed(x) # (B, d_model, H/P, W/P)
x = x.flatten(2).transpose(1, 2) # (B, num_patches, d_model)
# 添加位置编码
x = x + self.pos_embed
# SSM 处理
x = self.ssm(x)
return x

9. SSM 的挑战与局限#

9.1 主要挑战#

挑战描述当前解决方案
选择性有限Mamba 的选择性仍有改进空间改进扫描机制
长程依赖SSM 状态容量有限增大状态维度
训练稳定性大规模训练仍有挑战更好的初始化
硬件适配需要特定优化CUDA 内核优化

9.2 与 Transformer 的能力差距#

虽然 SSM 在长序列任务上表现出色,但在某些任务上仍落后于 Transformer:

任务SSMTransformer原因分析
短序列语言理解中等Attention 的全局建模更强
长序列建模中等SSM 的线性复杂度优势
精确检索中等需要复杂 KV 模式
代码生成待观察需要精确执行路径

9.3 状态维度的权衡#

# 状态维度 vs 性能 的权衡
"""
状态维度 N 太小:
- 计算快,内存省
- 但容量不足,无法建模长程依赖
状态维度 N 太大:
- 容量充足
- 但计算变慢,接近 Transformer
平衡点: N ≈ 16-64 对大多数任务足够
"""
# 实验数据参考
configs = {
'N=8': {'speed': 1.0, 'quality': 0.85, 'memory': 0.5},
'N=16': {'speed': 0.85, 'quality': 0.92, 'memory': 0.7},
'N=32': {'speed': 0.70, 'quality': 0.95, 'memory': 0.85},
'N=64': {'speed': 0.55, 'quality': 0.97, 'memory': 1.0},
}

10. 核心公式总结#

  1. 连续时间状态方程
x˙(t)=Ax(t)+Bu(t)\dot{x}(t) = \mathbf{A}x(t) + \mathbf{B}u(t)
  1. 离散化状态方程
xk+1=Aˉxk+Bˉukx_{k+1} = \mathbf{\bar{A}}x_k + \mathbf{\bar{B}}u_k

其中 Aˉ=eAΔ\mathbf{\bar{A}} = e^{\mathbf{A}\Delta}Bˉ=(A1(eAΔI))B\mathbf{\bar{B}} = (\mathbf{A}^{-1}(e^{\mathbf{A}\Delta} - I))\mathbf{B}

  1. SSM 输出
yk=Cxk+Duky_k = \mathbf{C}x_k + \mathbf{D}u_k
  1. Mamba 选择性机制
x=A(u)x+B(u)ux' = \mathbf{A}(u) \cdot x + \mathbf{B}(u) \cdot u

其中 A(u)=exp(ΔAinit)\mathbf{A}(u) = \exp(\Delta \cdot \mathbf{A}_{init})Δ\Delta 是输入相关的步长

  1. 卷积核计算
Ki=CAiB,i=0,1,,L1K_i = \mathbf{C}\mathbf{A}^i\mathbf{B}, \quad i = 0, 1, \ldots, L-1
  1. HiPPO 矩阵(Legendre)
An,k={2n+1k=n+1nk=n0otherwise\mathbf{A}_{n,k} = \begin{cases} \sqrt{2n+1} & k = n+1 \\ -\sqrt{n} & k = n \\ 0 & \text{otherwise} \end{cases}

11. 总结与展望#

SSM 的核心优势#

  1. 线性复杂度O(LN)O(L \cdot N) 而非 O(L2)O(L^2)
  2. 并行可训练:类似 CNN 的高效并行
  3. 恒定内存:推理时状态大小与序列长度无关
  4. 长程依赖:通过状态压缩捕捉长距离依赖
  5. 硬件友好:比 Attention 更易于硬件优化

SSM vs Transformer 的定位#

SSM 不是要完全替代 Transformer,而是为不同的场景提供更好的选择。

场景推荐架构
短序列 + 高精度Transformer
长序列 + 高效SSM (Mamba)
超长序列SSM + 滑动窗口
混合需求SSM + Transformer

未来方向#

  1. 更强的选择性:更智能的输入相关路由
  2. 多尺度 SSM:不同层使用不同状态维度
  3. 与 Diffusion 结合:SSM-Diffusion 生成模型
  4. 多模态 SSM:统一的文本、图像、音频建模
  5. 硬件协同设计:针对 SSM 特性的专用芯片

SSM 代表了深度学习序列建模的一个重要突破,它优雅地融合了 RNN 的效率、CNN 的并行性和 Transformer 的表达能力。随着研究的深入和硬件的优化,SSM 有望在更多场景中发挥重要作用。

参考资料#

  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., et al. (2019). “HiPPO: Recurrent Memory of Optimal Projections.” NeurIPS.
  4. Gu, A., et al. (2021). “Learnable Fourier Features for Time-Series Forecasting.” NeurIPS.
  5. Poli, M., et al. (2023). “Hyena Hierarchy: Towards Larger Convolutional Language Models.” ICML.
  6. Lieber, O., et al. (2024). “Jamba: A Hybrid Transformer-Mamba Language Model.” arXiv.
  7. Xue, F., et al. (2024). “Mamba-2: State Space Models at Scale.” arXiv.
  8. Zhao, B., et al. (2024). “Vision Mamba: Efficient Visual Representation Learning with Bidirectional State Space Model.” arXiv.

文章分享

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

深入理解状态空间模型 (SSM):RNN 与 Transformer 的优雅融合
https://aiattnstudio.link/posts/state-space-model/
作者
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标签