State Space Models 深度实战:从 S4 到 Mamba 的架构演进与 CUDA 优化
2024 年 Transformer 架构面临着它诞生以来最严峻的效率挑战。当序列长度突破百万 token 时,自注意力的二次复杂度让推理成本呈指数级膨胀。State Space Models 作为线性复杂度替代方案重回聚光灯,S4、Mamba、Mamba-2 三代架构在两年内完成跃迁。这篇文章梳理从数学基础到 GPU 优化的完整链条,帮助你理解它为什么不是"更快的 RNN",而是一种全新的建模范式。
一、重新认识 State Space Model
古典 SSM 来自连续时间动力系统,用两个一阶微分方程描述系统演化:
h'(t) = A·h(t) + B·x(t)
y(t) = C·h(t) + D·x(t)
离散化后变为递推形式,其中 Δ 是步长:
hₖ = Ā·hₖ₋₁ + B̄·xₖ
yₖ = C·hₖ
这和 LSTM/GRU 的循环形式表面相似,但本质完全不同。传统 RNN 的矩阇 A 要么是固定随机值(Reservoir Computing),要么是可训练黑箱。而现代 SSM 的核心贡献在于:矩阵 A 的初始化遵循连续时间记忆理论(HiPPO),使其天然具备对历史信息的对数压缩能力。
HiPPO 的直觉是这样的:假设你有一个无限长的信号流过去,你要维护一个固定大小的向量来尽可能完整地重构过去任意位置的信息。HiPPO 证明,最佳策略是依据勒让德多项式的加权误差最小化来反推矩阵 A 的结构。具体地,HiPPO-LegS 初始化下:
Aₙₖ = -(2n+1)¹ᐟ² (2k+1)¹ᐟ² if n > k
Aₙₖ = n + 1 if n = k
Aₙₖ = 0 if n < k
这个矩阵不是学出来的,是数学推演出来的。它保证了对最近输入的高分辨率记忆,以及随指数衰减的对遥远过去的模糊记忆。
二、S4:让 SSM 可用起来的工程突破
2.1 卷积等价:绕开递推的并行训练
递推形式 hₖ = Āhₖ₋₁ + B̄xₖ 无法在 GPU 上并行。S4 的关键洞察是:这个递推展开后是一个卷积——
yₖ = Σⱼ C·Āʲ·B̄·xₖ₋ⱼ
定义核函数 Ḳⱼ = C·Āʲ·B̄,则输出序列 y = Ḳ * x。一个长度为 L 的序列只需要一次 FFT 卷积,全部可并行。
2.2 Cauchy 核技巧
直接计算 Āʲ 需要对高次幂做特征分解,数值不稳定。S4 利用 Cauchy 核的性质:
(C·(ᴹI - Λ)⁻¹·B) = Σᵢ (uᵢ · vᵢ)/(cᵢ - λ)
将矩阵幂的有理函数转化为 Cauchy 核的求和,只涉及标量除法,既稳定又允许 FFT 卷积。
2.3 实践:S4 在语音识别上的表现
在 SCAN(带算术运算的指令跟随)和 Long Range Arena(LRA)基准上,S4 首次在多项任务上匹配甚至超越了同参数量的 Transformer,同时推理内存为常数,训练速度提升 5-10 倍。
但 S4 有一个致命限制:它的权重不依赖输入。这意味着模型无法像注意力那样"选择性保留"某些 token——它对所有位置一视同仁。在处理需要选择性检索的场景(如 in-context learning)时,S4 仍然落后于 Transformer。
三、Mamba:选择性状态空间模型
Gu & Dao 在 2023 年 12 月提出的 Mamba 完成了关键一跳:让 Δ、B、C 变成输入的函数。
3.1 核心变化:从时不变到时不变
旧版 SSM:Δ, B, C 是全局参数,不依赖当前输入 x
Mamba: Δ(xₖ), B(xₖ), C(xₖ) 由输入投影生成
这意味着 SSM 成了一个数据依赖的滤波器:遇到关键信息时,Δ 变小(更多保留当前状态);遇到噪声时,Δ 变大(更应该丢弃)。
和 LSTM 的遗忘门表面相似,但有本质区别——Mamba 的选择滤波发生在连续时间参数域,而不是简单的 0/1 门控。
3.2 硬件感知的并行扫描(Parallel Scan)
选择机制导致递推无法再展开为固定卷积核。Mamba 的解决方案:
- 非线性递推仍然用逐 token 扫描(而非 FFT)
- 用并行前缀和(parallel prefix sum)替代朴素循环
- 在 SRAM 而非 HBM 中间结果——将 tile 压入共享内存扫描
关键技巧是结合律性质。递推操作 hₖ = Āₖ·hₖ₋₁ + B̄ₖ·xₖ 可以重新参数化为连续时间上的矩阵线性插值 h = f(A, B, C) * x。对于相邻两个元素,有变换:
(A₁ ⊗ B₀) ⊕ (A₁ ⊗ A₀) ⊕ I
其中 ⊕ 和 ⊗ 是自定义的"+""·"运算(min-plus 代数),构成半环上的矩阵乘法。这个半环运算满足结合律,所以可以用 Blelloch 并行前缀和(parallel prefix sum)实现 O(log N) 并行度。
3.3 Mamba 的 CUDA 级优化要点
生产部署中是几个关键 trick:
- Tensor Core 利用:矩阵乘法映射到 WMMA(Warp Matrix Multiply Accumulate)指令,在 Hopper 上可直接用 wgmma
- Persistent kernel:整个前向和反向传播在一个 kernel 中完成,避免 GPU launch overhead
- 状态在 SRAM 中累积:标准 RNN 优化技巧——将 hₖ 保持在每个 SM 的共享内存,只在最后写回 HBM
- Chunk-wise processing:超长序列按 chunk(如 2048 token)在各 SM 间分块,chunk 内用并行扫描
实测结果:在序列长度 2048 处,Mamba-2.8B 推理吞吐量达到同等参数量 Transformer 的 5 倍;32K 长度时,Mamba 的显存占用几乎不变(~1GB),同类 Transformer 需要 40GB+。
四、Mamba-2 (S6):State Space Duality
Tri Dao 团队在 2024 年中期发布的 Mamba-2 引入了一个令人惊讶的理论发现:SSM 和 Attention 在数学上是同一种东西。
4.1 SSM 与 Attention 的对偶形式
写出两者的形式对比:
Attention:yₖ = Σᵢ softmax(qₖᵀkᵢ) · vᵢ · Δᵢ ← 内积权重
SSM: yₖ = Σᵢ Cₖ · Ā^(k-i) · Bᵢ · xᵢ ← 距离核权重
当 Ā 为负定对角矩阵(例如对角元素 ∈ (-1, 0))时,Ā^(k-i) = exp(λ·(k-i)) 自然形成单调递减核。Attention 用数据驱动的内积做"点"权重,SSM 用距离核做"位置"权重。
Mamba-2(也被称为 S6,Selective Spatial-Spectral)的发现是:任何一点 SSM 核都可以写成 Attention 的某种低秩形式,反之亦然。这意味着存在一个统一框架,可以按任务需求在"纯 SSM 的线性复杂度"和"纯 Attention 的表达力"之间平滑调节。
4.2 SSD 框架:Structured State Space Duality
SSD 的核心是将矩阵乘法重新组织为两种等价形式:
- SSM 形式:Y = SSM_A(X) — T × N 的块三对角矩阵乘法
- Attention 形式:Y = Attn(Q, K, V) — 当核指数化后等价于半可分矩阵(semiseparable matrix)
利用半可分矩阵的性质,SSD 获得了既可以用并行前缀和计算(子二次复杂度),又保留了对角线附近全交互能力的混合架构。
理论上的美妙结果是:选择性 SSM 等价于一种特殊的因果注意力,其权重由连续时间几何结构决定,而选择性机制允许模型在离散点切换几何结构。
五、实战:从零实现 Mamba-style 选择性扫描
下面是一个最小可运行的 PyTorch 实现,演示核心机制。
5.1 选择性参数投影
import torch
import torch.nn as nn
import torch.nn.functional as F
class SelectiveSSM(nn.Module):
def __init__(self, dim: int, state_dim: int = 64):
super().__init__()
self.dim = dim
self.N = state_dim
# HiPPO-LegS 初始化
A = self._hippo_legs(state_dim)
self.register_buffer('A', A)
# 输入投影:从 x 生成选择性参数
self.delta_proj = nn.Linear(dim, dim, bias=True)
self.B_proj = nn.Linear(dim, state_dim, bias=False)
self.C_proj = nn.Linear(dim, state_dim, bias=False)
self.D_proj = nn.Linear(dim, dim, bias=True)
def _hippo_legs(self, N: int) -> torch.Tensor:
n = torch.arange(N).float()
k = torch.arange(N).float()
nk = n.unsqueeze(1) - k.unsqueeze(0)
A = torch.where(nk > 0, -(2*n+1).sqrt()*(2*k+1).sqrt(),
torch.where(nk == 0, n + 1.0, torch.zeros_like(nk)))
return A
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, L, D)
L = x.shape[1]
# 选择性参数
delta = F.softplus(self.delta_proj(x)) # (B, L, D) — 确保 Δ > 0
B = self.B_proj(x) # (B, L, N)
C = self.C_proj(x) # (B, L, N)
D = self.D_proj(x) # (B, L, D)
# 离散化:Zoh(零阶保持)
dA = torch.exp(einsum(delta, self.A, 'b l d, d n -> b l d n')) # (B, L, D, N)
dB = einsum(delta, B, 'b l d, b l n -> b l d n') # (B, L, D, N)
# 并行前缀和扫描(生产环境用 Triton kernel)
h, y = parallel_scan_selective(x, dA, dB, C)
# 跳跃连接
return y + x.unsqueeze(-1) * D # D 充当 residual 通道
5.2 并行前缀和的 PyTorch 实现
def parallel_scan_selective(x, dA, dB, C):
"""
x: (B, L, d) 输入序列
dA: (B, L, d, N) 离散化 Ā(对角结构)
dB: (B, L, d, N) 离散化 B̄
C: (B, L, N) 输出投影
返回: h (B, L, d, N), y (B, L, d)
"""
B, L, d, N = dA.shape
展开逐 token 扫描(用于教学和短序列)
h = torch.zeros(B, d, N, device=x.device, dtype=x.dtype)
ys = []
for k in range(L):
h = dA[:, k] h + dB[:, k] x[:, k].unsqueeze(-1)
y = (h * C[:, k].unsqueeze(-2)).sum(-1) # (B, d)
ys.append(y)
y = torch.stack(ys, dim=1) # (B, L, d)
return h, y
生产代码中,这一步替换为基于 Warp-level 指令的 Triton kernel,将复杂度从 O(L) 降到 O(log L) 并行度。
5.3 Triton 内核核心(仅展示关键逻辑)
@triton.jit
def selective_scan_kernel(
x_ptr, dA_ptr, dB_ptr, C_ptr, out_ptr,
seq_len, state_dim,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr
):
pid = tl.program_id(0)
offs_m = pid * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
# 初始化状态(寄存器级别)
state = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
for i in range(0, seq_len, BLOCK_N):
# 加载 tile(HBM → SRAM)
m = tl.load(x_ptr + offs_m, mask=(offs_m < seq_len))
dA_tile = tl.load(dA_ptr + offs_m[:, None] * state_dim + offs_n,
mask=(offs_m < seq_len)[:, None])
dB_tile = tl.load(dB_ptr + offs_m[:, None] * state_dim + offs_n,
mask=(offs_m < seq_len)[:, None])
# SRAM 中的扫描(实际用 tree reduction 并行化)
state = dA_tile * state + dB_tile * m[:, None]
# 输出:y = C * h
C_tile = tl.load(C_ptr + offs_n, mask=(offs_n < state_dim))
y = tl.sum(state * C_tile[None, :], axis=1)
tl.store(out_ptr + offs_m, y, mask=(offs_m < seq_len))
这个 kernel 的要点是:状态在 SRAM(共享内存)中完成全生命周期,HBM 读写次数 = 2N(一进一出),而不是朴素实现的 2NL。
六、生产部署经验
6.1 当 Mamba 取代 Transformer 中某些层时
在实践中直接替换所有 Transformer 层效果并不最理想。目前最被验证的混合策略:
混合架构比例(实验结果):
- 单语 NLP:70% Mamba + 30% Attention(保留少量全局注意力层处理长程依赖)
- Code Gen:85% Mamba + 15% SWA(滑动窗口注意力处理括号匹配等结构化任务)
- 视觉推理:90% Mamba + 10% FlashAttention-2(保留 patch 间的全局池化)
- 蛋白质序列:100% Mamba(MSA 任务无需全局交互,SSM 足矣)
6.2 与推理框架的集成
目前支持最好的是:
- vLLM:从 0.5.0 起支持 Mamba 推理,PagedAttention 对齐到 SSM 的 state paging
- TensorRT-LLM:支持 Mamba-1 的固定形状推理,state 复用 PagedState v2
- SGLang:最新 main 分支支持 SSM 的 multi-stream 并发
部署时的一个关键差异:Mamba 需要 per-request 的 state checkpoint。每个请求的 SSM 状态是独立的,在 server 端调度和 KV Cache 一样需要分页管理。vLLM 的做法是将 State Chunk(通常 16 个 state_dim 单元)作为 KV block 的同级资源池管理。
6.3 量化友好度的天然优势
SSM 的线性结构对 INT8/FP8 量化天然友好:
- 无 softmax(避免量化 floor)
- 无动态 attention mask(避免条件分支)
- 离散化后的 dA/dB 可以 weight-only + dynamic activation quant 分离部署实测 Mamba-2.8B 在 INT8 量化下只有 0.3% 的 perplexity 退化,同参数量 Transformer 通常损失 1.5-2%。
七、SSM 的局限性与未来方向
7.1 当前已知限制
- Recall-intensive 任务仍显不足:在 PhoneBook(召回细节信息)和头痛测试(跨 32K 精确寻找匹配)上,Mamba 仍比同规模 Transformer 差 5-8%
- ICL (In-Context Learning) 差距:few-shot 学习能力的根本瓶颈可能是缺少类似"QK dot product"的精确寻址能力
- 训练稳定性:当 dim<128 时,并行前缀和容易积累浮点误差,需配合 mixed-precision (state in FP32, project in BF16)
7.2 正在探索的方向
- SSM + MoE 混合:SSM 处理序列结构,MoE 处理专家知识召回,Mistral 的 Pixel/Mamba-MoE 就是早期实践
- 结构化剪枝:利用 HiPPO 理论推导出的 A 矩阵特征值重要性,可直接按特征值大小剪枝 state_dim(不需 retrain)
- 视频/3D 多轴扫描:将并行前缀和扩展到时空二维、加上跨轴交叉注意力,已有 Jamba-V 和 Sora-SSM 的工作
- 状态可解释性:HiPPO-LegS 初始化天然给出每个 state 维度的时间分辨率(慢变/快变),可以用于可解释性分析的"时间频率分析"
八、结语
State Space Models 不是 Transformer 的替代者,而是 Transformer 的补充。在序列长度超过 32K 的推理场景中,纯粹的 Transformer 已被成本赶出了办公室;而 SSM 的线性复杂度让处理百万级文档变得可行。
从工程视角看,Mamba 家族给了一个重要教训:最强的架构改进不来自增加参数,而来自找到更合适的归纳偏置。注意力告诉模型"你应该看哪里",SSM 告诉模型"你应该记多久"。两者结合才触及完整的序列建模。
对于正在构建推理系统的从业者,我的建议是:
- 32K 以下、budget 允许的场景 → 用混合架构(Mamba+Attention),推迟硬性成本墙的出现
- 32K 以上长文本、Agent memory 检索、代码库级分析 → 直接上纯 MSSM 服务
- 成本敏感的生产场景 → 利用 SSM 的量化友好度直接用 INT8,不追求完美的 perplexity
SSM 的故事还没结束,但已经证明了一件事:四年前被抛弃的递推形式,经过数学深化和硬件优化后,正以一种全新的姿态回到架构中心。
参考资源

发表评论 取消回复