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) dtdx(t)=Ax(t)+Bu(t)
输出方程 (描述如何从状态产生输出):
y(t)=Cx(t)+Du(t) y(t) = \mathbf{C} x(t) + \mathbf{D} u(t) y(t)=Cx(t)+Du(t)
其中:
- x(t)x(t)x(t) 是状态向量 (state),维度为 NNN
- u(t)u(t)u(t) 是输入向量 (input)
- y(t)y(t)y(t) 是输出向量 (output)
- A\mathbf{A}A 是状态转移矩阵 (state transition matrix)
- B\mathbf{B}B 是输入矩阵 (input matrix)
- C\mathbf{C}C 是输出矩阵 (output matrix)
- D\mathbf{D}D 是直接馈通矩阵 (feedthrough matrix)
2.2 离散化
在计算机中,我们需要将连续时间模型离散化。假设采样间隔为 Δ\DeltaΔ:
xk+1=Adxk+Bduk x_{k+1} = \mathbf{A}_d x_k + \mathbf{B}_d u_k xk+1=Adxk+Bduk
yk=Cdxk+Dduk y_k = \mathbf{C}_d x_k + \mathbf{D}_d u_k yk=Cdxk+Dduk
零阶保持离散化:
Ad=eAΔ \mathbf{A}_d = e^{\mathbf{A} \Delta} Ad=eAΔ
Bd=(A−1(eAΔ−I))B \mathbf{B}_d = (\mathbf{A}^{-1}(e^{\mathbf{A}\Delta} - I)) \mathbf{B} Bd=(A−1(eAΔ−I))B
在实际实现中,我们使用更简洁的近似:
Ad≈AΔ+I \mathbf{A}_d \approx \mathbf{A} \Delta + I Ad≈AΔ+I
Bd≈BΔ \mathbf{B}_d \approx \mathbf{B} \Delta Bd≈BΔ
2.3 状态空间模型的参数
SSM 的参数是 (A,B,C,D)(\mathbf{A}, \mathbf{B}, \mathbf{C}, \mathbf{D})(A,B,C,D) 四个矩阵:
| 矩阵 | 形状 | 作用 |
|---|---|---|
| A\mathbf{A}A | (N,N)(N, N)(N,N) | 状态转移:控制历史信息如何传递到未来 |
| B\mathbf{B}B | (N,D)(N, D)(N,D) | 输入投影:将输入映射到状态空间 |
| C\mathbf{C}C | (H,N)(H, N)(H,N) | 状态输出:将状态映射到输出空间 |
| D\mathbf{D}D | (H,D)(H, D)(H,D) | 直接馈通:输入直接传递到输出(跳步连接) |
其中 NNN 是状态维度 (SSM 的隐藏状态大小),DDD 是输入维度 ,HHH 是输出维度。
2.4 状态维度的物理意义
状态维度 NNN 是 SSM 的核心超参数,它决定了模型的"记忆容量":
| N 值 | 特点 | 适用场景 |
|---|---|---|
| 小 (4-16) | 内存紧凑,计算快 | 简单模式 |
| 中 (32-64) | 平衡选择 | 一般序列任务 |
| 大 (128-256) | 长程依赖强 | 长序列、复杂任务 |
3. S4:结构化状态空间序列模型
3.1 S4 的核心创新
S4 (Structured State Space Sequence Model) 由 Albert Gu 和 Tri Dao 在 2021 年提出,是现代 SSM 的奠基之作。
S4 的三大创新:
- HiPPO 矩阵初始化 :用高阶多项色 Hirschberger (HiPPO) 框架初始化 A\mathbf{A}A 矩阵,捕捉长程依赖
- 对角+低秩结构 :将 A\mathbf{A}A 矩阵参数化为对角+低秩形式,支持高效计算
- 卷积计算:将 SSM 重写为循环卷积,支持并行训练
3.2 HiPPO:高效的矩阵初始化
S4 的关键洞察是:A\mathbf{A}A 矩阵的初始化对 SSM 的性能至关重要。
HiPPO 框架提出用特定的多项式基函数来初始化 A\mathbf{A}A:
对于Legendre 多项式 基:
An,k={2n+1if k=n+1−nif 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} An,k=⎩ ⎨ ⎧2n+1 −n 0if k=n+1if k=notherwise
python
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 可以重写为离散的线性卷积形式:
定义卷积核 K∈RL\mathbf{K} \in \mathbb{R}^LK∈RL:
yk=CAkBu0+CAk−1Bu1+⋯+CBuk y_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 yk=CAkBu0+CAk−1Bu1+⋯+CBuk
这相当于:
y=u∗K y = u * \mathbf{K} y=u∗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)K=(CB,CAB,CA2B,...)
python
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 的完整实现
python
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}A, B\mathbf{B}B, C\mathbf{C}C 矩阵是与输入无关的常数。这限制了 SSM 根据输入内容动态调整的能力。
Mamba 通过选择性扫描机制 (Selective Scan) 解决了这个问题。
4.2 选择性 SSM
Mamba 的核心改变是让 B\mathbf{B}B, C\mathbf{C}C(以及离散化的 Aˉ\mathbf{\bar{A}}Aˉ, Bˉ\mathbf{\bar{B}}Bˉ)变成输入相关的函数:
标准 SSM(输入无关) :
x′=Ax+Bu x' = \mathbf{A} x + \mathbf{B} u x′=Ax+Bu
Mamba SSM(输入相关) :
x′=A(u)x+B(u)u x' = \mathbf{A}(u) x + \mathbf{B}(u) u x′=A(u)x+B(u)u
其中 B(u)=LinearB(u)\mathbf{B}(u) = \text{Linear}_B(u)B(u)=LinearB(u),A(u)=A⋅Softplus(LinearA(u))\mathbf{A}(u) = \mathbf{A} \cdot \text{Softplus}(\text{Linear}_A(u))A(u)=A⋅Softplus(LinearA(u))
python
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 关键差异
| 特性 | S4 | Mamba |
|---|---|---|
| A, B, C 矩阵 | 固定,与输入无关 | 输入相关,可学习 |
| 扫描方式 | 并行卷积 | 顺序扫描(可并行优化) |
| 选择性 | 无 | 有 |
| 长序列建模 | 强 | 更强 |
| 因果建模 | 隐式 | 显式选择 |
| 计算效率 | 高(卷积) | 高(扫描 + 并行) |
5. SSM 与其他架构的对比
5.1 计算复杂度对比
| 架构 | 前向计算 | 内存 | 长序列适应性 |
|---|---|---|---|
| Transformer (Full) | O(L2)O(L^2)O(L2) | O(L2)O(L^2)O(L2) | 需位置编码外推 |
| Transformer (Flash Attention) | O(L2)O(L^2)O(L2) | O(L)O(L)O(L) | 有限 |
| RNN/LSTM | O(L)O(L)O(L) | O(1)O(1)O(1) | 梯度消失 |
| SSM (S4) | O(L⋅N)O(L \cdot N)O(L⋅N) | O(L⋅N)O(L \cdot N)O(L⋅N) | 强 |
| SSM (Mamba) | O(L⋅N)O(L \cdot N)O(L⋅N) | O(N)O(N)O(N) | 很强 |
其中 LLL 是序列长度,NNN 是状态维度(通常 N≪LN \ll LN≪L)。
5.2 并行化能力对比
Transformer: 完全并行(注意力矩阵可并行计算)
│
│███████████████████████████████████████
│
RNN: 顺序计算(无法并行)
│
│ →
│ →
│ →
│ →
│ →
│ →
│
SSM (S4): 可并行卷积
│
│███████████████ (FFT + 卷积)
│
SSM (Mamba): 并行扫描(并行扫描算法)
│
│▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓ (并行扫描)
5.3 线性注意力:SSM 的等价形式
有趣的是,SSM 可以被重写为线性注意力的形式!
标准线性注意力:
yi=∑j=1iexp(qi⋅kj)∑l=1iexp(qi⋅kl)vj y_i = \sum_{j=1}^{i} \frac{\exp(q_i \cdot k_j)}{\sum_{l=1}^{i} \exp(q_i \cdot k_l)} v_j yi=j=1∑i∑l=1iexp(qi⋅kl)exp(qi⋅kj)vj
当 exp(qi⋅kj)=ϕ(sj−1)⋅ψ(xj)\exp(q_i \cdot k_j) = \phi(s_{j-1}) \cdot \psi(x_j)exp(qi⋅kj)=ϕ(sj−1)⋅ψ(xj) 时,上式可以递归化为:
si=ϕ(si−1)⋅ψ(xi)+si−1 s_i = \phi(s_{i-1}) \cdot \psi(x_i) + s_{i-1} si=ϕ(si−1)⋅ψ(xi)+si−1
yi=θ(si) y_i = \theta(s_i) yi=θ(si)
这正是 SSM 的递归形式!其中 ϕ\phiϕ 对应 A\mathbf{A}A,ψ\psiψ 对应 B\mathbf{B}B,θ\thetaθ 对应 C\mathbf{C}C。
5.4 SSM 与 Transformer 的融合
现代架构经常将 SSM 与 Transformer 结合:
| 模型 | 融合方式 |
|---|---|
| MambaFormer | 交替使用 Mamba 层和 Transformer 层 |
| Jamba | Transformer 为主,插入 Mamba 层 |
| Mistral Mamba | 纯 Mamba,替换 Transformer |
| StripedHyena | 混合 SSM、Attention、卷积 |
python
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 的一个关键优势是状态压缩:
python
# 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 使用并行扫描来处理递归依赖:
python
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 系列
| 模型 | 开发者 | 特点 |
|---|---|---|
| S4 | UC Berkeley | 奠基之作 |
| S4-ND | UC Berkeley | 支持任意维度 |
| S4-D | UC Berkeley | 对角结构优化 |
| BlackMamba | Muse & Carper.ai | 量化友好 |
7.2 Mamba 系列
| 模型 | 开发者 | 特点 |
|---|---|---|
| Mamba | CMU & Tri Dao | 选择性 SSM |
| Mamba-2 | Tri Dao | 改进并行性 |
| Mistral Mamba | Mistral AI | 生产级 |
| Mamba-2-78M | MLC | 轻量高效 |
7.3 混合架构
| 模型 | 架构 | 开发者 |
|---|---|---|
| Jamba | Transformer + Mamba | AI21 Labs |
| MambaFormer | 交替 | 研究 |
| StripedHyena | 多专家混合 | Together AI |
| Raven | RWKV + SSM | 研究 |
7.4 开源模型对比
| 模型 | 参数量 | 上下文 | 类型 |
|---|---|---|---|
| Mamba-2.8B | 2.7B | 256K | 纯 Mamba |
| Mistral-Nemo-Mamba | 12B | 128K | 纯 Mamba |
| Jamba-Mini | 12B | 256K | 混合 |
| Falcon-Mamba | 7B | 2048 | 纯 Mamba |
| SeqGPT | 1B | 8K | 纯 Mamba |
8. SSM 的实际应用
8.1 时间序列预测
SSM 在时间序列预测中表现优异:
python
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 的强项:
python
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 Mamba | Vim | 图像分类 |
| SSM-UNet | SSM + U-Net | 分割 |
| DiS | Diffusion + SSM | 生成 |
python
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:
| 任务 | SSM | Transformer | 原因分析 |
|---|---|---|---|
| 短序列语言理解 | 中等 | 强 | Attention 的全局建模更强 |
| 长序列建模 | 强 | 中等 | SSM 的线性复杂度优势 |
| 精确检索 | 中等 | 强 | 需要复杂 KV 模式 |
| 代码生成 | 待观察 | 强 | 需要精确执行路径 |
9.3 状态维度的权衡
python
# 状态维度 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. 核心公式总结
-
连续时间状态方程 :
x˙(t)=Ax(t)+Bu(t) \dot{x}(t) = \mathbf{A}x(t) + \mathbf{B}u(t) x˙(t)=Ax(t)+Bu(t)
-
离散化状态方程 :
xk+1=Aˉxk+Bˉuk x_{k+1} = \mathbf{\bar{A}}x_k + \mathbf{\bar{B}}u_k xk+1=Aˉxk+Bˉuk
其中 Aˉ=eAΔ\mathbf{\bar{A}} = e^{\mathbf{A}\Delta}Aˉ=eAΔ,Bˉ=(A−1(eAΔ−I))B\mathbf{\bar{B}} = (\mathbf{A}^{-1}(e^{\mathbf{A}\Delta} - I))\mathbf{B}Bˉ=(A−1(eAΔ−I))B
-
SSM 输出 :
yk=Cxk+Duk y_k = \mathbf{C}x_k + \mathbf{D}u_k yk=Cxk+Duk
-
Mamba 选择性机制 :
x′=A(u)⋅x+B(u)⋅u x' = \mathbf{A}(u) \cdot x + \mathbf{B}(u) \cdot u x′=A(u)⋅x+B(u)⋅u
其中 A(u)=exp(Δ⋅Ainit)\mathbf{A}(u) = \exp(\Delta \cdot \mathbf{A}_{init})A(u)=exp(Δ⋅Ainit),Δ\DeltaΔ 是输入相关的步长
-
卷积核计算 :
Ki=CAiB,i=0,1,...,L−1 K_i = \mathbf{C}\mathbf{A}^i\mathbf{B}, \quad i = 0, 1, \ldots, L-1 Ki=CAiB,i=0,1,...,L−1
-
HiPPO 矩阵(Legendre) :
An,k={2n+1k=n+1−nk=n0otherwise \mathbf{A}_{n,k} = \begin{cases} \sqrt{2n+1} & k = n+1 \\ -\sqrt{n} & k = n \\ 0 & \text{otherwise} \end{cases} An,k=⎩ ⎨ ⎧2n+1 −n 0k=n+1k=notherwise
11. 总结与展望
SSM 的核心优势
- 线性复杂度 :O(L⋅N)O(L \cdot N)O(L⋅N) 而非 O(L2)O(L^2)O(L2)
- 并行可训练:类似 CNN 的高效并行
- 恒定内存:推理时状态大小与序列长度无关
- 长程依赖:通过状态压缩捕捉长距离依赖
- 硬件友好:比 Attention 更易于硬件优化
SSM vs Transformer 的定位
SSM 不是要完全替代 Transformer,而是为不同的场景提供更好的选择。
| 场景 | 推荐架构 |
|---|---|
| 短序列 + 高精度 | Transformer |
| 长序列 + 高效 | SSM (Mamba) |
| 超长序列 | SSM + 滑动窗口 |
| 混合需求 | SSM + Transformer |
SSM 代表了深度学习序列建模的一个重要突破,它优雅地融合了 RNN 的效率、CNN 的并行性和 Transformer 的表达能力。随着研究的深入和硬件的优化,SSM 有望在更多场景中发挥重要作用。