0. 为什么要固定大小的记忆状态
标准 softmax 注意力解码时必须缓存全部历史的 Key/Value,显存和单步延迟都是 O(n⋅d),上下文翻倍就一起翻倍。把历史压缩成一个固定大小的状态 S∈Rd×d,显存和单步开销变成 O(d2),与 n 无关。
代价是这个状态必须会写、会改、会忘。两条路线各给出一半答案:线性注意力造出固定状态,SSM 教它怎么遗忘。
1. 路线一:线性注意力
1.1 去掉 softmax,结合律才可用
裸注意力 (QK⊤)V 的两种括号化复杂度不同:先算 n×n 的 QK⊤ 是 O(dn2),先算 K⊤V∈Rd×d 则是 O(nd2)。矩阵乘法满足 (AB)C=A(BC),先算哪边是自由的。
softmax 挡在中间:它作用在 QK⊤ 之后、乘 V 之前,且是行内非线性。两个性质各堵死一条路——非线性使 softmax(QK⊤)V=Qsoftmax′(K⊤V),中间结果无法先合并;行内归一化的分母 ∑jeqi⋅kj 依赖该行所有 key,没有可以先结合掉的独立块。
所以平方复杂度不是矩阵乘法的错,是 softmax 的作用位置的错。
1.2 把非线性提前到 Q 和 K 各自身上
核技巧:只要相似度非负(Mercer 条件),就存在特征映射 ϕ(⋅) 使 sim(q,k)=ϕ(q)⊤ϕ(k)。非线性被吸收进 ϕ,分别作用于 q 和 k 各自身上,相似度本身变回线性内积,于是
(ϕ(Q)ϕ(K)⊤)V=ϕ(Q)(ϕ(K)⊤V)
线性注意力取 ϕ(x)=ELU(x)+1,取值范围 (0,∞) 恒正——这是核分解存在的条件,不是装饰。写成逐 token 形式:
Attn(q)=ϕ(q)⊤zϕ(q)⊤S,S=i∑ϕ(ki)vi⊤,z=i∑ϕ(ki)
S 是 KV 外积矩阵(关联记忆本体),z 是从 softmax 分母继承下来的归一化项。后续 DeltaNet/GDN/KDA 改用 RMSNorm,z 就退场了,只有 S 的更新规则一路演进下去。
两者都能写成递推,一步一读出:
St=St−1+ϕ(kt)vt⊤,ot=ϕ(qt)⊤ztϕ(qt)⊤St
这就是 fast weights 视角:S 是一块随输入不断被改写的「快速权重」,写入靠 ϕ(kt)vt⊤,读出靠 ϕ(qt)⊤St。
1.3 纯加性状态的代价:记忆碰撞
St=St−1+ϕ(kt)vt⊤ 只有叠加,没有删除。查询定义为右乘 Sk,把递推展开就能看出问题:
Stk=i∑vi(ki⊤k)
每个 value 的系数是它存入时的 key 与查询 key 的内积。key 完全匹配则完整取回,正交则不干扰——这是「按 key 相似度加权取回 value」。但同一个 key 先后写入两个不同的 value 时,两个系数都是 1,读出的是两者之和而不是最新那个。这就是记忆碰撞。
修法是 delta rule:写入前先查旧值,只写差值。
St=St−1+(vt−St−1kt)kt⊤
差值中的负项抵消掉旧 key 上的残余,等于先删除再写入,从而保证 Stkt≈vt。
1.4 手算一个例子:擦除到底发生了什么
取 dk=dv=2,S0=0,依次写入三个 token(关键在 k3=k1,同一个 key 写入新值):
k1=[1,0], v1=[1,0];k2=[0.6,0.8], v2=[0,1];k3=[1,0], v3=[2,0]
线性注意力(St=St−1+vtkt⊤,直接叠外积):
S1=[1000],S2=[10.600.8],S3=[30.600.8]
第三步的 v3k3⊤=[2000] 不清除 S2 里已有的旧值,直接往上叠,第一行变成 [3,0]——碰撞就在这一瞬发生。查询 k1 得 S3k1=[3,0.6]:旧值 v1 与新值 v3 完整叠加(两个系数都是 1),再混进 0.6 份 v2(k2⊤k1=0.6)。
delta rule(取 β=1,每步先查后写):
- t=1:S0k1=0,差值 u1=v1=[1,0],得 S1=[1000](与线性注意力相同);
- t=2:先查 S1k2=[0.6,0](k2 在第一维有 0.6 分量,部分命中旧记忆),差值 u2=v2−S1k2=[−0.6,1]。这个负分量就是擦除,它扣除 k2 方向上已存的 v1 残余,得 S2=[0.640.6−0.480.8];
- t=3:先查 S2k3=[0.64,0.6],差值 u3=v3−S2k3=[1.36,−0.6],得
S3=[20−0.480.8]
查询 k1 取 S3 第一列,得 [2,0]——精确返回最新的 v3。两个分量各自对应一次擦除:第一行的 2 是 u3 的 +1.36 把 t=2 碎掉的 0.64 补回到 2;第二行的 0 是 u3 的 −0.6 正好抵消 t=2 混进来的 0.6。对比线性注意力的 [3,0.6]:旧值和残余都还在里面。
DeltaNet 的雏形到此成立:固定大小的外积状态 + 先删旧再写新。
2. 路线二:SSM
2.1 连续 SSM 的定义
h′(t)=Ah(t)+Bx(t),y(t)=Ch(t)
其中 h(t)∈RN 是状态(历史的压缩),A 决定旧记忆如何衰减,B 是写入强度,C 是读出权重。要解决的问题:把微分方程改写成递推式 ht=Aˉht−1+Bˉxt,并求出 Aˉ,Bˉ。
用到矩阵指数 eM:=∑k≥0Mk/k! 的两条性质:dtdeAt=AeAt,以及 (eAt)−1=e−At。直觉上 eAΔ 就是「让系统按自身动力学自由演化 Δ 时间」的算子。
2.2 通解:积分因子法
定理:h′(t)=Ah(t)+Bx(t) 满足初值 h(t0) 的解为
h(t)=eA(t−t0)h(t0)+∫t0teA(t−s)Bx(s)ds
证明。移项得 h′(t)−Ah(t)=Bx(t),两边左乘积分因子 e−At:
e−Ath′(t)−e−AtAh(t)=e−AtBx(t)
左端恰是一个乘积的全导数,这是整个推导的机关:
dtd[e−Ath(t)]=e−Ath′(t)−e−AtAh(t)=e−AtBx(t)
从 t0 到 t 积分,再左乘 eAt(用 eAte−As=eA(t−s))即得结论。■
两项分别是旧记忆自由演化与区间内每一瞬输入贡献的叠加。
2.3 零阶保持(ZOH)离散化
假设采样区间内输入保持常值 x(tk+τ)=xk。在通解中取 t0=tk、t=tk+Δ,把 xk 提出积分号:
hk+1=AˉeAΔhk+Bˉ(∫0ΔeAsBds)xk
对级数逐项积分算出 Bˉ:
∫0ΔeAsds=k≥0∑(k+1)!AkΔk+1=A−1(eAΔ−I)
Aˉ=eΔA,Bˉ=A−1(eΔA−I)B
ZOH 在「输入确为分段常数」的假设下不是近似而是精确等价。Δ 很小时有 Aˉ≈I+ΔA、Bˉ≈ΔB(欧拉法),但 S4/Mamba 实现都用精确公式。
2.4 Aˉ=eΔA 就是遗忘门
若 A 负定(Mamba-2 取 A=−a⋅I,a>0),则
Aˉt=e−Δta∈(0,1)
Δt→∞ 时 Aˉt→0,清空旧记忆;Δt→0 时 Aˉt→1,冻结状态。Mamba 让 Δt=softplus(Linear(xt)),于是「步长」变成了看内容的遗忘门。
3. 两条路线在 Mamba-2 汇合
Mamba-2 把状态写成外积形式,矩阵 A 退化为标量衰减:
St=αtSt−1+vtkt⊤,αt=e−Δta
对照线性注意力的 St=St−1+ϕ(kt)vt⊤,只多了一个乘在旧状态上的 αt。线性注意力给出了外积状态的样子,SSM 给出了衰减门,从这一步起两者在数学上是同一个东西的两个记法(SSD 框架)。
此后 GDN 在写入侧加 delta rule,KDA 把标量门打开成逐通道门 Diag(αt) 并给衰减加下界。骨架始终是同一条:状态 ×(衰减/删除算子)+(写入项)。后续演化见《KDA 的来龙去脉》。
参考:
- Katharopoulos et al., Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention, arXiv:2006.16236
- Schlag et al., Linear Transformers Are Secretly Fast Weight Programmers, arXiv:2102.11174
- Gu et al., Efficiently Modeling Long Sequences with Structured State Spaces (S4), arXiv:2111.00396
- Gu & Dao, Mamba: Linear-Time Sequence Modeling with Selective State Spaces, arXiv:2312.00752
- Dao & Gu, Transformers are SSMs: Generalized Models and Efficient Algorithms Through Structured State Space Duality, arXiv:2405.21060