线性注意力与 SSM:两条技术路线的完整推导
线性注意力和 SSM(State Space Model,状态空间模型)是序列建模的两条技术路线,也是后来 KDA/GDN/DeltaNet 一族模型的两个源头:线性注意力 贡献了「外积记忆状态 + 结合律」,SSM 贡献了「衰减门 + 递推骨架」。本文把这两条路线的计算过程完整推一遍,每步都有推导、数值验算和工程上的理由。
0. 为什么要固定大小的记忆状态
标准 softmax 注意力在生成文本时,必须把每个历史 token 的 Key 和 Value 都存下来(KV cache)–因为每生成一个新 token,都要和前面所有 token 重新算一遍注意力。序列越长,这笔账越贵:
开销
softmax 注意力(KV cache)
固定状态(线性注意力一族)
解码显存
O ( n ⋅ d ) O(n \cdot d) O ( n ⋅ d ) (全部历史的 KV 都得缓存)
O ( d 2 ) O(d^2) O ( d 2 ) (一个固定矩阵 S,与 n 无关)
单步解码
O ( n ⋅ d ) O(n \cdot d) O ( n ⋅ d ) (扫全部历史)
O ( d 2 ) O(d^2) O ( d 2 ) (更新 S + 读出,两次矩阵乘)
预填充
O ( n 2 d ) O(n^2 d) O ( n 2 d )
O ( n d 2 ) O(n d^2) O ( n d 2 )
上下文长度翻倍
显存、延迟一起翻倍
完全不变
(单头视角,d 为特征维度。)
几十万 token 的上下文就能把 KV cache 撑到几十 GB。想摆脱这笔账,就得把「全部历史的 KV」压缩成一个固定大小的状态 ,而且这个状态要会写、会改、会忘 –两个源头各给出了一半答案:线性注意力先把固定状态造出来,SSM 教会它怎么遗忘。
1. 路线一:线性注意力–两个前提,缺一不可
线性注意力能成立,靠的是两件事:去掉 softmax 的非线性 和让 Q/K 先结合掉 。更准确地说,这是同一枚硬币的两面–softmax 必须作用在耦合后的 Q K ⊤ QK^\top Q K ⊤ 上,它既是非线性的来源,也是耦合的来源;把它换成可分离的 ϕ ( q ) ⊤ ϕ ( k ) \phi(q)^\top\phi(k) ϕ ( q ) ⊤ ϕ ( k ) ,非线性和耦合一起消失,结合律才重新可用。后面 KDA 一族的全部推导都建立在这个前提上,值得把两个条件一个一个讲透。
1.1 前提一:去掉 softmax,结合律才可用
先看没有 softmax 的裸注意力 A t t n ( Q , K , V ) = ( Q K ⊤ ) V \mathrm{Attn}(Q, K, V) = (QK^\top)V Attn ( Q , K , V ) = ( Q K ⊤ ) V ,复杂度账本:
括号化
计算路径
复杂度
主导项
( Q K ⊤ ) V (QK^\top)V ( Q K ⊤ ) V
先算出 n × n n \times n n × n 矩阵,再乘 V
O ( d ⋅ n 2 ) + O ( n ⋅ n d ) = O ( d n 2 ) O(d\cdot n^2) + O(n\cdot nd) = O(dn^2) O ( d ⋅ n 2 ) + O ( n ⋅ n d ) = O ( d n 2 )
n(平方)
Q ( K ⊤ V ) Q(K^\top V) Q ( K ⊤ V )
先算 K ⊤ V ∈ R d × d K^\top V \in \mathbb{R}^{d \times d} K ⊤ V ∈ R d × d (与 n 无关),再左乘 Q
O ( n ⋅ d 2 ) + O ( d ⋅ n d ) = O ( n d 2 ) O(n\cdot d^2) + O(d\cdot nd) = O(nd^2) O ( n ⋅ d 2 ) + O ( d ⋅ n d ) = O ( n d 2 )
d(线性)
同一个数学式,两种括号化,复杂度差一个 n 的幂–这就是结合律的诱惑 :矩阵乘法满足 ( A B ) C = A ( B C ) (AB)C = A(BC) ( A B ) C = A ( B C ) ,先算哪边是自由的。裸注意力(线性代数层面)本来就能线性化。
那 softmax 挡在哪? 标准注意力是 s o f t m a x ( Q K ⊤ ) V \mathrm{softmax}(QK^\top)V softmax ( Q K ⊤ ) V ,softmax 在 Q K ⊤ QK^\top Q K ⊤ 之后 、乘 V 之前 介入,并且它是行内非线性 :第 i 行(第 i 个 query 对所有 key 的打分)必须整体过 e x i / ∑ j e x j e^{x_i}/\sum_j e^{x_j} e x i / ∑ j e x j –指数逐项非线性 + 行内归一化耦合。两个性质各堵死一条路:
非线性破坏结合律 :softmax 不是线性算子,s o f t m a x ( Q K ⊤ ) V ≠ Q s o f t m a x ′ ( K ⊤ V ) \mathrm{softmax}(QK^\top)V \ne Q\,\mathrm{softmax}'(K^\top V) softmax ( Q K ⊤ ) V = Q softmax ′ ( K ⊤ V ) ,中间结果没法先合并;n × n n \times n n × n 矩阵必须先完整算出来;
归一化引入全局耦合 :分母 ∑ j e q i ⋅ k j \sum_j e^{q_i \cdot k_j} ∑ j e q i ⋅ k j 依赖该行所有 n 个 key,哪怕只想算一个 query 的输出,也得先把整行算完–信息在 key 维度上全局耦合,没有可以「先结合掉」的独立块。
所以平方复杂度不是矩阵乘法的错,是 softmax 的作用位置 的错:它非要等 Q 和 K 耦合完才动手。要线性化,就得把这个非线性从耦合点上搬走。
1.2 前提二:把非线性提前到 Q 和 K 各自身上
搬法就是核化(kernelization)。把注意力抽象成通用相似度函数 s i m ( q , k ) \mathrm{sim}(q, k) sim ( q , k ) –它不必是 softmax 的指数,多项式相似度、RBF 核都属此类。数学依据是核技巧 :只要 sim 非负(Mercer 条件),就存在特征映射 ϕ ( ⋅ ) \phi(\cdot) ϕ ( ⋅ ) 使
s i m ( q , k ) = ϕ ( q ) ⊤ ϕ ( k ) \mathrm{sim}(q, k) = \phi(q)^\top \phi(k)
sim ( q , k ) = ϕ ( q ) ⊤ ϕ ( k )
注意这个形式的本质:非线性被吸收进 ϕ \phi ϕ ,分别作用于 q 和 k 各自身上,相似度本身变成线性的内积 。softmax 的 e q ⋅ k e^{q \cdot k} e q ⋅ k 作用在耦合后的标量上;核化的 ϕ ( q ) ⊤ ϕ ( k ) \phi(q)^\top\phi(k) ϕ ( q ) ⊤ ϕ ( k ) 让 q 和 k 各自先做完所有非线性变换,最后只留一次线性内积。
线性注意力取 ϕ ( x ) = E L U ( x ) + 1 \phi(x) = \mathrm{ELU}(x) + 1 ϕ ( x ) = ELU ( x ) + 1 (Katharopoulos et al. 2020)。先把 ELU(Exponential Linear Unit,指数线性单元)本身说清楚,它是个分段函数:
E L U ( x ) = { x x > 0 e x − 1 x ≤ 0 \mathrm{ELU}(x) = \begin{cases} x & x > 0 \\ e^{x} - 1 & x \le 0 \end{cases}
ELU ( x ) = { x e x − 1 x > 0 x ≤ 0
正数区 :原样输出(和 ReLU 一样);
负数区 :输出 e x − 1 ∈ ( − 1 , 0 ) e^x - 1 \in (-1, 0) e x − 1 ∈ ( − 1 , 0 ) ,一条平滑曲线,越负越接近 -1 但永远不到 -1。
所以 E L U ( x ) + 1 \mathrm{ELU}(x) + 1 ELU ( x ) + 1 的取值范围是 ( 0 , ∞ ) (0, \infty) ( 0 , ∞ ) :恒正 。算两个具体值感受一下:x = 0 x = 0 x = 0 时 E L U ( 0 ) + 1 = 0 + 1 = 1 \mathrm{ELU}(0)+1 = 0 + 1 = 1 ELU ( 0 ) + 1 = 0 + 1 = 1 ;x = − 3 x = -3 x = − 3 时 e − 3 − 1 + 1 = e − 3 ≈ 0.05 e^{-3} - 1 + 1 = e^{-3} \approx 0.05 e − 3 − 1 + 1 = e − 3 ≈ 0.05 ,很小但仍是正数。这个恒正不是自选的装饰,是核分解存在的条件(Mercer 条件要求相似度非负):ϕ ( q ) ⊤ ϕ ( k ) \phi(q)^\top\phi(k) ϕ ( q ) ⊤ ϕ ( k ) 是 d 个正数乘正数再求和,结果必为正,才能扮演「打分」的角色。于是:
( ϕ ( Q ) ϕ ( K ) ⊤ ) V = ϕ ( Q ) ( ϕ ( K ) ⊤ V ) \big(\phi(Q)\,\phi(K)^\top\big)V = \phi(Q)\big(\phi(K)^\top V\big)
( ϕ ( Q ) ϕ ( K ) ⊤ ) V = ϕ ( Q ) ( ϕ ( K ) ⊤ V )
结合律在非线性世界里重新可用:先算 ϕ ( K ) ⊤ V ∈ R d × d \phi(K)^\top V \in \mathbb{R}^{d \times d} ϕ ( K ) ⊤ V ∈ R d × d ,复杂度回到 O ( n d 2 ) O(nd^2) O ( n d 2 ) 。两个前提到此汇成一句话:非线性提前,耦合消失,结合律重新可用 。写成逐 token 的归一化形式(softmax 的归一化也没丢,变成显式分母):
A t t n ( q ) = ∑ i ϕ ( q ) ⊤ ϕ ( k i ) v i ∑ i ϕ ( q ) ⊤ ϕ ( k i ) = ϕ ( q ) ⊤ S ϕ ( q ) ⊤ z , S = ∑ i ϕ ( k i ) v i ⊤ \mathrm{Attn}(q) = \frac{\sum_i \phi(q)^\top \phi(k_i)\, v_i}{\sum_i \phi(q)^\top \phi(k_i)} = \frac{\phi(q)^\top S}{\phi(q)^\top z}, \qquad S = \sum_{i} \phi(k_i)\, v_i^\top
Attn ( q ) = ∑ i ϕ ( q ) ⊤ ϕ ( k i ) ∑ i ϕ ( q ) ⊤ ϕ ( k i ) v i = ϕ ( q ) ⊤ z ϕ ( q ) ⊤ S , S = i ∑ ϕ ( k i ) v i ⊤
两个新符号 S 和 z 分别是:
S = ∑ i ϕ ( k i ) v i ⊤ ∈ R d × d S = \sum_i \phi(k_i)v_i^\top \in \mathbb{R}^{d \times d} S = ∑ i ϕ ( k i ) v i ⊤ ∈ R d × d :分子里的 KV 外积矩阵 ,存「key-value 关联」的记忆本体;
z = ∑ i ϕ ( k i ) ∈ R d z = \sum_i \phi(k_i) \in \mathbb{R}^{d} z = ∑ i ϕ ( k i ) ∈ R d :分母里的 key 累积和 ,就是 softmax 分母 ∑ i e q ⋅ k i \sum_i e^{q\cdot k_i} ∑ i e q ⋅ k i 的核化版–归一化因子。
先说清 z,因为后续论文里它最容易被略写 。z 在公式里承担的角色就是归一化:没有它,输出会随序列变长而无界增长。但后面 DeltaNet/GDN/KDA 的论文里,递推式往往只写 S S S 的更新,z 要么藏在一句「输出再过 RMSNorm」的描述里,要么干脆不提–因为实测发现把归一化换成 RMSNorm 效果更好,z 就被简化掉了。所以读者常遇到的困惑是:看 KDA 论文时公式里根本没有 z,翻早期线性注意力文献才发现 z 是这里从 softmax 分母继承下来的归一化项。一句话:z = 归一化分母的累积状态,S = 关联记忆本体 ;后续模型改用 RMSNorm 后 z 退场,但 S 的更新规则一路演进到 KDA。
分子分母都变成对 S 的查询,S 一遍扫过序列累积即可–复杂度对 n 线性。更关键的是 S 可以写成递推,一步一读出(z 同样逐步累积):
S t = S t − 1 + ϕ ( k t ) v t ⊤ , o t = ϕ ( q t ) ⊤ S t ϕ ( q t ) ⊤ z t S_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}
S t = S t − 1 + ϕ ( k t ) v t ⊤ , o t = ϕ ( q t ) ⊤ z t ϕ ( q t ) ⊤ S t
这个递推还有一个更形象的理解,叫 fast weights 视角(Schlag et al., Linear Transformers Are Secretly Fast Weight Programmers ):把 S 看作一块快速权重 –模型真正的参数(普通权重)训练完就固定了,而 S 每来一个 token 就被改写一次,相当于一个「随输入不断更新的小型权重」。写入靠 ϕ ( k t ) v t ⊤ \phi(k_t)v_t^\top ϕ ( k t ) v t ⊤ (每个 token 都在给这块权重编程),读出靠查询 ϕ ( q t ) ⊤ S t \phi(q_t)^\top S_t ϕ ( q t ) ⊤ S t (拿当前问题去这块权重里查答案)。所以这个视角下,序列处理 = 一边更新权重、一边用权重答题。
回过头总结一下,这两个前提各自换来了什么 :
前提
换来的直接好处
对后续模型的影响
抛弃 softmax 非线性
结合律可用,O ( n 2 ) → O ( n ) O(n^2) \to O(n) O ( n 2 ) → O ( n )
注意力矩阵 n × n n\times n n × n 消失,换成固定大小状态 S ∈ R d × d S \in \mathbb{R}^{d \times d} S ∈ R d × d
非线性提前到 Q/K 各自身上
ϕ ( q ) ⊤ ϕ ( k ) \phi(q)^\top\phi(k) ϕ ( q ) ⊤ ϕ ( k ) 可分离
S 可以递推累积 -> RNN 化 -> fast weights -> 一切后续演化的载体
注意代价也在表里:换来的 S 只会加法。这就引出了线性注意力最大的问题。
1.3 纯加性状态的问题:记忆碰撞
纯加性状态的核心问题 :state 是纯加性的(S t = S t − 1 + ϕ ( k t ) v t ⊤ S_t = S_{t-1} + \phi(k_t)v_t^\top S t = S t − 1 + ϕ ( k t ) v t ⊤ ),写入只有「叠加」,没有「删除」。序列长度远超 state 有效容量(d × d d \times d d × d 矩阵只能存这么多关联)时,不同 k → v k \to v k → v 关联互相干扰,旧信息永远赖在状态里,新信息无法覆盖–记忆像一个只进不出的仓库。
数值演示 。设 d k = d v = 2 d_k = d_v = 2 d k = d v = 2 ,S 从零开始,依次写入四个 token:
token
key
value
说明
t=1
k 1 = [ 1 , 0 ] k_1 = [1, 0] k 1 = [ 1 , 0 ]
v 1 = [ 1 , 0 ] v_1 = [1, 0] v 1 = [ 1 , 0 ]
t=2
k 2 = [ 0.6 , 0.8 ] k_2 = [0.6, 0.8] k 2 = [ 0.6 , 0.8 ]
v 2 = [ 0 , 1 ] v_2 = [0, 1] v 2 = [ 0 , 1 ]
t=3
k 3 = [ 1 , 0 ] k_3 = [1, 0] k 3 = [ 1 , 0 ]
v 3 = [ 2 , 0 ] v_3 = [2, 0] v 3 = [ 2 , 0 ]
和 k 1 k_1 k 1 同一个 key,写入新值
t=4
k 4 = [ 0 , 1 ] k_4 = [0, 1] k 4 = [ 0 , 1 ]
v 4 = [ 3 , 1 ] v_4 = [3, 1] v 4 = [ 3 , 1 ]
线性注意力这边,把每一步都算出来 。写入规则是 S t = S t − 1 + v t k t ⊤ S_t = S_{t-1} + v_t k_t^\top S t = S t − 1 + v t k t ⊤ ,每个 v k ⊤ v k^\top v k ⊤ 是一个 2×2 外积。
为什么是 v k ⊤ v k^\top v k ⊤ 而不是 k v ⊤ k v^\top k v ⊤ ? 外积的顺序决定「谁进谁出」。v k ⊤ v k^\top v k ⊤ 是一个从 key 空间到 value 空间的线性映射:key 从右边进,value 从左边出 –直接验证:( v k ⊤ ) k = v ( k ⊤ k ) (v k^\top)\,k = v\,(k^\top k) ( v k ⊤ ) k = v ( k ⊤ k ) ,查询 key 与存储 key 的内积落在系数上,value 原样输出。如果反过来存 k v ⊤ k v^\top k v ⊤ 还从右边乘,得到 k ( v ⊤ k ) k\,(v^\top k) k ( v ⊤ k ) –变成「拿 key 查、按 value 相似度返回 key」,角色整个反了,查不出想要的东西。要让 k v ⊤ k v^\top k v ⊤ 布局工作,查询必须从左边进:行向量 q ⊤ S = ∑ i ( q ⊤ k i ) v i ⊤ q^\top S = \sum_i (q^\top k_i)\,v_i^\top q ⊤ S = ∑ i ( q ⊤ k i ) v i ⊤ ,结论完全一样。所以两种写法都存在:§1.2 的 S = ∑ i ϕ ( k i ) v i ⊤ S = \sum_i \phi(k_i)v_i^\top S = ∑ i ϕ ( k i ) v i ⊤ 配左乘读出 ϕ ( q ) ⊤ S \phi(q)^\top S ϕ ( q ) ⊤ S 、KDA 论文的 S S S 配读出 S ⊤ q S^\top q S ⊤ q ,是 k v ⊤ k v^\top k v ⊤ 约定;本文手算例子和 DeltaNet 一节用 v k ⊤ v k^\top v k ⊤ 约定、右乘读出 S q S\,q S q 。两个约定下的 S 互为转置,全部结论不变–只是读论文时看到外积顺序翻转不要慌,先看它查询是从哪边乘的。
t=1 :v 1 k 1 ⊤ = [ 1 0 ] [ 1 0 ] = [ 1 0 0 0 ] v_1 k_1^\top = \begin{bmatrix}1\\0\end{bmatrix}\begin{bmatrix}1&0\end{bmatrix} = \begin{bmatrix}1&0\\0&0\end{bmatrix} v 1 k 1 ⊤ = [ 1 0 ] [ 1 0 ] = [ 1 0 0 0 ] ,所以
S 1 = [ 1 0 0 0 ] S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}
S 1 = [ 1 0 0 0 ]
t=2 :v 2 k 2 ⊤ = [ 0 1 ] [ 0.6 0.8 ] = [ 0 0 0.6 0.8 ] v_2 k_2^\top = \begin{bmatrix}0\\1\end{bmatrix}\begin{bmatrix}0.6&0.8\end{bmatrix} = \begin{bmatrix}0&0\\0.6&0.8\end{bmatrix} v 2 k 2 ⊤ = [ 0 1 ] [ 0.6 0.8 ] = [ 0 0.6 0 0.8 ] ,叠加后
S 2 = [ 1 0 0.6 0.8 ] S_2 = \begin{bmatrix}1&0\\0.6&0.8\end{bmatrix}
S 2 = [ 1 0.6 0 0.8 ]
t=3 (关键一步,同一个 key 写入新值):v 3 k 3 ⊤ = [ 2 0 ] [ 1 0 ] = [ 2 0 0 0 ] v_3 k_3^\top = \begin{bmatrix}2\\0\end{bmatrix}\begin{bmatrix}1&0\end{bmatrix} = \begin{bmatrix}2&0\\0&0\end{bmatrix} v 3 k 3 ⊤ = [ 2 0 ] [ 1 0 ] = [ 2 0 0 0 ] ,注意它不清除 S 2 S_2 S 2 里已有的 [ 1 0 0 0 ] \begin{bmatrix}1&0\\0&0\end{bmatrix} [ 1 0 0 0 ] ,直接往上叠 :
S 3 = [ 3 0 0.6 0.8 ] S_3 = \begin{bmatrix}3&0\\0.6&0.8\end{bmatrix}
S 3 = [ 3 0.6 0 0.8 ]
第一行变成了 [3, 0]–旧值 1 和新值 2 加在一起,这就是碰撞发生的瞬间。
t=4 :v 4 k 4 ⊤ = [ 3 1 ] [ 0 1 ] = [ 0 3 0 1 ] v_4 k_4^\top = \begin{bmatrix}3\\1\end{bmatrix}\begin{bmatrix}0&1\end{bmatrix} = \begin{bmatrix}0&3\\0&1\end{bmatrix} v 4 k 4 ⊤ = [ 3 1 ] [ 0 1 ] = [ 0 0 3 1 ] ,叠加后
S 4 = [ 3 3 0.6 1.8 ] S_4 = \begin{bmatrix}3&3\\0.6&1.8\end{bmatrix}
S 4 = [ 3 0.6 3 1.8 ]
接下来演示「查询」–先说这个操作是什么 。写入是每个 token 做一次 S + = v k ⊤ S \mathrel{+}= v k^\top S + = v k ⊤ ,把四步叠起来,S 4 S_4 S 4 其实就是四个外积之和:
S 4 = v 1 k 1 ⊤ + v 2 k 2 ⊤ + v 3 k 3 ⊤ + v 4 k 4 ⊤ S_4 = v_1k_1^\top + v_2k_2^\top + v_3k_3^\top + v_4k_4^\top
S 4 = v 1 k 1 ⊤ + v 2 k 2 ⊤ + v 3 k 3 ⊤ + v 4 k 4 ⊤
查询的定义:拿一个向量 k k k 去右乘 S,即 S k S\,k S k 。为什么这样就是「查询」?把上式代入 S 4 k 1 S_4 k_1 S 4 k 1 :
S 4 k 1 = v 1 ( k 1 ⊤ k 1 ) + v 2 ( k 2 ⊤ k 1 ) + v 3 ( k 3 ⊤ k 1 ) + v 4 ( k 4 ⊤ k 1 ) S_4 k_1 = v_1(k_1^\top k_1) + v_2(k_2^\top k_1) + v_3(k_3^\top k_1) + v_4(k_4^\top k_1)
S 4 k 1 = v 1 ( k 1 ⊤ k 1 ) + v 2 ( k 2 ⊤ k 1 ) + v 3 ( k 3 ⊤ k 1 ) + v 4 ( k 4 ⊤ k 1 )
每个 token 的 value 前面多了一个系数–它存入时的 key 与查询 key 的内积 。key 完全匹配(内积=1)value 完整取回;key 正交(内积=0)完全不干扰;部分对齐就按比例混进来一份。这就是「按 key 相似度加权取回 value」,也正是注意力打分的雏形(softmax 注意力把内积换成 e q ⋅ k e^{q\cdot k} e q ⋅ k ,思路相同)。
代入数字 。k 1 ⊤ k 1 = 1 k_1^\top k_1 = 1 k 1 ⊤ k 1 = 1 ,k 2 ⊤ k 1 = 0.6 k_2^\top k_1 = 0.6 k 2 ⊤ k 1 = 0.6 ,k 3 ⊤ k 1 = 1 k_3^\top k_1 = 1 k 3 ⊤ k 1 = 1 ,k 4 ⊤ k 1 = 0 k_4^\top k_1 = 0 k 4 ⊤ k 1 = 0 :
S 4 k 1 = 1 × v 1 + 0.6 × v 2 + 1 × v 3 + 0 × v 4 = [ 1 , 0 ] + [ 0 , 0.6 ] + [ 2 , 0 ] = [ 3 0.6 ] S_4 k_1 = 1{\times}v_1 + 0.6{\times}v_2 + 1{\times}v_3 + 0{\times}v_4 = [1,0] + [0, 0.6] + [2,0] = \begin{bmatrix}3\\0.6\end{bmatrix}
S 4 k 1 = 1 × v 1 + 0.6 × v 2 + 1 × v 3 + 0 × v 4 = [ 1 , 0 ] + [ 0 , 0.6 ] + [ 2 , 0 ] = [ 3 0.6 ]
(直接用矩阵乘验证:取 S 4 S_4 S 4 第一列,同样是 [ 3 , 0.6 ] ⊤ [3, 0.6]^\top [ 3 , 0.6 ] ⊤ 。)
拆开看这个结果:[3.0, 0.6] = 旧值 v 1 = [ 1 , 0 ] v_1=[1,0] v 1 = [ 1 , 0 ] + 新值 v 3 = [ 2 , 0 ] v_3=[2,0] v 3 = [ 2 , 0 ] 各自完整叠加(两个「系数 1」都命中),再加 0.6 份 v 2 v_2 v 2 (k 2 k_2 k 2 和 k 1 k_1 k 1 内积 0.6,部分对齐,混进来 0.6 个 [0,1])。想读「key=[1,0] 对应的最新值」,读出来的却是一锅大杂烩。
delta rule 这边,同样四步 。写入规则换成「先减旧值再加新值」:S t = S t − 1 + ( v t − S t − 1 k t ) k t ⊤ S_t = S_{t-1} + (v_t - S_{t-1}k_t)k_t^\top S t = S t − 1 + ( v t − S t − 1 k t ) k t ⊤ (取 β = 1 \beta = 1 β = 1 ):
t=1 :S 0 k 1 = 0 S_0 k_1 = 0 S 0 k 1 = 0 ,写入 v 1 k 1 ⊤ v_1 k_1^\top v 1 k 1 ⊤ ,得 S 1 = [ 1 0 0 0 ] S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix} S 1 = [ 1 0 0 0 ] (和线性注意力相同);
t=2 :先查旧值 S 1 k 2 = [ 0.6 , 0 ] ⊤ S_1 k_2 = [0.6, 0]^\top S 1 k 2 = [ 0.6 , 0 ] ⊤ (k 2 = [ 0.6 , 0.8 ] k_2 = [0.6, 0.8] k 2 = [ 0.6 , 0.8 ] 在第一维有 0.6 的分量,所以部分命中第一列)。写入差值 u 2 = v 2 − S 1 k 2 = [ − 0.6 , 1 ] u_2 = v_2 - S_1k_2 = [-0.6, 1] u 2 = v 2 − S 1 k 2 = [ − 0.6 , 1 ] ,外积 u 2 k 2 ⊤ = [ − 0.36 − 0.48 0.6 0.8 ] u_2 k_2^\top = \begin{bmatrix}-0.36&-0.48\\0.6&0.8\end{bmatrix} u 2 k 2 ⊤ = [ − 0.36 0.6 − 0.48 0.8 ] 。注意左上角是负数 –它在把 k 2 k_2 k 2 方向上已存的 v 1 v_1 v 1 残余往回抠,不是盲目叠加。得 S 2 = [ 0.64 − 0.48 0.6 0.8 ] S_2 = \begin{bmatrix}0.64&-0.48\\0.6&0.8\end{bmatrix} S 2 = [ 0.64 0.6 − 0.48 0.8 ] ;
t=3 (同一个 key 写新值):先查旧值 S 2 k 3 = [ 0.64 , 0.6 ] ⊤ S_2 k_3 = [0.64, 0.6]^\top S 2 k 3 = [ 0.64 , 0.6 ] ⊤ 。理想情况这里应该读出 [ 1 , 0 ] [1, 0] [ 1 , 0 ] (t=1 写入的 v 1 v_1 v 1 ),但读出的被 t=2 的写入带着偏了。写入差值 u 3 = v 3 − S 2 k 3 = [ 1.36 , − 0.6 ] u_3 = v_3 - S_2k_3 = [1.36, -0.6] u 3 = v 3 − S 2 k 3 = [ 1.36 , − 0.6 ] ,外积只动第一列方向,把 k 3 k_3 k 3 方向修正到 v 3 v_3 v 3 。得 S 3 = [ 2 − 0.48 0 0.8 ] S_3 = \begin{bmatrix}2&-0.48\\0&0.8\end{bmatrix} S 3 = [ 2 0 − 0.48 0.8 ] –第一行回到 2,新值到位;
t=4 :先查 S 3 k 4 = [ − 0.48 , 0.8 ] ⊤ S_3 k_4 = [-0.48, 0.8]^\top S 3 k 4 = [ − 0.48 , 0.8 ] ⊤ ,写入差值 u 4 = v 4 − S 3 k 4 = [ 3.48 , 0.2 ] u_4 = v_4 - S_3k_4 = [3.48, 0.2] u 4 = v 4 − S 3 k 4 = [ 3.48 , 0.2 ] ,只修正 k 4 k_4 k 4 方向。得 S 4 = [ 2 3 0 1 ] S_4 = \begin{bmatrix}2&3\\0&1\end{bmatrix} S 4 = [ 2 0 3 1 ] 。
最后查询 k 1 = [ 1 , 0 ] k_1 = [1, 0] k 1 = [ 1 , 0 ] :线性注意力给出 [ 3.0 , 0.6 ] [3.0, 0.6] [ 3.0 , 0.6 ] (新旧叠加 + 残余污染),delta rule 给出 [ 2.0 , 0.0 ] [2.0, 0.0] [ 2.0 , 0.0 ] –精确返回最新的 v 3 v_3 v 3 。
同一个 key 写两次,线性注意力读出的是两个 value 的叠加 ,delta rule 读出的是新的那个 –这就是「记忆碰撞」(memory collision)。修复它(delta 写入、衰减门)就是后续模型的故事,见《KDA 的来龙去脉》 。
2. 路线二:SSM–从微分方程到遗忘门
线性注意力造出了固定状态,但这个状态只会加法。SSM 是从另一个领域来的–控制理论的状态空间模型–它恰好能精确回答「状态该怎么衰减」,最终和线性注意力在 Mamba-2 处汇合。
2.1 连续 SSM 的定义
单输入单输出(SISO)线性时不变(LTI)状态空间模型:
{ h ′ ( t ) = A h ( t ) + B x ( t ) (状态方程) y ( t ) = C h ( t ) + D x ( t ) (输出方程) \begin{cases}
h'(t) = A\,h(t) + B\,x(t) & \text{(状态方程)}\\[4pt]
y(t) = C\,h(t) + D\,x(t) & \text{(输出方程)}
\end{cases}
{ h ′ ( t ) = A h ( t ) + B x ( t ) y ( t ) = C h ( t ) + D x ( t ) ( 状态方程 ) ( 输出方程 )
符号
维度
含义
x ( t ) ∈ R x(t) \in \mathbb{R} x ( t ) ∈ R
标量
输入信号
h ( t ) ∈ R N h(t) \in \mathbb{R}^N h ( t ) ∈ R N
N 维向量
状态 :到 t 为止历史的压缩
y ( t ) ∈ R y(t) \in \mathbb{R} y ( t ) ∈ R
标量
输出
A ∈ R N × N A \in \mathbb{R}^{N\times N} A ∈ R N × N
矩阵
状态转移:旧记忆如何演化/衰减
B ∈ R N × 1 B \in \mathbb{R}^{N\times 1} B ∈ R N × 1
向量
输入到状态的写入强度
C ∈ R 1 × N C \in \mathbb{R}^{1\times N} C ∈ R 1 × N
向量
状态到输出的读出权重
D ∈ R D \in \mathbb{R} D ∈ R
标量
直通项(深度学习中常设 0,以下略去)
要解决的问题 :神经网络处理的是离散序列 x 1 , x 2 , … x_1, x_2, \dots x 1 , x 2 , … ,需要把微分方程改写为递推式 h t = A ˉ h t − 1 + B ˉ x t h_t = \bar A h_{t-1} + \bar B x_t h t = A ˉ h t − 1 + B ˉ x t ,并求出 A ˉ , B ˉ \bar A, \bar B A ˉ , B ˉ 与 A , B A, B A , B 的精确关系。
2.2 预备知识:矩阵指数
定义 (对标量指数的泰勒级数的直接推广):
e M : = ∑ k = 0 ∞ M k k ! = I + M + M 2 2 ! + M 3 3 ! + ⋯ e^{M} := \sum_{k=0}^{\infty} \frac{M^k}{k!} = I + M + \frac{M^2}{2!} + \frac{M^3}{3!} + \cdots
e M := k = 0 ∑ ∞ k ! M k = I + M + 2 ! M 2 + 3 ! M 3 + ⋯
本推导用到的三条性质 :
微分 :d d t e A t = A e A t = e A t A \dfrac{d}{dt}e^{At} = A\,e^{At} = e^{At}A d t d e A t = A e A t = e A t A (与标量 e a t e^{at} e a t 求导完全平行);
交换性 :A A A 与 e A t e^{At} e A t 可交换(因为 e A t e^{At} e A t 是 A A A 的幂级数);
逆 :( e A t ) − 1 = e − A t (e^{At})^{-1} = e^{-At} ( e A t ) − 1 = e − A t (由 e A t e − A t = e A ( t − t ) = I e^{At}e^{-At} = e^{A(t-t)} = I e A t e − A t = e A ( t − t ) = I )。
直觉:e A Δ e^{A\Delta} e A Δ 是「让系统按自身动力学自由演化 Δ 时间」的算子。A 的特征值实部为负时,它就是各种速度衰减的混合。
2.3 齐次方程的解(无输入情形)
命题 :若 x ( t ) ≡ 0 x(t) \equiv 0 x ( t ) ≡ 0 ,则 h ′ ( t ) = A h ( t ) h'(t) = Ah(t) h ′ ( t ) = A h ( t ) 的解为
h ( t ) = e A ( t − t 0 ) h ( t 0 ) h(t) = e^{A(t - t_0)}\, h(t_0)
h ( t ) = e A ( t − t 0 ) h ( t 0 )
证明 :直接验证满足方程与初值。令 h ( t ) = e A ( t − t 0 ) h ( t 0 ) h(t) = e^{A(t-t_0)}h(t_0) h ( t ) = e A ( t − t 0 ) h ( t 0 ) ,则
h ′ ( t ) = d d t [ e A ( t − t 0 ) ] h ( t 0 ) = A e A ( t − t 0 ) h ( t 0 ) = A h ( t ) . ■ h'(t) = \frac{d}{dt}\Big[e^{A(t-t_0)}\Big]h(t_0) = A\,e^{A(t-t_0)}h(t_0) = A\,h(t). \quad\blacksquare
h ′ ( t ) = d t d [ e A ( t − t 0 ) ] h ( t 0 ) = A e A ( t − t 0 ) h ( t 0 ) = A h ( t ) . ■
含义 :没有输入时,旧记忆按 e A Δ e^{A\Delta} e A Δ 自然演化–这是通解第一项的来源。
2.4 通解推导:积分因子法
定理 :h ′ ( t ) = A h ( t ) + B x ( t ) h'(t) = A h(t) + B x(t) h ′ ( t ) = A h ( t ) + B x ( t ) 满足初值 h ( t 0 ) h(t_0) h ( t 0 ) 的解为
h ( t ) = e A ( t − t 0 ) h ( t 0 ) + ∫ t 0 t e A ( t − s ) B x ( s ) d s h(t) = e^{A(t-t_0)}\,h(t_0) + \int_{t_0}^{t} e^{A(t-s)}\,B\,x(s)\,ds
h ( t ) = e A ( t − t 0 ) h ( t 0 ) + ∫ t 0 t e A ( t − s ) B x ( s ) d s
证明 (积分因子法,四步):
第 1 步:移项。 把含 h 的项移到左边:
h ′ ( t ) − A h ( t ) = B x ( t ) h'(t) - A\,h(t) = B\,x(t)
h ′ ( t ) − A h ( t ) = B x ( t )
第 2 步:乘积分因子 e − A t e^{-At} e − A t 。 两边左乘 e − A t e^{-At} e − A t :
e − A t h ′ ( t ) − e − A t A h ( t ) = e − A t B x ( t ) e^{-At}h'(t) - e^{-At}A\,h(t) = e^{-At}B\,x(t)
e − A t h ′ ( t ) − e − A t A h ( t ) = e − A t B x ( t )
第 3 步:识别乘积导数。 由矩阵指数的微分性质,
d d t [ e − A t h ( t ) ] = − e − A t A h ( t ) + e − A t h ′ ( t ) \frac{d}{dt}\Big[e^{-At}h(t)\Big] = -e^{-At}A\,h(t) + e^{-At}h'(t)
d t d [ e − A t h ( t ) ] = − e − A t A h ( t ) + e − A t h ′ ( t )
恰好等于左边。于是方程变成:
d d t [ e − A t h ( t ) ] = e − A t B x ( t ) \frac{d}{dt}\Big[e^{-At}h(t)\Big] = e^{-At}B\,x(t)
d t d [ e − A t h ( t ) ] = e − A t B x ( t )
这一步是整个推导的「机关」:乘以 e − A t e^{-At} e − A t 后,左端塌缩成一个全导数,方程立刻可积。
第 4 步:两边积分并整理。 从 t 0 t_0 t 0 到 t t t 积分:
e − A t h ( t ) − e − A t 0 h ( t 0 ) = ∫ t 0 t e − A s B x ( s ) d s e^{-At}h(t) - e^{-At_0}h(t_0) = \int_{t_0}^{t} e^{-As}B\,x(s)\,ds
e − A t h ( t ) − e − A t 0 h ( t 0 ) = ∫ t 0 t e − A s B x ( s ) d s
两边左乘 e A t e^{At} e A t (注意 e A t e − A t 0 = e A ( t − t 0 ) e^{At}e^{-At_0} = e^{A(t-t_0)} e A t e − A t 0 = e A ( t − t 0 ) ,且 e A t e − A s = e A ( t − s ) e^{At}e^{-As} = e^{A(t-s)} e A t e − A s = e A ( t − s ) ):
h ( t ) = e A ( t − t 0 ) h ( t 0 ) + ∫ t 0 t e A ( t − s ) B x ( s ) d s ■ h(t) = e^{A(t-t_0)}h(t_0) + \int_{t_0}^{t} e^{A(t-s)}B\,x(s)\,ds \quad\blacksquare
h ( t ) = e A ( t − t 0 ) h ( t 0 ) + ∫ t 0 t e A ( t − s ) B x ( s ) d s ■
两项的物理含义 :旧记忆自由演化 + 区间内每一瞬输入贡献的叠加 (叠加原理)。
2.5 零阶保持(ZOH)离散化
零阶保持假设 :在采样区间 [ t k , t k + Δ ) [t_k,\ t_k + \Delta) [ t k , t k + Δ ) 内,输入保持为采样值:
x ( t k + τ ) = x k , ∀ τ ∈ [ 0 , Δ ) x(t_k + \tau) = x_k, \qquad \forall\,\tau \in [0, \Delta)
x ( t k + τ ) = x k , ∀ τ ∈ [ 0 , Δ )
推导 :在通解中取 t 0 = t k t_0 = t_k t 0 = t k ,t = t k + Δ t = t_k + \Delta t = t k + Δ :
h k + 1 = e A Δ h k + ∫ 0 Δ e A ( Δ − τ ) B x ( t k + τ ) d τ h_{k+1} = e^{A\Delta}h_k + \int_{0}^{\Delta} e^{A(\Delta-\tau)}B\,x(t_k + \tau)\,d\tau
h k + 1 = e A Δ h k + ∫ 0 Δ e A ( Δ − τ ) B x ( t k + τ ) d τ
由 ZOH 假设,x ( t k + τ ) = x k x(t_k+\tau) = x_k x ( t k + τ ) = x k 是常数,提出积分号:
h k + 1 = e A Δ ⏟ A ˉ h k + ( ∫ 0 Δ e A ( Δ − τ ) B d τ ) ⏟ B ˉ x k h_{k+1} = \underbrace{e^{A\Delta}}_{\bar A}\,h_k + \underbrace{\left(\int_0^{\Delta} e^{A(\Delta-\tau)}B\,d\tau\right)}_{\bar B}\,x_k
h k + 1 = A ˉ e A Δ h k + B ˉ ( ∫ 0 Δ e A ( Δ − τ ) B d τ ) x k
计算 B ˉ \bar B B ˉ 的积分 (标量情形最直观;矩阵情形在 A 可逆时同样成立)。换元 s = Δ − τ s = \Delta - \tau s = Δ − τ :
B ˉ = ∫ 0 Δ e A s B d s \bar B = \int_0^{\Delta} e^{As}\,B\,ds
B ˉ = ∫ 0 Δ e A s B d s
对级数逐项积分:
∫ 0 Δ e A s d s = ∫ 0 Δ ∑ k = 0 ∞ ( A s ) k k ! d s = ∑ k = 0 ∞ A k Δ k + 1 ( k + 1 ) ! = A − 1 ( e A Δ − I ) \int_0^{\Delta} e^{As}\,ds = \int_0^{\Delta}\sum_{k=0}^{\infty}\frac{(As)^k}{k!}ds = \sum_{k=0}^{\infty}\frac{A^k\Delta^{k+1}}{(k+1)!} = A^{-1}\big(e^{A\Delta} - I\big)
∫ 0 Δ e A s d s = ∫ 0 Δ k = 0 ∑ ∞ k ! ( A s ) k d s = k = 0 ∑ ∞ ( k + 1 )! A k Δ k + 1 = A − 1 ( e A Δ − I )
最后一步验证:A − 1 ( e A Δ − I ) = A − 1 ∑ k ≥ 1 ( A Δ ) k k ! = ∑ k ≥ 1 A k − 1 Δ k k ! = ∑ j ≥ 0 A j Δ j + 1 ( j + 1 ) ! A^{-1}(e^{A\Delta} - I) = A^{-1}\sum_{k\ge1}\frac{(A\Delta)^k}{k!} = \sum_{k\ge1}\frac{A^{k-1}\Delta^k}{k!} = \sum_{j\ge0}\frac{A^j\Delta^{j+1}}{(j+1)!} A − 1 ( e A Δ − I ) = A − 1 ∑ k ≥ 1 k ! ( A Δ ) k = ∑ k ≥ 1 k ! A k − 1 Δ k = ∑ j ≥ 0 ( j + 1 )! A j Δ j + 1 ✓
结论(ZOH 离散化公式) :
A ˉ = e Δ A , B ˉ = A − 1 ( e Δ A − I ) B \boxed{\ \bar{A} = e^{\Delta A}, \qquad \bar{B} = A^{-1}\big(e^{\Delta A} - I\big)\,B\ }
A ˉ = e Δ A , B ˉ = A − 1 ( e Δ A − I ) B
2.6 小步长近似
当 Δ 很小时,e Δ A = I + Δ A + O ( Δ 2 ) e^{\Delta A} = I + \Delta A + O(\Delta^2) e Δ A = I + Δ A + O ( Δ 2 ) ,于是
B ˉ = A − 1 ( Δ A + O ( Δ 2 ) ) B = Δ B + O ( Δ 2 ) \bar B = A^{-1}\big(\Delta A + O(\Delta^2)\big)B = \Delta B + O(\Delta^2)
B ˉ = A − 1 ( Δ A + O ( Δ 2 ) ) B = Δ B + O ( Δ 2 )
即 B ˉ ≈ Δ B \bar B \approx \Delta B B ˉ ≈ Δ B 。同理 A ˉ ≈ I + Δ A \bar A \approx I + \Delta A A ˉ ≈ I + Δ A (这等价于欧拉法)。
注意 :近似只在 Δ -> 0 时可靠。S4/Mamba 的实际实现都用精确的 ZOH 公式;且 Mamba 中 Δ 是逐 token 变化的,每步现算 e Δ t A e^{\Delta_t A} e Δ t A 。
2.7 数值验证:完整算一遍
取标量系统 A = − 1 A = -1 A = − 1 ,B = 2 B = 2 B = 2 ,C = 1 C = 1 C = 1 ,Δ = 0.5 \Delta = 0.5 Δ = 0.5 。
第 1 步:离散化。
A ˉ = e − 0.5 ≈ 0.6065 , B ˉ = e − 0.5 − 1 − 1 × 2 = ( 1 − 0.6065 ) × 2 ≈ 0.7869 \bar A = e^{-0.5} \approx 0.6065, \qquad \bar B = \frac{e^{-0.5}-1}{-1}\times 2 = (1 - 0.6065)\times 2 \approx 0.7869
A ˉ = e − 0.5 ≈ 0.6065 , B ˉ = − 1 e − 0.5 − 1 × 2 = ( 1 − 0.6065 ) × 2 ≈ 0.7869
第 2 步:递归。 h 0 = 0 h_0 = 0 h 0 = 0 ,输入 x = [ 1 , 1 ] x = [1, 1] x = [ 1 , 1 ] :
k
x k x_k x k
h k = 0.6065 h k − 1 + 0.7869 x k h_k = 0.6065\,h_{k-1} + 0.7869\,x_k h k = 0.6065 h k − 1 + 0.7869 x k
1
1
0.7869
2
1
0.6065×0.7869 + 0.7869 ≈ 1.2642
第 3 步:与连续精确解对拍。 对 x ≡ 1 x \equiv 1 x ≡ 1 的常数输入,连续方程 h ′ = − h + 2 h' = -h + 2 h ′ = − h + 2 的解为 h ( t ) = 2 + ( h 0 − 2 ) e − t h(t) = 2 + (h_0 - 2)e^{-t} h ( t ) = 2 + ( h 0 − 2 ) e − t 。在 t = 1 t = 1 t = 1 (即两步后):
h ( 1 ) = 2 − 2 e − 1 ≈ 2 − 0.7358 = 1.2642 ✓ h(1) = 2 - 2e^{-1} \approx 2 - 0.7358 = 1.2642 \quad✓
h ( 1 ) = 2 − 2 e − 1 ≈ 2 − 0.7358 = 1.2642 ✓
离散递归与连续精确解逐步完全一致 –这正是 ZOH 离散化的性质:在 ZOH 假设成立(输入确为分段常数)时,它不是近似,而是精确等价 。
2.8 卷积形式:递归展开 = 一维卷积
先把「卷积」本身说清楚 (离散因果时序卷积)。输入时序序列 u 1 , u 2 , … , u t u_1, u_2, \dots, u_t u 1 , u 2 , … , u t ,卷积核(权重模板)K 0 , K 1 , K 2 , … K_0, K_1, K_2, \dots K 0 , K 1 , K 2 , … ,其中 K k K_k K k 是与当前时刻相隔 k 步的历史信息的权重 。因果 (causal):只能看过去,看不到未来。形式化:
y t = ∑ τ = 1 t K t − τ ⋅ u τ y_t = \sum_{\tau=1}^{t} K_{t-\tau} \cdot u_\tau
y t = τ = 1 ∑ t K t − τ ⋅ u τ
最简单的数字例子 。输入 u = [ 10 , 20 , 30 ] u = [10, 20, 30] u = [ 10 , 20 , 30 ] ,衰减卷积核 K = [ 1 , 0.5 , 0.25 ] K = [1,\ 0.5,\ 0.25] K = [ 1 , 0.5 , 0.25 ] (越远权重越小):
y 1 = K 0 u 1 = 1 × 10 = 10 y 2 = K 1 u 1 + K 0 u 2 = 0.5 × 10 + 1 × 20 = 25 y 3 = K 2 u 1 + K 1 u 2 + K 0 u 3 = 0.25 × 10 + 0.5 × 20 + 1 × 30 = 42.5 \begin{aligned}
y_1 &= K_0 u_1 = 1\times10 = 10 \\
y_2 &= K_1 u_1 + K_0 u_2 = 0.5\times10 + 1\times20 = 25 \\
y_3 &= K_2 u_1 + K_1 u_2 + K_0 u_3 = 0.25\times10 + 0.5\times20 + 1\times30 = 42.5
\end{aligned}
y 1 y 2 y 3 = K 0 u 1 = 1 × 10 = 10 = K 1 u 1 + K 0 u 2 = 0.5 × 10 + 1 × 20 = 25 = K 2 u 1 + K 1 u 2 + K 0 u 3 = 0.25 × 10 + 0.5 × 20 + 1 × 30 = 42.5
每一步都把全部历史按距离加权求和 ,权重模板固定、随距离翻牌–这就是卷积。y 1 y_1 y 1 只用了 u 1 u_1 u 1 (因果),三个输出用的是同一套 K(时不变)。带着这两个属性看 SSM 的卷积形式,就是把 K k K_k K k 具体化为 C A ˉ k B ˉ C\bar A^k \bar B C A ˉ k B ˉ 。
命题 :离散 SSM h t = A ˉ h t − 1 + B ˉ x t h_t = \bar A h_{t-1} + \bar B x_t h t = A ˉ h t − 1 + B ˉ x t ,y t = C h t y_t = C h_t y t = C h t (设 h 0 = 0 h_0 = 0 h 0 = 0 )等价于
y = K ˉ ∗ x , K ˉ = ( C B ˉ , C A ˉ B ˉ , C A ˉ 2 B ˉ , … , C A ˉ L − 1 B ˉ ) y = \bar K * x, \qquad \bar K = \big(C\bar B,\ C\bar A\bar B,\ C\bar A^2\bar B,\ \dots,\ C\bar A^{L-1}\bar B\big)
y = K ˉ ∗ x , K ˉ = ( C B ˉ , C A ˉ B ˉ , C A ˉ 2 B ˉ , … , C A ˉ L − 1 B ˉ )
证明 (直接展开):
h t = ∑ i = 1 t A ˉ t − i B ˉ x i h_t = \sum_{i=1}^{t} \bar A^{\,t-i}\bar B\, x_i
h t = i = 1 ∑ t A ˉ t − i B ˉ x i
(归纳:h 1 = B ˉ x 1 h_1 = \bar B x_1 h 1 = B ˉ x 1 ;设 h t − 1 = ∑ i ≤ t − 1 A ˉ t − 1 − i B ˉ x i h_{t-1} = \sum_{i\le t-1}\bar A^{t-1-i}\bar B x_i h t − 1 = ∑ i ≤ t − 1 A ˉ t − 1 − i B ˉ x i ,则 h t = A ˉ h t − 1 + B ˉ x t = ∑ i ≤ t − 1 A ˉ t − i B ˉ x i + B ˉ x t h_t = \bar A h_{t-1} + \bar B x_t = \sum_{i\le t-1}\bar A^{t-i}\bar B x_i + \bar B x_t h t = A ˉ h t − 1 + B ˉ x t = ∑ i ≤ t − 1 A ˉ t − i B ˉ x i + B ˉ x t ✓)
代入输出方程:
y t = ∑ i = 1 t C A ˉ t − i B ˉ x i = ∑ j = 0 t − 1 K ˉ j x t − j ■ y_t = \sum_{i=1}^{t} C\bar A^{\,t-i}\bar B\, x_i = \sum_{j=0}^{t-1} \bar K_j\, x_{t-j} \quad\blacksquare
y t = i = 1 ∑ t C A ˉ t − i B ˉ x i = j = 0 ∑ t − 1 K ˉ j x t − j ■
数值验证 (§2.7 的系统,输入 x = [ 2 , 4 , 0 , 8 ] x=[2,4,0,8] x = [ 2 , 4 , 0 , 8 ] ):
卷积核:K ˉ j = 0.6065 j × 0.7869 \bar K_j = 0.6065^j \times 0.7869 K ˉ j = 0.606 5 j × 0.7869 ,即 [ 0.7869 , 0.4772 , 0.2894 , 0.1755 ] [0.7869,\ 0.4772,\ 0.2894,\ 0.1755] [ 0.7869 , 0.4772 , 0.2894 , 0.1755 ]
y 4 = 0.7869 × 8 + 0.4772 × 0 + 0.2894 × 4 + 0.1755 × 2 = 6.295 + 1.158 + 0.351 ≈ 7.804 y_4 = 0.7869{\times}8 + 0.4772{\times}0 + 0.2894{\times}4 + 0.1755{\times}2 = 6.295 + 1.158 + 0.351 \approx 7.804
y 4 = 0.7869 × 8 + 0.4772 × 0 + 0.2894 × 4 + 0.1755 × 2 = 6.295 + 1.158 + 0.351 ≈ 7.804
递归验证:h 1 = 1.574 , h 2 = 4.102 , h 3 = 2.488 , h 4 = 0.6065 × 2.488 + 0.7869 × 8 ≈ 1.509 + 6.295 = 7.804 h_1=1.574,\ h_2=4.102,\ h_3=2.488,\ h_4=0.6065{\times}2.488+0.7869{\times}8 \approx 1.509+6.295 = 7.804 h 1 = 1.574 , h 2 = 4.102 , h 3 = 2.488 , h 4 = 0.6065 × 2.488 + 0.7869 × 8 ≈ 1.509 + 6.295 = 7.804 ✓
推论 :LTI 系统拥有双形式–训练用卷积(FFT,并行),推理用递归(O(1)/token)。Mamba 使 B , C , Δ B, C, \Delta B , C , Δ 依赖输入后,K ˉ \bar K K ˉ 不再固定,卷积形式失效,改由并行扫描承担训练侧并行。
2.9 从 Δ 到遗忘门:通往 Mamba 的最后一环
回看 A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A :若 A 为负定(如 Mamba-2 取 A = − a ⋅ I A = -a\cdot I A = − a ⋅ I ,a > 0 a>0 a > 0 ),则
A ˉ t = e − Δ t a ∈ ( 0 , 1 ) \bar A_t = e^{-\Delta_t a} \in (0, 1)
A ˉ t = e − Δ t a ∈ ( 0 , 1 )
就是一个衰减门 :
Δ t → ∞ \Delta_t \to \infty Δ t → ∞ :A ˉ t → 0 \bar A_t \to 0 A ˉ t → 0 – 清空旧记忆,只看当前输入;
Δ t → 0 \Delta_t \to 0 Δ t → 0 :A ˉ t → 1 \bar A_t \to 1 A ˉ t → 1 – 冻结状态,忽略当前输入。
Mamba 让 Δ t = s o f t p l u s ( L i n e a r ( x t ) ) \Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t)) Δ t = softplus ( Linear ( x t )) ,离散化公式每步现算,于是「步长」变成了看内容的标量遗忘门 。
2.10 附录:双线性变换离散化对比
除 ZOH 外,S4 论文实际使用的是双线性变换(Tustin 方法) :用梯形法则近似积分,
A ˉ = ( I − Δ 2 A ) − 1 ( I + Δ 2 A ) , B ˉ = ( I − Δ 2 A ) − 1 Δ B \bar A = \Big(I - \frac{\Delta}{2}A\Big)^{-1}\Big(I + \frac{\Delta}{2}A\Big), \qquad
\bar B = \Big(I - \frac{\Delta}{2}A\Big)^{-1}\Delta B
A ˉ = ( I − 2 Δ A ) − 1 ( I + 2 Δ A ) , B ˉ = ( I − 2 Δ A ) − 1 Δ B
方法
性质
适用
ZOH
ZOH 假设下精确 ;保稳定性;公式含 e Δ A e^{\Delta A} e Δ A
Mamba 系列:Δ t \Delta_t Δ t 逐 token 变化,每步现算矩阵指数(S4D/Mamba 的对角 A 使 e Δ t A e^{\Delta_t A} e Δ t A 就是逐元素指数)
双线性
保稳定性;避免矩阵指数;把左半平面解析映射到单位圆内
S4 原论文:A 为 HiPPO 矩阵(非对角),双线性把它变成有理函数,配合 Cauchy 核计算
一句话收束:离散化方法的选择是被 A 的结构决定的 –对角 A 配 ZOH(指数按元素算),非对角 HiPPO A 配双线性(变成多项式比值,可借 Cauchy 核求解)。
3. Mamba 一族:从 S4 到 Mamba-2
上面这条 SSM 线,从深度学习视角走过了三个关键站点:
S4(2021) :把连续 SSM 正经地搬进序列建模。A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A 、B ˉ = A − 1 ( e Δ A − I ) B \bar B = A^{-1}(e^{\Delta A} - I)B B ˉ = A − 1 ( e Δ A − I ) B 都是固定矩阵 (与输入无关),于是 LTI 成立、递归 = 卷积,训练用 FFT 卷积、推理用 O ( 1 ) O(1) O ( 1 ) 递归(双形式,§2.8)。A 用 HiPPO 矩阵初始化(记住长历史的专门构造),离散化用双线性变换。限制也很明显:参数不随输入变,「记忆怎么衰减」在推理前就定死了。
Mamba / S6(2023) :选择性改造。让 B , C , Δ B, C, \Delta B , C , Δ 都依赖输入(Δ t = s o f t p l u s ( L i n e a r ( x t ) ) \Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t)) Δ t = softplus ( Linear ( x t )) ),每步现算 e Δ t A e^{\Delta_t A} e Δ t A –「步长」变成了看内容的遗忘门(§2.9)。代价是 LTI 被打破、固定卷积核 K ˉ \bar K K ˉ 失效,训练侧改用并行扫描(associative scan);A 限制为对角结构使 e Δ t A e^{\Delta_t A} e Δ t A 逐元素可算。这一步是「SSM 学会看内容」的关键。
Mamba-2 / SSD(2024) :与线性注意力形式统一。状态改写成外积形式 S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_t S_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤ ,矩阵 A A A 退化为标量衰减 α t = e − Δ t a \alpha_t = e^{-\Delta_t a} α t = e − Δ t a –从这一步起,SSM 与线性注意力在数学上就是同一个东西的两个记法 (SSD 框架)。训练用 chunkwise:chunk 内 Q K ⊤ QK^\top Q K ⊤ 小注意力 + chunk 间递归,比纯 scan 的 GEMM 利用率高得多。
对照第 1 节的线性注意力递推 S t = S t − 1 + ϕ ( k t ) v t ⊤ S_t = S_{t-1} + \phi(k_t)v_t^\top S t = S t − 1 + ϕ ( k t ) v t ⊤ :Mamba-2 只多了一个乘在旧状态上的标量 α t \alpha_t α t 。两条路线在 Mamba-2 处汇合 :线性注意力给出了外积状态的样子,SSM 给出了衰减门。
4. 全链路一图流
h ′ = A h + B x ⏟ 连续 SSM → 积分因子法 h ( t + Δ ) = e A Δ h ( t ) + ∫ 0 Δ e A ( Δ − τ ) B x d τ ⏟ 通解:旧记忆衰减 + 输入累积 → ZOH h t = A ˉ h t − 1 + B ˉ x t ⏟ 离散递归(推理用) → 展开 y = K ˉ ∗ x ⏟ 卷积(训练用) \underbrace{h' = Ah + Bx}_{\text{连续 SSM}}
\xrightarrow{\text{积分因子法}}
\underbrace{h(t+\Delta) = e^{A\Delta}h(t) + \int_0^\Delta e^{A(\Delta-\tau)}Bx\,d\tau}_{\text{通解:旧记忆衰减 + 输入累积}}
\xrightarrow{\text{ZOH}}
\underbrace{h_t = \bar A h_{t-1} + \bar B x_t}_{\text{离散递归(推理用)}}
\xrightarrow{\text{展开}}
\underbrace{y = \bar K * x}_{\text{卷积(训练用)}}
连续 SSM h ′ = A h + B x 积分因子法 通解 : 旧记忆衰减 + 输入累积 h ( t + Δ ) = e A Δ h ( t ) + ∫ 0 Δ e A ( Δ − τ ) B x d τ ZOH 离散递归 ( 推理用 ) h t = A ˉ h t − 1 + B ˉ x t 展开 卷积 ( 训练用 ) y = K ˉ ∗ x
线性注意力侧:s o f t m a x ( Q K ⊤ ) V \mathrm{softmax}(QK^\top)V softmax ( Q K ⊤ ) V --核化–> ϕ ( Q ) ( ϕ ( K ) ⊤ V ) \phi(Q)(\phi(K)^\top V) ϕ ( Q ) ( ϕ ( K ) ⊤ V ) --逐 token–> S t = S t − 1 + ϕ ( k t ) v t ⊤ S_t = S_{t-1} + \phi(k_t)v_t^\top S t = S t − 1 + ϕ ( k t ) v t ⊤ 。
两线在 Mamba-2 汇合:S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_t S_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤ 。
此后:GDN 在写入侧加 delta rule(可删除的写入);KDA 把标量门打开成逐通道门 D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 并给衰减加下界。骨架始终是同一条:状态 ×(衰减/删除算子)+(写入项) 。后续演化见《KDA 的来龙去脉》 。
总结
线性注意力 :去掉 softmax 的耦合非线性 + 核化把非线性提前,Q ( K ⊤ V ) Q(K^\top V) Q ( K ⊤ V ) 结合律重新可用,O ( n 2 ) → O ( n ) O(n^2) \to O(n) O ( n 2 ) → O ( n ) ;代价是纯加性状态,只会叠加不会删除(记忆碰撞);
SSM 一条链推完 :齐次解 e A t h 0 e^{At}h_0 e A t h 0 -> 积分因子法得通解 -> ZOH 假设下提出 x k x_k x k 得离散递归,A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A 、B ˉ = A − 1 ( e Δ A − I ) B \bar B = A^{-1}(e^{\Delta A} - I)B B ˉ = A − 1 ( e Δ A − I ) B ;
ZOH 不是近似 :输入分段常数时离散递归与连续解逐步相等(§2.7 数值对拍 1.2642 = 1.2642);
递归 = 卷积 :LTI 的双形式是 S4 训练(FFT 卷积)与推理(O(1) 递归)各取所长的根基;
A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A 就是遗忘门 :这是从控制论到 Mamba/KDA 的思想连续性;
两条路线在 Mamba-2 汇合 :外积状态(线性注意力)+ 标量衰减门(SSM)= S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_t S_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤ ,此后 delta rule 和逐通道门的演化都在这条骨架上进行。
参考 :
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