跳转至

Linear / RNN 路线:砍 softmax,固定 state 替代 KV

更新日期:2026-04-27

核心原理:把 attention 的 \(\text{softmax}(QK^\top) V\) 改成 递推形式:

\[ S_t = f(S_{t-1}, k_t, v_t),\qquad o_t = g(q_t, S_t) \]

其中 \(S \in \mathbb{R}^{d \times d}\)(或 per-channel 向量)是固定大小 state,把所有历史压进去。KV cache 不再存 —— 只存当前 \(S\)

本文专注原理、公式与复杂度。所有变体的 PyTorch / Triton 实测见 A2-1.1 attention kernel,bench 数字(softmax decode vs GLA / Mamba-2 / KDA / Lightning)也在那一篇。


一、两条同源路线(看公式就懂)

历史上 从两个不同入口 走到同一个数学对象:

起点 路线 思想 代表演化
Attention 视角 Pure Linear Attention 把 softmax 拆掉,利用 \((QK^\top)V = Q(K^\top V)\) 结合律换矩阵乘顺序 Katharopoulos → GLA → DeltaNet → KDA → Lightning
控制论视角 SSM (State Space Model) 把 attention 看成 RNN,用状态空间方程 \(\dot h = Ah + Bx,\, y = Ch\) 描述 S4 → S5 → H3 → Mamba (S6) → Mamba-2
flowchart LR
    subgraph LA["Pure Linear Attention(attention 视角)"]
        L1["1. Linear Attn<br/>无 gate"] -->|加 forget gate| L2["2. GLA"]
        L2 -->|+ delta rule + DPLR| L3["3. DeltaNet / KDA"]
        L2 -->|改归一化方式| L4["4. Lightning<br/>(transnormer)"]
    end
    subgraph SS["SSM(控制论视角)"]
        S1["1. S4<br/>HiPPO + conv"] -->|MIMO + scan| S2["2. S5"]
        S2 -->|+ 乘法 gate| S3["3. H3"]
        S3 -->|selective + CUDA| S4["4. Mamba (S6)"]
        S4 -->|SSD framework| S5["5. Mamba-2"]
    end
    SSD(("Mamba-2 SSD 定理<br/>SSM = Linear Attention 同构"))
    L3 -.-> SSD
    L4 -.-> SSD
    S5 -.-> SSD

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    classDef tag fill:#f5f3eb,stroke:#bdb9ab,color:#1a1a1a;
    class L1,L2,L3,L4,S1,S2,S3,S4,S5 stage
    class SSD tag

收敛点Mamba-2 (Dao & Gu, ICML 2024) 证明了带特定结构的 SSM 状态转移 ≡ 带特定 mask 的 linear attention。两条线殊途同归。

根本约束(两条线都逃不掉):state 容量 = \(O(d^2)\) floats,长 ctx 上"recall 任意早期 token"必然弱于 softmax。这不是工程问题,是表达力上限。所以 frontier 全用 hybrid(linear 担成本、softmax 担精度)。


二、Pure Linear Attention 路线(attention 视角)

2.1 朴素 Linear Attention — 起点(Katharopoulos+ 2020)

参考 Transformers are RNNs (Katharopoulos et al., ICML 2020)

公式

把 softmax 换成 kernel feature map \(\phi\)(典型 \(\phi(x) = \text{ELU}(x) + 1\),保非负):

\[ \text{Attn}(Q, K, V)_t \;=\; \frac{\phi(q_t)^\top \sum_{\tau \le t} \phi(k_\tau) v_\tau^\top}{\phi(q_t)^\top \sum_{\tau \le t} \phi(k_\tau)} \]

关键 trick:分子分母都是累加和,可以递推:

\[ S_t = S_{t-1} + \phi(k_t) v_t^\top,\quad z_t = z_{t-1} + \phi(k_t),\quad o_t = \frac{\phi(q_t)^\top S_t}{\phi(q_t)^\top z_t} \]

完整 PyTorch reference(最简单的版本)

def linear_attn_step(q_t, k_t, v_t, S_prev, z_prev, eps=1e-6):
    """朴素 Linear Attention 一步.
    q_t, k_t, v_t: [d]
    S_prev:        [d, d]   累积 K^T V
    z_prev:        [d]      累积 K
    """
    phi_q = F.elu(q_t) + 1                  # 非负 feature map
    phi_k = F.elu(k_t) + 1
    S_new = S_prev + torch.outer(phi_k, v_t)         # [d, d]
    z_new = z_prev + phi_k                            # [d]
    num   = phi_q @ S_new                             # [d]
    denom = phi_q @ z_new + eps                       # scalar
    return num / denom, S_new, z_new

复杂度 + 限制

  • State size: \(S \in \mathbb{R}^{d \times d}\) + \(z \in \mathbb{R}^d\)\(O(d^2)\) 字节,与 N 无关
  • Per step: \(O(d^2)\) FLOPs(外积 + 矩阵-向量)
  • 致命问题: 没有遗忘 —— 旧 token 的贡献无限累积,长 ctx 后 state 被淹没;早期实验上 quality 显著差于 softmax,所以 2020-2022 之间无人重视

2.2 GLA — 加 forget gate(Yang+ 2023)

参考 GLA paper (Yang et al., 2023)

公式(diff vs §2.1)

\[ \boxed{S_t = \underbrace{\text{diag}(g_t)}_{\text{新增}} S_{t-1} + k_t v_t^\top},\qquad o_t = q_t^\top S_t \]

其中 \(g_t \in (0,1)^d\)data-dependent gate(每个 key-dim 独立衰减),来自 \(g_t = \sigma(W_g x_t)\)

改进了什么

  • 遗忘机制:每步给 state 的每行 \(i\)\(g_t[i] \in (0, 1)\),旧信息按指数衰减。\(g_t \to 1\) 时保留全部,\(g_t \to 0\) 时清零
  • 细粒度per-key-dim gate 比标量 gate 表达力强(每维独立决定衰减率)
  • 去掉 normalization:实践中 LayerNorm 输出 + RMS 已经够稳,省掉 \(z_t\) 简化

还差什么

  • gate 只能"减弱"旧记忆,不能精确覆盖早期 (k, v) 对应的 value(如要记新事实替换旧事实,做不到)
  • 训练时 \(\text{diag}(g_t)\) 是逐步累乘 —— chunk-wise 并行需要小心处理

2.3 DeltaNet → KDA — 加 delta rule + DPLR(Moonshot 2025-10)

参考 Kimi Linear / KDA (Moonshot, 2025)

公式(diff vs §2.2)

\[ \boxed{S_t = \underbrace{(I - \beta_t v_t k_t^\top)}_{\text{rank-1 减项}} S_{t-1} + \alpha_t v_t k_t^\top} \]

其中 \(\alpha_t, \beta_t \in (0, 1)^d\) 都是 per-dim 向量 gate(不是 scalar)。

转移矩阵从 diag → DPLR

路线 转移矩阵 \(A_t\) 表达力
GLA \(\text{diag}(g_t)\)(对角) 各 row 独立衰减
KDA \(I - \beta_t v_t k_t^\top\)Diagonal Plus Low-Rank, DPLR 特殊变体) 能在 \(k_t\) 方向减去旧 v(即"覆盖")

I - β v k^T 是 rank-1 修正:作用在 \(S_{t-1}\) 上等于"沿 \(k_t\) 方向把旧 value 减掉一部分"。配合 \(+\alpha_t v_t k_t^\top\) 写入新 value,实现 delta rule:state 中关于 \(k_t\) 的 value 被精确替换

改进了什么

  • 能覆盖旧记忆:解决 GLA "只能衰减不能精确替换"的痛点(典型场景:知识更新、上下文中事实改写)
  • per-dim gate\(\alpha, \beta\) 双向量,更细粒度
  • 跟 attention 表达力差距进一步缩小

工程代价

  • DPLR 转移矩阵的 chunk-wise 并行更复杂(不是简单 cumprod),需要 fused kernel
  • Kimi 自己写了 Kimi-Linear KDA kernel 才让训练效率追上

2.4 Lightning Attention — 改 normalization(MiniMax M1, 2025-06)

参考 MiniMax-M1 (2025) + Transnormer (Yang et al., 2024)

公式(另一种 diff vs §2.1)

\[ \boxed{S_t = S_{t-1} + k_t v_t^\top},\qquad o_t = \text{Norm}\!\big(q_t^\top S_t\big) \]

改进点

  • 没有 gate(state 直接累加,无衰减)
  • 靠 RMSNorm / 输出 normalization 控范围(替代 softmax 的 normalization 角色)
  • \(Q, K\) 走 SiLU 或 ELU+1 feature map 保非负

跟 GLA 的本质区别

GLA 试图"在 state 内部做范围控制"(gate 衰减),Lightning 选择"放任 state 累加,输出端 normalize"。后者更简单,但需要更激进的 hybrid(MiniMax-M1 用 1:7 = softmax : Lightning)来补长 ctx 精度。

2.5 Pure Linear Attention 路线总结

模型 转移矩阵 \(A_t\) Gate 形式 能否覆盖旧 v 训练并行 论文
Linear Attn \(I\)(无) \(O(N d^2)\) chunk scan Katharopoulos 2020
GLA \(\text{diag}(g_t)\) per-dim scalar 否(只能衰减) chunk-wise cumprod Yang 2023
DeltaNet/KDA \(I - \beta_t v_t k_t^\top\) (DPLR) \(\alpha, \beta\) 双向量 chunk-wise (复杂) Moonshot 2025
Lightning \(I\)(无 gate) 否(靠 norm 控) 简单(线性累加) MiniMax 2025

演化方向:复杂度↑、表达力↑、训练并行↓。所以 DeltaNet/KDA 是当前最强(Kimi Linear 论文报告 hybrid 下可超 full attention),但 kernel 开发成本最高。


三、SSM 路线(控制论视角)

3.1 起点:连续状态空间方程

控制论老朋友。系统的 internal state \(h(t) \in \mathbb{R}^N\) 由 ODE 演化:

\[ \dot{h}(t) = A\, h(t) + B\, x(t),\qquad y(t) = C\, h(t) \]
  • \(A \in \mathbb{R}^{N \times N}\):state transition(决定记忆衰减/振荡)
  • \(B \in \mathbb{R}^{N \times 1}\):input projection
  • \(C \in \mathbb{R}^{1 \times N}\):output projection

离散化(Zero-Order Hold 或 Euler)后变成 RNN:

\[ h_t = \bar A\, h_{t-1} + \bar B\, x_t,\qquad y_t = C\, h_t \]

其中 \(\bar A = e^{A \Delta}\), \(\bar B = (\bar A - I) A^{-1} B\)(步长 \(\Delta\))。

跟 Pure Linear Attention 的差别

Pure Linear Attention SSM
state \(S \in \mathbb{R}^{d \times d}\)(外积形式) \(h \in \mathbb{R}^N\)(向量,per-channel)
输入 \(k_t v_t^\top\)(外积) \(\bar B \cdot x_t\)(线性)
衰减 gate \(g_t\)(数据相关,分类讨论) \(\bar A\)早期 SSM 数据无关
输出 \(q_t^\top S_t\)(再投影) \(C^\top h_t\)(直接读出)

关键差别:早期 SSM (\(A, B, C\)) 是数据无关的常数矩阵,区别于 attention 的 \(Q, K, V\) 完全数据相关。这决定了 SSM 演化的方向就是逐步把参数变得数据相关

3.2 S4 — HiPPO + 卷积形式(Gu+ 2021)

参考 Efficiently Modeling Long Sequences with Structured State Spaces (Gu et al., NeurIPS 2021)

公式

\(h_t = \bar A h_{t-1} + \bar B x_t,\quad y_t = C h_t\),其中 \(A, B, C, \Delta\) 都是 learned 但数据无关 的参数。

关键创新 1:HiPPO 初始化

\(A\) 不是随机初始化,而是 HiPPO matrix(High-Order Polynomial Projection Operators,Gu et al. 2020),一种特殊的下三角矩阵:

\[ A_{nk} = -\begin{cases} (2n+1)^{1/2} (2k+1)^{1/2} & n > k \\ n+1 & n = k \\ 0 & n < k \end{cases} \]

作用:让 state \(h_t\) 在 Legendre / Chebyshev 多项式基上保留整段历史的最优投影。这是 S4 第一次让 SSM 在长序列上 work 的关键 —— 之前的线性 RNN 都因为初始化烂导致梯度爆炸 / 消失。

关键创新 2:卷积形式

\(y_t = C \bar A^{t-1} \bar B x_1 + C \bar A^{t-2} \bar B x_2 + \cdots = (\bar K * x)_t\)

其中 \(\bar K = (C\bar B, C\bar A \bar B, C\bar A^2 \bar B, \ldots)\)卷积核。训练时一次性算完整段 \(\bar K\),用 FFT 卷积 → \(O(N \log N)\)。这让 S4 训练能跟 Transformer 速度持平。

限制

  • \(A, B, C\) 数据无关:所有 token 用同一组转移矩阵,无法"针对内容选择性记忆"。比如知识更新场景,S4 没有 attention 那种"我看到这个 token 就强力记住"的能力
  • HiPPO 是数学瑰宝但工程黑箱:很多人不知道为什么这样初始化 work,跟 LSTM 一样依赖经验
  • 离散化 \(\Delta\) 是固定步长:不同 frequency 的信号需要不同 \(\Delta\),固定步长损失自适应能力

3.3 S5 — MIMO + simplified parallel scan(Smith+ 2022)

参考 Simplified State Space Layers for Sequence Modeling (Smith et al., 2022)

改进 vs S4

  • MIMO (Multi-Input Multi-Output):S4 是 SISO(每个通道独立),S5 改成 \(A \in \mathbb{R}^{N \times N}, B \in \mathbb{R}^{N \times d}, C \in \mathbb{R}^{d \times N}\),多通道直接耦合
  • Diagonal \(A\):把 HiPPO 转化为对角矩阵(同等表达力但便宜得多)
  • Parallel scan:用 Blelloch's prefix scan algorithm 替代 FFT 卷积。\(O(N)\) 工作 + \(O(\log N)\) 深度,更适合 GPU 并行

限制

data-independent。所以即使训练快了,长 ctx 选择性记忆问题没解。

3.4 H3 — SSM + 乘法 gate(Fu+ 2022)

参考 Hungry Hungry Hippos (Fu et al., 2022)

改进思路

试图用 SSM 模拟 attention 的功能:把 SSM 的输出跟一个"shift" SSM 的输出做乘法 gate(对应 attention 的 query × key 乘积),再加一个输出投影。结构类似:

\[ y = \text{SSM}_{C}(\sigma(\text{SSM}_{B}(x)) \odot x) \]

限制

H3 的 quality 离 attention 还有差距 —— 乘法 gate 模拟出"选择性"了一点,但没解决 \(A, B, C\) 数据无关的核心问题。这导致 H3 论文 demo 都是合成任务,没在真 LM scale 上击败 attention。

H3 的历史价值:它是最后一个还试图保留 SSM 数据无关性的模型。下一步 Mamba 直接放弃这一约束。

3.5 Mamba (S6) — Selective + hardware-aware(Gu & Dao 2023)

参考 Mamba (Gu & Dao 2023)

关键突破:让 \(\bar A, \bar B, C, \Delta\) 全部数据相关

\[ \boxed{\bar A_t = e^{\Delta_t \cdot A},\quad \bar B_t = \Delta_t \cdot B_t,\quad C_t = C_t(x_t),\quad \Delta_t = \Delta_t(x_t)} \]

每个 token 的 \(\bar A_t\) 由当前输入 \(x_t\) 决定 → selective state space = LSTM 级别的"该忘就忘"。这是 SSM 第一次拿到真正的"内容相关记忆"。

代价:转移矩阵随时间变化 → 不能用卷积(FFT 形式失效)。需要 sequential scan。

关键突破 2:hardware-aware kernel

Mamba 论文的 50% 篇幅在讲 CUDA 工程(不是数学):写了 selective_scan_cuda 把 selective scan 写成 fused kernel,让 GPU 利用率追上 Transformer。S4/S5 数学上更早,但卡在 kernel 没人写出来;Mamba 一次性把数学 + 工程都做完才破圈。

限制

  • 单纯 SSM 长 ctx 精确召回不足(NIAH benchmark 拉胯)
  • 训练 selective scan 依赖手写 CUDA → 不是标准 PyTorch op,调试困难

3.6 Mamba-2 — SSD framework + chunk matmul(Dao & Gu 2024)

参考 Mamba-2 / SSD (Dao & Gu, ICML 2024)

核心定理:SSM ≡ Linear Attention(结构化)

带特定结构(1-Semiseparable matrix = scalar × identity + low rank)的 SSM 状态转移,等价于 带特定 mask 的 linear attention:

  • SSM 视角:\(h_t = \bar A_t h_{t-1} + \bar B_t x_t \;\Rightarrow\; h_t = \sum_{\tau \le t} \big(\prod_{\tau < i \le t} \bar A_i\big) \bar B_\tau x_\tau\)
  • \((\prod \bar A_i) \bar B_\tau\) 当作 \(k_\tau\)\(x_\tau\) 当作 \(v_\tau\) → 正好是 \(S_t = \sum_\tau k_\tau v_\tau^\top\) 的递推

两条路线(Pure Linear Attention 和 SSM)数学上殊途同归

完整 PyTorch reference(最复杂的版本)

class Mamba2Block(nn.Module):
    """简化版 Mamba-2 block,演示 selective state + SSD chunk-wise 训练路径。
    Production 用 selective_scan_cuda + SSD chunk-matmul kernel。
    """
    def __init__(self, d_model, d_state=128, d_conv=4, expand=2):
        super().__init__()
        self.d_inner = expand * d_model            # 扩展内部维(典型 2×)
        self.d_state = d_state                     # SSM state 维度

        # 输入投影 + 1D conv (short-range mixing)
        self.in_proj = nn.Linear(d_model, 2 * self.d_inner)
        self.conv = nn.Conv1d(self.d_inner, self.d_inner, d_conv,
                              groups=self.d_inner, padding=d_conv - 1)

        # selective: Δ, B, C 数据相关 (Mamba vs S4 的关键差别)
        self.x_proj = nn.Linear(self.d_inner, d_state * 2 + 1)  # → B, C, Δ
        self.dt_proj = nn.Linear(1, self.d_inner)               # Δ broadcast

        # A 是 learned but stable (负实数对角)
        A = -torch.arange(1, d_state + 1, dtype=torch.float).repeat(self.d_inner, 1)
        self.A_log = nn.Parameter(torch.log(-A))

        self.out_proj = nn.Linear(self.d_inner, d_model)

    def forward(self, x, h_prev=None):
        B, S, D = x.shape
        x_in, res = self.in_proj(x).chunk(2, dim=-1)             # [B, S, d_inner]

        # 1D conv mixing
        x_in = self.conv(x_in.transpose(1, 2))[:, :, :S].transpose(1, 2)
        x_in = F.silu(x_in)

        # Selective: Δ, B, C 全部由 input 决定
        proj = self.x_proj(x_in)
        delta_raw, B_param, C_param = proj.split(
            [1, self.d_state, self.d_state], dim=-1)
        delta = F.softplus(self.dt_proj(delta_raw))              # [B, S, d_inner]

        # Discretize
        A = -torch.exp(self.A_log)                               # [d_inner, d_state]
        A_bar = torch.exp(delta.unsqueeze(-1) * A)               # [B, S, d_inner, d_state]
        B_bar = delta.unsqueeze(-1) * B_param.unsqueeze(-2)      # [B, S, d_inner, d_state]

        # 序列扫描 (production 用 fused selective_scan_cuda 或 SSD chunk-matmul)
        h = torch.zeros(B, self.d_inner, self.d_state, device=x.device)
        ys = []
        for t in range(S):
            h = A_bar[:, t] * h + B_bar[:, t] * x_in[:, t].unsqueeze(-1)
            y_t = (h * C_param[:, t].unsqueeze(-2)).sum(-1)      # [B, d_inner]
            ys.append(y_t)
        y = torch.stack(ys, dim=1)
        return self.out_proj(y * F.silu(res)), h

工程优势(vs Mamba-1)

  1. 训练统一:用 SSD chunk-wise matmul 实现,吃 Tensor Core,比 selective-scan 快 2-8×
  2. 架构借鉴:linear attention 的 multi-head / FlashAttention 优化都能套到 SSM
  3. 同时存在两种部署形态:训练 chunk matmul,推理 decode 一步一更新 \(h\)

限制

  • 长 ctx 精确召回仍弱于 softmax → 必须 hybrid(Nemotron-H 用 ~8% softmax + 92% Mamba-2/MLP 才能追上 full attn 的 NIAH 性能)

3.7 SSM 路线总结

模型 \(A\) 是否数据相关 \(B, C\) 是否数据相关 训练形式 关键贡献 论文
S4 (2021) FFT 卷积 \(O(N \log N)\) HiPPO 初始化让 SSM 第一次 work Gu 2021
S5 (2022) Parallel scan \(O(N)\) + diag \(A\) 简化 + GPU 友好 Smith 2022
H3 (2022) 同 S4 + 乘法 gate 试图用乘法门模拟 attention,未完成 Fu 2022
Mamba (S6) (2023) \(\bar A_t = e^{\Delta_t A}\)(间接) Sequential scan,依赖 CUDA kernel Selective + hardware-aware Gu & Dao 2023
Mamba-2 (2024) Chunk matmul (SSD) SSM ≡ Linear Attention 同构定理 Dao & Gu 2024

演化方向:参数从数据无关 → 全部数据相关;训练从 FFT → scan → chunk matmul(越来越像 attention 的形式)。


四、SSD 定理:两路汇合

4.1 直观证明

把 Mamba 的 \(h_t = \bar A_t h_{t-1} + \bar B_t x_t\) 一路展开:

\[ h_t = \sum_{\tau \le t} \underbrace{\Big(\prod_{\tau < i \le t} \bar A_i\Big)}_{=: \tilde k_\tau} \bar B_\tau x_\tau \]

\(\tilde k_\tau\)(包含累乘衰减)当作"key",\(x_\tau\) 当作"value",输出 \(y_t = C_t h_t\) 就是:

\[ y_t = C_t \sum_{\tau \le t} \tilde k_\tau \bar B_\tau x_\tau \;\;\equiv\;\; q_t^\top \sum_{\tau \le t} k_\tau v_\tau^\top \]

这就是 linear attention 的递推形式。所以 Mamba ≡ 一种特殊参数化的 GLA。

4.2 实践意义

  1. 训练算法可互换:SSD chunk-matmul(Mamba-2 用)和 GLA chunk-wise scan(KDA 用)数学上等价,工程上选哪个都行
  2. 架构借鉴:linear attention 那边的 multi-head、FlashAttention-style IO 优化、Triton 实现都能直接搬到 SSM
  3. 两条线汇合后:frontier 的"linear vs SSM"之争其实是伪命题,本质都是结构化线性 RNN,差别在初始化和参数化细节

五、统一对比表(公式 + 复杂度 + 区别)

模型 公式(state 更新) State size Per-step Train parallel 论文
Linear Attn \(S_t = S_{t-1} + k_t v_t^\top\) \(d^2\) \(O(d^2)\) chunk scan Katharopoulos 2020
GLA \(S_t = \text{diag}(g_t) S_{t-1} + k_t v_t^\top\) \(d^2\) \(O(d^2)\) chunk + cumprod Yang 2023
DeltaNet/KDA \(S_t = (I - \beta_t v_t k_t^\top) S_{t-1} + \alpha_t v_t k_t^\top\) \(d^2\) \(O(d^2)\) chunk DPLR Moonshot 2025
Lightning \(S_t = S_{t-1} + k_t v_t^\top\)(输出端 norm) \(d^2\) \(O(d^2)\) chunk scan MiniMax 2025
S4 \(h_t = \bar A h_{t-1} + \bar B x_t\)\(\bar A\) 数据无关) \(N\) per channel \(O(d \cdot N)\) FFT \(O(N \log N)\) Gu 2021
S5 同 S4 + diag \(A\) + MIMO \(N\) \(O(d \cdot N)\) parallel scan \(O(N)\) Smith 2022
H3 SSM + 乘法 gate \(N\) \(O(d \cdot N)\) 同 S4 Fu 2022
Mamba (S6) \(h_t = \bar A_t h_{t-1} + \bar B_t x_t\)全 selective \(N\) \(O(d \cdot N)\) selective_scan_cuda Gu & Dao 2023
Mamba-2 同 Mamba + SSD 1-SS 结构 \(N\) \(O(d \cdot N)\) chunk matmul Dao & Gu 2024

注意:

  • Pure Linear Attention 系列 state 是 \(\mathbb{R}^{d \times d}\) matrix(每个 head 一份)
  • SSM 系列 state 是 \(\mathbb{R}^N\) vector(per-channel;总 size 跟 head dim 不直接对应)
  • Mamba-2 SSD 定理证明二者等价(可互换参数化)

bench 实测对比见 kernel.md §3-§6(GLA / Mamba-2 / KDA / Lightning vs softmax decode 的 RTX 4070 Ti SUPER 实测)。


六、我的判断 — 两条线哪个赢了

学术上:Mamba-2 的 SSD 定理已经定理化「同一类东西」。所以「Pure Linear Attention vs SSM」从数学角度是伪命题。

产业上SSM 线(Mamba 系)出圈得更彻底,三个非技术原因:

  1. 品牌叙事:Mamba 论文恰好踩中"Transformer killer"节点,传播力远超学术上更早的 linear attention 工作
  2. HiPPO 起点:S4 的 HiPPO 初始化让 SSM 一开始就接近"恒等映射",loss 曲线顺;linear attention 从随机 init 起步,早期 loss spike 多
  3. CUDA kernel 先发:Mamba 第一天就给 selective_scan_cuda,研究者复制粘贴就能用;GLA 系的 chunk-wise kernel 等到 Kimi Linear 的 KDA 才补齐(晚了将近一年)

Frontier hybrid 配方

模型 linear 层 hybrid 配比 长 ctx 表现
Nemotron-H / Nano-2 Mamba-2 ~8% softmax + 92% Mamba+MLP 接近 full attn
Jamba (AI21) Mamba-1 ~12% softmax + 88% Mamba+MoE 中规中矩
Kimi Linear KDA (DeltaNet++) 25% MLA + 75% KDA 长 ctx + RL 强项
MiniMax-M1 Lightning 12.5% softmax + 87.5% Lightning 1M ctx 商用

给个人选型的建议

  • 复现 / 研究:选 Mamba-2(生态最成熟、kernel 现成、论文链完整)
  • 长 ctx 产品:选 GLA 系(KDA / Lightning,跟 attention hybrid 更平滑)
  • 不要纠结哪条线"赢了"——它们是同一类东西,工程细节决定收益

参考文献

Pure Linear Attention 系

SSM 系

生产 kernel 库

hybrid 模型


上级 · A2 注意力机制全景