状态空间模型(State Space Model, SSM)近年来在长序列建模领域重新受到关注,其线性复杂度的推理优势让它成为Transformer的有力竞争者。然而经典SSM基于线性时不变(LTI)假设,状态转移矩阵和输入投影矩阵在序列所有位置共享同一组参数,导致模型无法根据当前输入的内容动态调整记忆策略。当序列长度达到数万甚至数十万时,早期关键信息会被后续无关 token 逐渐淹没,产生类似RNN的长程遗忘问题。选择性状态空间(Selective SSM)通过让参数成为输入的函数,从根本上改变了信息筛选机制。本文将系统探讨选择性SSM的原理、并行扫描算法的实现思路,以及它们如何协同解决长序列遗忘难题。

我们先从一个直观的视角理解问题本质。传统SSM的离散化形式可以写成:h_t = A * h_{t-1} + B * x_t,y_t = C * h_t。这里 A、B、C 对所有时间步固定,意味着模型无法辨别当前 token 是重要的人名、日期还是无意义的填充词。相比之下,人类阅读长文档时会自动忽略连词和标点,而把注意力集中在专有名词和数字上。选择性SSM正是借鉴了这种机制,将 B、C 甚至步长 Δ 设计为输入 x_t 的函数,让模型学会“什么时候该记住,什么时候该忘记”。这种设计并非简单地增加参数数量,而是让状态更新规则本身具备了上下文感知能力。
从线性时不变到输入依赖:选择性状态空间的数学重构
经典SSM的核心在于线性时不变性质,它允许使用卷积或FFT进行高效计算,但代价是表达能力受限。为了引入选择性,Mamba论文提出了一个优雅的方案:让离散化步长 Δ_t、输入投影矩阵 B_t 和输出矩阵 C_t 都依赖于当前输入。具体来说,对于输入序列 x ∈ R^{B×L×D},我们通过线性层和激活函数生成参数:Δ_t = softplus(Linear_Δ(x_t) + bias_Δ),B_t = Linear_B(x_t),C_t = Linear_C(x_t)。状态转移矩阵 A 可以保持固定但通常设置为可学习参数,因为它主要负责长期衰减模式,而步长 Δ 则控制当前输入对状态的瞬时影响程度。
离散化过程使用零阶保持(ZOH)方法:A_bar_t = exp(Δ_t * A),B_bar_t = (Δ_t * A)^{-1} (exp(Δ_t * A) - I) * Δ_t * B_t。当 Δ_t 较大时,A_bar 趋近于零矩阵,历史状态被大幅遗忘,当前输入的影响增强;当 Δ_t 较小时,A_bar 接近单位矩阵,历史状态得以保留。这种通过连续参数 Δ 来控制有效记忆长度的方法,使得模型能够在不同位置动态调整感受野。例如在处理代码时,遇到函数定义的关键字需要长程记忆,而遇到注释则可以选择性忽略,这完全由输入数据本身的特征驱动。
一个值得注意的实现细节是,A 矩阵通常被参数化为对角矩阵或低秩矩阵。对角化后 A_bar 的计算变成逐元素操作,避免了矩阵指数和求逆的高昂代价。在实际代码中,Mamba 将 A 设置为可学习的实数向量(长度为隐藏维度),并对 Δ 施加 softplus 保证非负性。这样每个隐藏通道拥有独立的衰减速率,而输入依赖的 Δ 则给所有通道一个全局缩放因子,两者结合既保留了参数的灵活性又避免了过拟合风险。
并行扫描:让选择性SSM摆脱递归瓶颈
选择性SSM虽然解决了遗忘问题,但破坏了线性时不变性质,使得原先依赖卷积的高效计算不再适用。如果按递归方式逐时间步计算,时间复杂度为 O(L),这在长序列上虽然优于 Transformer 的 O(L²),但仍然难以充分利用GPU并行性,且无法在训练时通过并行化加速。并行扫描(Parallel Scan)算法给出了答案:它利用结合律将线性递归转化为可以在对数深度内完成的并行操作。
考虑状态更新方程 h_t = A_bar_t * h_{t-1} + B_bar_t * x_t,定义二元操作 o = (a1, b1) ⊕ (a2, b2) = (a2 * a1, a2 * b1 + b2),则可以用一个序列 s_t = (A_bar_t, B_bar_t * x_t) 表示局部变换,整个序列的扫描结果就是按顺序对所有 s_t 做 ⊕ 运算。由于 ⊕ 满足结合律,可以将序列分成两半分别扫描,再合并结果。递归地将序列二分,每层合并的复杂度为 O(L),总层数为 log L,因此总时间复杂度为 O(L log L)。在GPU上,每一层的合并操作可以完全并行,实际运行时间通常接近 O(L) 的量级。
def parallel_scan(a_bar, b_bar_x):
"""
a_bar: shape (L, N) 每个时间步的状态衰减因子
b_bar_x: shape (L, N) 每个时间步的输入投影
返回: h, shape (L, N) 所有时间步的隐藏状态
"""
L, N = a_bar.shape
# 初始化扫描元素为 (a, b) 对
elements = [(a_bar[i], b_bar_x[i]) for i in range(L)]
# 迭代合并,每次将相邻两个元素合并
step = 1
while step < L:
new_elements = []
for i in range(0, L, 2 * step):
if i + step < L:
left_a, left_b = elements[i]
right_a, right_b = elements[i + step]
# 合并操作: (a2*a1, a2*b1 + b2)
merged_a = right_a * left_a
merged_b = right_a * left_b + right_b
new_elements.append((merged_a, merged_b))
else:
new_elements.append(elements[i])
elements = new_elements
step *= 2
# 最后一个元素包含完整的扫描结果,需要展开所有中间状态
# 实际实现中会使用更高效的并行树结构,这里给出简化示意
h = [None] * L
curr_a = 1.0
curr_b = 0.0
# 从后往前提取中间状态(伪代码,实际用树结构避免O(L))
for t in range(L):
# 此处仅为示意,详细实现需反向传播结合律
pass
return h
上述代码仅为概念演示,实际工程中并行扫描采用 Blelloch 扫描或 Hillis-Steele 扫描算法,它们在GPU上通过共享内存和 warp shuffle 实现极高的吞吐量。Mamba 的官方实现使用了自定义 CUDA 内核,将扫描过程分解为多个块内扫描和跨块扫描,在 A100 上处理长度 100K 的序列仅需毫秒级计算时间。反向传播同样利用扫描的结合性,用类似方式计算梯度,避免了逐时间步回传的大量中间状态存储,显存占用与序列长度无关,这使得选择性SSM能够训练超长序列而不会像 Transformer 那样因注意力矩阵导致显存爆炸。
与Transformer及传统RNN的对比:遗忘问题的量化分析
Transformer 通过自注意力机制为每个 token 分配全局感受野,理论上不存在遗忘问题,但它的计算复杂度和显存需求随序列长度平方增长,实际处理长序列时被迫使用窗口注意力或稀疏注意力,间接引入了遗忘。选择性SSM在保持线性复杂度的同时,通过内容感知的状态更新实现了类似注意力的信息筛选能力。实验表明,在语言建模任务上,Mamba(基于选择性SSM)在相同参数量下超越了 Transformer++ 架构,并且在处理长度为 100K 的 DNA 序列和音频波形时表现出稳定的长程依赖捕捉能力。
传统RNN(LSTM/GRU)通过门控机制缓解遗忘,但其状态更新是逐时间步串行的,无法在训练时并行化,且门控参数与输入无关,表达能力受限。选择性SSM可以看作一种广义的门控线性递归,门控信号本身由输入驱动,并且通过并行扫描绕开了递归的串行瓶颈。从信息论角度看,选择性SSM的状态向量可以动态调整保留哪些维度的信息,相当于在每个时间步对记忆空间进行重分配,而 LSTM 只能通过固定门控比例来调节记忆强度。
下表总结了三种架构在长序列建模上的关键差异:
| 特性 | Transformer | LSTM/GRU | 选择性SSM |
|---|---|---|---|
| 长程遗忘问题 | 无(全注意力)但窗口化后存在 | 有(梯度消失/信息覆盖) | 通过输入依赖门控显著缓解 |
| 训练并行性 | 高(注意力矩阵并行) | 低(时间步串行) | 高(并行扫描) |
| 推理复杂度 | O(L²) 或窗口化 O(L·W) | O(L) | O(L) |
| 状态更新规则 | 无显式状态,KV缓存 | 固定门控,输入无关 | 输入依赖,动态调整 |
从实际应用角度看,选择性SSM已经扩展到多模态、强化学习和时间序列预测领域。例如在基因组建模中,DNA 序列长度可达百万碱基对,Transformer 完全无法处理,而选择性SSM可以高效训练并捕获远距离调控元件之间的关联。在音频生成任务中,选择性SSM替代自回归Transformer后,生成速度提升数倍且音质不降。这些案例证明解决遗忘问题不仅是理论上的优雅,更是真实场景的迫切需求。
值得注意的是,选择性SSM并非要完全取代Transformer。许多最新工作将二者结合,例如在选择性SSM的隐藏层之间插入稀疏注意力模块,或者用混合架构同时捕捉局部纹理和全局依赖。这种趋势表明,未来的长序列建模将更多地采用多种机制的有机融合,而选择性状态空间与并行扫描构成的核心技术栈,无疑是其中最具潜力的基础组件之一。