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. 基础理论