PyTorch新手也能懂:手把手拆解Mamba-minimal中的selective_scan实现

在深度学习领域,状态空间模型(State Space Models, SSM)正逐渐成为处理序列数据的新范式。而Mamba作为SSM家族的最新成员,凭借其选择性扫描机制(selective_scan)在长序列建模任务中展现出惊人潜力。本文将带您深入Mamba-minimal实现中最核心的 selective_scan 函数,用PyTorch初学者的视角逐行解析这个看似复杂实则精妙的设计。

1. 理解状态空间模型的基础

状态空间模型本质上描述了一个动态系统的演变过程。在离散时间步中,系统状态$x_k$和输出$y_k$由以下方程决定:

x_k = A * x_{k-1} + B * u_k
y_k = C * x_k + D * u_k

其中:

  • $A$是状态转移矩阵
  • $B$是输入矩阵
  • $C$是输出矩阵
  • $D$是前馈矩阵
  • $u_k$是当前输入

Mamba的创新之处在于让这些矩阵参数 动态依赖于输入数据 ,而传统SSM(如S4)使用固定参数。这种数据依赖性使得模型能够根据输入内容自适应地调整状态转移方式。

2. selective_scan的输入参数解析

让我们先看看 selective_scan 函数的完整签名:

def selective_scan(self, u, delta, A, B, C, D):

各参数含义及维度如下表所示:

参数 维度 说明
u (b, l, d_in) 输入序列,b为batch大小,l为序列长度
delta (b, l, d_in) 数据依赖的时间步长参数
A (d_in, n) 状态转移矩阵
B (b, l, n) 输入矩阵(数据依赖)
C (b, l, n) 输出矩阵(数据依赖)
D (d_in) 前馈矩阵

关键点在于:

  • 与传统SSM不同,B和C矩阵 每个时间步都有不同值
  • delta参数控制着离散化的时间步长,也是输入依赖的

3. 离散化过程详解

Mamba采用两种离散化方法的组合:

  1. **零阶保持(ZOH)**用于状态矩阵A
  2. 前向欧拉 用于输入矩阵B

对应的离散化公式实现如下:

deltaA = torch.exp(einsum(delta, A, 'b l d_in, d_in n -> b l d_in n'))
deltaB_u = einsum(delta, B, u, 'b l d_in, b l n, b l d_in -> b l d_in n')

这里使用了爱因斯坦求和约定(einsum)进行高效张量运算。让我们拆解这两行代码:

3.1 状态矩阵的ZOH离散化

deltaA 的计算对应ZOH离散化中的$e^{ΔA}$项:

  • delta 形状:(b, l, d_in)
  • A 形状:(d_in, n)
  • 通过einsum在d_in维度上做乘法,然后取指数

3.2 输入矩阵的欧拉离散化

deltaB_u 对应欧拉离散化的$ΔB·u$项:

  • 将delta、B和u三个张量在多个维度上相乘
  • 结果形状为(b, l, d_in, n)

提示:einsum操作虽然高效,但对初学者可能不太直观。可以想象它是在指定维度上进行"乘法+求和"的组合操作。

4. 选择性扫描的核心循环

真正的扫描过程由一个简单的for循环实现:

x = torch.zeros((b, d_in, n), device=deltaA.device)
ys = []
for i in range(l):
    x = deltaA[:, i] * x + deltaB_u[:, i]  # 状态更新
    y = einsum(x, C[:, i, :], 'b d_in n, b n -> b d_in')  # 输出计算
    ys.append(y)

这个循环实现了以下功能:

  1. 初始化状态x为零
  2. 对每个时间步i:
    • 更新状态:$x = e^{ΔA}x + ΔB·u$
    • 计算输出:$y = Cx$
  3. 收集所有时间步的输出

值得注意的是,这与原始论文的并行实现不同:

  • 这里使用 顺序扫描 ,时间复杂度O(l)
  • 论文使用 并行算法 ,时间复杂度O(log l)
  • 教学实现更易理解,但实际应用应使用CUDA优化版本

5. 完整流程与输出处理

扫描完成后,我们需要处理输出结果:

y = torch.stack(ys, dim=1)  # 将列表转为张量 (b, l, d_in)
y = y + u * D  # 添加前馈连接
return y

最后一步 y + u * D 体现了状态空间模型的完整输出方程:

  • $y = Cx + Du$
  • D矩阵提供了从输入到输出的直接路径
  • 这种跳跃连接有助于梯度传播

6. 与理论公式的对照

让我们将代码与离散化理论公式做对比:

前向欧拉离散化(用于B矩阵)

x_k = (I + Δ_k A)x_{k-1} + Δ_k B u_k

零阶保持离散化(用于A矩阵)

x_k = e^{Δ_k A}x_{k-1} + (Δ_k A)^{-1}(e^{Δ_k A} - I)Δ_k B u_k

在Mamba-minimal实现中:

  • 对A矩阵采用完整ZOH离散化(包含指数项)
  • 对B矩阵做了简化,相当于只保留一阶泰勒展开
  • 这种混合策略在效果和效率间取得了平衡

7. 实际应用中的注意事项

在您自己的项目中实现或修改selective_scan时,需要注意:

  1. 数值稳定性

    • 指数运算可能导致数值爆炸
    • 实际实现可能需要添加归一化
  2. 初始化策略

    • A矩阵初始化为对数空间
    • 使用 softplus 确保delta为正
  3. 性能考量

    • 序列较长时,顺序扫描会成为瓶颈
    • 实际应用应参考论文的并行实现

以下是一个简化的初始化示例:

# A矩阵的初始化(对数空间)
A = repeat(torch.arange(1, args.d_state + 1), 'n -> d n', d=args.d_inner)
self.A_log = nn.Parameter(torch.log(A))

# delta的处理
delta = F.softplus(self.dt_proj(delta))  # 确保为正

理解selective_scan的实现是掌握Mamba架构的关键。虽然这个简化版本牺牲了部分效率,但它清晰地展现了选择性状态空间的核心思想——通过输入相关的参数实现动态序列建模。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐