线性注意力与 SSM:两条技术路线的完整推导

0. 为什么要固定大小的记忆状态

标准 softmax 注意力解码时必须缓存全部历史的 Key/Value,显存和单步延迟都是 O(nd)O(n \cdot d),上下文翻倍就一起翻倍。把历史压缩成一个固定大小的状态 SRd×dS \in \mathbb{R}^{d \times d},显存和单步开销变成 O(d2)O(d^2),与 nn 无关。

代价是这个状态必须会写、会改、会忘。两条路线各给出一半答案:线性注意力造出固定状态,SSM 教它怎么遗忘。


1. 路线一:线性注意力

1.1 去掉 softmax,结合律才可用

裸注意力 (QK)V(QK^\top)V 的两种括号化复杂度不同:先算 n×nn \times nQKQK^\topO(dn2)O(dn^2),先算 KVRd×dK^\top V \in \mathbb{R}^{d\times d} 则是 O(nd2)O(nd^2)。矩阵乘法满足 (AB)C=A(BC)(AB)C = A(BC),先算哪边是自由的。

softmax 挡在中间:它作用在 QKQK^\top 之后、乘 VV 之前,且是行内非线性。两个性质各堵死一条路——非线性使 softmax(QK)VQsoftmax(KV)\mathrm{softmax}(QK^\top)V \ne Q\,\mathrm{softmax}'(K^\top V),中间结果无法先合并;行内归一化的分母 jeqikj\sum_j e^{q_i \cdot k_j} 依赖该行所有 key,没有可以先结合掉的独立块。

所以平方复杂度不是矩阵乘法的错,是 softmax 的作用位置的错。

1.2 把非线性提前到 Q 和 K 各自身上

核技巧:只要相似度非负(Mercer 条件),就存在特征映射 ϕ()\phi(\cdot) 使 sim(q,k)=ϕ(q)ϕ(k)\mathrm{sim}(q, k) = \phi(q)^\top \phi(k)非线性被吸收进 ϕ\phi,分别作用于 q 和 k 各自身上,相似度本身变回线性内积,于是

(ϕ(Q)ϕ(K))V=ϕ(Q)(ϕ(K)V)\big(\phi(Q)\,\phi(K)^\top\big)V = \phi(Q)\big(\phi(K)^\top V\big)

线性注意力取 ϕ(x)=ELU(x)+1\phi(x) = \mathrm{ELU}(x) + 1,取值范围 (0,)(0, \infty) 恒正——这是核分解存在的条件,不是装饰。写成逐 token 形式:

Attn(q)=ϕ(q)Sϕ(q)z,S=iϕ(ki)vi,z=iϕ(ki)\mathrm{Attn}(q) = \frac{\phi(q)^\top S}{\phi(q)^\top z}, \qquad S = \sum_{i} \phi(k_i)\, v_i^\top,\quad z = \sum_i \phi(k_i)

SS 是 KV 外积矩阵(关联记忆本体),zz 是从 softmax 分母继承下来的归一化项。后续 DeltaNet/GDN/KDA 改用 RMSNorm,zz 就退场了,只有 SS 的更新规则一路演进下去。

两者都能写成递推,一步一读出:

St=St1+ϕ(kt)vt,ot=ϕ(qt)Stϕ(qt)ztS_t = S_{t-1} + \phi(k_t)\, v_t^\top, \qquad o_t = \frac{\phi(q_t)^\top S_t}{\phi(q_t)^\top z_t}

这就是 fast weights 视角:SS 是一块随输入不断被改写的「快速权重」,写入靠 ϕ(kt)vt\phi(k_t)v_t^\top,读出靠 ϕ(qt)St\phi(q_t)^\top S_t

1.3 纯加性状态的代价:记忆碰撞

St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top 只有叠加,没有删除。查询定义为右乘 SkS\,k,把递推展开就能看出问题:

Stk=ivi(kik)S_t\,k = \sum_i v_i\,(k_i^\top k)

每个 value 的系数是它存入时的 key 与查询 key 的内积。key 完全匹配则完整取回,正交则不干扰——这是「按 key 相似度加权取回 value」。但同一个 key 先后写入两个不同的 value 时,两个系数都是 1,读出的是两者之和而不是最新那个。这就是记忆碰撞。

修法是 delta rule:写入前先查旧值,只写差值。

St=St1+(vtSt1kt)ktS_t = S_{t-1} + \big(v_t - S_{t-1}k_t\big)k_t^\top

差值中的负项抵消掉旧 key 上的残余,等于先删除再写入,从而保证 StktvtS_t k_t \approx v_t

1.4 手算一个例子:擦除到底发生了什么

dk=dv=2d_k = d_v = 2S0=0S_0 = 0,依次写入三个 token(关键在 k3=k1k_3 = k_1,同一个 key 写入新值):

k1=[1,0], v1=[1,0];k2=[0.6,0.8], v2=[0,1];k3=[1,0], v3=[2,0]k_1 = [1, 0],\ v_1 = [1, 0];\qquad k_2 = [0.6, 0.8],\ v_2 = [0, 1];\qquad k_3 = [1, 0],\ v_3 = [2, 0]

线性注意力St=St1+vtktS_t = S_{t-1} + v_tk_t^\top,直接叠外积):

S1=[1000],S2=[100.60.8],S3=[300.60.8]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix},\qquad S_2 = \begin{bmatrix}1&0\\0.6&0.8\end{bmatrix},\qquad S_3 = \begin{bmatrix}3&0\\0.6&0.8\end{bmatrix}

第三步的 v3k3=[2000]v_3k_3^\top = \begin{bmatrix}2&0\\0&0\end{bmatrix} 不清除 S2S_2 里已有的旧值,直接往上叠,第一行变成 [3,0][3, 0]——碰撞就在这一瞬发生。查询 k1k_1S3k1=[3,0.6]S_3k_1 = [3, 0.6]:旧值 v1v_1 与新值 v3v_3 完整叠加(两个系数都是 1),再混进 0.6 份 v2v_2k2k1=0.6k_2^\top k_1 = 0.6)。

delta rule(取 β=1\beta = 1,每步先查后写):

  • t=1t=1S0k1=0S_0k_1 = 0,差值 u1=v1=[1,0]u_1 = v_1 = [1, 0],得 S1=[1000]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}(与线性注意力相同);
  • t=2t=2:先查 S1k2=[0.6,0]S_1k_2 = [0.6, 0]k2k_2 在第一维有 0.6 分量,部分命中旧记忆),差值 u2=v2S1k2=[0.6,1]u_2 = v_2 - S_1k_2 = [-0.6, 1]。这个负分量就是擦除,它扣除 k2k_2 方向上已存的 v1v_1 残余,得 S2=[0.640.480.60.8]S_2 = \begin{bmatrix}0.64&-0.48\\0.6&0.8\end{bmatrix}
  • t=3t=3:先查 S2k3=[0.64,0.6]S_2k_3 = [0.64, 0.6],差值 u3=v3S2k3=[1.36,0.6]u_3 = v_3 - S_2k_3 = [1.36, -0.6],得

S3=[20.4800.8]S_3 = \begin{bmatrix}2&-0.48\\0&0.8\end{bmatrix}

查询 k1k_1S3S_3 第一列,得 [2,0][2, 0]——精确返回最新的 v3v_3。两个分量各自对应一次擦除:第一行的 2 是 u3u_3+1.36+1.36t=2t{=}2 碎掉的 0.64 补回到 2;第二行的 0 是 u3u_30.6-0.6 正好抵消 t=2t{=}2 混进来的 0.6。对比线性注意力的 [3,0.6][3, 0.6]:旧值和残余都还在里面。

DeltaNet 的雏形到此成立:固定大小的外积状态 + 先删旧再写新。


2. 路线二:SSM

2.1 连续 SSM 的定义

h(t)=Ah(t)+Bx(t),y(t)=Ch(t)h'(t) = A\,h(t) + B\,x(t), \qquad y(t) = C\,h(t)

其中 h(t)RNh(t) \in \mathbb{R}^N 是状态(历史的压缩),AA 决定旧记忆如何衰减,BB 是写入强度,CC 是读出权重。要解决的问题:把微分方程改写成递推式 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,并求出 Aˉ,Bˉ\bar A, \bar B

用到矩阵指数 eM:=k0Mk/k!e^{M} := \sum_{k\ge0} M^k/k! 的两条性质:ddteAt=AeAt\frac{d}{dt}e^{At} = A\,e^{At},以及 (eAt)1=eAt(e^{At})^{-1} = e^{-At}。直觉上 eAΔe^{A\Delta} 就是「让系统按自身动力学自由演化 Δ\Delta 时间」的算子。

2.2 通解:积分因子法

定理h(t)=Ah(t)+Bx(t)h'(t) = A h(t) + B x(t) 满足初值 h(t0)h(t_0) 的解为

h(t)=eA(tt0)h(t0)+t0teA(ts)Bx(s)dsh(t) = e^{A(t-t_0)}\,h(t_0) + \int_{t_0}^{t} e^{A(t-s)}\,B\,x(s)\,ds

证明。移项得 h(t)Ah(t)=Bx(t)h'(t) - A\,h(t) = B\,x(t),两边左乘积分因子 eAte^{-At}

eAth(t)eAtAh(t)=eAtBx(t)e^{-At}h'(t) - e^{-At}A\,h(t) = e^{-At}B\,x(t)

左端恰是一个乘积的全导数,这是整个推导的机关:

ddt[eAth(t)]=eAth(t)eAtAh(t)=eAtBx(t)\frac{d}{dt}\Big[e^{-At}h(t)\Big] = e^{-At}h'(t) - e^{-At}A\,h(t) = e^{-At}B\,x(t)

t0t_0tt 积分,再左乘 eAte^{At}(用 eAteAs=eA(ts)e^{At}e^{-As} = e^{A(t-s)})即得结论。\blacksquare

两项分别是旧记忆自由演化区间内每一瞬输入贡献的叠加

2.3 零阶保持(ZOH)离散化

假设采样区间内输入保持常值 x(tk+τ)=xkx(t_k + \tau) = x_k。在通解中取 t0=tkt_0 = t_kt=tk+Δt = t_k + \Delta,把 xkx_k 提出积分号:

hk+1=eAΔAˉhk+(0ΔeAsBds)Bˉxkh_{k+1} = \underbrace{e^{A\Delta}}_{\bar A}\,h_k + \underbrace{\left(\int_0^{\Delta} e^{As}B\,ds\right)}_{\bar B}\,x_k

对级数逐项积分算出 Bˉ\bar B

0ΔeAsds=k0AkΔk+1(k+1)!=A1(eAΔI)\int_0^{\Delta} e^{As}\,ds = \sum_{k\ge0}\frac{A^k\Delta^{k+1}}{(k+1)!} = A^{-1}\big(e^{A\Delta} - I\big)

 Aˉ=eΔA,Bˉ=A1(eΔAI)B \boxed{\ \bar{A} = e^{\Delta A}, \qquad \bar{B} = A^{-1}\big(e^{\Delta A} - I\big)\,B\ }

ZOH 在「输入确为分段常数」的假设下不是近似而是精确等价Δ\Delta 很小时有 AˉI+ΔA\bar A \approx I + \Delta ABˉΔB\bar B \approx \Delta B(欧拉法),但 S4/Mamba 实现都用精确公式。

2.4 Aˉ=eΔA\bar A = e^{\Delta A} 就是遗忘门

AA 负定(Mamba-2 取 A=aIA = -a\cdot Ia>0a>0),则

Aˉt=eΔta(0,1)\bar A_t = e^{-\Delta_t a} \in (0, 1)

Δt\Delta_t \to \inftyAˉt0\bar A_t \to 0,清空旧记忆;Δt0\Delta_t \to 0Aˉt1\bar A_t \to 1,冻结状态。Mamba 让 Δt=softplus(Linear(xt))\Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t)),于是「步长」变成了看内容的遗忘门。


3. 两条路线在 Mamba-2 汇合

Mamba-2 把状态写成外积形式,矩阵 AA 退化为标量衰减:

St=αtSt1+vtkt,αt=eΔtaS_t = \alpha_t S_{t-1} + v_tk_t^\top, \qquad \alpha_t = e^{-\Delta_t a}

对照线性注意力的 St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,只多了一个乘在旧状态上的 αt\alpha_t线性注意力给出了外积状态的样子,SSM 给出了衰减门,从这一步起两者在数学上是同一个东西的两个记法(SSD 框架)。

此后 GDN 在写入侧加 delta rule,KDA 把标量门打开成逐通道门 Diag(αt)\mathrm{Diag}(\alpha_t) 并给衰减加下界。骨架始终是同一条:状态 ×(衰减/删除算子)+(写入项)。后续演化见《KDA 的来龙去脉》


参考