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

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

线性注意力和 SSM(State Space Model,状态空间模型)是序列建模的两条技术路线,也是后来 KDA/GDN/DeltaNet 一族模型的两个源头:线性注意力贡献了「外积记忆状态 + 结合律」,SSM 贡献了「衰减门 + 递推骨架」。本文把这两条路线的计算过程完整推一遍,每步都有推导、数值验算和工程上的理由。

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

标准 softmax 注意力在生成文本时,必须把每个历史 token 的 Key 和 Value 都存下来(KV cache)–因为每生成一个新 token,都要和前面所有 token 重新算一遍注意力。序列越长,这笔账越贵:

开销 softmax 注意力(KV cache) 固定状态(线性注意力一族)
解码显存 O(nd)O(n \cdot d)(全部历史的 KV 都得缓存) O(d2)O(d^2)(一个固定矩阵 S,与 n 无关)
单步解码 O(nd)O(n \cdot d)(扫全部历史) O(d2)O(d^2)(更新 S + 读出,两次矩阵乘)
预填充 O(n2d)O(n^2 d) O(nd2)O(n d^2)
上下文长度翻倍 显存、延迟一起翻倍 完全不变

(单头视角,d 为特征维度。)

几十万 token 的上下文就能把 KV cache 撑到几十 GB。想摆脱这笔账,就得把「全部历史的 KV」压缩成一个固定大小的状态,而且这个状态要会写、会改、会忘–两个源头各给出了一半答案:线性注意力先把固定状态造出来,SSM 教会它怎么遗忘。


1. 路线一:线性注意力–两个前提,缺一不可

线性注意力能成立,靠的是两件事:去掉 softmax 的非线性让 Q/K 先结合掉。更准确地说,这是同一枚硬币的两面–softmax 必须作用在耦合后的 QKQK^\top 上,它既是非线性的来源,也是耦合的来源;把它换成可分离的 ϕ(q)ϕ(k)\phi(q)^\top\phi(k),非线性和耦合一起消失,结合律才重新可用。后面 KDA 一族的全部推导都建立在这个前提上,值得把两个条件一个一个讲透。

1.1 前提一:去掉 softmax,结合律才可用

先看没有 softmax 的裸注意力 Attn(Q,K,V)=(QK)V\mathrm{Attn}(Q, K, V) = (QK^\top)V,复杂度账本:

括号化 计算路径 复杂度 主导项
(QK)V(QK^\top)V 先算出 n×nn \times n 矩阵,再乘 V O(dn2)+O(nnd)=O(dn2)O(d\cdot n^2) + O(n\cdot nd) = O(dn^2) n(平方)
Q(KV)Q(K^\top V) 先算 KVRd×dK^\top V \in \mathbb{R}^{d \times d}(与 n 无关),再左乘 Q O(nd2)+O(dnd)=O(nd2)O(n\cdot d^2) + O(d\cdot nd) = O(nd^2) d(线性)

同一个数学式,两种括号化,复杂度差一个 n 的幂–这就是结合律的诱惑:矩阵乘法满足 (AB)C=A(BC)(AB)C = A(BC),先算哪边是自由的。裸注意力(线性代数层面)本来就能线性化。

那 softmax 挡在哪? 标准注意力是 softmax(QK)V\mathrm{softmax}(QK^\top)V,softmax 在 QKQK^\top 之后、乘 V 之前介入,并且它是行内非线性:第 i 行(第 i 个 query 对所有 key 的打分)必须整体过 exi/jexje^{x_i}/\sum_j e^{x_j}–指数逐项非线性 + 行内归一化耦合。两个性质各堵死一条路:

  1. 非线性破坏结合律:softmax 不是线性算子,softmax(QK)VQsoftmax(KV)\mathrm{softmax}(QK^\top)V \ne Q\,\mathrm{softmax}'(K^\top V),中间结果没法先合并;n×nn \times n 矩阵必须先完整算出来;
  2. 归一化引入全局耦合:分母 jeqikj\sum_j e^{q_i \cdot k_j} 依赖该行所有 n 个 key,哪怕只想算一个 query 的输出,也得先把整行算完–信息在 key 维度上全局耦合,没有可以「先结合掉」的独立块。

所以平方复杂度不是矩阵乘法的错,是 softmax 的作用位置的错:它非要等 Q 和 K 耦合完才动手。要线性化,就得把这个非线性从耦合点上搬走。

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

搬法就是核化(kernelization)。把注意力抽象成通用相似度函数 sim(q,k)\mathrm{sim}(q, k)–它不必是 softmax 的指数,多项式相似度、RBF 核都属此类。数学依据是核技巧:只要 sim 非负(Mercer 条件),就存在特征映射 ϕ()\phi(\cdot) 使

sim(q,k)=ϕ(q)ϕ(k)\mathrm{sim}(q, k) = \phi(q)^\top \phi(k)

注意这个形式的本质:非线性被吸收进 ϕ\phi,分别作用于 q 和 k 各自身上,相似度本身变成线性的内积。softmax 的 eqke^{q \cdot k} 作用在耦合后的标量上;核化的 ϕ(q)ϕ(k)\phi(q)^\top\phi(k) 让 q 和 k 各自先做完所有非线性变换,最后只留一次线性内积。

线性注意力取 ϕ(x)=ELU(x)+1\phi(x) = \mathrm{ELU}(x) + 1(Katharopoulos et al. 2020)。先把 ELU(Exponential Linear Unit,指数线性单元)本身说清楚,它是个分段函数:

ELU(x)={xx>0ex1x0\mathrm{ELU}(x) = \begin{cases} x & x > 0 \\ e^{x} - 1 & x \le 0 \end{cases}

  • 正数区:原样输出(和 ReLU 一样);
  • 负数区:输出 ex1(1,0)e^x - 1 \in (-1, 0),一条平滑曲线,越负越接近 -1 但永远不到 -1。

所以 ELU(x)+1\mathrm{ELU}(x) + 1 的取值范围是 (0,)(0, \infty):恒正。算两个具体值感受一下:x=0x = 0ELU(0)+1=0+1=1\mathrm{ELU}(0)+1 = 0 + 1 = 1;x=3x = -3e31+1=e30.05e^{-3} - 1 + 1 = e^{-3} \approx 0.05,很小但仍是正数。这个恒正不是自选的装饰,是核分解存在的条件(Mercer 条件要求相似度非负):ϕ(q)ϕ(k)\phi(q)^\top\phi(k) 是 d 个正数乘正数再求和,结果必为正,才能扮演「打分」的角色。于是:

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

结合律在非线性世界里重新可用:先算 ϕ(K)VRd×d\phi(K)^\top V \in \mathbb{R}^{d \times d},复杂度回到 O(nd2)O(nd^2)两个前提到此汇成一句话:非线性提前,耦合消失,结合律重新可用。写成逐 token 的归一化形式(softmax 的归一化也没丢,变成显式分母):

Attn(q)=iϕ(q)ϕ(ki)viiϕ(q)ϕ(ki)=ϕ(q)Sϕ(q)z,S=iϕ(ki)vi\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

两个新符号 S 和 z 分别是:

  • S=iϕ(ki)viRd×dS = \sum_i \phi(k_i)v_i^\top \in \mathbb{R}^{d \times d}:分子里的 KV 外积矩阵,存「key-value 关联」的记忆本体;
  • z=iϕ(ki)Rdz = \sum_i \phi(k_i) \in \mathbb{R}^{d}:分母里的 key 累积和,就是 softmax 分母 ieqki\sum_i e^{q\cdot k_i} 的核化版–归一化因子。

先说清 z,因为后续论文里它最容易被略写。z 在公式里承担的角色就是归一化:没有它,输出会随序列变长而无界增长。但后面 DeltaNet/GDN/KDA 的论文里,递推式往往只写 SS 的更新,z 要么藏在一句「输出再过 RMSNorm」的描述里,要么干脆不提–因为实测发现把归一化换成 RMSNorm 效果更好,z 就被简化掉了。所以读者常遇到的困惑是:看 KDA 论文时公式里根本没有 z,翻早期线性注意力文献才发现 z 是这里从 softmax 分母继承下来的归一化项。一句话:z = 归一化分母的累积状态,S = 关联记忆本体;后续模型改用 RMSNorm 后 z 退场,但 S 的更新规则一路演进到 KDA。

分子分母都变成对 S 的查询,S 一遍扫过序列累积即可–复杂度对 n 线性。更关键的是 S 可以写成递推,一步一读出(z 同样逐步累积):

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 视角(Schlag et al., Linear Transformers Are Secretly Fast Weight Programmers):把 S 看作一块快速权重–模型真正的参数(普通权重)训练完就固定了,而 S 每来一个 token 就被改写一次,相当于一个「随输入不断更新的小型权重」。写入靠 ϕ(kt)vt\phi(k_t)v_t^\top(每个 token 都在给这块权重编程),读出靠查询 ϕ(qt)St\phi(q_t)^\top S_t(拿当前问题去这块权重里查答案)。所以这个视角下,序列处理 = 一边更新权重、一边用权重答题。

回过头总结一下,这两个前提各自换来了什么:

前提 换来的直接好处 对后续模型的影响
抛弃 softmax 非线性 结合律可用,O(n2)O(n)O(n^2) \to O(n) 注意力矩阵 n×nn\times n 消失,换成固定大小状态 SRd×dS \in \mathbb{R}^{d \times d}
非线性提前到 Q/K 各自身上 ϕ(q)ϕ(k)\phi(q)^\top\phi(k) 可分离 S 可以递推累积 -> RNN 化 -> fast weights -> 一切后续演化的载体

注意代价也在表里:换来的 S 只会加法。这就引出了线性注意力最大的问题。

1.3 纯加性状态的问题:记忆碰撞

纯加性状态的核心问题:state 是纯加性的(St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top),写入只有「叠加」,没有「删除」。序列长度远超 state 有效容量(d×dd \times d 矩阵只能存这么多关联)时,不同 kvk \to v 关联互相干扰,旧信息永远赖在状态里,新信息无法覆盖–记忆像一个只进不出的仓库。

数值演示。设 dk=dv=2d_k = d_v = 2,S 从零开始,依次写入四个 token:

token key value 说明
t=1 k1=[1,0]k_1 = [1, 0] v1=[1,0]v_1 = [1, 0]
t=2 k2=[0.6,0.8]k_2 = [0.6, 0.8] v2=[0,1]v_2 = [0, 1]
t=3 k3=[1,0]k_3 = [1, 0] v3=[2,0]v_3 = [2, 0] k1k_1 同一个 key,写入新值
t=4 k4=[0,1]k_4 = [0, 1] v4=[3,1]v_4 = [3, 1]

线性注意力这边,把每一步都算出来。写入规则是 St=St1+vtktS_t = S_{t-1} + v_t k_t^\top,每个 vkv k^\top 是一个 2×2 外积。

为什么是 vkv k^\top 而不是 kvk v^\top? 外积的顺序决定「谁进谁出」。vkv k^\top 是一个从 key 空间到 value 空间的线性映射:key 从右边进,value 从左边出–直接验证:(vk)k=v(kk)(v k^\top)\,k = v\,(k^\top k),查询 key 与存储 key 的内积落在系数上,value 原样输出。如果反过来存 kvk v^\top 还从右边乘,得到 k(vk)k\,(v^\top k)–变成「拿 key 查、按 value 相似度返回 key」,角色整个反了,查不出想要的东西。要让 kvk v^\top 布局工作,查询必须从左边进:行向量 qS=i(qki)viq^\top S = \sum_i (q^\top k_i)\,v_i^\top,结论完全一样。所以两种写法都存在:§1.2 的 S=iϕ(ki)viS = \sum_i \phi(k_i)v_i^\top 配左乘读出 ϕ(q)S\phi(q)^\top S、KDA 论文的 SS 配读出 SqS^\top q,是 kvk v^\top 约定;本文手算例子和 DeltaNet 一节用 vkv k^\top 约定、右乘读出 SqS\,q。两个约定下的 S 互为转置,全部结论不变–只是读论文时看到外积顺序翻转不要慌,先看它查询是从哪边乘的。

t=1:v1k1=[10][10]=[1000]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},所以

S1=[1000]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}

t=2:v2k2=[01][0.60.8]=[000.60.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},叠加后

S2=[100.60.8]S_2 = \begin{bmatrix}1&0\\0.6&0.8\end{bmatrix}

t=3(关键一步,同一个 key 写入新值):v3k3=[20][10]=[2000]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}注意它不清除 S2S_2 里已有的 [1000]\begin{bmatrix}1&0\\0&0\end{bmatrix},直接往上叠:

S3=[300.60.8]S_3 = \begin{bmatrix}3&0\\0.6&0.8\end{bmatrix}

第一行变成了 [3, 0]–旧值 1 和新值 2 加在一起,这就是碰撞发生的瞬间。

t=4:v4k4=[31][01]=[0301]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},叠加后

S4=[330.61.8]S_4 = \begin{bmatrix}3&3\\0.6&1.8\end{bmatrix}

接下来演示「查询」–先说这个操作是什么。写入是每个 token 做一次 S+=vkS \mathrel{+}= v k^\top,把四步叠起来,S4S_4 其实就是四个外积之和:

S4=v1k1+v2k2+v3k3+v4k4S_4 = v_1k_1^\top + v_2k_2^\top + v_3k_3^\top + v_4k_4^\top

查询的定义:拿一个向量 kk右乘 S,即 SkS\,k。为什么这样就是「查询」?把上式代入 S4k1S_4 k_1:

S4k1=v1(k1k1)+v2(k2k1)+v3(k3k1)+v4(k4k1)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)

每个 token 的 value 前面多了一个系数–它存入时的 key 与查询 key 的内积。key 完全匹配(内积=1)value 完整取回;key 正交(内积=0)完全不干扰;部分对齐就按比例混进来一份。这就是「按 key 相似度加权取回 value」,也正是注意力打分的雏形(softmax 注意力把内积换成 eqke^{q\cdot k},思路相同)。

代入数字k1k1=1k_1^\top k_1 = 1,k2k1=0.6k_2^\top k_1 = 0.6,k3k1=1k_3^\top k_1 = 1,k4k1=0k_4^\top k_1 = 0:

S4k1=1×v1+0.6×v2+1×v3+0×v4=[1,0]+[0,0.6]+[2,0]=[30.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}

(直接用矩阵乘验证:取 S4S_4 第一列,同样是 [3,0.6][3, 0.6]^\top。)

拆开看这个结果:[3.0, 0.6] = 旧值 v1=[1,0]v_1=[1,0] + 新值 v3=[2,0]v_3=[2,0] 各自完整叠加(两个「系数 1」都命中),再加 0.6 份 v2v_2(k2k_2k1k_1 内积 0.6,部分对齐,混进来 0.6 个 [0,1])。想读「key=[1,0] 对应的最新值」,读出来的却是一锅大杂烩。

delta rule 这边,同样四步。写入规则换成「先减旧值再加新值」:St=St1+(vtSt1kt)ktS_t = S_{t-1} + (v_t - S_{t-1}k_t)k_t^\top(取 β=1\beta = 1):

  • t=1:S0k1=0S_0 k_1 = 0,写入 v1k1v_1 k_1^\top,得 S1=[1000]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}(和线性注意力相同);
  • t=2:先查旧值 S1k2=[0.6,0]S_1 k_2 = [0.6, 0]^\top(k2=[0.6,0.8]k_2 = [0.6, 0.8] 在第一维有 0.6 的分量,所以部分命中第一列)。写入差值 u2=v2S1k2=[0.6,1]u_2 = v_2 - S_1k_2 = [-0.6, 1],外积 u2k2=[0.360.480.60.8]u_2 k_2^\top = \begin{bmatrix}-0.36&-0.48\\0.6&0.8\end{bmatrix}。注意左上角是负数–它在把 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=3(同一个 key 写新值):先查旧值 S2k3=[0.64,0.6]S_2 k_3 = [0.64, 0.6]^\top。理想情况这里应该读出 [1,0][1, 0](t=1 写入的 v1v_1),但读出的被 t=2 的写入带着偏了。写入差值 u3=v3S2k3=[1.36,0.6]u_3 = v_3 - S_2k_3 = [1.36, -0.6],外积只动第一列方向,把 k3k_3 方向修正到 v3v_3。得 S3=[20.4800.8]S_3 = \begin{bmatrix}2&-0.48\\0&0.8\end{bmatrix}–第一行回到 2,新值到位;
  • t=4:先查 S3k4=[0.48,0.8]S_3 k_4 = [-0.48, 0.8]^\top,写入差值 u4=v4S3k4=[3.48,0.2]u_4 = v_4 - S_3k_4 = [3.48, 0.2],只修正 k4k_4 方向。得 S4=[2301]S_4 = \begin{bmatrix}2&3\\0&1\end{bmatrix}

最后查询 k1=[1,0]k_1 = [1, 0]:线性注意力给出 [3.0,0.6][3.0, 0.6](新旧叠加 + 残余污染),delta rule 给出 [2.0,0.0][2.0, 0.0]精确返回最新的 v3v_3

同一个 key 写两次,线性注意力读出的是两个 value 的叠加,delta rule 读出的是新的那个–这就是「记忆碰撞」(memory collision)。修复它(delta 写入、衰减门)就是后续模型的故事,见《KDA 的来龙去脉》


2. 路线二:SSM–从微分方程到遗忘门

线性注意力造出了固定状态,但这个状态只会加法。SSM 是从另一个领域来的–控制理论的状态空间模型–它恰好能精确回答「状态该怎么衰减」,最终和线性注意力在 Mamba-2 处汇合。

2.1 连续 SSM 的定义

单输入单输出(SISO)线性时不变(LTI)状态空间模型:

{h(t)=Ah(t)+Bx(t)(状态方程)y(t)=Ch(t)+Dx(t)(输出方程)\begin{cases} h'(t) = A\,h(t) + B\,x(t) & \text{(状态方程)}\\[4pt] y(t) = C\,h(t) + D\,x(t) & \text{(输出方程)} \end{cases}

符号 维度 含义
x(t)Rx(t) \in \mathbb{R} 标量 输入信号
h(t)RNh(t) \in \mathbb{R}^N N 维向量 状态:到 t 为止历史的压缩
y(t)Ry(t) \in \mathbb{R} 标量 输出
ARN×NA \in \mathbb{R}^{N\times N} 矩阵 状态转移:旧记忆如何演化/衰减
BRN×1B \in \mathbb{R}^{N\times 1} 向量 输入到状态的写入强度
CR1×NC \in \mathbb{R}^{1\times N} 向量 状态到输出的读出权重
DRD \in \mathbb{R} 标量 直通项(深度学习中常设 0,以下略去)

要解决的问题:神经网络处理的是离散序列 x1,x2,x_1, x_2, \dots,需要把微分方程改写为递推式 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,并求出 Aˉ,Bˉ\bar A, \bar BA,BA, B 的精确关系。

2.2 预备知识:矩阵指数

定义(对标量指数的泰勒级数的直接推广):

eM:=k=0Mkk!=I+M+M22!+M33!+e^{M} := \sum_{k=0}^{\infty} \frac{M^k}{k!} = I + M + \frac{M^2}{2!} + \frac{M^3}{3!} + \cdots

本推导用到的三条性质:

  1. 微分:ddteAt=AeAt=eAtA\dfrac{d}{dt}e^{At} = A\,e^{At} = e^{At}A(与标量 eate^{at} 求导完全平行);
  2. 交换性:AAeAte^{At} 可交换(因为 eAte^{At}AA 的幂级数);
  3. :(eAt)1=eAt(e^{At})^{-1} = e^{-At}(由 eAteAt=eA(tt)=Ie^{At}e^{-At} = e^{A(t-t)} = I)。

直觉:eAΔe^{A\Delta} 是「让系统按自身动力学自由演化 Δ 时间」的算子。A 的特征值实部为负时,它就是各种速度衰减的混合。

2.3 齐次方程的解(无输入情形)

命题:若 x(t)0x(t) \equiv 0,则 h(t)=Ah(t)h'(t) = Ah(t) 的解为

h(t)=eA(tt0)h(t0)h(t) = e^{A(t - t_0)}\, h(t_0)

证明:直接验证满足方程与初值。令 h(t)=eA(tt0)h(t0)h(t) = e^{A(t-t_0)}h(t_0),则

h(t)=ddt[eA(tt0)]h(t0)=AeA(tt0)h(t0)=Ah(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

含义:没有输入时,旧记忆按 eAΔe^{A\Delta} 自然演化–这是通解第一项的来源。

2.4 通解推导:积分因子法

定理: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

证明(积分因子法,四步):

第 1 步:移项。 把含 h 的项移到左边:

h(t)Ah(t)=Bx(t)h'(t) - A\,h(t) = B\,x(t)

第 2 步:乘积分因子 eAte^{-At} 两边左乘 eAte^{-At}:

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

第 3 步:识别乘积导数。 由矩阵指数的微分性质,

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

恰好等于左边。于是方程变成:

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

这一步是整个推导的「机关」:乘以 eAte^{-At} 后,左端塌缩成一个全导数,方程立刻可积。

第 4 步:两边积分并整理。t0t_0tt 积分:

eAth(t)eAt0h(t0)=t0teAsBx(s)dse^{-At}h(t) - e^{-At_0}h(t_0) = \int_{t_0}^{t} e^{-As}B\,x(s)\,ds

两边左乘 eAte^{At}(注意 eAteAt0=eA(tt0)e^{At}e^{-At_0} = e^{A(t-t_0)},且 eAteAs=eA(ts)e^{At}e^{-As} = e^{A(t-s)}):

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 \quad\blacksquare

两项的物理含义:旧记忆自由演化 + 区间内每一瞬输入贡献的叠加(叠加原理)。

2.5 零阶保持(ZOH)离散化

零阶保持假设:在采样区间 [tk, tk+Δ)[t_k,\ t_k + \Delta) 内,输入保持为采样值:

x(tk+τ)=xk,τ[0,Δ)x(t_k + \tau) = x_k, \qquad \forall\,\tau \in [0, \Delta)

推导:在通解中取 t0=tkt_0 = t_k,t=tk+Δt = t_k + \Delta:

hk+1=eAΔhk+0ΔeA(Δτ)Bx(tk+τ)dτh_{k+1} = e^{A\Delta}h_k + \int_{0}^{\Delta} e^{A(\Delta-\tau)}B\,x(t_k + \tau)\,d\tau

由 ZOH 假设,x(tk+τ)=xkx(t_k+\tau) = x_k 是常数,提出积分号:

hk+1=eAΔAˉhk+(0ΔeA(Δτ)Bdτ)Bˉxkh_{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

计算 Bˉ\bar B 的积分(标量情形最直观;矩阵情形在 A 可逆时同样成立)。换元 s=Δτs = \Delta - \tau:

Bˉ=0ΔeAsBds\bar B = \int_0^{\Delta} e^{As}\,B\,ds

对级数逐项积分:

0ΔeAsds=0Δk=0(As)kk!ds=k=0AkΔk+1(k+1)!=A1(eAΔ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)

最后一步验证:A1(eAΔI)=A1k1(AΔ)kk!=k1Ak1Δkk!=j0AjΔ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)!}

结论(ZOH 离散化公式):

 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\ }

2.6 小步长近似

当 Δ 很小时,eΔA=I+ΔA+O(Δ2)e^{\Delta A} = I + \Delta A + O(\Delta^2),于是

Bˉ=A1(Δ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ˉΔB\bar B \approx \Delta B。同理 AˉI+ΔA\bar A \approx I + \Delta A(这等价于欧拉法)。

注意:近似只在 Δ -> 0 时可靠。S4/Mamba 的实际实现都用精确的 ZOH 公式;且 Mamba 中 Δ 是逐 token 变化的,每步现算 eΔtAe^{\Delta_t A}

2.7 数值验证:完整算一遍

取标量系统 A=1A = -1,B=2B = 2,C=1C = 1,Δ=0.5\Delta = 0.5

第 1 步:离散化。

Aˉ=e0.50.6065,Bˉ=e0.511×2=(10.6065)×20.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

第 2 步:递归。 h0=0h_0 = 0,输入 x=[1,1]x = [1, 1]:

k xkx_k hk=0.6065hk1+0.7869xkh_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 步:与连续精确解对拍。x1x \equiv 1 的常数输入,连续方程 h=h+2h' = -h + 2 的解为 h(t)=2+(h02)eth(t) = 2 + (h_0 - 2)e^{-t}。在 t=1t = 1(即两步后):

h(1)=22e120.7358=1.2642h(1) = 2 - 2e^{-1} \approx 2 - 0.7358 = 1.2642 \quad✓

离散递归与连续精确解逐步完全一致–这正是 ZOH 离散化的性质:在 ZOH 假设成立(输入确为分段常数)时,它不是近似,而是精确等价

2.8 卷积形式:递归展开 = 一维卷积

先把「卷积」本身说清楚(离散因果时序卷积)。输入时序序列 u1,u2,,utu_1, u_2, \dots, u_t,卷积核(权重模板)K0,K1,K2,K_0, K_1, K_2, \dots,其中 KkK_k与当前时刻相隔 k 步的历史信息的权重因果(causal):只能看过去,看不到未来。形式化:

yt=τ=1tKtτuτy_t = \sum_{\tau=1}^{t} K_{t-\tau} \cdot u_\tau

最简单的数字例子。输入 u=[10,20,30]u = [10, 20, 30],衰减卷积核 K=[1, 0.5, 0.25]K = [1,\ 0.5,\ 0.25](越远权重越小):

y1=K0u1=1×10=10y2=K1u1+K0u2=0.5×10+1×20=25y3=K2u1+K1u2+K0u3=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}

每一步都把全部历史按距离加权求和,权重模板固定、随距离翻牌–这就是卷积。y1y_1 只用了 u1u_1(因果),三个输出用的是同一套 K(时不变)。带着这两个属性看 SSM 的卷积形式,就是把 KkK_k 具体化为 CAˉkBˉC\bar A^k \bar B

命题:离散 SSM ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,yt=Chty_t = C h_t(设 h0=0h_0 = 0)等价于

y=Kˉx,Kˉ=(CBˉ, CAˉBˉ, CAˉ2Bˉ, , CAˉL1Bˉ)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)

证明(直接展开):

ht=i=1tAˉtiBˉxih_t = \sum_{i=1}^{t} \bar A^{\,t-i}\bar B\, x_i

(归纳:h1=Bˉx1h_1 = \bar B x_1;设 ht1=it1Aˉt1iBˉxih_{t-1} = \sum_{i\le t-1}\bar A^{t-1-i}\bar B x_i,则 ht=Aˉht1+Bˉxt=it1AˉtiBˉxi+Bˉxth_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 ✓)

代入输出方程:

yt=i=1tCAˉtiBˉxi=j=0t1Kˉjxtjy_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

数值验证(§2.7 的系统,输入 x=[2,4,0,8]x=[2,4,0,8]):

卷积核:Kˉj=0.6065j×0.7869\bar K_j = 0.6065^j \times 0.7869,即 [0.7869, 0.4772, 0.2894, 0.1755][0.7869,\ 0.4772,\ 0.2894,\ 0.1755]

y4=0.7869×8+0.4772×0+0.2894×4+0.1755×2=6.295+1.158+0.3517.804y_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

递归验证:h1=1.574, h2=4.102, h3=2.488, h4=0.6065×2.488+0.7869×81.509+6.295=7.804h_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

推论:LTI 系统拥有双形式–训练用卷积(FFT,并行),推理用递归(O(1)/token)。Mamba 使 B,C,ΔB, C, \Delta 依赖输入后,Kˉ\bar K 不再固定,卷积形式失效,改由并行扫描承担训练侧并行。

2.9 从 Δ 到遗忘门:通往 Mamba 的最后一环

回看 Aˉ=eΔA\bar A = e^{\Delta A}:若 A 为负定(如 Mamba-2 取 A=aIA = -a\cdot I,a>0a>0),则

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

就是一个衰减门:

  • Δt\Delta_t \to \infty:Aˉt0\bar A_t \to 0 – 清空旧记忆,只看当前输入;
  • Δt0\Delta_t \to 0:Aˉt1\bar A_t \to 1 – 冻结状态,忽略当前输入。

Mamba 让 Δt=softplus(Linear(xt))\Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t)),离散化公式每步现算,于是「步长」变成了看内容的标量遗忘门

2.10 附录:双线性变换离散化对比

除 ZOH 外,S4 论文实际使用的是双线性变换(Tustin 方法):用梯形法则近似积分,

Aˉ=(IΔ2A)1(I+Δ2A),Bˉ=(IΔ2A)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

方法 性质 适用
ZOH ZOH 假设下精确;保稳定性;公式含 eΔAe^{\Delta A} Mamba 系列:Δt\Delta_t 逐 token 变化,每步现算矩阵指数(S4D/Mamba 的对角 A 使 eΔtAe^{\Delta_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}Bˉ=A1(eΔAI)B\bar B = A^{-1}(e^{\Delta A} - I)B 都是固定矩阵(与输入无关),于是 LTI 成立、递归 = 卷积,训练用 FFT 卷积、推理用 O(1)O(1) 递归(双形式,§2.8)。A 用 HiPPO 矩阵初始化(记住长历史的专门构造),离散化用双线性变换。限制也很明显:参数不随输入变,「记忆怎么衰减」在推理前就定死了。

Mamba / S6(2023):选择性改造。让 B,C,ΔB, C, \Delta 都依赖输入(Δt=softplus(Linear(xt))\Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t))),每步现算 eΔtAe^{\Delta_t A}–「步长」变成了看内容的遗忘门(§2.9)。代价是 LTI 被打破、固定卷积核 Kˉ\bar K 失效,训练侧改用并行扫描(associative scan);A 限制为对角结构使 eΔtAe^{\Delta_t A} 逐元素可算。这一步是「SSM 学会看内容」的关键。

Mamba-2 / SSD(2024):与线性注意力形式统一。状态改写成外积形式 St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top,矩阵 AA 退化为标量衰减 αt=eΔta\alpha_t = e^{-\Delta_t a}从这一步起,SSM 与线性注意力在数学上就是同一个东西的两个记法(SSD 框架)。训练用 chunkwise:chunk 内 QKQK^\top 小注意力 + chunk 间递归,比纯 scan 的 GEMM 利用率高得多。

对照第 1 节的线性注意力递推 St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top:Mamba-2 只多了一个乘在旧状态上的标量 αt\alpha_t两条路线在 Mamba-2 处汇合:线性注意力给出了外积状态的样子,SSM 给出了衰减门。


4. 全链路一图流

h=Ah+Bx连续 SSM积分因子法h(t+Δ)=eAΔh(t)+0ΔeA(Δτ)Bxdτ通解:旧记忆衰减 + 输入累积ZOHht=Aˉht1+Bˉxt离散递归(推理用)展开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{卷积(训练用)}}

线性注意力侧:softmax(QK)V\mathrm{softmax}(QK^\top)V --核化–> ϕ(Q)(ϕ(K)V)\phi(Q)(\phi(K)^\top V) --逐 token–> St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top

两线在 Mamba-2 汇合:St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top

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

总结

  • 线性注意力:去掉 softmax 的耦合非线性 + 核化把非线性提前,Q(KV)Q(K^\top V) 结合律重新可用,O(n2)O(n)O(n^2) \to O(n);代价是纯加性状态,只会叠加不会删除(记忆碰撞);
  • SSM 一条链推完:齐次解 eAth0e^{At}h_0 -> 积分因子法得通解 -> ZOH 假设下提出 xkx_k 得离散递归,Aˉ=eΔA\bar A = e^{\Delta A}Bˉ=A1(eΔAI)B\bar B = A^{-1}(e^{\Delta A} - I)B;
  • ZOH 不是近似:输入分段常数时离散递归与连续解逐步相等(§2.7 数值对拍 1.2642 = 1.2642);
  • 递归 = 卷积:LTI 的双形式是 S4 训练(FFT 卷积)与推理(O(1) 递归)各取所长的根基;
  • Aˉ=eΔA\bar A = e^{\Delta A} 就是遗忘门:这是从控制论到 Mamba/KDA 的思想连续性;
  • 两条路线在 Mamba-2 汇合:外积状态(线性注意力)+ 标量衰减门(SSM)= St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top,此后 delta rule 和逐通道门的演化都在这条骨架上进行。

参考: