跳转至

SSM 与混合架构:Mamba、Jamba、RWKV

更新日期:2026-04-14

本文目标:理解 SSM 与 Transformer 的本质区别,为什么混合架构有价值,以及工程上怎么选。


一、为什么需要 Transformer 的替代

Transformer 的问题 影响 SSM 如何解决
注意力 O(n²) 复杂度 长序列训练/推理极慢 O(n) 线性复杂度
KV Cache 线性增长 长序列推理内存爆炸 固定大小状态,恒定内存
推理速度随上下文增长 每个新 token 都要扫描全部 KV 恒定推理速度
SSM 的代价 影响
无法精确回忆远距离细节 "大海捞针"类任务差
训练吞吐量不如 FlashAttention 优化的 Transformer 实际训练可能更慢
生态不成熟 工具链、预训练权重少

二、Mamba 核心原理

2.1 从连续 SSM 到离散 Mamba

# 连续状态空间模型 (线性 ODE):
# dh/dt = A·h + B·x    (状态更新)
# y = C·h + D·x        (观测/输出)
# A: 状态转移矩阵 [N, N], N=状态维度
# B: 输入矩阵 [N, 1]
# C: 输出矩阵 [1, N]

# 离散化 (Zero-Order Hold):
# h_t = A_bar · h_{t-1} + B_bar · x_t
# y_t = C · h_t
# 其中 A_bar = exp(Δ·A), B_bar = (A_bar - I)·A^{-1}·B
# Δ: 步长参数

# Mamba 的关键创新: 让 B, C, Δ 依赖输入!
# → "选择性" SSM: 模型可以选择记住/忘记什么

2.2 完整实现

class MambaBlock:
    def __init__(self, d_model, d_state=16, d_conv=4, expand=2):
        self.d_inner = d_model * expand

        # 输入投影
        self.in_proj = Linear(d_model, 2 * self.d_inner)  # x → (z, x_proj)

        # 1D 卷积 (局部上下文)
        self.conv = Conv1d(self.d_inner, self.d_inner, d_conv, groups=self.d_inner)

        # SSM 参数 (从输入动态生成)
        self.x_proj = Linear(self.d_inner, d_state + d_state + 1)  # → B, C, Δ

        # 固定参数 A (初始化为对角负矩阵)
        A = -exp(arange(1, d_state+1).float())  # 负值 → 衰减
        self.A_log = Parameter(log(-A))

        self.D = Parameter(ones(self.d_inner))

        # 输出投影
        self.out_proj = Linear(self.d_inner, d_model)

    def forward(self, x):
        B, S, D = x.shape

        # 输入投影: 分为 x 和 gate z
        xz = self.in_proj(x)  # [B, S, 2*d_inner]
        x_proj, z = xz.chunk(2, dim=-1)

        # 1D 卷积 (捕捉局部上下文)
        x_conv = self.conv(x_proj.transpose(1,2)).transpose(1,2)  # [B, S, d_inner]
        x_conv = silu(x_conv)

        # 从输入动态生成 SSM 参数
        ssm_params = self.x_proj(x_conv)  # [B, S, 2N+1]
        B_t = ssm_params[:, :, :self.d_state]       # [B, S, N]
        C_t = ssm_params[:, :, self.d_state:2*self.d_state]  # [B, S, N]
        delta = softplus(ssm_params[:, :, -1])       # [B, S] — 步长, > 0

        A = -exp(self.A_log)  # [d_inner, N]

        # === 核心: 选择性扫描 (selective scan) ===
        y = selective_scan(x_conv, delta, A, B_t, C_t)

        # 门控 + 输出
        y = y * silu(z)  # gate
        return self.out_proj(y)

def selective_scan(x, delta, A, B, C):
    # x: [B, S, d_inner]
    # 递推计算
    B_size, S, d_inner = x.shape
    N = A.shape[1]  # 状态维度

    h = zeros(B_size, d_inner, N)  # 隐状态
    outputs = []

    for t in range(S):
        # 离散化 A
        dA = exp(delta[:, t, :, None] * A)  # [B, d_inner, N]
        dB = delta[:, t, :, None] * B[:, t, None, :]  # [B, d_inner, N]

        # 状态更新
        h = dA  h + dB  x[:, t, :, None]  # [B, d_inner, N]

        # 读出
        y_t = (h * C[:, t, None, :]).sum(-1)  # [B, d_inner]
        outputs.append(y_t)

    return stack(outputs, dim=1)  # [B, S, d_inner]

    # 训练时: 用并行扫描 (parallel scan) 替代 for 循环 → GPU 友好
    # 推理时: 用递推模式 → O(1) per token, 恒定速度

2.3 选择性机制的直觉

非选择性 (传统 SSM): A, B, C 对所有 token 相同 → 模型无法根据内容决定记忆/遗忘 → 等价于线性时不变系统

选择性 (Mamba): B, C, Δ 由当前输入决定 Δ 大 → exp(Δ·A) 接近 0 → 遗忘旧状态, 记忆新输入 Δ 小 → exp(Δ·A) 接近 1 → 保持旧状态, 忽略新输入

直觉: 看到重要信息时 Δ 增大 → "记住这个" 看到不重要信息时 Δ 很小 → "跳过这个"


三、RWKV

对比 Mamba RWKV
基础 状态空间模型 RNN + 线性注意力
核心操作 选择性扫描 WKV (Weighted Key-Value) 机制
训练并行化 并行扫描 Chunk-wise 并行
状态维度 d_inner × d_state d_model × d_model
最大模型 Mamba-2 (研究) RWKV-6 14B (实际部署)
社区 学术为主 有活跃开源社区

四、混合架构

4.1 Jamba (AI21, 2024)

class JambaModel:
    def __init__(self, n_layers=80):
        self.layers = []
        for i in range(n_layers):
            if i % 8 in [0, 4]:
                # 每 8 层中 2 层用 Transformer (精确检索)
                if i % 16 == 0:
                    self.layers.append(TransformerBlock(use_moe=True))   # + MoE
                else:
                    self.layers.append(TransformerBlock(use_moe=False))  # Dense
            else:
                # 其余 6 层用 Mamba (高效序列建模)
                self.layers.append(MambaBlock())

    # 效果:
    # - 256K 上下文窗口 (单 A100 80GB!)
    # - 质量接近纯 Transformer
    # - 推理吞吐量更高 (Mamba 层无 KV Cache)

4.2 为什么混合比纯 SSM 好

4.3 混合比例选择


五、追问延伸

问题 方向
Mamba-2 有什么改进? SSD (Structured State-space Duality): 将 SSM 和注意力统一
混合架构的 Scaling Law 是什么样? 研究较少, 但初步显示与纯 Transformer 类似
SSM 能替代 Transformer 吗? 短期不能, 但混合使用越来越常见
训练混合架构用什么框架? Mamba 有 PyTorch 实现, 但 Megatron 支持有限

参考链接

SSM 与混合架构:Mamba、Jamba、RWKV

更新日期:2026-04-14

本文目标:理解 SSM 与 Transformer 的本质区别,为什么混合架构有价值,以及工程上怎么选。


一、为什么需要 Transformer 的替代

Transformer 的问题 影响 SSM 如何解决
注意力 O(n²) 复杂度 长序列训练/推理极慢 O(n) 线性复杂度
KV Cache 线性增长 长序列推理内存爆炸 固定大小状态,恒定内存
推理速度随上下文增长 每个新 token 都要扫描全部 KV 恒定推理速度
SSM 的代价 影响
无法精确回忆远距离细节 "大海捞针"类任务差
训练吞吐量不如 FlashAttention 优化的 Transformer 实际训练可能更慢
生态不成熟 工具链、预训练权重少

二、Mamba 核心原理

2.1 从连续 SSM 到离散 Mamba

连续状态空间模型(线性 ODE):

  • 状态更新:dh/dt = A·h + B·x

  • 观测/输出:y = C·h + D·x

  • 其中 A: [N, N] 状态转移矩阵,B: [N, 1] 输入矩阵,C: [1, N] 输出矩阵,N 为状态维度

离散化(Zero-Order Hold):

  • h_t = A_bar · h_{t-1} + B_bar · x_t

  • y_t = C · h_t

  • 其中 A_bar = exp(Delta·A)B_bar = (A_bar - I)·A^{-1}·B,Delta 为步长参数

Mamba 的关键创新:让 B、C、Delta 依赖输入(而非固定参数),从而实现"选择性"SSM——模型可以根据内容决定记住还是忘记。

2.2 完整实现

class MambaBlock:
    def __init__(self, d_model, d_state=16, d_conv=4, expand=2):
        self.d_inner = d_model * expand

        # 输入投影
        self.in_proj = Linear(d_model, 2 * self.d_inner)  # x → (z, x_proj)

        # 1D 卷积 (局部上下文)
        self.conv = Conv1d(self.d_inner, self.d_inner, d_conv, groups=self.d_inner)

        # SSM 参数 (从输入动态生成)
        self.x_proj = Linear(self.d_inner, d_state + d_state + 1)  # → B, C, Δ

        # 固定参数 A (初始化为对角负矩阵)
        A = -exp(arange(1, d_state+1).float())  # 负值 → 衰减
        self.A_log = Parameter(log(-A))

        self.D = Parameter(ones(self.d_inner))

        # 输出投影
        self.out_proj = Linear(self.d_inner, d_model)

    def forward(self, x):
        B, S, D = x.shape

        # 输入投影: 分为 x 和 gate z
        xz = self.in_proj(x)  # [B, S, 2*d_inner]
        x_proj, z = xz.chunk(2, dim=-1)

        # 1D 卷积 (捕捉局部上下文)
        x_conv = self.conv(x_proj.transpose(1,2)).transpose(1,2)  # [B, S, d_inner]
        x_conv = silu(x_conv)

        # 从输入动态生成 SSM 参数
        ssm_params = self.x_proj(x_conv)  # [B, S, 2N+1]
        B_t = ssm_params[:, :, :self.d_state]       # [B, S, N]
        C_t = ssm_params[:, :, self.d_state:2*self.d_state]  # [B, S, N]
        delta = softplus(ssm_params[:, :, -1])       # [B, S] — 步长, > 0

        A = -exp(self.A_log)  # [d_inner, N]

        # === 核心: 选择性扫描 (selective scan) ===
        y = selective_scan(x_conv, delta, A, B_t, C_t)

        # 门控 + 输出
        y = y * silu(z)  # gate
        return self.out_proj(y)

def selective_scan(x, delta, A, B, C):
    # x: [B, S, d_inner]
    # 递推计算
    B_size, S, d_inner = x.shape
    N = A.shape[1]  # 状态维度

    h = zeros(B_size, d_inner, N)  # 隐状态
    outputs = []

    for t in range(S):
        # 离散化 A
        dA = exp(delta[:, t, :, None] * A)  # [B, d_inner, N]
        dB = delta[:, t, :, None] * B[:, t, None, :]  # [B, d_inner, N]

        # 状态更新
        h = dA  h + dB  x[:, t, :, None]  # [B, d_inner, N]

        # 读出
        y_t = (h * C[:, t, None, :]).sum(-1)  # [B, d_inner]
        outputs.append(y_t)

    return stack(outputs, dim=1)  # [B, S, d_inner]

    # 训练时: 用并行扫描 (parallel scan) 替代 for 循环 → GPU 友好
    # 推理时: 用递推模式 → O(1) per token, 恒定速度

2.3 选择性机制的直觉

选择性机制的直觉——Delta 控制遗忘与记忆的平衡:

  • Delta 大exp(Delta·A) 接近 0 → 遗忘旧状态,记忆新输入("记住这个")

  • Delta 小exp(Delta·A) 接近 1 → 保持旧状态,忽略新输入("跳过这个")


三、RWKV

对比 Mamba RWKV
基础 状态空间模型 RNN + 线性注意力
核心操作 选择性扫描 WKV (Weighted Key-Value) 机制
训练并行化 并行扫描 Chunk-wise 并行
状态维度 d_inner × d_state d_model × d_model
最大模型 Mamba-2 (研究) RWKV-6 14B (实际部署)
社区 学术为主 有活跃开源社区

四、混合架构

4.1 Jamba (AI21, 2024)

class JambaModel:
    def __init__(self, n_layers=80):
        self.layers = []
        for i in range(n_layers):
            if i % 8 in [0, 4]:
                # 每 8 层中 2 层用 Transformer (精确检索)
                if i % 16 == 0:
                    self.layers.append(TransformerBlock(use_moe=True))   # + MoE
                else:
                    self.layers.append(TransformerBlock(use_moe=False))  # Dense
            else:
                # 其余 6 层用 Mamba (高效序列建模)
                self.layers.append(MambaBlock())

    # 效果:
    # - 256K 上下文窗口 (单 A100 80GB!)
    # - 质量接近纯 Transformer
    # - 推理吞吐量更高 (Mamba 层无 KV Cache)

4.2 为什么混合比纯 SSM 好

4.3 混合比例选择


五、追问延伸

问题 方向
Mamba-2 有什么改进? SSD (Structured State-space Duality): 将 SSM 和注意力统一
混合架构的 Scaling Law 是什么样? 研究较少, 但初步显示与纯 Transformer 类似
SSM 能替代 Transformer 吗? 短期不能, 但混合使用越来越常见
训练混合架构用什么框架? Mamba 有 PyTorch 实现, 但 Megatron 支持有限

参考链接


上级 · A. 基础理论