Back to blog

S4: Efficiently Modeling Long Sequences with Structured State Spaces

序列建模的一个核心目标是设计一个单一的、有原则的模型,能够处理跨模态和任务的序列数据,特别是长程依赖(LRDs)。然而,传统模型包括 RNNs、CNNs 和 Transformers 的专门变体仍然难以扩展到 10000 步以上的超长序列。

S4: Efficiently Modeling Long Sequences with Structured State Spaces

一、论文概述

项目内容
标题Efficiently Modeling Long Sequences with Structured State Spaces (S4)
作者Albert Gu, Karan Goel, Christopher Ré
机构Stanford University
论文https://arxiv.org/abs/2111.00396
代码https://github.com/HazyResearch/state-spaces
发布2021-10-31
许可-
领域cs.LG (Machine Learning)

二、核心思想

问题定义

序列建模的一个核心目标是设计一个单一的、有原则的模型,能够处理跨模态和任务的序列数据,特别是长程依赖(LRDs)。然而,传统模型包括 RNNs、CNNs 和 Transformers 的专门变体仍然难以扩展到 10000 步以上的超长序列。

现有方法的局限性:

  1. RNNs:梯度消失/爆炸,难以捕捉长程依赖
  2. CNNs:感受野有限,需要大量层来覆盖长序列
  3. Transformers:O(L²) 复杂度,内存和计算瓶颈
  4. LSSL:理论上可行,但计算和内存需求过高

解决方案概述

S4(Structured State Space sequence model) 基于状态空间模型(SSM)的新参数化:

SSM 基础: x′(t)=Ax(t)+Bu(t)x'(t) = Ax(t) + Bu(t) y(t)=Cx(t)+Du(t)y(t) = Cx(t) + Du(t)

关键创新:

  • 对 A 进行低秩修正,使其可稳定对角化
  • 将 SSM 计算归约为 Cauchy 核计算
  • 高效地在递归和卷积表示之间切换

核心成果

  • 顺序 CIFAR-10 上 91% 准确率(无数据增强),与更大 2-D ResNet 相当
  • 在图像和语言建模任务上大幅缩小与 Transformers 的差距,生成速度快 60×
  • Long Range Arena 每个任务上达到 SOTA,包括解决 16k 长度的 Path-X 任务
  • 与其他竞争方法一样高效

三、技术架构

SSM 概览

状态空间模型(SSM):

  • 参数化:矩阵 A, B, C, D
  • 映射输入信号 u(t) 到输出 y(t) 通过潜在状态 x(t)
  • 可以作为递归或卷积计算

HiPPO 矩阵:

  • 专门设计用于处理长程依赖
  • 连续时间记忆的数学框架
  • 但直接使用计算和内存需求过高

核心公式

SSM 离散化

连续 SSM 通过步长 Δ 离散化: Aˉ=(I−Δ2A)−1(I+Δ2A)\bar{A} = (I - \frac{\Delta}{2}A)^{-1}(I + \frac{\Delta}{2}A) Bˉ=(I−Δ2A)−1ΔB\bar{B} = (I - \frac{\Delta}{2}A)^{-1}\Delta B Cˉ=C\bar{C} = C

递归表示

xk=Aˉxk−1+Bˉukx_k = \bar{A}x_{k-1} + \bar{B}u_k yk=Cˉxky_k = \bar{C}x_k

卷积表示

y=Kˉ∗uy = \bar{K} * u Kˉ=(CˉBˉ,CˉAˉBˉ,...,CˉAˉL−1Bˉ)\bar{K} = (\bar{C}\bar{B}, \bar{C}\bar{A}\bar{B}, ..., \bar{C}\bar{A}^{L-1}\bar{B})

NPLR 参数化

关键定理:所有 HiPPO 矩阵可以分解为 Normal Plus Low-Rank (NPLR) 形式: A=VΛV∗−PQ∗=V(Λ−(V∗P)(V∗Q)∗)V∗A = V\Lambda V^* - PQ^* = V(\Lambda - (V^*P)(V^*Q)^*)V^*

其中 V 是酉矩阵,Λ 是对角矩阵,P, Q 是低秩因子。

S4 卷积核算法

Algorithm 1(概要):

  1. 截断 SSM 生成函数 (SSMGF) 到长度 L
  2. 应用 Woodbury 恒等式处理低秩项
  3. 归约为 Cauchy 核计算
  4. 在单位根上评估 SSMGF
  5. 应用逆 FFT 得到卷积核

计算复杂度:Õ(N + L) 操作和 O(N + L) 空间

复杂度对比

模型参数训练计算训练空间训练并行推理计算
ConvolutionLHÕLH(B+H)BLHYesLH²
RecurrenceH²BLH²BLHNoH²
AttentionH²B(L²H+LH²)B(L²+HL)YesL²H+H²
S4H²BH(ÕH+ÕL)BLHYesH²

S4 的优势:

  • 训练时像卷积一样并行
  • 推理时像递归一样高效
  • 结合了两者的优势

核心组件

组件说明关键参数
SSM状态空间模型A, B, C, D 矩阵
HiPPO长程依赖矩阵专门设计的 A 矩阵
NPLR参数化低秩修正 + 对角化
Cauchy Kernel计算核心高效数值算法
S4 Layer网络层H 个独立 SSM + 线性混合

四、核心创新

创新点说明理论/实验依据
NPLR 参数化低秩修正使 A 可稳定对角化解决 LSSL 数值不稳定性
Cauchy 核归约SSM 计算归约为 Cauchy 核Õ(N+L) 复杂度
递归+卷积统一高效切换两种表示训练并行+推理高效
HiPPO 矩阵应用连续时间记忆理论数学上处理 LRDs
深度 S4 架构类似 depthwise-separable CNN全局卷积核

五、代码实现分析

技术栈

  • 框架:PyTorch
  • 核心库:pykeops(内存高效核操作)
  • GPU:A100 (40GB)
  • 算法:Cauchy 核计算,FFT

关键实现细节

  1. NPLR 参数化:

    • 将 HiPPO 矩阵分解为 VΛV* - PQ*
    • 低秩因子 r=1 或 r=2
    • 支持稳定对角化
  2. Cauchy 核计算:

    • 当前实现使用朴素 O(NL) 算法
    • GPU 上易于并行
    • 使用 pykeops 库
    • 有更快的近线性算法可用
  3. 深度 S4 架构:

    • H 个独立 SSM 副本
    • 位置-wise 线性层混合特征
    • 非线性激活函数
    • 类似 depthwise-separable CNN
  4. HiPPO 参数学习:

    • A 矩阵初始化为 HiPPO 矩阵
    • 学习率降低(最大 0.001)提高稳定性
    • Δ 等参数可学习

六、实验结果

效率基准

S4 vs LSSL:

维度LSSL 训练 (ms)S4 训练 (ms)比率LSSL 内存 (MB)S4 内存 (MB)比率
1289.324.771.9×222.15.342.0×
25620.63.076.7×168512.6133×
512140.74.7529.6×1314033.5392×

S4 vs 高效 Transformers:

方法长度 1024 速度长度 1024 内存长度 4096 速度长度 4096 内存
Transformer1×1×1×1×
Performer1.23×0.43×3.79×0.086×
Linear Trans.1.58×0.37×5.35×0.067×
S41.58×0.43×5.19×0.091×

Long Range Arena 基准

模型ListOpsTextRetrievalImagePathfinderPath-X平均
Transformer36.3764.2757.4642.4471.40~random53.66
Reformer37.2756.1053.4038.0768.50~random50.56
BigBird36.0564.0259.2940.8374.87~random54.17
Performer18.0165.4053.8242.7777.05~random51.18
S4 (original)58.3576.0287.0987.2686.0588.1080.48
S4 (updated)59.6086.8290.9088.6594.2096.3586.09

关键发现:

  • S4 在所有 6 个 LRA 任务上达到 SOTA
  • Path-X 任务(16k 长度):S4 达到 88.10-96.35%,所有其他方法接近随机
  • 平均准确率大幅领先(80.48-86.09 vs 最佳其他 59.37)

顺序 CIFAR-10

  • S4 达到 91% 准确率
  • 无数据增强或辅助损失
  • 与更大 2-D ResNet 相当

图像和语言建模

  • 大幅缩小与 Transformers 的差距
  • 生成速度快 60×
  • 在多种模态上表现有效

Path-X 可视化

任务:判断图像中的标记是否由路径连接

观察:

  • SSM 卷积核可视化为 128×128 图像
  • 第一层:局部特征检测
  • 最后层:全局路径理解
  • 成功解决 16k 长度任务

七、相关工作

状态空间模型

  • SSMs:控制理论基础模型
  • HiPPO:连续时间记忆框架
  • LSSL:Linear State Space Layer(S4 的前身)

长程依赖模型

  • Orthogonal/Lipschitz RNNs:对抗梯度消失
  • Dilated Convolutions:增加感受野
  • 高效 Transformers:降低 O(L²) 复杂度

数值方法

  • Cauchy 核:数值分析中的经典问题
  • Fast Multipole Method (FMM):近线性算法
  • Woodbury 恒等式:低秩修正

八、总结

核心贡献

  1. NPLR 参数化:解决 LSSL 数值不稳定性
  2. Cauchy 核归约:Õ(N+L) 高效计算
  3. 递归+卷积统一:训练并行+推理高效
  4. S4 模型:通用序列建模解决方案
  5. LRA SOTA:包括解决 Path-X 16k 任务

技术影响

  • 序列建模:统一处理多种模态
  • 长程依赖:数学上优雅的解决方案
  • 效率:训练和推理都高效
  • 后续工作:启发 Mamba、S5 等一系列 SSM 研究

局限性

  1. 语言建模:与 Transformers 仍有差距
  2. 实现优化:Cauchy 核有更快算法未实现
  3. 硬件依赖:需要 GPU 优化
  4. 模型规模:实验限于相对较小规模

未来方向

  • 与 Transformers 组合
  • 音频数据预训练和生成
  • 高维数据(图像、视频)扩展
  • 更高效 Cauchy 核实现

九、参考资源