KDA 的来龙去脉:从线性注意力到 Kimi Delta Attention
KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。本文不展开 K3 的完整架构,只讨论一条演进路线:KDA 之前的几个模型(线性注意力、Mamba、DeltaNet/GDN)分别解决了什么问题,KDA 又在它们的基础上改了什么 。
演进脉络如下:线性注意力先构造出「固定大小的记忆」,但该记忆只能累加、无法清理;Mamba 引入遗忘机制,但衰减是全局统一的;DeltaNet 支持对单条记忆的精确改写;GDN 将「擦除」与「改写」结合;KDA 进一步让每个通道独立决定衰减速度。每一步均给出公式推导、数值验算与工程动机。
0. 出发点:固定大小的记忆
标准 softmax 注意力的 KV cache 随序列长度线性膨胀,几十万 token 的上下文可使其达到数十 GB 量级,显存与带宽同时成为瓶颈。K3 把「序列维度缩放」列为头号工程目标,93 层里 69 层换成 KDA,就是把 KV cache 换成一个固定大小的状态矩阵 :O ( n ⋅ d ) O(n \cdot d) O ( n ⋅ d ) 变 O ( d 2 ) O(d^2) O ( d 2 ) ,上下文长度翻倍时状态大小保持不变。
但固定大小是有代价的。一个 d v × d k d_v \times d_k d v × d k 的矩阵要承载任意长度的序列,记忆必须可写、可改、可遗忘 :若只能写入而不能清理,序列一长信息就会相互干扰。如何设计这块「可自我管理的固定记忆」,是本文的主线。
线性注意力与 SSM 两条路线的完整推导见前置篇《线性注意力与 SSM:两条技术路线的完整推导》 ,本文直接使用其结论。
记号约定:两套写法互为转置
先把记号钉死。文献里状态矩阵有两种摆法,内容完全等价、互为转置 ,但混用会让推导看起来「中途换了个式子」:
主约定(§1.1 与 §3 起使用,与 KDA 论文一致)
转置约定(部分实现与 chunkwise 推导常用)
状态形状
S t ∈ R d k × d v S_t \in \mathbb{R}^{d_k \times d_v} S t ∈ R d k × d v
S ^ t ∈ R d v × d k \hat S_t \in \mathbb{R}^{d_v \times d_k} S ^ t ∈ R d v × d k
读出
o t = S t ⊤ q t o_t = S_t^\top q_t o t = S t ⊤ q t
o t = S ^ t q t o_t = \hat S_t q_t o t = S ^ t q t
当前 key 的旧内容
S t − 1 ⊤ k t S_{t-1}^\top k_t S t − 1 ⊤ k t
S ^ t − 1 k t \hat S_{t-1} k_t S ^ t − 1 k t
写入项
k t v t ⊤ k_t v_t^\top k t v t ⊤ (左 key 右 value)
v t k t ⊤ v_t k_t^\top v t k t ⊤ (左 value 右 key)
擦除算子
左乘:( I − β t k t k t ⊤ ) S t − 1 (I - \beta_t k_t k_t^\top)S_{t-1} ( I − β t k t k t ⊤ ) S t − 1
右乘:S ^ t − 1 ( I − β t k t k t ⊤ ) \hat S_{t-1}(I - \beta_t k_t k_t^\top) S ^ t − 1 ( I − β t k t k t ⊤ )
两者的严格关系就是一次转置 S ^ t = S t ⊤ \hat S_t = S_t^\top S ^ t = S t ⊤ 。由于擦除算子 H t = I − β t k t k t ⊤ H_t = I - \beta_t k_t k_t^\top H t = I − β t k t k t ⊤ 是对称矩阵 (H t ⊤ = H t H_t^\top = H_t H t ⊤ = H t ,因为 ( k t k t ⊤ ) ⊤ = k t k t ⊤ (k_tk_t^\top)^\top = k_tk_t^\top ( k t k t ⊤ ) ⊤ = k t k t ⊤ ),转置可以直接穿过它:
( H t S t − 1 + β t k t v t ⊤ ) ⊤ = S t − 1 ⊤ H t ⊤ + β t v t k t ⊤ = S ^ t − 1 H t + β t v t k t ⊤ \big(H_t S_{t-1} + \beta_t k_t v_t^\top\big)^{\!\top} = S_{t-1}^\top H_t^\top + \beta_t v_t k_t^\top = \hat S_{t-1} H_t + \beta_t v_t k_t^\top
( H t S t − 1 + β t k t v t ⊤ ) ⊤ = S t − 1 ⊤ H t ⊤ + β t v t k t ⊤ = S ^ t − 1 H t + β t v t k t ⊤
所以后文若看到 k v ⊤ k v^\top k v ⊤ 与 v k ⊤ v k^\top v k ⊤ 互换、擦除算子从左边跑到右边,那是切换了约定,不是等式变了 。GDN 数值验算与 chunkwise 一节为对齐论文公式会改用转置约定,届时会再次点明。
各模型按出场顺序(统一写成主约定):
线性注意力 :固定状态 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 ⊤ ,只支持累加写入(同一个 key 写两次,读出的是两个 value 的叠加,即记忆碰撞);
Mamba-2 :引入标量衰减门 S t = α t S t − 1 + k t v t ⊤ S_t = \alpha_t S_{t-1} + k_t v_t^\top S t = α t S t − 1 + k t v t ⊤ ,具备遗忘能力,但衰减是全局统一的;
DeltaNet :改为差值写入 S t = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ S_t = (I - \beta_tk_tk_t^\top)S_{t-1} + \beta_t k_t v_t^\top S t = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ ,支持定点改写;
GDN :衰减门与差值写入结合 S t = α t ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ S_t = \alpha_t(I - \beta_tk_tk_t^\top)S_{t-1} + \beta_t k_t v_t^\top S t = α t ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ ;
KDA :将标量门拆为逐通道门 Diag ( α t ) \operatorname{Diag}(\bm{\alpha}_t) Diag ( α t ) ,各维度独立决定衰减速度,并加数值下界,用于 K3。
1. 演进主线:从 delta rule 到 GDN
前置篇给出两个结论:线性注意力提供了固定状态,但只能累加写入、存在记忆碰撞;SSM 推进到 Mamba-2 提供了标量衰减门 α t \alpha_t α t ,但衰减是全局统一的。本节讨论这两者如何组合成 GDN(Gated DeltaNet)——先看写入规则怎么从加法变成差值(§1.1),再把遗忘门拼进来(§1.2),然后手算验证(§1.3)并给出可并行的 chunkwise 形式(§1.4)。KDA 对 GDN 的改动从 §3 开始。
1.1 DeltaNet:从加法到差值写入
线性注意力得到的 S 只能不断叠加写入,修正这一点正是 DeltaNet 的动机。解法:写入差值,不写全值 。写入前,先用当前 key 把状态里已存的内容读一遍:
v o l d = S t − 1 ⊤ k t ( 当前 key 指向的旧内容 ) , u t = β t ( v t − v o l d ) v_{\mathrm{old}} = S_{t-1}^\top k_t \quad (\text{当前 key 指向的旧内容}), \qquad u_t = \beta_t (v_t - v_{\mathrm{old}})
v old = S t − 1 ⊤ k t ( 当前 key 指向的旧内容 ) , u t = β t ( v t − v old )
S t = S t − 1 + k t u t ⊤ S_t = S_{t-1} + k_t u_t^\top
S t = S t − 1 + k t u t ⊤
符号
形状
含义
S t S_t S t
[ d k , d v ] [d_k, d_v] [ d k , d v ]
状态矩阵(主约定)
k t k_t k t
[ d k ] [d_k] [ d k ]
当前 key,经 L2 归一化后 ∣ k t ∣ 2 = 1 |k_t|_2 = 1 ∣ k t ∣ 2 = 1
v o l d v_{\mathrm{old}} v old
[ d v ] [d_v] [ d v ]
当前 key 从状态里读出的已有内容
β t \beta_t β t
标量 ∈ ( 0 , 1 ) \in (0,1) ∈ ( 0 , 1 )
内容替换强度
u t u_t u t
[ d v ] [d_v] [ d v ]
实际写入的差值
严格改写:差值写入 = 先删后写 。上式右侧看不出「擦除」在哪里,把 u t u_t u t 代入展开即可,每一步只用外积的结合律 k t ( k t ⊤ S t − 1 ) = ( k t k t ⊤ ) S t − 1 k_t(k_t^\top S_{t-1}) = (k_tk_t^\top)S_{t-1} k t ( k t ⊤ S t − 1 ) = ( k t k t ⊤ ) S t − 1 :
S t = S t − 1 + k t u t ⊤ 定义 = S t − 1 + k t [ β t ( v t − v o l d ) ] ⊤ 代入 u t = S t − 1 + β t k t v t ⊤ − β t k t v o l d ⊤ 转置展开 = S t − 1 + β t k t v t ⊤ − β t k t ( S t − 1 ⊤ k t ) ⊤ 代入 v o l d = S t − 1 ⊤ k t = S t − 1 + β t k t v t ⊤ − β t k t k t ⊤ S t − 1 ( S t − 1 ⊤ k t ) ⊤ = k t ⊤ S t − 1 = ( I − β t k t k t ⊤ ) S t − 1 ⏟ 删除项:沿 k t 方向擦除 + β t k t v t ⊤ ⏟ 新值项 提取公因式 S t − 1 \begin{aligned}
S_t &= S_{t-1} + k_t u_t^\top && \text{定义} \\[2pt]
&= S_{t-1} + k_t\big[\beta_t(v_t - v_{\mathrm{old}})\big]^\top && \text{代入 } u_t \\[2pt]
&= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t v_{\mathrm{old}}^\top && \text{转置展开} \\[2pt]
&= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t \big(S_{t-1}^\top k_t\big)^{\!\top} && \text{代入 } v_{\mathrm{old}} = S_{t-1}^\top k_t \\[2pt]
&= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t k_t^\top S_{t-1} && (S_{t-1}^\top k_t)^\top = k_t^\top S_{t-1} \\[2pt]
&= \underbrace{\big(I - \beta_t k_t k_t^\top\big) S_{t-1}}_{\text{删除项:沿 } k_t \text{ 方向擦除}} + \underbrace{\beta_t k_t v_t^\top}_{\text{新值项}} && \text{提取公因式 } S_{t-1}
\end{aligned}
S t = S t − 1 + k t u t ⊤ = S t − 1 + k t [ β t ( v t − v old ) ] ⊤ = S t − 1 + β t k t v t ⊤ − β t k t v old ⊤ = S t − 1 + β t k t v t ⊤ − β t k t ( S t − 1 ⊤ k t ) ⊤ = S t − 1 + β t k t v t ⊤ − β t k t k t ⊤ S t − 1 = 删除项:沿 k t 方向擦除 ( I − β t k t k t ⊤ ) S t − 1 + 新值项 β t k t v t ⊤ 定义 代入 u t 转置展开 代入 v old = S t − 1 ⊤ k t ( S t − 1 ⊤ k t ) ⊤ = k t ⊤ S t − 1 提取公因式 S t − 1
关键的一步是第五行:v o l d v_{\mathrm{old}} v old 自己就是由 S t − 1 S_{t-1} S t − 1 算出来的,所以「减去旧值」这个动作必然能写成一个作用在 S t − 1 S_{t-1} S t − 1 上的线性算子 ,而不是一个额外的加项。提取公因式之后,− β t k t k t ⊤ -\beta_t k_tk_t^\top − β t k t k t ⊤ 就并入了单位阵,变成擦除算子。预测误差 v t − v o l d v_t - v_{\mathrm{old}} v t − v old 就是 delta,Delta Rule 由此得名 :每次写入的从来不是新值本身,而是新值与旧值之差。
两端均为 d k × d v d_k \times d_v d k × d v ,维度自洽;若换成转置约定,同一个式子写作 S ^ t = S ^ t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤ \hat S_t = \hat S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_t v_tk_t^\top S ^ t = S ^ t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤ 。
写入差值这一形式并非经验设计,它等价于对回归损失做一步梯度下降 。上面是先有「写差值」这个想法、再发现它等于先删后写;反方向走一遍会看到,差值根本不是选的,而是推出来的。
第一步:把记忆当成一个在线回归问题 。状态 S S S 要承担的职责是一张查询表:拿钥匙 k t k_t k t 来,应当取出 v t v_t v t ,即希望 S ⊤ k t ≈ v t S^\top k_t \approx v_t S ⊤ k t ≈ v t 。把这个愿望写成当前样本上的平方损失:
L ( S ) = 1 2 ∥ S ⊤ k t − v t ∥ 2 \mathcal{L}(S) = \tfrac{1}{2}\big\|S^\top k_t - v_t\big\|^2
L ( S ) = 2 1 S ⊤ k t − v t 2
S ⊤ k t S^\top k_t S ⊤ k t 是「用现在的记忆查 k t k_t k t 能取出的值」,v t v_t v t 是「应该取到的值」,两者之差就是残差 ;系数 1 2 \tfrac12 2 1 纯为求导后消掉 2。
第二步:求梯度,得到残差与钥匙的外积 。先用最熟的一维情形建立直觉:1 2 ( w x − y ) 2 \tfrac12(wx-y)^2 2 1 ( w x − y ) 2 对 w w w 的导数是 ( w x − y ) x (wx-y)\,x ( w x − y ) x ——误差乘输入 。矩阵版一模一样,只是乘法变成外积。记残差 r t = S ⊤ k t − v t ∈ R d v r_t = S^\top k_t - v_t \in \mathbb{R}^{d_v} r t = S ⊤ k t − v t ∈ R d v ,取扰动 Δ \Delta Δ 算一阶项:
L ( S + Δ ) − L ( S ) = ⟨ r t , Δ ⊤ k t ⟩ + O ( ∥ Δ ∥ 2 ) = ⟨ k t r t ⊤ , Δ ⟩ + O ( ∥ Δ ∥ 2 ) \mathcal{L}(S+\Delta) - \mathcal{L}(S) = \big\langle r_t,\ \Delta^\top k_t \big\rangle + O(\|\Delta\|^2)
= \big\langle k_t r_t^\top,\ \Delta \big\rangle + O(\|\Delta\|^2)
L ( S + Δ ) − L ( S ) = ⟨ r t , Δ ⊤ k t ⟩ + O ( ∥Δ ∥ 2 ) = ⟨ k t r t ⊤ , Δ ⟩ + O ( ∥Δ ∥ 2 )
与 Δ \Delta Δ 配对的那个矩阵就是梯度:
∇ S L = k t ( S ⊤ k t − v t ) ⊤ = k t r t ⊤ ∈ R d k × d v \nabla_S \mathcal{L} = k_t\,\big(S^\top k_t - v_t\big)^{\!\top} = k_t r_t^\top \ \in\ \mathbb{R}^{d_k \times d_v}
∇ S L = k t ( S ⊤ k t − v t ) ⊤ = k t r t ⊤ ∈ R d k × d v
形状与 S S S 一致,可直接用于更新。注意差值 S ⊤ k t − v t S^\top k_t - v_t S ⊤ k t − v t 是自己冒出来的 ——平方损失的梯度天然就长成「残差 ⊗ 钥匙」的样子。这就是「为何恰好是差值」的答案:没人规定写入要用差值,是平方损失的梯度只能是差值。
第三步:以 β t \beta_t β t 为步长走一步 SGD ,展开重新归项:
S t = S t − 1 − β t ∇ S L ∣ S = S t − 1 一步梯度下降 = S t − 1 − β t k t ( S t − 1 ⊤ k t − v t ) ⊤ 代入梯度 = S t − 1 − β t k t k t ⊤ S t − 1 + β t k t v t ⊤ 展开 = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ 提公因式 \begin{aligned}
S_t &= S_{t-1} - \beta_t \nabla_S\mathcal{L}\big|_{S = S_{t-1}} && \text{一步梯度下降} \\[2pt]
&= S_{t-1} - \beta_t k_t\big(S_{t-1}^\top k_t - v_t\big)^{\!\top} && \text{代入梯度} \\[2pt]
&= S_{t-1} - \beta_t k_t k_t^\top S_{t-1} + \beta_t k_t v_t^\top && \text{展开} \\[2pt]
&= \big(I - \beta_t k_t k_t^\top\big)S_{t-1} + \beta_t k_t v_t^\top && \text{提公因式}
\end{aligned}
S t = S t − 1 − β t ∇ S L S = S t − 1 = S t − 1 − β t k t ( S t − 1 ⊤ k t − v t ) ⊤ = S t − 1 − β t k t k t ⊤ S t − 1 + β t k t v t ⊤ = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ 一步梯度下降 代入梯度 展开 提公因式
结果与前面从「写差值」出发得到的式子逐字相同 。两条路径交汇于同一个等式,于是 β t \beta_t β t 的角色也明确了:它就是学习率 。同样,「先删后写」不是设计直觉,而是展开式里必然的两项:− β t k t k t ⊤ -\beta_t k_tk_t^\top − β t k t k t ⊤ 并入单位阵成为擦除,+ β t k t v t ⊤ +\beta_t k_tv_t^\top + β t k t v t ⊤ 就是写入。
验证:写完立即读一次 。取 ∥ k t ∥ 2 = 1 \|k_t\|_2 = 1 ∥ k t ∥ 2 = 1 (L2Norm 之后),用同一个 k t k_t k t 查新状态:
S t ⊤ k t = [ ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ ] ⊤ k t = ( 1 − β t ) ⏟ 保留旧值 S t − 1 ⊤ k t + β t ⏟ 接纳新值 v t S_t^\top k_t = \big[(I-\beta_tk_tk_t^\top)S_{t-1} + \beta_tk_tv_t^\top\big]^{\!\top}k_t
= \underbrace{(1-\beta_t)}_{\text{保留旧值}}\,S_{t-1}^\top k_t + \underbrace{\beta_t}_{\text{接纳新值}} v_t
S t ⊤ k t = [ ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ ] ⊤ k t = 保留旧值 ( 1 − β t ) S t − 1 ⊤ k t + 接纳新值 β t v t
读出结果是旧值与新值的凸组合 ,β t \beta_t β t 就是插值系数:β t = 1 \beta_t = 1 β t = 1 时 S t ⊤ k t = v t S_t^\top k_t = v_t S t ⊤ k t = v t ,该样本的损失一步降到 0,即完全覆写;β t = 0.5 \beta_t = 0.5 β t = 0.5 时新旧各一半;β t = 0 \beta_t = 0 β t = 0 时不学。这也解释了 L2Norm 为何是前提:只有 ∥ k t ∥ = 1 \|k_t\| = 1 ∥ k t ∥ = 1 时上式才是干净的插值,否则系数会变成 1 − β t ∥ k t ∥ 2 1 - \beta_t\|k_t\|^2 1 − β t ∥ k t ∥ 2 ,可能跌出 [ 0 , 1 ] [0,1] [ 0 , 1 ] 。
一句话总结这条链:记忆 = 在线回归 → 平方损失 → 梯度 = 残差 ⊗ 钥匙 → 一步 SGD = 先删后写 。它有严格出处:Widrow-Hoff 1960 年的 delta rule / LMS 规则(名字里的 delta 正是指误差 δ \delta δ )在联想记忆矩阵上的应用。后面 GDN 与 KDA 做的事,就是在这个损失里再加一项正则 ,把「别忘了旧账」量化进目标(见 §1.2 的在线学习统一视角)。
DeltaNet 只改写入规则,两个前置条件一律保留:特征映射 ϕ \phi ϕ (实现里常取 ϕ = L 2 N o r m \phi = \mathrm{L2Norm} ϕ = L2Norm ,即下文的单位球约束)、外积状态、逐步递推形式完全没动,变的只有一处:+ k v ⊤ ⟶ + k ( v − S ⊤ k ) ⊤ +\,k v^\top \longrightarrow +\,k(v - S^\top k)^\top + k v ⊤ ⟶ + k ( v − S ⊤ k ) ⊤ 。
直观类比:state 是白板,key 是指针 。纯加性写入相当于在白板上不断叠加便签,写满后内容互相遮盖;DeltaNet 先擦除指针 k t k_t k t 指向的区域,再写入新内容。β t = 0 \beta_t = 0 β t = 0 :完全不写入,状态不动;β t = 1 \beta_t = 1 β t = 1 :完全替换,指针指向的内容被整体替换。这种「先擦后写」的机制使状态能够覆盖错误记忆,这是 DeltaNet 与 GDN 这条路线的核心改进。
关键整理:写成转移矩阵形式 。定义 H t = I − β t k t k t ⊤ H_t = I - \beta_t k_t k_t^\top H t = I − β t k t k t ⊤ ,状态更新变成
S t = H t S t − 1 + β t k t v t ⊤ S_t = H_t\, S_{t-1} + \beta_t k_t v_t^\top
S t = H t S t − 1 + β t k t v t ⊤
这个形式之所以非常关键,是因为它暴露了 DeltaNet 和 SSM 的同构:H t H_t H t 正是随输入变化的状态转移矩阵 (SSM 里是 A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A ),β t k t v t ⊤ \beta_t k_t v_t^\top β t k t v t ⊤ 正是写入项(SSM 里是 B ˉ x t \bar B x_t B ˉ x t )。
H t H_t H t 到底是什么:秩 1 与特征值
H t = I − β t k t k t ⊤ H_t = I - \beta_t k_t k_t^\top H t = I − β t k t k t ⊤ 是全文出现频率最高的矩阵,值得把它彻底拆开。它由两块拼成:单位阵,减去一个秩 1 矩阵 。
先看「秩」 。矩阵的秩 = 它的列里真正不同的方向 有几个,形式定义是线性无关列的最大个数,直觉是这个矩阵作为变换、输出能铺满几维空间。单位阵 I ∈ R d × d I \in \mathbb{R}^{d\times d} I ∈ R d × d 的 d d d 列是 d d d 根坐标轴,谁也不沾谁,秩 = d = d = d (满秩);而 k k ⊤ k k^\top k k ⊤ 的所有列都挤在同一条线上,秩 = 1 = 1 = 1 。
为什么 k k ⊤ kk^\top k k ⊤ 的列都是 k k k 的倍数 。k k ⊤ kk^\top k k ⊤ 的 ( i , j ) (i,j) ( i , j ) 元素是 k i k j k_ik_j k i k j ,于是它的第 j j j 列是
( k 1 k j k 2 k j ⋮ k d k j ) = k j ⋅ ( k 1 k 2 ⋮ k d ) = k j ⋅ k \begin{pmatrix} k_1 k_j \\ k_2 k_j \\ \vdots \\ k_d k_j \end{pmatrix}
= k_j \cdot \begin{pmatrix} k_1 \\ k_2 \\ \vdots \\ k_d \end{pmatrix} = k_j \cdot k
k 1 k j k 2 k j ⋮ k d k j = k j ⋅ k 1 k 2 ⋮ k d = k j ⋅ k
第 j j j 列 = 标量 k j k_j k j 乘同一个向量 k k k ,换 j j j 只换倍率、方向永远是 k k k 。取 k = ( 1 , 2 ) ⊤ k = (1,2)^\top k = ( 1 , 2 ) ⊤ 验算:
k k ⊤ = ( 1 2 ) ( 1 2 ) = ( 1 2 2 4 ) k k^\top = \begin{pmatrix}1\\2\end{pmatrix}\begin{pmatrix}1&2\end{pmatrix} = \begin{pmatrix}1&2\\2&4\end{pmatrix}
k k ⊤ = ( 1 2 ) ( 1 2 ) = ( 1 2 2 4 )
第二列 ( 2 , 4 ) (2,4) ( 2 , 4 ) 恰是第一列 ( 1 , 2 ) (1,2) ( 1 , 2 ) 的 2 倍——形式上是 2 × 2 2\times2 2 × 2 矩阵,实际只携带一个方向的信息,故行列式为 0、不可逆。反过来也成立:任何秩 1 矩阵都能写成某个外积 u v ⊤ uv^\top u v ⊤ ,所以「秩 1 矩阵」与「外积」基本同义。
作为变换,它把一切压到 k k k 那条线上 。作用在任意向量 x x x 上,用结合律:
( k k ⊤ ) x = k ( k ⊤ x ) = ( k ⊤ x ) ⋅ k (k k^\top)\, x = k\,(k^\top x) = (k^\top x)\cdot k
( k k ⊤ ) x = k ( k ⊤ x ) = ( k ⊤ x ) ⋅ k
k ⊤ x k^\top x k ⊤ x 是标量,即 x x x 在 k k k 方向的投影长度;乘回 k k k 得到一个沿 k k k 的向量。不管输入什么,输出永远落在 k k k 张成的一维直线上 ——这就是「只张出一维」的几何含义,整个 d d d 维空间被拍扁成一条线。
再看「特征值」 。若存在非零向量 x x x 使 M x = λ x Mx = \lambda x M x = λ x ,即 x x x 经变换后方向不变(或恰好反向)、只被缩放 λ \lambda λ 倍,则 x x x 是特征向量 (eigenvector)、λ \lambda λ 是特征值 (eigenvalue)。多数向量过一个矩阵会又转又缩,特征向量是躺在矩阵「主轴」上的特例,变换对它们只是纯缩放;特征值就是各主轴上的缩放倍率。
用 k = ( 1 , 2 ) ⊤ k=(1,2)^\top k = ( 1 , 2 ) ⊤ 的例子算:沿 k k k 方向,( k k ⊤ ) k = k ( k ⊤ k ) = ∥ k ∥ 2 k = 5 k (kk^\top)k = k(k^\top k) = \|k\|^2 k = 5k ( k k ⊤ ) k = k ( k ⊤ k ) = ∥ k ∥ 2 k = 5 k ,故 λ 1 = ∥ k ∥ 2 = 5 \lambda_1 = \|k\|^2 = 5 λ 1 = ∥ k ∥ 2 = 5 ;与 k k k 垂直的 x = ( 2 , − 1 ) ⊤ x = (2,-1)^\top x = ( 2 , − 1 ) ⊤ 满足 k ⊤ x = 0 k^\top x = 0 k ⊤ x = 0 ,于是 ( k k ⊤ ) x = k ⋅ 0 = 0 = 0 ⋅ x (kk^\top)x = k\cdot 0 = 0 = 0\cdot x ( k k ⊤ ) x = k ⋅ 0 = 0 = 0 ⋅ x ,故 λ 2 = 0 \lambda_2 = 0 λ 2 = 0 。2 × 2 2\times2 2 × 2 恰好两个特征值 5 与 0,正对应「秩 1 把正交方向压没、只在 k k k 方向放大 ∥ k ∥ 2 \|k\|^2 ∥ k ∥ 2 倍」。
放回 H t H_t H t 。结构一目了然:I I I 让所有方向原样保留,减去 β t k t k t ⊤ \beta_t k_tk_t^\top β t k t k t ⊤ 只在 k t k_t k t 这一个方向上动刀 ,其余 d − 1 d-1 d − 1 个方向碰都不碰。特征值分两种:
H t k t = k t − β t k t ( k t ⊤ k t ) = ( 1 − β t ∥ k t ∥ 2 ) k t , H t x = x ( ∀ x ⊥ k t ) H_t k_t = k_t - \beta_t k_t(k_t^\top k_t) = \big(1 - \beta_t\|k_t\|^2\big)k_t, \qquad
H_t x = x \quad (\forall\, x \perp k_t)
H t k t = k t − β t k t ( k t ⊤ k t ) = ( 1 − β t ∥ k t ∥ 2 ) k t , H t x = x ( ∀ x ⊥ k t )
配合 L2Norm(∥ k t ∥ = 1 \|k_t\| = 1 ∥ k t ∥ = 1 )就是:沿 k t k_t k t 缩放 1 − β t 1-\beta_t 1 − β t ,正交补方向特征值为 1,即特征值落在 [ 1 − β t , 1 ] ⊂ ( 0 , 1 ] [1-\beta_t,\ 1] \subset (0,1] [ 1 − β t , 1 ] ⊂ ( 0 , 1 ] ——不放大、不翻转、不发散,这就是数值稳定性的特征值表述。反例也很直白:若不做归一化,取 k = ( 1 , 2 ) ⊤ k = (1,2)^\top k = ( 1 , 2 ) ⊤ 、β = 0.6 \beta = 0.6 β = 0.6 ,则 1 − β ∥ k ∥ 2 = 1 − 3 = − 2 1-\beta\|k\|^2 = 1-3 = -2 1 − β ∥ k ∥ 2 = 1 − 3 = − 2 ,特征值变号且模长大于 1,反复作用必然发散。
H t H_t H t 还有两个后文要用的性质:对称 (H t ⊤ = H t H_t^\top = H_t H t ⊤ = H t ,这是前面两套约定能靠转置互换的原因);以及这类「I I I 减秩 1」结构与**豪斯霍尔德变换(Householder transformation)**同型,QR 分解用它做反射消元。区别在于 Householder 取 β = 2 / ∥ k ∥ 2 \beta = 2/\|k\|^2 β = 2/∥ k ∥ 2 ,∥ k ∥ = 1 \|k\|=1 ∥ k ∥ = 1 时是精确反射(保长);DeltaNet 的 β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 ) 是「部分反射」,软化成可学习的写入强度。
秩 1 也解释了为何删除项开销低:擦除单一方向不需要满秩运算,一次外积即可。更进一步,S ← H t S + β t k t v t ⊤ S \leftarrow H_tS + \beta_t k_tv_t^\top S ← H t S + β t k t v t ⊤ 每步只给状态加一个秩 1 矩阵(秩 1 更新 ),chunk 内 C C C 步就是 C C C 个秩 1 更新的连乘叠加——而 WY / UT 变换正是数值线性代数里专门处理「一串秩 1 更新如何打包」的经典工具,后文 chunkwise 一节的源头就在这里。
为什么特征值这个概念到处出现 。因为它回答了迭代系统最关心的问题:一个变换反复作用很多次之后会怎样 。M M M 作用 n n n 次,在特征向量方向上就是 λ n \lambda^n λ n :∣ λ ∣ > 1 |\lambda| > 1 ∣ λ ∣ > 1 的方向爆炸,∣ λ ∣ < 1 |\lambda| < 1 ∣ λ ∣ < 1 的方向衰减消失,λ = 1 \lambda = 1 λ = 1 的方向保持不变。本文这条线上的约束几乎都在围着它转:
出现位置
特征值/缩放倍率的约束
目的
SSM/Mamba 的离散化 A ˉ \bar A A ˉ
特征值(或对角衰减因子)模长 ≤ 1 \le 1 ≤ 1
状态不发散
GDN / KDA 的 α ∈ ( 0 , 1 ) \alpha \in (0,1) α ∈ ( 0 , 1 )
直接强制衰减算子逐方向缩放 < 1 < 1 < 1
可控遗忘
H t = I − β t k t k t ⊤ H_t = I-\beta_tk_tk_t^\top H t = I − β t k t k t ⊤
特征值 ∈ [ 1 − β t , 1 ] \in [1-\beta_t, 1] ∈ [ 1 − β t , 1 ]
擦除是收缩的
1 / Γ 1/\Gamma 1/Γ 溢出(K3 加下界的原因)
Γ = ∏ α \Gamma = \prod\alpha Γ = ∏ α 连乘趋 0,倒数爆炸
与「$
一句话:特征向量是矩阵的「自然方向」,特征值是每个自然方向上的缩放倍率 ;看懂这两个数,就看懂了矩阵反复作用后的长期行为。
把递归展开,消掉时间依赖 。记 B t = β t k t v t ⊤ B_t = \beta_t k_t v_t^\top B t = β t k t v t ⊤ ,递推 S t = H t S t − 1 + B t S_t = H_t S_{t-1} + B_t S t = H t S t − 1 + B t 逐层代入(主约定下 H H H 在左侧,越晚的时间步越靠外):
S 1 = H 1 S 0 + B 1 S_1 = H_1 S_0 + B_1
S 1 = H 1 S 0 + B 1
S 2 = H 2 H 1 S 0 + H 2 B 1 + B 2 S_2 = H_2H_1 S_0 + H_2B_1 + B_2
S 2 = H 2 H 1 S 0 + H 2 B 1 + B 2
S 3 = H 3 H 2 H 1 S 0 + H 3 H 2 B 1 + H 3 B 2 + B 3 S_3 = H_3H_2H_1 S_0 + H_3H_2B_1 + H_3B_2 + B_3
S 3 = H 3 H 2 H 1 S 0 + H 3 H 2 B 1 + H 3 B 2 + B 3
S 4 = H 4 H 3 H 2 H 1 S 0 + H 4 H 3 H 2 B 1 + H 4 H 3 B 2 + H 4 B 3 + B 4 S_4 = H_4H_3H_2H_1 S_0 + H_4H_3H_2B_1 + H_4H_3B_2 + H_4B_3 + B_4
S 4 = H 4 H 3 H 2 H 1 S 0 + H 4 H 3 H 2 B 1 + H 4 H 3 B 2 + H 4 B 3 + B 4
规律如下:
S t = ( ∏ i = t 1 H i ) S 0 + ∑ i ≤ t ( ∏ j = t i + 1 H j ) B i S_t = \Big(\prod_{i=t}^{1} H_i\Big) S_0 + \sum_{i \le t} \Big(\prod_{j=t}^{i+1} H_j\Big) B_i
S t = ( i = t ∏ 1 H i ) S 0 + i ≤ t ∑ ( j = t ∏ i + 1 H j ) B i
其中 ∏ i = t 1 H i = H t H t − 1 ⋯ H 1 \prod_{i=t}^{1}H_i = H_tH_{t-1}\cdots H_1 ∏ i = t 1 H i = H t H t − 1 ⋯ H 1 表示按时间倒序左乘。展开后递归被彻底消掉了 :每个 S t S_t S t 都是初始状态、各步写入项与转移矩阵连乘的线性组合,只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化方案之一——换个括号方式(比如二叉树式两两合并),H H H 的连乘与写入项的累积可以在 O ( log n ) O(\log n) O ( log n ) 深度内并行完成。这就是 parallel scan / associative scan(并行扫描/结合扫描) 类方法的核心思想,也是 chunkwise 并行化和 SSM 训练并行(如 Mamba 的 selective scan)共同的理论根基。
到这里两条路线可以拼在一起了:DeltaNet 的 S t = H t S t − 1 + β t k t v t ⊤ S_t = H_t S_{t-1} + \beta_t k_t v_t^\top S t = H t S t − 1 + β t k t v t ⊤ 与 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 结构完全同构,差的只有一件事——H t H_t H t 的遗忘是「沿 k t k_t k t 方向删一块」,没有 SSM 那种全通道的指数衰减。
1.2 GDN:把遗忘门与定点改写拼在一起
Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。论文的核心洞察:gating 和 delta rule 是互补 的两种记忆管理机制:
机制
比喻
能力
缺陷
Gating(Mamba-2)
板擦
大面积擦除,即全局衰减
无法定点修改
Delta rule(DeltaNet)
铅笔
定点覆写某个 key 的关联
无法快速清空
序列长度超过状态容量时,记忆碰撞必然发生;GDN 以板擦与铅笔的组合来管理这块固定大小的白板。主约定下写作
S t = α t ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ S_t = \alpha_t\big(I - \beta_tk_tk_t^\top\big)S_{t-1} + \beta_t k_t v_t^\top
S t = α t ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤
从这里开始切换到转置约定 (S ^ = S ⊤ \hat S = S^\top S ^ = S ⊤ ,读出 o t = S ^ t q t o_t = \hat S_tq_t o t = S ^ t q t ),以便与 GDN 论文及 chunkwise 推导逐项对齐;§1.2 余下部分与 §1.3、§1.4 的数值验算和 chunkwise 公式全部使用它。同一个式子转置后是
S ^ t = S ^ t − 1 ( α t ( I − β t k t k t ⊤ ) ) + β t v t k t ⊤ \hat S_t = \hat S_{t-1}\big(\alpha_t(I - \beta_tk_tk_t^\top)\big) + \beta_t v_t k_t^\top
S ^ t = S ^ t − 1 ( α t ( I − β t k t k t ⊤ ) ) + β t v t k t ⊤
为减少符号负担,转置约定内部仍把状态记作 S t S_t S t (即下文 S t S_t S t 指 S ^ t \hat S_t S ^ t ,形状 d v × d k d_v \times d_k d v × d k ,读出为 S t q t S_tq_t S t q t 、旧读出为 S t − 1 k t S_{t-1}k_t S t − 1 k t )。
其中 α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 ) 是数据相关的标量门(Mamba-2 的参数化:α = exp ( − S o f t p l u s ( L i n e a r ( x t ) ) ) \alpha = \exp(-\mathrm{Softplus}(\mathrm{Linear}(x_t))) α = exp ( − Softplus ( Linear ( x t ))) ,在 log 空间计算以保证数值稳定)。通过三种极限情形理解这个式子:
极限
行为
对应模型
α t → 1 \alpha_t \to 1 α t → 1
纯 delta rule,只定点改写
DeltaNet
β t → 1 \beta_t \to 1 β t → 1 ,k ⊥ 已有记忆
退化为 S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_tS_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤
Mamba-2
α t → 0 \alpha_t \to 0 α t → 0
整表清零再写入(硬重置)
新能力:两者都做不到
几何解释 :( I − β k k ⊤ ) (I - \beta kk^\top) ( I − β k k ⊤ ) 是广义 Householder 反射,沿 k k k 方向压缩状态;标量 α \alpha α 则将整个状态矩阵均匀缩小。前者是定向 操作,后者是全局 操作,两者作用于不同自由度,因此可以叠加。
注意作用顺序 :擦除量是 β t ( S t − 1 k t ) \beta_t(S_{t-1}k_t) β t ( S t − 1 k t ) ,用的是未衰减 的旧读出;而 delta 对照的是衰减后 的旧值(v t − α t S t − 1 k t v_t - \alpha_tS_{t-1}k_t v t − α t S t − 1 k t )。α \alpha α 乘的是整个 ( I − β k k ⊤ ) (I-\beta kk^\top) ( I − β k k ⊤ ) ,这一点在实现时极易出错(把 α 只乘到擦除项上,chunkwise 形式会与递归形式对不上)。读取侧同样随时间累积衰减:token 在时间步 x x x 写入,在 x + t x+t x + t 读取时已经被 α x α x + 1 … α x + t \alpha_x\alpha_{x+1}\dots\alpha_{x+t} α x α x + 1 … α x + t 衰减过。实现中通过 γ r / γ i \gamma^r/\gamma^i γ r / γ i 项修正——分子分母都是 α \alpha α 连乘,相除就是区间衰减,本质是乘法形式的前缀和(prefix-sum) ,与 SSM 的 A ˉ \bar A A ˉ 连乘完全同源。
统一视角 。至此四个模型均已出现,它们本质上是同一个在线优化问题的闭式解 ,差别只在目标函数:
Linear Attn ( 纯加性 ) → Mamba-2 ( 全局遗忘门 ) ↘ DeltaNet ( 定点覆写 ) ↗ GDN ( 遗忘门 + 定点覆写 ) → KDA ( 逐通道门 + 下界 ) \text{Linear Attn}\,(\text{纯加性}) \to \text{Mamba-2}\,(\text{全局遗忘门}) \searrow \\
\text{DeltaNet}\,(\text{定点覆写}) \nearrow \;\; \text{GDN}\,(\text{遗忘门} + \text{定点覆写}) \to \text{KDA}\,(\text{逐通道门} + \text{下界})
Linear Attn ( 纯加性 ) → Mamba-2 ( 全局遗忘门 ) ↘ DeltaNet ( 定点覆写 ) ↗ GDN ( 遗忘门 + 定点覆写 ) → KDA ( 逐通道门 + 下界 )
先说清这个优化问题本身。把每个时间步看成一次在线学习:已有状态 S t − 1 S_{t-1} S t − 1 ,新到一对样本 ( k t , v t ) (k_t, v_t) ( k t , v t ) ,需要解出新状态 S t S_t S t 。所有四个模型的目标函数都是下面这个形式,自变量是 S t S_t S t ,S t − 1 S_{t-1} S t − 1 、k t k_t k t 、v t v_t v t 均为已知量:
L ( S t ) = ∥ S t − A t ∥ F 2 ⏟ 正则项 − 2 ⟨ S t k t , u t ⟩ ⏟ 拟合项 \mathcal{L}(S_t) = \underbrace{\|S_t - A_t\|_F^2}_{\text{正则项}} - \underbrace{2\langle S_tk_t,\ u_t\rangle}_{\text{拟合项}}
L ( S t ) = 正则项 ∥ S t − A t ∥ F 2 − 拟合项 2 ⟨ S t k t , u t ⟩
两项各自度量的是:
正则项 ∥ S t − A t ∥ F 2 \|S_t - A_t\|_F^2 ∥ S t − A t ∥ F 2 :新状态与锚点 A t A_t A t 的 Frobenius 距离,即所有矩阵元素差的平方和。它惩罚状态的改动量,对应记忆保留 。锚点取 S t − 1 S_{t-1} S t − 1 表示要求尽量不动,取 α t S t − 1 \alpha_tS_{t-1} α t S t − 1 表示允许先按 α t \alpha_t α t 收缩再比较,即容忍遗忘;
拟合项 − 2 ⟨ S t k t , u t ⟩ -2\langle S_tk_t,\ u_t\rangle − 2 ⟨ S t k t , u t ⟩ :用 k t k_t k t 检索新状态得到 S t k t S_tk_t S t k t ,再与写入目标 u t u_t u t 做内积。前面的负号使内积越大、损失越小,即要求检索结果朝 u t u_t u t 的方向对齐 ,对应关联学习。
该目标对 S t S_t S t 是二次的,∇ S t L = 2 ( S t − A t ) − 2 u t k t ⊤ = 0 \nabla_{S_t}\mathcal{L} = 2(S_t - A_t) - 2u_tk_t^\top = 0 ∇ S t L = 2 ( S t − A t ) − 2 u t k t ⊤ = 0 ,因此闭式解统一为
S t = A t + u t k t ⊤ S_t = A_t + u_tk_t^\top
S t = A t + u t k t ⊤
于是四个模型的差别可以完全归结为两个量的选择:锚点 A t A_t A t 决定怎么遗忘,写入目标 u t u_t u t 决定怎么写入 。
模型
锚点 A t A_t A t
写入目标 u t u_t u t
在线学习目标
状态更新的闭式解
Linear Attn
S t − 1 S_{t-1} S t − 1
v t v_t v t
∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩ |S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle ∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩
S t = S t − 1 + v t k t ⊤ S_t = S_{t-1} + v_tk_t^\top S t = S t − 1 + v t k t ⊤
Mamba-2
α t S t − 1 \alpha_tS_{t-1} α t S t − 1
v t v_t v t
∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩ |S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle ∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩
S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_tS_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤
DeltaNet
S t − 1 S_{t-1} S t − 1
β t ( v t − S t − 1 k t ) \beta_t(v_t - S_{t-1}k_t) β t ( v t − S t − 1 k t )
∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − S t − 1 k t ) ⟩ |S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - S_{t-1}k_t)\rangle ∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − S t − 1 k t )⟩
S t = S t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤ S_t = S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top S t = S t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤
GDN
α t S t − 1 \alpha_tS_{t-1} α t S t − 1
β t ( v t − α t S t − 1 k t ) \beta_t(v_t - \alpha_tS_{t-1}k_t) β t ( v t − α t S t − 1 k t )
∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − α t S t − 1 k t ) ⟩ |S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - \alpha_tS_{t-1}k_t)\rangle ∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − α t S t − 1 k t )⟩
S t = S t − 1 ( α t ( I − β t k t k t ⊤ ) ) + β t v t k t ⊤ S_t = S_{t-1}(\alpha_t(I-\beta_tk_tk_t^\top)) + \beta_tv_tk_t^\top S t = S t − 1 ( α t ( I − β t k t k t ⊤ )) + β t v t k t ⊤
KDA
S t − 1 D i a g ( α t ) S_{t-1}\mathrm{Diag}(\bm\alpha_t) S t − 1 Diag ( α t ) 逐通道收缩
β t ( v t − S t − 1 D i a g ( α t ) k t ) \beta_t(v_t - S_{t-1}\mathrm{Diag}(\bm\alpha_t)k_t) β t ( v t − S t − 1 Diag ( α t ) k t )
GDN 的逐通道化
S t = S t − 1 D i a g ( α t ) ( I − β t k t k t ⊤ ) + β t v t k t ⊤ S_t = S_{t-1}\mathrm{Diag}(\bm\alpha_t)(I-\beta_tk_tk_t^\top) + \beta_tv_tk_t^\top S t = S t − 1 Diag ( α t ) ( I − β t k t k t ⊤ ) + β t v t k t ⊤
把 A t A_t A t 与 u t u_t u t 代入 S t = A t + u t k t ⊤ S_t = A_t + u_tk_t^\top S t = A t + u t k t ⊤ 即可得到最后一列。以 DeltaNet 为例:u t = β t ( v t − S t − 1 k t ) u_t = \beta_t(v_t - S_{t-1}k_t) u t = β t ( v t − S t − 1 k t ) ,代入后 S t = S t − 1 + β t ( v t − S t − 1 k t ) k t ⊤ S_t = S_{t-1} + \beta_t(v_t - S_{t-1}k_t)k_t^\top S t = S t − 1 + β t ( v t − S t − 1 k t ) k t ⊤ ,展开即 S t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤ S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top S t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤ ,与前面用梯度下降推出的结果一致。
这张表读法很简单:损失里只有两样东西 ——v t v_t v t 是 ground truth(这个 key 本来该存的值),S t k t S_tk_t S t k t 是当前记忆对它的预测(这个 key 实际读出来的值),拟合项衡量两者是否一致;正则项则约束新状态别离锚点太远,即该保留多少旧记忆。四个模型的差别只在于:拟合项拿什么当 target(整份 v t v_t v t ,还是残差 v t − A t k t v_t - A_tk_t v t − A t k t ),正则项拿什么当锚点(固定的 S t − 1 S_{t-1} S t − 1 ,还是可收缩的 α t S t − 1 \alpha_tS_{t-1} α t S t − 1 ) 。GDN 在两处都取强化版本,KDA 再把锚点的标量收缩换成逐通道的 D i a g ( α t ) \mathrm{Diag}(\bm\alpha_t) Diag ( α t ) 。
说到底,这就是把「k → v k \to v k → v 这条记忆是否还对得上」写成了一个损失函数:对不上就修(拟合项),但别为了修这一条把整张表推翻(正则项)。
KDA 的最后一步由此确定。Kimi Linear(后演化为 KDA)在 Gated DeltaNet 基础上的核心改进是细粒度门控 :不再是每个注意力头一个标量 α \alpha α ,而是每个通道一个独立衰减值:
α t ∈ ( 0 , 1 ) d k ( channel-wise 遗忘 ) \alpha_t \in (0,1)^{d_k} \quad (\text{channel-wise 遗忘})
α t ∈ ( 0 , 1 ) d k ( channel-wise 遗忘 )
作用是模型可以对不同维度做不同程度的记忆衰减 :部分通道 α \alpha α 接近 1,保留长期信息;部分通道 α \alpha α 接近 0,快速遗忘。类比来说,Gated DeltaNet 的标量门相当于总开关,KDA 的逐通道门相当于每个通道各有一个独立调节旋钮:同一个状态内,慢通道承载长程依赖,快通道负责局部上下文,记忆容量按维度重新分配。再加上衰减下界(log \log log -decay 限制在 ( g min , 0 ) (g_{\min}, 0) ( g m i n , 0 ) )以保证数值稳定,即得到 KDA 的完整递推式,见下一小节。
下面先梳理 GDN(Gated DeltaNet)的机制:手工验算一遍、推导 chunkwise 并行形式,最后对照 KDA 分析其继承与改动。
1.3 GDN 数值验算:手算一遍
设定 d k = d v = 2 d_k = d_v = 2 d k = d v = 2 ,C = 3 C = 3 C = 3 个 token,S 0 = 0 S_0 = 0 S 0 = 0 。数据刻意构造为**k 1 = k 2 = e 1 k_1 = k_2 = e_1 k 1 = k 2 = e 1 ,以制造 key 碰撞**;query 取 Q = K Q = K Q = K ,即每一步都用当前 token 自己的 key 去查(q 1 = q 2 = e 1 q_1 = q_2 = e_1 q 1 = q 2 = e 1 、q 3 = e 2 q_3 = e_2 q 3 = e 2 ),这样读出结果能直接对照「这个地址此刻存的是什么」:
Q = K = [ 1 0 1 0 0 1 ] , V = [ 1 0 2 0 3 1 ] , α = [ 0.8 , 0.5 , 0.9 ] , β = [ 1 , 1 , 0.6 ] Q = K = \begin{bmatrix}1&0\\1&0\\0&1\end{bmatrix},\ V = \begin{bmatrix}1&0\\2&0\\3&1\end{bmatrix},\ \alpha = [0.8,\,0.5,\,0.9],\ \beta = [1,\,1,\,0.6]
Q = K = 1 1 0 0 0 1 , V = 1 2 3 0 0 1 , α = [ 0.8 , 0.5 , 0.9 ] , β = [ 1 , 1 , 0.6 ]
本节沿用 §1.2 声明的转置约定 :S ∈ R d v × d k S \in \mathbb{R}^{d_v \times d_k} S ∈ R d v × d k ,读出 o t = S t q t o_t = S_tq_t o t = S t q t ,旧读出为 S t − 1 k t S_{t-1}k_t S t − 1 k t 。
顺序递归 ,作为对照基准。累积衰减 γ j = ∏ i ≤ j α i = [ 0.8 , 0.4 , 0.36 ] \gamma_j = \prod_{i\le j}\alpha_i = [0.8, 0.4, 0.36] γ j = ∏ i ≤ j α i = [ 0.8 , 0.4 , 0.36 ] 。
t=1 (k = e 1 k=e_1 k = e 1 , v = [ 1 , 0 ] v=[1,0] v = [ 1 , 0 ] , α = 0.8 \alpha=0.8 α = 0.8 , β = 1 \beta=1 β = 1 ):S 0 k = 0 S_0k = 0 S 0 k = 0 ,无旧记忆可删:
S 1 = 0.8 ⋅ 0 + 1 ⋅ ( [ 1 , 0 ] − 0 ) e 1 ⊤ = [ 1 0 0 0 ] , o 1 = S 1 q 1 = S 1 e 1 = [ 1 , 0 ] S_1 = 0.8\cdot 0 + 1\cdot([1,0]-0)\,e_1^\top = \begin{bmatrix}1&0\\0&0\end{bmatrix}, \qquad o_1 = S_1q_1 = S_1e_1 = [1,0]
S 1 = 0.8 ⋅ 0 + 1 ⋅ ([ 1 , 0 ] − 0 ) e 1 ⊤ = [ 1 0 0 0 ] , o 1 = S 1 q 1 = S 1 e 1 = [ 1 , 0 ]
t=2 (k = e 1 k=e_1 k = e 1 , v = [ 2 , 0 ] v=[2,0] v = [ 2 , 0 ] , α = 0.5 \alpha=0.5 α = 0.5 , β = 1 \beta=1 β = 1 ):旧读出 S 1 k 2 = [ 1 , 0 ] = v 1 S_1k_2 = [1,0] = v_1 S 1 k 2 = [ 1 , 0 ] = v 1 ——碰撞发生 。擦除 − 1 ⋅ [ 1 , 0 ] e 1 ⊤ -1\cdot[1,0]e_1^\top − 1 ⋅ [ 1 , 0 ] e 1 ⊤ 把 v 1 v_1 v 1 完全擦掉;写入 delta = [ 2 , 0 ] − 0.5 ⋅ [ 1 , 0 ] = [ 1.5 , 0 ] = [2,0] - 0.5\cdot[1,0] = [1.5, 0] = [ 2 , 0 ] − 0.5 ⋅ [ 1 , 0 ] = [ 1.5 , 0 ] :
S 2 = 0.5 S 1 − [ 1 , 0 ] e 1 ⊤ + [ 1.5 , 0 ] e 1 ⊤ = [ 2 0 0 0 ] S_2 = 0.5S_1 - [1,0]e_1^\top + [1.5,0]e_1^\top = \begin{bmatrix}2&0\\0&0\end{bmatrix}
S 2 = 0.5 S 1 − [ 1 , 0 ] e 1 ⊤ + [ 1.5 , 0 ] e 1 ⊤ = [ 2 0 0 0 ]
此时读出 o 2 = S 2 q 2 = S 2 e 1 = [ 2 , 0 ] o_2 = S_2q_2 = S_2e_1 = [2,0] o 2 = S 2 q 2 = S 2 e 1 = [ 2 , 0 ] ,即 v 2 v_2 v 2 ,v 1 v_1 v 1 已被覆写(对照线性注意力:纯加性写入会得到 [3,0] 的叠加结果)。
t=3 (k = e 2 k=e_2 k = e 2 , v = [ 3 , 1 ] v=[3,1] v = [ 3 , 1 ] , α = 0.9 \alpha=0.9 α = 0.9 , β = 0.6 \beta=0.6 β = 0.6 ):旧读出 S 2 k 3 = 0 S_2k_3 = 0 S 2 k 3 = 0 ,正交无碰撞。写入 0.6 ⋅ [ 3 , 1 ] e 2 ⊤ 0.6\cdot[3,1]e_2^\top 0.6 ⋅ [ 3 , 1 ] e 2 ⊤ ,同时全表再乘 α \alpha α (此处也乘了已写入的 e 1 e_1 e 1 行):
S 3 = [ 1.8 1.8 0 0.6 ] , o 3 = S 3 q 3 = S 3 e 2 = [ 1.8 , 0.6 ] S_3 = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}, \qquad o_3 = S_3q_3 = S_3e_2 = [1.8,\,0.6]
S 3 = [ 1.8 0 1.8 0.6 ] , o 3 = S 3 q 3 = S 3 e 2 = [ 1.8 , 0.6 ]
检查点:为什么 S 3 [ 0 , 0 ] = 1.8 S_3[0,0] = 1.8 S 3 [ 0 , 0 ] = 1.8 不是 2.0? t=3 的 α 3 = 0.9 \alpha_3=0.9 α 3 = 0.9 作用在整张表 上:第 1 行 2.0 × 0.9 = 1.8 2.0\times0.9 = 1.8 2.0 × 0.9 = 1.8 。这就是 gating 与 delta 的交互:即使 token 3 的 key 与 e 1 e_1 e 1 正交,它的遗忘门仍然衰减了 e 1 e_1 e 1 通道上的记忆。逐通道门(KDA)与下界衰减都是围绕这一约束做文章。
1.4 GDN 的 Chunkwise 并行形式
推理用递归(O ( 1 ) O(1) O ( 1 ) /token),但训练/prefill 必须并行。GDN 的贡献是把 gating 并入 DeltaNet 的 WY 表示 chunkwise 框架。设 chunk 大小 C,chunk 入口状态 S 0 S_0 S 0 ,目标:一次矩阵乘算出整个 chunk 的 O 和 chunk 出口状态 S C S_C S C 。
第一步:部分展开递归 :
S r = γ r S 0 P r ⏟ F r + ∑ i = 1 r γ r γ i u ~ i k i ⊤ ⏟ G r S_r = \underbrace{\gamma_r S_0 P_r}_{F_r} + \underbrace{\sum_{i=1}^r\frac{\gamma_r}{\gamma_i}\,\tilde u_i k_i^\top}_{G_r}
S r = F r γ r S 0 P r + G r i = 1 ∑ r γ i γ r u ~ i k i ⊤
γ r = ∏ j ≤ r α j \gamma_r = \prod_{j\le r}\alpha_j γ r = ∏ j ≤ r α j (α \alpha α 是标量,可提到矩阵连乘外面);
P r = ∏ i ≤ r ( I − β i k i k i ⊤ ) P_r = \prod_{i\le r}(I - \beta_ik_ik_i^\top) P r = ∏ i ≤ r ( I − β i k i k i ⊤ ) :纯 Householder 连乘,与 gating 无关 ;
u ~ i \tilde u_i u ~ i :吸收了 β \beta β 和衰减修正的「伪 value」。
第二步:WY 表示——秩 1 连乘压缩成两个小矩阵 。
先说清为什么非做不可。P r P_r P r 是 C C C 个 d k × d k d_k\times d_k d k × d k 矩阵逐个相乘,O ( C d k 3 ) O(C\,d_k^3) O ( C d k 3 ) 且严格顺序 ——比原递推还贵,并行化直接失败。输出侧同理,每个 o c o_c o c 都要「前 c c c 个擦除矩阵的乘积」。
关键观察:这种乘积永远不膨胀 。每个因子都是「单位阵减秩 1」,这类矩阵连乘的结果仍是「单位阵减一个低秩矩阵」,秩不超过因子个数:
P r : = ∏ i = 1 r ( I − β i k i k i ⊤ ) = I − ∑ i ≤ r w i k i ⊤ P_r := \prod_{i=1}^{r}(I - \beta_i k_i k_i^\top) = I - \sum_{i\le r} w_i k_i^\top
P r := i = 1 ∏ r ( I − β i k i k i ⊤ ) = I − i ≤ r ∑ w i k i ⊤
证明(对 r r r 归纳) 。r = 0 r=0 r = 0 时 P 0 = I P_0 = I P 0 = I ,空和成立。设 P r − 1 = I − ∑ i < r w i k i ⊤ P_{r-1} = I - \sum_{i<r} w_i k_i^\top P r − 1 = I − ∑ i < r w i k i ⊤ 已成立,右乘第 r r r 个因子:
P r = P r − 1 ( I − β r k r k r ⊤ ) = ( I − ∑ i < r w i k i ⊤ ) − β r ( I − ∑ i < r w i k i ⊤ ) k r k r ⊤ = I − ∑ i < r w i k i ⊤ − β r k r k r ⊤ + β r ∑ i < r w i ( k i ⊤ k r ) ⏟ 标量 k r ⊤ \begin{aligned}
P_r &= P_{r-1}\big(I - \beta_r k_r k_r^\top\big) \\[2pt]
&= \Big(I - \sum_{i<r} w_i k_i^\top\Big) - \beta_r\Big(I - \sum_{i<r} w_i k_i^\top\Big)k_r k_r^\top \\[2pt]
&= I - \sum_{i<r} w_i k_i^\top - \beta_r k_r k_r^\top + \beta_r \sum_{i<r} w_i \underbrace{(k_i^\top k_r)}_{\text{标量}} k_r^\top
\end{aligned}
P r = P r − 1 ( I − β r k r k r ⊤ ) = ( I − i < r ∑ w i k i ⊤ ) − β r ( I − i < r ∑ w i k i ⊤ ) k r k r ⊤ = I − i < r ∑ w i k i ⊤ − β r k r k r ⊤ + β r i < r ∑ w i 标量 ( k i ⊤ k r ) k r ⊤
第三、四项都以 k r ⊤ k_r^\top k r ⊤ 结尾,合并同类项:
P r = I − ∑ i < r w i k i ⊤ − β r ( k r − ∑ i < r w i ( k i ⊤ k r ) ) ⏟ 记作 w r k r ⊤ = I − ∑ i ≤ r w i k i ⊤ P_r = I - \sum_{i<r} w_i k_i^\top - \underbrace{\beta_r\Big(k_r - \sum_{i<r} w_i (k_i^\top k_r)\Big)}_{\text{记作 } w_r} k_r^\top
= I - \sum_{i\le r} w_i k_i^\top
P r = I − i < r ∑ w i k i ⊤ − 记作 w r β r ( k r − i < r ∑ w i ( k i ⊤ k r ) ) k r ⊤ = I − i ≤ r ∑ w i k i ⊤
归纳完成。注意 w r w_r w r 不是定义出来的技巧,而是「乘积保持低秩」这个要求逼出来的 ——要让结果保持 I − ∑ i w i k i ⊤ I - \sum_i w_ik_i^\top I − ∑ i w i k i ⊤ 的形式,括号里那一坨只能是 w r w_r w r :
w r = β r ( k r − ∑ i < r w i ( k i ⊤ k r ) ) w_r = \beta_r\Big(k_r - \sum_{i<r}w_i\,(k_i^\top k_r)\Big)
w r = β r ( k r − i < r ∑ w i ( k i ⊤ k r ) )
value 侧同理。把第 i i i 次写入 β i v i k i ⊤ \beta_iv_ik_i^\top β i v i k i ⊤ 穿过它之后所有的擦除矩阵与衰减,追踪一遍即得
u ~ r = β r ( v r − ∑ i < r u ~ i γ r γ i ( k i ⊤ k r ) ) \tilde u_r = \beta_r\Big(v_r - \sum_{i<r}\tilde u_i\,\tfrac{\gamma_r}{\gamma_i}(k_i^\top k_r)\Big)
u ~ r = β r ( v r − i < r ∑ u ~ i γ i γ r ( k i ⊤ k r ) )
两条递归的直觉 。w r w_r w r 是修正后的擦除向量 :第 r r r 步本想擦除 k r k_r k r 方向,但若 k r k_r k r 与之前的 k i k_i k i 有重叠(k i ⊤ k r ≠ 0 k_i^\top k_r \ne 0 k i ⊤ k r = 0 ),连乘展开时前面的擦除项已经顺带擦过这部分,w r w_r w r 把已擦的量减掉以避免重复擦除 ,减法权重恰是重叠度 k i ⊤ k r k_i^\top k_r k i ⊤ k r 。u ~ r \tilde u_r u ~ r 是同一修正的 value 版,唯一差别是那个衰减比 γ r / γ i \gamma_r/\gamma_i γ r / γ i :第 i i i 次写入到第 r r r 步时已多衰减 ∏ s = i + 1 r α s = γ r / γ i \prod_{s=i+1}^{r}\alpha_s = \gamma_r/\gamma_i ∏ s = i + 1 r α s = γ r / γ i 倍,故其干扰要按此比例打折——越早的写入衰减越多、干扰越小 。
为什么 γ \gamma γ 只出现在 u ~ \tilde u u ~ 里 (即 gating 并入的位置):GDN 的衰减是标量 ,标量与一切矩阵可交换,擦除连乘里的衰减可整体提到外面变成总因子 γ r \gamma_r γ r ,所以 w w w 的递归里看不见 γ \gamma γ ;但 value 侧每次写入的「存活时长」不同,这个相对 衰减无法外提,只能以比值留在递归里。对照 KDA:衰减变成向量后与擦除不可交换 ,γ \gamma γ 再也提不出去,只能渗进内积本身(见 §4.1 的 M c i M_{ci} M c i )。
第三步:UT 变换——递归变成一次下三角方程求解 。
两条递归里 w r w_r w r 只依赖 w i < r w_{i<r} w i < r ,是严格下三角依赖 ,因此可以整体写成矩阵方程。把 w r w_r w r 的递归移项:
w r + ∑ i < r β r ( k i ⊤ k r ) w i = β r k r w_r + \sum_{i<r} \beta_r(k_i^\top k_r)\,w_i = \beta_r k_r
w r + i < r ∑ β r ( k i ⊤ k r ) w i = β r k r
令 W ∈ R C × d k W \in \mathbb{R}^{C\times d_k} W ∈ R C × d k 的第 r r r 行为 w r ⊤ w_r^\top w r ⊤ ,并定义严格下三角矩阵
L r i = { β r ( k i ⊤ k r ) , i < r 0 , i ≥ r 即 L = s t r i c t L o w e r ( d i a g ( β ) K K ⊤ ) L_{ri} = \begin{cases}\beta_r\,(k_i^\top k_r), & i < r\\ 0, & i \ge r\end{cases}
\qquad\text{即}\qquad L = \mathrm{strictLower}\big(\mathrm{diag}(\beta)\,KK^\top\big)
L r i = { β r ( k i ⊤ k r ) , 0 , i < r i ≥ r 即 L = strictLower ( diag ( β ) K K ⊤ )
则 C C C 个方程可一次写成 ( I + L ) W = d i a g ( β ) K (I + L)\,W = \mathrm{diag}(\beta)\,K ( I + L ) W = diag ( β ) K ,于是
W = ( I + L ) − 1 d i a g ( β ) K = T plain K , T plain = [ I + s t r i c t L o w e r ( d i a g ( β ) K K ⊤ ) ] − 1 d i a g ( β ) W = (I+L)^{-1}\,\mathrm{diag}(\beta)\,K = T_{\text{plain}}K, \qquad T_{\text{plain}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)KK^\top)\big]^{-1}\mathrm{diag}(\beta)
W = ( I + L ) − 1 diag ( β ) K = T plain K , T plain = [ I + strictLower ( diag ( β ) K K ⊤ ) ] − 1 diag ( β )
求这个逆很便宜,原因是 L L L 幂零 。严格下三角矩阵满足 L C = 0 L^C = 0 L C = 0 ,所以 Neumann 级数有限项精确截断 :
( I + L ) − 1 = I − L + L 2 − ⋯ + ( − 1 ) C − 1 L C − 1 (I+L)^{-1} = I - L + L^2 - \cdots + (-1)^{C-1}L^{C-1}
( I + L ) − 1 = I − L + L 2 − ⋯ + ( − 1 ) C − 1 L C − 1
既不需要迭代、也不存在收敛性问题,一次前代法(forward substitution)即可,代价 O ( C 2 ) O(C^2) O ( C 2 ) ——相对于省下的 O ( C d k 3 ) O(C\,d_k^3) O ( C d k 3 ) 连乘完全可以忽略。这就是 UT 变换 :( I + L ) (I+L) ( I + L ) 是单位下三角 (Unit Triangular,对角为 1,因为 w r w_r w r 完整依赖自己),求解它即「UT」名字的来源。U ~ \tilde U U ~ 侧只需把内积换成带衰减比的版本:
U ~ = T gated V , T gated = [ I + s t r i c t L o w e r ( d i a g ( β ) ( Γ ⊙ K K ⊤ ) ) ] − 1 d i a g ( β ) \tilde U = T_{\text{gated}}V, \qquad T_{\text{gated}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)(\Gamma\odot KK^\top))\big]^{-1}\mathrm{diag}(\beta)
U ~ = T gated V , T gated = [ I + strictLower ( diag ( β ) ( Γ ⊙ K K ⊤ )) ] − 1 diag ( β )
所以 UT 不是额外发明的东西,它就是这两条递归的矩阵形态 。
乘法链 → 加法链:这才是并行的真正来源 。整件事的本质是把擦除矩阵的连乘 ∏ r ( I − β r k r k r ⊤ ) \prod_r(I - \beta_rk_rk_r^\top) ∏ r ( I − β r k r k r ⊤ ) 换成了求和 I − ∑ r w r k r ⊤ I - \sum_r w_rk_r^\top I − ∑ r w r k r ⊤ ,代价是求和项不再是原始的 k r k_r k r 而是带修正的 w r w_r w r ,修正系数由那个小三角求解预先算清。准确地说:把大的顺序依赖 (C C C 个 d k × d k d_k\times d_k d k × d k 矩阵依次相乘)换成了小的顺序依赖 (C × C C\times C C × C 三角求解)加一堆可任意并行的加法 。
求和为什么就是胜利:加法可交换、可结合,因而可任意分组——树形归约、分块、稠密 matmul,GPU 的全部并行性都在奖励「求和结构」;而连乘与递推必须一步一步来。
这个手法在本文这条技术路线上出现了至少四次,难度递增但模式相同:
场景
恒等式
把什么变成了求和
log 空间衰减
∏ s α s = exp ( ∑ s g s ) \prod_s \alpha_s = \exp(\sum_s g_s) ∏ s α s = exp ( ∑ s g s )
累积衰减 → cumsum(最字面的一个)
SSM 的卷积形式
x t = ∑ s A ˉ t − s B ˉ u s x_t = \sum_s \bar A^{t-s}\bar B u_s x t = ∑ s A ˉ t − s B ˉ u s
顺序递推 → 对历史的加权和,可用卷积/FFT
Mamba-2 的 SSD
递推 ≡ \equiv ≡ 半可分矩阵,块内 ( C B ⊤ ) ⊙ L (CB^\top)\odot L ( C B ⊤ ) ⊙ L
选择性递推 → 下三角掩码 × 求和
DeltaNet/GDN/KDA 的 WY-UT
∏ ( I − β k k ⊤ ) = I − ∑ w i k i ⊤ \prod(I-\beta kk^\top) = I - \sum w_ik_i^\top ∏ ( I − β k k ⊤ ) = I − ∑ w i k i ⊤
秩 1 连乘 → 外积求和(本节)
一个反向的注脚:状态本身 S = ∑ i u ~ i k i ⊤ S = \sum_i \tilde u_ik_i^\top S = ∑ i u ~ i k i ⊤ 也是求和——推理时每步加一个秩 1,训练时 WY 把顺序过程也变成求和。累积是语义,求和是算法 ,这条路线的美学是自洽的。
辨析:因果掩码 ≠ UT 变换 。两者容易混,因为碰巧都是下三角,但要管的完全是两件事:
因果掩码 (tril):把 score 矩阵的严格上三角直接置零,一行掩码操作,没有任何「变换」可言。它管的是「第 i 个输出不许看未来的 token」;
UT 变换 (求 ( I + L ) − 1 (I+L)^{-1} ( I + L ) − 1 ):处理的是块内历史写入之间的相互影响 。delta rule 每次写入都是「先读、再改」,块内第 2 次写入读到了第 1 次的结果,第 3 次读到前两次,历史写入相互耦合,UT 变换负责解耦。
两者都是下三角不是巧合,是同一个原因:因果性 。位置 i 的写入只能影响 i 之后的位置,所以「干扰系数矩阵」天然下三角;对角线天然是 1(自己的写入自己完整可见),于是要逆的矩阵恰好是单位 下三角——这就是「UT」(Unit Triangular)名字的由来。一句话:因果性决定了它是三角的,但做它的目的是解耦,不是掩码 。
为什么必须做 :并行化要求把 C 次顺序写入合并为一次外积累加 ∑ i v ~ i k i ⊤ \sum_i \tilde v_i k_i^\top ∑ i v ~ i k i ⊤ 。如果直接用原始 v i v_i v i 累加,重叠部分会被重复计算——第 1 次写入的内容会透过后续写入的「先读」环节被间接再写一遍。UT 变换算出每个 v i v_i v i 该扣除多少,使等式精确成立。不做的代价:要么结果错,要么退回逐 token 循环。
这类下三角变换是个大家族 。「顺序递推 ↔ 三角矩阵求逆」是个通用模式,KDA 的 UT 只是其中一员:
家族成员
三角结构
与 UT 的关系
三角方程组求解(前代/回代)
LU、Cholesky 分解之后的三角系统
UT 变换的计算过程即一次前代法,两者为同一算法
Householder QR 的紧凑 WY 表示
反射连乘 ∏ ( I − β k k ⊤ ) = I + Y T Y ⊤ \prod(I - \beta kk^\top) = I + YTY^\top ∏ ( I − β k k ⊤ ) = I + Y T Y ⊤ ,T T T 三角
「WY」「UT」两个名字的学术出处(Schreiber-Van Loan 1989);DeltaNet 把同样的打包思想借到 delta rule
因果卷积的逆(去卷积)
因果线性系统 = 下三角 Toeplitz,其逆也是下三角 Toeplitz
信号处理经典:「用三角逆矩阵解顺序依赖」
Mamba-2/SSD 的半可分矩阵
块内注意力 ( C B ⊤ ) ⊙ L (CB^\top)\odot L ( C B ⊤ ) ⊙ L ,L L L 下三角衰减
同一枚硬币另一面:顺序 SSM 递推等价于带结构下三角矩阵,KDA 的 A q k A^{qk} A q k 衰减注意力项完全是这个结构
幂零矩阵 Neumann 级数
( I + L ) − 1 = I − L + L 2 − ⋯ (I+L)^{-1} = I - L + L^2 - \cdots ( I + L ) − 1 = I − L + L 2 − ⋯
严格下三角矩阵幂零(L C = 0 L^C=0 L C = 0 )故有限项精确截断,即上面 UT 能精确且便宜算出的数学原因
归纳一条通则:凡是「顺序执行的因果更新」,在分块并行化时都会转化为一个下三角矩阵的求逆或求解问题 ,SSM、delta rule、因果卷积、QR 分解都属于这一模式。KDA chunkwise 中的 WY 表示与 UT 变换,分别是这个模式在「写入打包」与「依赖解耦」上的具体化。
Householder 与 WY 的出处(1958 / 1989)
WY 与 UT 都是数值线性代数的经典老物件,被 DeltaNet 系列「考古」出来复用。既然本文反复用到,把家谱交代清楚。
Householder 变换(1958)就是关于一个超平面的镜像反射 。给定单位法向量 u u u ,反射矩阵为 H = I − 2 u u ⊤ H = I - 2uu^\top H = I − 2 u u ⊤ 。把任意 x x x 拆成沿 u u u 的分量与平行镜面的分量,反射即把法向分量翻号:
H x = x − 2 u ( u ⊤ x ) = x − 2 ( u ⊤ x ) u Hx = x - 2u(u^\top x) = x - 2(u^\top x)\,u
H x = x − 2 u ( u ⊤ x ) = x − 2 ( u ⊤ x ) u
三个性质直接从这个结构来:对称 (H ⊤ = H H^\top = H H ⊤ = H )、正交 (H ⊤ H = I H^\top H = I H ⊤ H = I ,反射保长)、对合 (H 2 = I H^2 = I H 2 = I ,照两次镜子回到原样)。其特征值是一个 − 1 -1 − 1 (法向翻转)与 d − 1 d-1 d − 1 个 + 1 +1 + 1 (镜面内不动)。
它与 delta rule 的擦除矩阵是同族对象 。对比 I − 2 u u ⊤ I - 2uu^\top I − 2 u u ⊤ 与 I − β t k t k t ⊤ I - \beta_tk_tk_t^\top I − β t k t k t ⊤ :取 β t ∥ k t ∥ 2 = 2 \beta_t\|k_t\|^2 = 2 β t ∥ k t ∥ 2 = 2 时两者完全相同。结合 §1.1 算过的特征值 1 − β ∥ k ∥ 2 1 - \beta\|k\|^2 1 − β ∥ k ∥ 2 :
β \beta β (取 ∣ k ∣ = 1 |k|=1 ∣ k ∣ = 1 )
沿 k k k 的特征值
性质
β → 0 \beta \to 0 β → 0
→ 1 \to 1 → 1
几乎不动,不擦除
β = 1 \beta = 1 β = 1
0 0 0
该方向完全清零(投影)
β = 2 \beta = 2 β = 2
− 1 -1 − 1
正交反射,即 Householder
所以 delta rule 的擦除矩阵可以理解为没照到底的半面镜子 :β ∈ ( 0 , 1 ) \beta \in (0,1) β ∈ ( 0 , 1 ) 只做收缩而非翻转,代价是不再正交,好处是「擦除强度」成了可学习的连续量。Householder 的主战场是 QR 分解 :对第 1 列选一面镜子把对角线以下全照成 0,再对第 2 列选一面……m m m 面镜子依次照完得到上三角 R R R ,镜子之积即正交阵 Q Q Q 。
WY 表示(1989,Schreiber & Van Loan)解决的是「一串镜子怎么存」 。QR 做完后 Q = H 1 H 2 ⋯ H m Q = H_1H_2\cdots H_m Q = H 1 H 2 ⋯ H m 是 m m m 个 n × n n\times n n × n 反射之积,每次用它都重新连乘既贵又顺序。他们的观察是
Q = H 1 H 2 ⋯ H m = I + W Y ⊤ , W , Y ∈ R n × m Q = H_1H_2\cdots H_m = I + WY^\top, \qquad W, Y \in \mathbb{R}^{n\times m}
Q = H 1 H 2 ⋯ H m = I + W Y ⊤ , W , Y ∈ R n × m
m m m 个大方阵之积压缩成两个瘦矩阵 (m ≪ n m \ll n m ≪ n ),用的时候两次矩阵乘即可:Q x = x + W ( Y ⊤ x ) Qx = x + W(Y^\top x) Q x = x + W ( Y ⊤ x ) 。推导与上面 w r w_r w r 的归纳一模一样 :每个因子是「I I I 加秩 1」,连乘时秩只累加不膨胀,归纳地合并出第 j j j 个修正向量。进一步可写成 Q = I − Y T Y ⊤ Q = I - YTY^\top Q = I − Y T Y ⊤ ,其中 T T T 是 m × m m\times m m × m 上三角 矩阵,由一个小递归算出——这就是 UT 变换名字的出处 。这套东西在 LAPACK 的 QR 例程(xGEQRF / xORMQR)底层已经跑了几十年。
把家谱摆出来:
数值线性代数(1958 / 1989)
DeltaNet / GDN / KDA(2020s)
基本砖块
I − 2 u u ⊤ I - 2uu^\top I − 2 u u ⊤ (反射)
I − β k k ⊤ I - \beta kk^\top I − β k k ⊤ (擦除)
要打包的对象
m m m 面镜子之积 Q Q Q
chunk 内 C C C 次擦除之积 P C P_C P C
紧凑形式
I + W Y ⊤ I + WY^\top I + W Y ⊤ (或 I − Y T Y ⊤ I - YTY^\top I − Y T Y ⊤ )
I − ∑ i w i k i ⊤ I - \sum_i w_ik_i^\top I − ∑ i w i k i ⊤
修正向量递归
逐个镜子扣重叠
w r = β r ( k r − ∑ i < r w i k i ⊤ k r ) w_r = \beta_r(k_r - \sum_{i<r}w_i\,k_i^\top k_r) w r = β r ( k r − ∑ i < r w i k i ⊤ k r )
小三角因子
T T T (上三角)
( I + L ) − 1 (I+L)^{-1} ( I + L ) − 1 (单位下三角)
目的
Q Q Q 的应用变成 BLAS-3
顺序写入变成稠密 matmul
一句话总括:Householder 变换是一类「I I I 减秩 1」的反射矩阵,WY 表示是把这类矩阵的连乘压缩成两个瘦矩阵的经典技巧 ;DeltaNet 的作者发现 delta rule 的擦除矩阵恰是同一族对象,于是把这个三十多年前的打包技巧搬了过来,让 chunk 内的顺序写入得以并行——GDN 与 KDA 一路继承,本文的 UT 变换就是 compact-WY 里那个三角因子的计算。
回到公式。上面 T gated T_{\text{gated}} T gated 里的衰减感知掩码即 Γ i j = γ i / γ j ( i > j ) \Gamma_{ij} = \gamma_i/\gamma_j\ (i>j) Γ ij = γ i / γ j ( i > j ) ;W W W 与 U ~ \tilde U U ~ 只差内积处的一个 γ \gamma γ 比。下三角求解(C × C C\times C C × C ,实践中 C = 64 C = 64 C = 64 )用前代法完成。
第四步:chunk 输出与出口状态 。记号沿用论文的箭头约定:( ⋅ ) ← r = γ r ( ⋅ ) r \overleftarrow{(\cdot)}_r = \gamma_r(\cdot)_r ( ⋅ ) r = γ r ( ⋅ ) r (衰减到 chunk 首端),( ⋅ ) → r = γ C γ r ( ⋅ ) r \overrightarrow{(\cdot)}_r = \frac{\gamma_C}{\gamma_r}(\cdot)_r ( ⋅ ) r = γ r γ C ( ⋅ ) r (衰减到 chunk 末端)。
输出 (两项:读 chunk 外旧状态 + chunk 内交互):
O = Q ← S 0 ⊤ + ( Q K ⊤ ⊙ Γ causal ) ( U ~ − W ← S 0 ⊤ ) O = \overleftarrow{Q}\,S_0^\top + \big(QK^\top\odot\Gamma_{\text{causal}}\big)\big(\tilde U - \overleftarrow W S_0^\top\big)
O = Q S 0 ⊤ + ( Q K ⊤ ⊙ Γ causal ) ( U ~ − W S 0 ⊤ )
出口状态 :
S C = γ C S 0 + ( U ~ → − W → S 0 ⊤ ) ⊤ K ( U ~ → r = γ C γ r u ~ r ) S_C = \gamma_C S_0 + \big(\overrightarrow{\tilde U} - \overrightarrow W S_0^\top\big)^\top K \qquad(\overrightarrow{\tilde U}_r = \tfrac{\gamma_C}{\gamma_r}\tilde u_r)
S C = γ C S 0 + ( U ~ − W S 0 ⊤ ) ⊤ K ( U ~ r = γ r γ C u ~ r )
结构读法:
输出的第二项是「chunk 内小注意力」:Q K ⊤ ⊙ Γ causal QK^\top\odot\Gamma_{\text{causal}} Q K ⊤ ⊙ Γ causal 就是带衰减的因果注意力矩阵 ,attend 的对象不是 V 而是修正后的伪 value U ~ − W ← S 0 ⊤ \tilde U - \overleftarrow WS_0^\top U ~ − W S 0 ⊤ (后者是「旧状态在这个 chunk 里该被擦掉的部分」);
出口状态 = 旧状态整体衰减 γ C \gamma_C γ C + 修正量加权写入(权重 γ C / γ i \gamma_C/\gamma_i γ C / γ i :越早写入衰减越多);
全部计算都是 C × C C\times C C × C 、C × d k C\times d_k C × d k 、C × d v C\times d_v C × d v 的稠密矩阵乘——Tensor Core 友好 ;chunk 间只传一个 d v × d k d_v\times d_k d v × d k 矩阵。
数值验算:用上面顺序递归的例子核对 chunkwise 形式 。
Step 1 :K K ⊤ = [ 1 1 0 1 1 0 0 0 1 ] KK^\top = \begin{bmatrix}1&1&0\\1&1&0\\0&0&1\end{bmatrix} K K ⊤ = 1 1 0 1 1 0 0 0 1 ,Γ strict ⊙ K K ⊤ \Gamma_{\text{strict}}\odot KK^\top Γ strict ⊙ K K ⊤ 只有 (2,1) 处非零 = γ 2 / γ 1 × 1 = 0.5 = \gamma_2/\gamma_1\times1 = 0.5 = γ 2 / γ 1 × 1 = 0.5 。
Step 2 :解两个下三角方程(β = [ 1 , 1 , 0.6 ] \beta = [1,1,0.6] β = [ 1 , 1 , 0.6 ] ):
T plain = [ 1 0 0 − 1 1 0 0 0 0.6 ] , T gated = [ 1 0 0 − 0.5 1 0 0 0 0.6 ] T_{\text{plain}} = \begin{bmatrix}1&0&0\\-1&1&0\\0&0&0.6\end{bmatrix}, \qquad T_{\text{gated}} = \begin{bmatrix}1&0&0\\-0.5&1&0\\0&0&0.6\end{bmatrix}
T plain = 1 − 1 0 0 1 0 0 0 0.6 , T gated = 1 − 0.5 0 0 1 0 0 0 0.6
对比唯一差别 (2,1):-1 → -0.5,正是 γ 2 / γ 1 = 0.5 \gamma_2/\gamma_1 = 0.5 γ 2 / γ 1 = 0.5 的衰减 ——t=2 时 v 1 v_1 v 1 已被衰减一半,擦除它的需求也减半。
Step 3 :W = T plain K = [ 1 0 0 0 0 0.6 ] W = T_{\text{plain}}K = \begin{bmatrix}1&0\\0&0\\0&0.6\end{bmatrix} W = T plain K = 1 0 0 0 0 0.6 ,U ~ = T gated V = [ 1 0 1.5 0 1.8 0.6 ] \tilde U = T_{\text{gated}}V = \begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix} U ~ = T gated V = 1 1.5 1.8 0 0 0.6 。
看 u ~ \tilde u u ~ 的第二行:$v_2 - 0.5\cdot\tilde u_1(k_1^\top k_2) = [2,0] - 0.5[1,0] = $ [1.5, 0] ,与顺序递归中计算的 delta 完全一致 ✓(W 2 = 0 W_2 = 0 W 2 = 0 是因为 w 1 w_1 w 1 已将 e 1 e_1 e 1 方向完全擦除,β = 1 \beta=1 β = 1 时无需重复擦除)。
Step 4 (S 0 = 0 S_0=0 S 0 = 0 ,修正项消失):U ~ → = γ C γ i u ~ i \overrightarrow{\tilde U} = \frac{\gamma_C}{\gamma_i}\tilde u_i U ~ = γ i γ C u ~ i 逐行乘 [ 0.36 / 0.8 , 0.36 / 0.4 , 1 ] = [ 0.45 , 0.9 , 1 ] [0.36/0.8, 0.36/0.4, 1] = [0.45, 0.9, 1] [ 0.36/0.8 , 0.36/0.4 , 1 ] = [ 0.45 , 0.9 , 1 ] :
U ~ → = [ 0.45 0 1.35 0 1.8 0.6 ] , S C = U ~ → ⊤ K = [ 1.8 1.8 0 0.6 ] ✓ \overrightarrow{\tilde U} = \begin{bmatrix}0.45&0\\1.35&0\\1.8&0.6\end{bmatrix}, \qquad S_C = \overrightarrow{\tilde U}^\top K = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}\ \checkmark
U ~ = 0.45 1.35 1.8 0 0 0.6 , S C = U ~ ⊤ K = [ 1.8 0 1.8 0.6 ] ✓
与顺序递归的 S 3 S_3 S 3 完全一致 。
Step 5 :Γ causal = [ 1 0 0 0.5 1 0 0.45 0.9 1 ] \Gamma_{\text{causal}} = \begin{bmatrix}1&0&0\\0.5&1&0\\0.45&0.9&1\end{bmatrix} Γ causal = 1 0.5 0.45 0 1 0.9 0 0 1 ,O = ( Q K ⊤ ⊙ Γ causal ) U ~ O = (QK^\top\odot\Gamma_{\text{causal}})\tilde U O = ( Q K ⊤ ⊙ Γ causal ) U ~ :
O = [ 1 0 0 0.5 1 0 0 0 1 ] [ 1 0 1.5 0 1.8 0.6 ] = [ 1 0 2 0 1.8 0.6 ] ✓ O = \begin{bmatrix}1&0&0\\0.5&1&0\\0&0&1\end{bmatrix}\begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix} = \begin{bmatrix}1&0\\2&0\\1.8&0.6\end{bmatrix}\ \checkmark
O = 1 0.5 0 0 1 0 0 0 1 1 1.5 1.8 0 0 0.6 = 1 2 1.8 0 0 0.6 ✓
与顺序递归的 o 1 = [ 1 , 0 ] o_1=[1,0] o 1 = [ 1 , 0 ] 、o 2 = [ 2 , 0 ] o_2=[2,0] o 2 = [ 2 , 0 ] 、o 3 = [ 1.8 , 0.6 ] o_3=[1.8,0.6] o 3 = [ 1.8 , 0.6 ] 逐步精确一致 (含非零 S 0 S_0 S 0 的一般情形同样对拍通过)。
2. GDN 在网络里长什么样
§1 讲的都是状态怎么更新,属于 token mixer 内部的数学。但 q , k , v , α , β q, k, v, \alpha, \beta q , k , v , α , β 这些量本身从哪来、算完的 o ~ \tilde o o ~ 又怎么变成层输出,还没交代。本节补上这一层:GDN 采用 Llama 式宏架构,把 attention 换成 gated delta token mixer:
1 2 3 4 5 6 7 x ─┬─ W_q/W_k ─ ShortConv ─ SiLU ─ L2Norm ──> q, k ├─ W_v ──── ShortConv ─ SiLU ───────────> v ├─ W_α(线性投影, log 空间负 Softplus)──> α ∈ (0,1) 标量 ├─ W_β(线性投影 + Sigmoid)────────────> β ∈ (0,1) │ 递归/分块计算 gated delta rule ──> õ ├─ W_g ─ SiLU ─────────────────────────> 输出门 └─ y = W_o( Sigmoid(W_g x) ⊙ RMSNorm(õ) ), 再接 SwiGLU MLP + 残差
要点:q/k 依次经过 ShortConv、SiLU 与 L2Norm,分别提供局部上下文与归一化后的稳定擦除;α \alpha α 与 β \beta β 是标量,只需线性投影,不经过卷积;输出门沿用 Mamba 的 SiLU 门设计,K3 将其升级为满秩。
2.1 ShortConv
ShortConv(Short Convolution,短因果深度卷积)出自 Kimi Linear(K3 技术报告引用 [64]),位于 Q/K/V 线性投影之后、进入 KDA 递归之前 ——顺序是 x → W q / k / v → S h o r t C o n v → … x \to W_{q/k/v} \to \mathrm{ShortConv} \to \dots x → W q / k / v → ShortConv → … ,即卷积是「KDA 前置」而非「投影前置」。它只捕捉当前 token 之前少量的局部上下文,不引入未来 token 信息;且在 Kimi Linear / FLA 实现里卷积不是裸用的,后面紧跟 SiLU:s i l u ( c o n v ( x ) ) \mathrm{silu}(\mathrm{conv}(x)) silu ( conv ( x )) 。
输入单头向量 x t ∈ R d x_t \in \mathbb{R}^d x t ∈ R d ,卷积核窗口大小 W(标准取 4),深度可分离(depthwise)且因果:
S h o r t C o n v ( x ) t = ∑ j = 0 W − 1 w j ⊙ x t − j \mathrm{ShortConv}(x)_t = \sum_{j=0}^{W-1} w_j \odot x_{t-j}
ShortConv ( x ) t = j = 0 ∑ W − 1 w j ⊙ x t − j
记号
含义
w j ∈ R d w_j \in \mathbb{R}^d w j ∈ R d
每通道独立卷积权重——depthwise,每个特征通道一套独立卷积核
⊙ \odot ⊙
逐元素相乘
因果约束
t − j ≥ 0 t - j \ge 0 t − j ≥ 0 ;t < j t < j t < j 时零填充,看不到未来 token
三个性质决定了它为什么放在这个位置:
时序维度滑动,只混合最近 W 个历史 token ,开销远小于全局注意力;
depthwise :w k w_k w k 与 x t − k x_{t-k} x t − k 逐通道相乘,通道间不混——局部上下文的注入不破坏各通道独立的衰减语义(与 D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 逐通道门配套);
因果 :t − j ≥ 0 t-j \ge 0 t − j ≥ 0 保证只看历史,t < j t<j t < j 时零填充,LLM 自回归必备。
其作用是补足线性 RNN 缺失的局部建模能力:递归状态只携带压缩后的全局历史,最近若干 token 的精细局部模式(词内字符、短程搭配)由这层卷积负责。
2.2 Swish 与 SiLU
ShortConv 之后、L2Norm 之前是 SiLU,即 Swish——严格说 Swish 带可学参数时为 x ⋅ σ ( β x ) x \cdot \sigma(\beta x) x ⋅ σ ( β x ) ,论文和实现里固定 β = 1 \beta = 1 β = 1 ,两个名字就此等价:
S w i s h ( x ) = x ⋅ σ ( x ) , σ ( x ) = 1 1 + e − x \mathrm{Swish}(x) = x \cdot \sigma(x), \qquad \sigma(x) = \frac{1}{1 + e^{-x}}
Swish ( x ) = x ⋅ σ ( x ) , σ ( x ) = 1 + e − x 1
计算分两步:其一,对输入的每个元素计算 σ ( x ) = 1 / ( 1 + exp ( − x ) ) \sigma(x) = 1/(1+\exp(-x)) σ ( x ) = 1/ ( 1 + exp ( − x )) ;其二,原输入 x 与 sigmoid 结果逐元素相乘 得到输出。Swish = 输入 × 输入自己的 sigmoid 门——自门控 (self-gated):门控信号不是外部来的,是输入自身,免参数。
此处使用它的原因有两点:其一,平滑且非单调(负区间存在小幅下凹,x → − ∞ x \to -\infty x → − ∞ 时输出趋于 0,而非 ReLU 的硬截断),梯度处处非零,深层网络训练更稳定;其二,门控形式与整个 block 的「信息通过与抑制」语义一致,Q/K/V 进入递归状态前先经过一次自门控,相当于对局部卷积混合后的特征做一次软筛选,随后由 L2Norm 归一化到单位范数。K3 将输出侧的这道门升级为满秩输入相关门,其来源即此处的 SiLU。
2.3 L2Norm
Q/K 支路的最后一步。对单头向量 z ∈ R d k \bm{z} \in \mathbb{R}^{d_k} z ∈ R d k ,逐通道 L2 标准化到单位范数:
L 2 N o r m ( z ) = z ∥ z ∥ 2 + ϵ , ∥ z ∥ 2 = ∑ i = 1 d k z i 2 \mathrm{L2Norm}(\bm{z}) = \frac{\bm{z}}{\|\bm{z}\|_2 + \epsilon}, \qquad \|\bm{z}\|_2 = \sqrt{\sum_{i=1}^{d_k} z_i^2}
L2Norm ( z ) = ∥ z ∥ 2 + ϵ z , ∥ z ∥ 2 = i = 1 ∑ d k z i 2
记号
含义
ϵ \epsilon ϵ
极小防除零常数(10 − 6 10^{-6} 1 0 − 6 左右)
操作维度
每个注意力头独立归一化,跨头不共享统计
为什么只有 Q/K 使用、V 不使用?回到 §1.1 的结论:归一化后 ∥ k t ∥ 2 = 1 \|k_t\|_2 = 1 ∥ k t ∥ 2 = 1 ,于是 delta rule 的擦除矩阵 I − β k k ⊤ I - \beta k k^\top I − β k k ⊤ 的特征值落在 [ 1 − β , 1 ] [1-\beta,\ 1] [ 1 − β , 1 ] 区间内,写入步长因此稳定,β = 1 \beta = 1 β = 1 时为精确的保长反射。Q 也做归一化 ,目的是使读出 o = S ⊤ q o = S^\top q o = S ⊤ q 的尺度可控:query 与 key 的范数均为 1,内积才具备可比性,chunkwise 公式中 Q K ⊤ QK^\top Q K ⊤ 的 score 也被限制在 [ − 1 , 1 ] [-1, 1] [ − 1 , 1 ] 内。V 不做归一化,因为写入内容的幅值本身携带信息,且 β t k t v t ⊤ \beta_t k_t v_t^\top β t k t v t ⊤ 的稳定性由 k 侧保证,与 v 的尺度无关。这与 softmax 注意力中 QK-norm 防止 logit 爆炸的动机同源,但在线性 RNN 中它还额外承担擦除操作 的数值稳定性,因此归一化并非可选优化,而是 delta rule 正常工作的前提。
至此 Q/K 支路的四步全部交代完毕:Linear 投影、ShortConv(局部混合)、SiLU(软筛选)、L2Norm(归一化)。这几步均为因果、逐通道操作,不涉及跨头统计,与 D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 的逐通道语义一致。三者各自负责一部分数值稳定性:ShortConv 负责局部混合,L2Norm 负责幅度稳定,逐通道衰减差 γ i − γ j ≤ 0 \gamma_i - \gamma_j \le 0 γ i − γ j ≤ 0 负责指数项不溢出 。这条全链路的数值稳定性是 KDA 能在 BF16 下训练的前提。
混合架构 :GDN + 滑窗注意力(H1)或 Mamba-2 + GDN + SWA(H2)交错堆叠,互补长短程建模能力,与 K3 的「3 KDA + 1 MLA」思路一致,GDN 论文可视为这一混合范式的早期工作。
3. 从 GDN 到 KDA:K3 继承了什么、改了什么
维度
GDN
KDA(Kimi Linear → K3)
动机
更新式
S t − 1 ( α t ( I − β k k ⊤ ) ) + β v k ⊤ S_{t-1}(\alpha_t(I-\beta kk^\top)) + \beta vk^\top S t − 1 ( α t ( I − β k k ⊤ )) + β v k ⊤
( I − β k k ⊤ ) D i a g ( α t ) S t − 1 + β v k ⊤ (I-\beta kk^\top)\,\mathrm{Diag}(\alpha_t)S_{t-1} + \beta vk^\top ( I − β k k ⊤ ) Diag ( α t ) S t − 1 + β v k ⊤
α 从标量 → 逐通道向量
遗忘门粒度
标量(整头同一衰减率)
向量 α t ∈ ( 0 , 1 ) d k \alpha_t\in(0,1)^{d_k} α t ∈ ( 0 , 1 ) d k
长期记忆通道与短期工作区分离
门作用顺序
α 在外乘整个更新
Diag(α) 在内侧先衰减 S,再删写
与逐通道参数化配套(KCP 推导需要)
α 参数化
负 Softplus,值域 ( − ∞ , 0 ) (-\infty,0) ( − ∞ , 0 )
K3:缩放 sigmoid,下界 g min = − 5 g_{\min}=-5 g m i n = − 5
1 / Γ 1/\Gamma 1/Γ 有界 → 对角 tile 全走 Tensor Core
chunkwise
WY + UT + 衰减箭头
同框架,Γ 从标量比变成向量累积比
表达力↑ 数值难度↑(K3 用下界解决)
输出门
低秩/简单门
K3:输入相关全秩 sigmoid 门
逐通道调节读出
两处容易忽略的改动:
作用顺序变了 :GDN 是 S ( α ( I − β k k ⊤ ) ) S(\alpha(I-\beta kk^\top)) S ( α ( I − β k k ⊤ )) (α 吸进 Householder 连乘),KDA 是 ( I − β k k ⊤ ) D i a g ( α ) S (I-\beta kk^\top)\mathrm{Diag}(\alpha)S ( I − β k k ⊤ ) Diag ( α ) S (α 先作用、再删写)。KDA 的逐通道门是矩阵 D i a g ( α ) \mathrm{Diag}(\alpha) Diag ( α ) ,与 Householder 不交换,顺序成为实质设计选择——它使 KCP 的「段转移分解」成为可能。
衰减在 Γ 处的复杂化 :GDN 的 γ 是标量连乘,Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j 只是数;KDA 的 Γ 是向量 逐元素累积,chunk 公式里 K/Γ、Q⊙Γ 等运算随之复杂化,数值范围问题(1 / Γ 1/\Gamma 1/Γ 溢出)由此而生——K3 的下界衰减正是对 GDN→KDA 这一步引入的新问题的修复。
3.1 SSM 路线的结论
SSM 从控制论的状态方程出发,经 ZOH 离散化、S4、Mamba 到 Mamba-2,推导过程见前置篇第 2 节。本文需要的只是其结论:乘在旧状态上的标量衰减系数 α t \alpha_t α t ,即线性注意力一族对「状态如何遗忘」的回答,但它是全局衰减、不区分通道。它与 DeltaNet 的定点覆写在 GDN 处结合,GDN 又被 KDA 逐通道化。整体骨架始终不变:
状态 × ( 衰减/删除算子 ) + ( 写入项 ) \text{状态} \times (\text{衰减/删除算子}) + (\text{写入项})
状态 × ( 衰减 / 删除算子 ) + ( 写入项 )
3.2 KDA 的递推公式
本节回到主约定 (S t ∈ R d k × d v \mathbf{S}_t \in \mathbb{R}^{d_k\times d_v} S t ∈ R d k × d v ,读出 S t ⊤ q t \mathbf{S}_t^\top\bm q_t S t ⊤ q t ),与 KDA 论文写法一致:
S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤ \mathbf{S}_t = \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) \operatorname{Diag}(\bm{\alpha}_t)\, \mathbf{S}_{t-1} + \beta_t \bm{k}_t \bm{v}_t^\top
S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤
o ~ t = S t ⊤ q t \tilde{\bm{o}}_t = \mathbf{S}_t^\top \bm{q}_t
o ~ t = S t ⊤ q t
符号说明 ,沿用论文定义:
符号
含义
α t ∈ R d k \bm{\alpha}_t \in \mathbb{R}^{d_k} α t ∈ R d k
逐通道一维保留因子向量(channel-wise one-step retention factor)
Diag ( α t ) \operatorname{Diag}(\bm{\alpha}_t) Diag ( α t )
向量转对角矩阵算子,把通道级衰减系数变成矩阵乘
I \mathbf{I} I
同维度单位矩阵
β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 )
delta rule 写入强度
该式包含三步操作 ,按从右往左的顺序:
通道衰减 :Diag ( α t ) S t − 1 \operatorname{Diag}(\bm{\alpha}_t)\mathbf{S}_{t-1} Diag ( α t ) S t − 1 等价于对 S t − 1 \mathbf{S}_{t-1} S t − 1 的每一行分别乘以对应通道的 α \alpha α 系数——逐通道缩放历史记忆。普通线性注意力/SSM 大多全局统一衰减(α t \alpha_t α t 是标量),KDA 给每一个 key 通道分配独立衰减权重:不同语义通道可以选择更快/更慢遗忘 S t − 1 \mathbf{S}_{t-1} S t − 1 ;
方向擦除 :( I − β t k t k t ⊤ ) \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) ( I − β t k t k t ⊤ ) 再对衰减后的状态做当前 token 的定点擦除(秩 1 Householder 型,详见 §1.1);
新信息写入 :叠加 β t k t v t ⊤ \beta_t \bm{k}_t \bm{v}_t^\top β t k t v t ⊤ 。
三步合起来:先通道衰减,再方向擦除,最后写入 ——带通道精细遗忘的 delta 递推。作用顺序是设计选择:Diag 与 Householder 不交换,这个顺序使 KCP 的段转移分解成为可能。
记号约定 :论文里大写 Diag ( ⋅ ) \operatorname{Diag}(\cdot) Diag ( ⋅ ) 是「向量 → 对角矩阵」算子;小写 diag ( ⋅ ) \operatorname{diag}(\cdot) diag ( ⋅ ) 有时指反向操作(输入矩阵、提取对角线为向量)。本文全程大写表示向量转对角矩阵。
3.3 为什么是逐通道门:动机与在线学习视角
标量门的表达力瓶颈 。GDN 的 α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 ) 是标量:每一步遗忘时,所有 key 通道以同一个比例 衰减——要么一起记住,要么一起忘记。但不同通道承担的角色不同:
有的通道在存「长期主题」(希望 α ≈ 1 \alpha \approx 1 α ≈ 1 ,几乎不遗忘);
有的通道在存「临时指针」(希望快速衰减,腾出容量)。
KDA 的核心改动只有一处:把标量换成向量 α t ∈ ( 0 , 1 ) d k \bm{\alpha}_t \in (0,1)^{d_k} α t ∈ ( 0 , 1 ) d k (β t \beta_t β t 仍是标量,这是 KDA 的选择而非必须)。直觉:S 的第 j 列对应 key 空间的第 j 个通道,D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 作用上去就是给每一列配一个独立的遗忘速度 。GDN 是 KDA 在 D i a g ( α t ) = α t I \mathrm{Diag}(\alpha_t) = \alpha_t I Diag ( α t ) = α t I 时的特例。
在线学习视角:逐通道权重衰减的 delta rule 。与 §1.2 的在线学习表同构,只需把正则项换成逐通道版。每步给定新样本 ( k t , v t ) (k_t, v_t) ( k t , v t ) ,希望新状态 S 满足两点:其一,拟合新样本 S k t ≈ v t Sk_t \approx v_t S k t ≈ v t ;其二,不过度偏离「衰减后的历史记忆」S ≈ S t − 1 D i a g ( α t ) S \approx S_{t-1}\mathrm{Diag}(\alpha_t) S ≈ S t − 1 Diag ( α t ) (以下记 D t = D i a g ( α t ) D_t = \mathrm{Diag}(\alpha_t) D t = Diag ( α t ) ,α t = exp ( g t ) \alpha_t = \exp(g_t) α t = exp ( g t ) ,g t ∈ R < 0 d k g_t \in \mathbb{R}_{<0}^{d_k} g t ∈ R < 0 d k ):
L t ( S ) = 1 2 ∥ S k t − v t ∥ 2 ⏟ 拟合 + 1 2 ∥ S − S t − 1 D t ∥ F 2 ⏟ 逐通道正则 \mathcal{L}_t(S) = \underbrace{\tfrac12 \|S k_t - v_t\|^2}_{\text{拟合}} + \underbrace{\tfrac12 \|S - S_{t-1} D_t\|_F^2}_{\text{逐通道正则}}
L t ( S ) = 拟合 2 1 ∥ S k t − v t ∥ 2 + 逐通道正则 2 1 ∥ S − S t − 1 D t ∥ F 2
从 S t − 1 D t S_{t-1}D_t S t − 1 D t 出发对第一项做一步梯度下降(步长 β t \beta_t β t ):
S t = S t − 1 D t − β t ( S t − 1 D t k t − v t ) k t ⊤ = S t − 1 D t ( I − β t k t k t ⊤ ) + β t v t k t ⊤ S_t = S_{t-1}D_t - \beta_t\,(S_{t-1}D_t\, k_t - v_t)\,k_t^\top = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top
S t = S t − 1 D t − β t ( S t − 1 D t k t − v t ) k t ⊤ = S t − 1 D t ( I − β t k t k t ⊤ ) + β t v t k t ⊤
正是递推式(转置约定下)。逐通道正则的含义 :第 j 列的「信任区域」宽度正比于 α t ( j ) \alpha_t^{(j)} α t ( j ) ——α \alpha α 小的通道,旧记忆先被压缩,新写入覆盖几乎没有阻力(快遗忘);α ≈ 1 \alpha \approx 1 α ≈ 1 的通道,旧记忆原样进入下一步,delta rule 只做精细增量(慢遗忘)。这就是「细粒度记忆控制」的优化论表述:KDA = 逐通道权重衰减 + delta rule。
实现注意:作用顺序不可颠倒 。D t D_t D t 作用于整个 S t − 1 S_{t-1} S t − 1 、先于擦除项 ——擦除时检索用的也是已衰减 的状态:
S t = S t − 1 D t ⏟ 先衰减 ( I − β t k t k t ⊤ ) ⏟ 再擦除/写入 S_t = \underbrace{S_{t-1} D_t}_{\text{先衰减}}\underbrace{(I - \beta_t k_t k_t^\top)}_{\text{再擦除/写入}}
S t = 先衰减 S t − 1 D t 再擦除 / 写入 ( I − β t k t k t ⊤ )
写成 S t − 1 ( I − β k k ⊤ ) D t S_{t-1}(I-\beta kk^\top)D_t S t − 1 ( I − β k k ⊤ ) D t (顺序颠倒)或只衰减单位阵部分都是错的——D t D_t D t 与 ( I − β k k ⊤ ) (I - \beta kk^\top) ( I − β k k ⊤ ) 不可交换 ,顺序错了结果就错(与 §1.2「α 乘整个括号」的警告同源,但逐通道化之后错误更隐蔽)。FLA 参考实现里对应 S = S * g.exp() 之后立刻用衰减后的 S 做检索 v - k^T S。
3.4 下界衰减(Lower-bounded decay)
KDA 的衰减参数化是一个关键改进。Kimi Linear 使用无界的负 Softplus 映射 g = − e A Softplus ( z ) ∈ ( − ∞ , 0 ) g = -e^A \text{Softplus}(z) \in (-\infty, 0) g = − e A Softplus ( z ) ∈ ( − ∞ , 0 ) ,而 K3 改用有界缩放 sigmoid :
g t h = g min ⋅ Sigmoid ( e A h z t h ) ∈ ( g min , 0 ) g_t^h = g_{\min} \cdot \text{Sigmoid}(e^{A_h} z_t^h) \in (g_{\min}, 0)
g t h = g m i n ⋅ Sigmoid ( e A h z t h ) ∈ ( g m i n , 0 )
其中 g min = − 5 g_{\min} = -5 g m i n = − 5 固定,A h A_h A h 是可学习的每头对数尺度。这意味着每个保留因子满足 α > e − 5 ≈ 6.7 × 10 − 3 \alpha > e^{-5} \approx 6.7 \times 10^{-3} α > e − 5 ≈ 6.7 × 1 0 − 3 ,16 token 块的累积对数衰减落在 ( − 80 , 0 ) (-80, 0) ( − 80 , 0 ) 内,对应的重缩放因子小于 e 80 e^{80} e 80 ,在 BF16 动态范围内。
计算收益 :有限范围使得因果对角块和离对角块都能用密集 Tensor Core 矩阵乘法,消除了 Kimi Linear 中需要的 position-pair 对角计算路径。
1 / Γ 1/\Gamma 1/Γ 溢出:向量门控引入的数值问题
为什么这个修复在 GDN 上不必要、在 KDA 上变成刚需?GDN 的标量衰减在 chunkwise 里只以比值 出现(∏ s = j + 1 i α s ≤ 1 \prod_{s=j+1}^{i}\alpha_s \le 1 ∏ s = j + 1 i α s ≤ 1 ,i ≥ j i \ge j i ≥ j ),天然安全。KDA 把 Γ \Gamma Γ 变成向量累积 Γ t = exp ( ∑ s ≤ t g s ) \Gamma_t = \exp(\sum_{s\le t} g_s) Γ t = exp ( ∑ s ≤ t g s ) ,逐通道独立:第 4 步式的推导可以刻意只用 i ≥ j i \ge j i ≥ j 的差(指数 ≤ 0 \le 0 ≤ 0 ,安全);但只要换一种等价写法——把状态「反归一化」回 chunk 起点、或把衰减从 key 上整体外提(论文公式 (4) 的 K / Γ K/\Gamma K /Γ 因式分解正是这种写法)——就得真的算出
Γ t − 1 = exp ( − ∑ s ≤ t g s ) (逐通道) \Gamma_t^{-1} = \exp\Big(-\sum_{s \le t} g_s\Big) \quad \text{(逐通道)}
Γ t − 1 = exp ( − s ≤ t ∑ g s ) ( 逐通道 )
衰减越快的通道,− γ t -\gamma_t − γ t 越大,1 / Γ 1/\Gamma 1/Γ 呈指数膨胀。具体量级如下:
恒定 α \alpha α
t = 100 t=100 t = 100
t = 500 t=500 t = 500
t = 1000 t=1000 t = 1000
t = 2000 t=2000 t = 2000
0.9
3.8 × 10 4 3.8\times10^{4} 3.8 × 1 0 4
7.6 × 10 22 7.6\times10^{22} 7.6 × 1 0 22
5.7 × 10 45 5.7\times10^{45} 5.7 × 1 0 45
3.3 × 10 91 3.3\times10^{91} 3.3 × 1 0 91
0.5
1.3 × 10 30 1.3\times10^{30} 1.3 × 1 0 30
3.3 × 10 150 3.3\times10^{150} 3.3 × 1 0 150
1.1 × 10 301 1.1\times10^{301} 1.1 × 1 0 301
溢出 float64
0.1
10 100 10^{100} 1 0 100
溢出
溢出
溢出
float64 上界约 e 709 ≈ 1.8 × 10 308 e^{709} \approx 1.8\times10^{308} e 709 ≈ 1.8 × 1 0 308 ;BF16 训练下数十步即会溢出。标量推广为向量后,衰减率的动态范围被放大 ,数值溢出由个别情形变为普遍现象,因此必须对衰减率本身设置边界。
Kimi Linear 的规避方案 :在对数空间算相对衰减(减法代替除法,不溢出),并把每个 chunk 再切成 16 token 的二级瓦片。效果:瓦片之间 (非对角块)可以安全交给 Tensor Core 稠密矩阵乘;但瓦片内部 (对角块)衰减可能极端,仍需按位置对显式计算,position-pair 路径无法组织成大矩阵乘,Tensor Core 利用率低,成为块内主要瓶颈。
K3 的解决方式 :不修改公式,而是修改参数化,使溢出在数学上不可能发生(负 Softplus 允许 α \alpha α 任意接近 0,即「一步清零」;缩放 sigmoid 不允许)。
表达力为何不受损 :快通道在瓦片内仍可将记忆衰减到 e − 80 ≈ 10 − 35 e^{-80} \approx 10^{-35} e − 80 ≈ 1 0 − 35 (数值上接近零,但并非严格为零);长期遗忘靠多步连乘(每步乘 0.0067 0.0067 0.0067 ),数十步后衰减量已足够小,不需要单步清零的能力。以一个衰减下界换取全链路 Tensor Core 化,是有利的取舍。
3.5 满秩门控(Full-rank gate)
K3 将 KDA 的输出门从低秩参数化改为输入相关的满秩投影 。在递推输出经过 head-wise RMSNorm 后,应用数据相关的输出门控:
y t = W o [ Sigmoid ( W g x t ) ⊙ RMSNorm ( o ~ t ) ] y_t = W_o [\text{Sigmoid}(W_g x_t) \odot \text{RMSNorm}(\tilde{o}_t)]
y t = W o [ Sigmoid ( W g x t ) ⊙ RMSNorm ( o ~ t )]
满秩门控允许每个 token 独立调制从循环状态读取的通道。
4. KDA 的 Chunkwise 并行形式
递推形式在推理阶段是优势——O ( 1 ) O(1) O ( 1 ) 状态更新;但在训练与 prefill 阶段成为瓶颈:每个 token 的状态依赖前一个 token,顺序循环使 GPU 的数千个核心无法并行工作。Chunkwise 并行化 把序列切成长度 C 的块:块内矩阵运算并行,块间只传状态。先看通用的复杂度框架,再看 KDA 逐通道门带来的推导细节。
通用框架 。Chunk 内 :token 交互用带衰减的因果注意力直接算,O ( C 2 d ) O(C^2 d) O ( C 2 d ) ——C 是常数(64/128),对序列长度 N 线性。Chunk 间 :每个 chunk 对状态做一次递推更新,chunk 内外积先归约成固定大小的 state 增量,块间只传 d × d d \times d d × d 的 state。总计算量:
O ( N ⋅ C ⋅ d ) ⏟ chunk 内注意力 + O ( N / C ⋅ d 2 ) ⏟ chunk 间状态更新 = 2 N d 2 ( 固定项 ) + 2 N C d \underbrace{O\big(N \cdot C \cdot d\big)}_{\text{chunk 内注意力}} + \underbrace{O\big(N/C \cdot d^2\big)}_{\text{chunk 间状态更新}} \;=\; 2Nd^2(\text{固定项}) + 2NCd
chunk 内注意力 O ( N ⋅ C ⋅ d ) + chunk 间状态更新 O ( N / C ⋅ d 2 ) = 2 N d 2 ( 固定项 ) + 2 N C d
C
形态
复杂度
C = 1 C = 1 C = 1
每个 token 一个 chunk,chunk 内注意力消失
纯线性注意力递推,FLOPs 最少但不一定最快(GPU 对小矩阵乘利用率低)
C = N C = N C = N
整个序列一个 chunk,递推消失
标准 O ( N 2 ) O(N^2) O ( N 2 ) 注意力
实践中 C 取 64 或 128:足够小以控制 chunk 间项的开销,足够大以让 C × C C\times C C × C 注意力矩阵填满 Tensor Core 的 tile。这与 S4 时期「训练用卷积、推理用递归」的双形式思路一致:选择性打破 LTI 后,chunkwise 就是卷积的继任者。
4.1 逐通道衰减下的 WY 表示与 UT 变换
记 chunk 起点传入状态 S [ 0 ] S_{[0]} S [ 0 ] (以下用局部下标 i = 1.. C i = 1..C i = 1.. C ;乘积按时间倒序)。核心记号是累积 log 衰减 :
γ i = ∑ s = 1 i g s ∈ R d k , Γ i ← j = diag ( exp ( γ i − γ j ) ) ( i ≥ j ) \gamma_i = \sum_{s=1}^{i} g_s \in \mathbb{R}^{d_k}, \qquad \Gamma_{i \leftarrow j} = \operatorname{diag}\big(\exp(\gamma_i - \gamma_j)\big) \quad (i \ge j)
γ i = s = 1 ∑ i g s ∈ R d k , Γ i ← j = diag ( exp ( γ i − γ j ) ) ( i ≥ j )
Γ i ← j \Gamma_{i\leftarrow j} Γ i ← j 是「从第 j 步衰减到第 i 步」的逐通道算子。两条性质:i ≥ j i \ge j i ≥ j 时 γ i − γ j ≤ 0 \gamma_i - \gamma_j \le 0 γ i − γ j ≤ 0 逐分量成立(g s < 0 g_s < 0 g s < 0 ),元素都在 ( 0 , 1 ] (0,1] ( 0 , 1 ] ,数值安全;反向 Γ j ← i − 1 \Gamma_{j\leftarrow i}^{-1} Γ j ← i − 1 的元素 ≥ 1 \ge 1 ≥ 1 且随 chunk 长度指数增长 ——这是 1 / Γ 1/\Gamma 1/Γ 爆炸的根源,先记住这个观察。
第一步:衰减 KKT 矩阵 M 。类比 GDN chunkwise 里的普通 k i ⊤ k j k_i^\top k_j k i ⊤ k j ,逐通道版需要带衰减的 key-key 内积:
M c i = ( k c ⊙ e γ c − γ i ) ⊤ k i , 1 ≤ i < c ≤ C M_{ci} = \big(k_c \odot e^{\gamma_c - \gamma_i}\big)^\top k_i, \qquad 1 \le i < c \le C
M c i = ( k c ⊙ e γ c − γ i ) ⊤ k i , 1 ≤ i < c ≤ C
含义:k i k_i k i 写入的记忆衰减到第 c 步时,与 k c k_c k c 的重叠程度。Hadamard 积 ⊙ \odot ⊙ 作用在 key 通道维——正是逐通道衰减出现的位置。
第二步:UT 变换 。构造严格下三角 L L L :L c i = β i M c i ( c > i ) L_{ci} = \beta_i M_{ci}\ (c > i) L c i = β i M c i ( c > i ) ,求 T = ( I + L ) − 1 T = (I+L)^{-1} T = ( I + L ) − 1 (幂零,有限项截断;实践中不显式求逆,前代法逐行解),再每列乘 β j \beta_j β j 得 A ^ = T diag ( β ) \hat{A} = T\,\operatorname{diag}(\beta) A ^ = T diag ( β ) 。FLA 实现中的三行循环,数学上即求解 ( I + L ) T = I (I+L)T = I ( I + L ) T = I 。
与 GDN 对照 :GDN 这一步的 M 是普通 k c ⊤ k i k_c^\top k_i k c ⊤ k i (标量衰减被拆成比值吸收进别的项);KDA 里衰减「长」在 M 内部 ,无法外提——这是向量门控带来的结构性变化,KDA chunkwise 推导的核心难点。
第三步:WY 表示 :
W = A ^ ( e γ i ⊙ k i ) i = 1.. C ∈ R C × d k , U = A ^ V ∈ R C × d v , V ~ = U − W S [ 0 ] ⊤ W = \hat{A}\,\big(e^{\gamma_i} \odot k_i\big)_{i=1..C} \in \mathbb{R}^{C \times d_k}, \qquad
U = \hat{A}\,V \in \mathbb{R}^{C \times d_v}, \qquad \tilde{V} = U - W\,S_{[0]}^\top
W = A ^ ( e γ i ⊙ k i ) i = 1.. C ∈ R C × d k , U = A ^ V ∈ R C × d v , V ~ = U − W S [ 0 ] ⊤
伪值 V ~ \tilde V V ~ 的含义 。UT 变换产出 U 和 W,由它们定义 V ~ : = U − W S \tilde V := U - WS V ~ := U − W S 。三个符号的角色如下:
符号
形状
角色
S [ 0 ] S_{[0]} S [ 0 ]
d k × d v d_k \times d_v d k × d v
历史记忆 :chunk 之前所有 token 写入的内容,块内计算时固定不变
W W W
C × d k C \times d_k C × d k
历史读取算子 :每行是某位置的衰减 key 修正组合;W S WS W S 表示该位置能从历史记忆中读到的内容
U U U
C × d v C \times d_v C × d v
块内互扣后的写入目标 :A ^ \hat A A ^ 作用于 V,块内写入之间的重叠已扣除
V ~ \tilde V V ~
C × d v C \times d_v C × d v
实际写入的净增量 = 写入目标减去(历史记忆已有部分 + 块内其他位置已写部分)
称其为「伪」值的原因:它们并非真实的 v(手算例子中 v ~ 2 = ( − 2 , 2 ) ≠ v 2 \tilde v_2 = (-2,2) \ne v_2 v ~ 2 = ( − 2 , 2 ) = v 2 ),而是扣除全部重叠后可直接累加、无需再做修正 的增量。这与单步 delta rule 一致:单步写入为 v t − S t − 1 ⊤ k t v_t - S_{t-1}^\top k_t v t − S t − 1 ⊤ k t ,即目标值减去历史读数;V ~ \tilde V V ~ 是其 chunk 并行版本,每行给出该位置真正新增 的部分。重叠在 V ~ \tilde V V ~ 中已扣除完毕,后续的块内注意力 A V ~ A\tilde V A V ~ 与状态更新才能以单次矩阵乘完成。
第四步:输出与跨 chunk 状态 :
A c j q k = ( q c ⊙ e γ c − γ j ) ⊤ k j ( j ≤ c ) , o c = S [ 0 ] ( q c ⊙ e γ c ) ⏟ 块间:衰减后的 query 检索历史状态 + ∑ j ≤ c A c j q k v ~ j ⏟ 块内:衰减注意力 A^{qk}_{cj} = \big(q_c \odot e^{\gamma_c - \gamma_j}\big)^\top k_j \ (j \le c), \qquad o_c = \underbrace{S_{[0]}\,(q_c \odot e^{\gamma_c})}_{\text{块间:衰减后的 query 检索历史状态}} + \underbrace{\sum_{j \le c} A^{qk}_{cj}\,\tilde v_j}_{\text{块内:衰减注意力}}
A c j q k = ( q c ⊙ e γ c − γ j ) ⊤ k j ( j ≤ c ) , o c = 块间:衰减后的 query 检索历史状态 S [ 0 ] ( q c ⊙ e γ c ) + 块内 : 衰减注意力 j ≤ c ∑ A c j q k v ~ j
S [ C ] = S [ 0 ] Γ C ← 0 + ∑ c = 1 C v ~ c ( k c ⊙ e γ C − γ c ) ⊤ S_{[C]} = S_{[0]}\,\Gamma_{C\leftarrow 0} + \sum_{c=1}^{C} \tilde v_c \big(k_c \odot e^{\gamma_C - \gamma_c}\big)^\top
S [ C ] = S [ 0 ] Γ C ← 0 + c = 1 ∑ C v ~ c ( k c ⊙ e γ C − γ c ) ⊤
出口状态读法:旧状态整体按整 chunk 累积衰减缩小;每个伪值以「写入时刻衰减到块末」的 key 为地址写入。所有指数都是 ≤ 0 \le 0 ≤ 0 的差,整条链路没有一个 ≥ 1 \ge 1 ≥ 1 的因子 ——刻意保持,原因见下界衰减一节。
4.2 数值验算(d k = d v = 2 d_k = d_v = 2 d k = d v = 2 ,C = 2)
零初始状态,每步恒定衰减 α = ( 0.5 , 0.25 ) \alpha = (0.5, 0.25) α = ( 0.5 , 0.25 ) (即 g = ( ln 0.5 , ln 0.25 ) g = (\ln 0.5, \ln 0.25) g = ( ln 0.5 , ln 0.25 ) ),β 取 1:
k 1 = ( 1 1 ) , v 1 = ( 2 0 ) ; k 2 = ( 1 2 ) , v 2 = ( 0 2 ) ; q 1 = q 2 = ( 1 1 ) k_1 = \binom{1}{1},\ v_1 = \binom{2}{0};\quad k_2 = \binom{1}{2},\ v_2 = \binom{0}{2};\quad q_1 = q_2 = \binom{1}{1}
k 1 = ( 1 1 ) , v 1 = ( 0 2 ) ; k 2 = ( 2 1 ) , v 2 = ( 2 0 ) ; q 1 = q 2 = ( 1 1 )
递推式 。第 1 步(S 0 = 0 S_0 = 0 S 0 = 0 ):S 1 = v 1 k 1 ⊤ = ( 2 2 0 0 ) S_1 = v_1k_1^\top = \begin{pmatrix}2&2\\0&0\end{pmatrix} S 1 = v 1 k 1 ⊤ = ( 2 0 2 0 ) ,o 1 = ( 4 , 0 ) o_1 = (4,0) o 1 = ( 4 , 0 ) 。第 2 步,D = d i a g ( 0.5 , 0.25 ) D = \mathrm{diag}(0.5, 0.25) D = diag ( 0.5 , 0.25 ) ,先衰减 S 1 D = ( 1 0.5 0 0 ) S_1D = \begin{pmatrix}1&0.5\\0&0\end{pmatrix} S 1 D = ( 1 0 0.5 0 ) (第 1 列 ×0.5、第 2 列 ×0.25——逐通道在动 ),再擦除写入:
S 1 D ( I − k 2 k 2 ⊤ ) = ( − 1 − 3.5 0 0 ) , S 2 = ( − 1 − 3.5 2 4 ) , o 2 = ( − 4.5 , 6 ) S_1D(I - k_2k_2^\top) = \begin{pmatrix}-1&-3.5\\0&0\end{pmatrix}, \qquad S_2 = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix}, \qquad o_2 = (-4.5,\ 6)
S 1 D ( I − k 2 k 2 ⊤ ) = ( − 1 0 − 3.5 0 ) , S 2 = ( − 1 2 − 3.5 4 ) , o 2 = ( − 4.5 , 6 )
Chunkwise 。e γ 1 = ( 0.5 , 0.25 ) e^{\gamma_1} = (0.5, 0.25) e γ 1 = ( 0.5 , 0.25 ) ,e γ 2 = ( 0.25 , 0.0625 ) e^{\gamma_2} = (0.25, 0.0625) e γ 2 = ( 0.25 , 0.0625 ) ,e γ 2 − γ 1 = ( 0.5 , 0.25 ) e^{\gamma_2-\gamma_1} = (0.5, 0.25) e γ 2 − γ 1 = ( 0.5 , 0.25 ) 。衰减 KKT:M 21 = ( k 2 ⊙ e γ 2 − γ 1 ) ⊤ k 1 = ( 0.5 , 0.5 ) ⋅ ( 1 , 1 ) = 1 M_{21} = (k_2 \odot e^{\gamma_2-\gamma_1})^\top k_1 = (0.5, 0.5)\cdot(1,1) = 1 M 21 = ( k 2 ⊙ e γ 2 − γ 1 ) ⊤ k 1 = ( 0.5 , 0.5 ) ⋅ ( 1 , 1 ) = 1 。UT:L = ( 0 0 1 0 ) L = \begin{pmatrix}0&0\\1&0\end{pmatrix} L = ( 0 1 0 0 ) ,T = I − L = ( 1 0 − 1 1 ) T = I - L = \begin{pmatrix}1&0\\-1&1\end{pmatrix} T = I − L = ( 1 − 1 0 1 ) ,A ^ = T \hat A = T A ^ = T 。WY:
W = A ^ ( 0.5 0.25 0.25 0.125 ) = ( 0.5 0.25 − 0.25 − 0.125 ) , U = A ^ ( 2 0 0 2 ) = ( 2 0 − 2 2 ) W = \hat A\begin{pmatrix}0.5&0.25\\0.25&0.125\end{pmatrix} = \begin{pmatrix}0.5&0.25\\-0.25&-0.125\end{pmatrix}, \quad U = \hat A\begin{pmatrix}2&0\\0&2\end{pmatrix} = \begin{pmatrix}2&0\\-2&2\end{pmatrix}
W = A ^ ( 0.5 0.25 0.25 0.125 ) = ( 0.5 − 0.25 0.25 − 0.125 ) , U = A ^ ( 2 0 0 2 ) = ( 2 − 2 0 2 )
伪值(S [ 0 ] = 0 ⇒ V ~ = U S_{[0]}=0 \Rightarrow \tilde V = U S [ 0 ] = 0 ⇒ V ~ = U ):v ~ 1 = ( 2 , 0 ) \tilde v_1 = (2,0) v ~ 1 = ( 2 , 0 ) ,v ~ 2 = ( − 2 , 2 ) \tilde v_2 = (-2,2) v ~ 2 = ( − 2 , 2 ) 。注意 v ~ 2 ≠ v 2 \tilde v_2 \ne v_2 v ~ 2 = v 2 :因为 k 2 k_2 k 2 与衰减后的 k 1 k_1 k 1 写入重叠(M 21 = 1 ≠ 0 M_{21} = 1 \ne 0 M 21 = 1 = 0 ),WY 把 v 2 v_2 v 2 修正为扣除重叠后真正的新增——单步 delta rule 的 chunk 版样子。衰减注意力与输出:
A q k = ( 2 0 0.75 3 ) , o 1 = 2 v ~ 1 = ( 4 , 0 ) ✓ , o 2 = 0.75 v ~ 1 + 3 v ~ 2 = ( − 4.5 , 6 ) ✓ A^{qk} = \begin{pmatrix}2&0\\0.75&3\end{pmatrix}, \qquad o_1 = 2\tilde v_1 = (4,0)\ \checkmark, \qquad o_2 = 0.75\,\tilde v_1 + 3\,\tilde v_2 = (-4.5,\ 6)\ \checkmark
A q k = ( 2 0.75 0 3 ) , o 1 = 2 v ~ 1 = ( 4 , 0 ) ✓ , o 2 = 0.75 v ~ 1 + 3 v ~ 2 = ( − 4.5 , 6 ) ✓
块末状态:S [ 2 ] = v ~ 1 ( k 1 ⊙ e γ 2 − γ 1 ) ⊤ + v ~ 2 k 2 ⊤ = ( − 1 − 3.5 2 4 ) S_{[2]} = \tilde v_1(k_1 \odot e^{\gamma_2-\gamma_1})^\top + \tilde v_2 k_2^\top = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix} S [ 2 ] = v ~ 1 ( k 1 ⊙ e γ 2 − γ 1 ) ⊤ + v ~ 2 k 2 ⊤ = ( − 1 2 − 3.5 4 ) ✓ \checkmark ✓ 与递推式逐项一致。(随机对拍:非零初始状态、随机门控下递推 vs chunkwise 最大误差 10 − 9 10^{-9} 1 0 − 9 量级。)
4.3 与论文公式 (4) 的对照
K3 报告(沿用 Kimi Linear)用乘积记号:γ i → j = ∏ r = i j α r = exp ( ∑ r = i j g r ) \gamma_{i\to j} = \prod_{r=i}^{j}\alpha_r = \exp(\sum_{r=i}^j g_r) γ i → j = ∏ r = i j α r = exp ( ∑ r = i j g r ) ;Γ ∈ R C × d k \Gamma \in \mathbb{R}^{C\times d_k} Γ ∈ R C × d k 把各步 γ \gamma γ 按行堆叠(注意 Γ \Gamma Γ 本身不是对角阵 ,每个位置 r r r 的 d i a g ( γ r ) \mathrm{diag}(\gamma_r) diag ( γ r ) 才是遗忘对角阵,Γ \Gamma Γ 是 C 个对角阵的打包)。论文状态约定 S ∈ R d k × d v S \in \mathbb{R}^{d_k \times d_v} S ∈ R d k × d v (本文的转置),其公式 (4):
A = T r i l [ ( Q ⊙ Γ ) ( K / Γ ) ⊤ ] , O = ( Γ ⊙ Q ) S ⏟ 块间 + A V ~ ⏟ 块内 A = \mathrm{Tril}\big[(Q\odot\Gamma)(K/\Gamma)^\top\big], \qquad O = \underbrace{(\Gamma\odot Q)S}_{\text{块间}} + \underbrace{A\,\tilde V}_{\text{块内}}
A = Tril [ ( Q ⊙ Γ ) ( K /Γ ) ⊤ ] , O = 块间 ( Γ ⊙ Q ) S + 块内 A V ~
为什么 A 能这样拆 :看 ( i , j ) (i,j) ( i , j ) 元素 A i j = ∑ d q i , d γ i , d ⋅ k j , d / γ j , d = ( q i ⊙ γ j → i ) ⊤ k j A_{ij} = \sum_d q_{i,d}\gamma_{i,d}\cdot k_{j,d}/\gamma_{j,d} = (q_i \odot \gamma_{j\to i})^\top k_j A ij = ∑ d q i , d γ i , d ⋅ k j , d / γ j , d = ( q i ⊙ γ j → i ) ⊤ k j ,与 A q k A^{qk} A q k 逐元素相同——逐对位置的衰减比值被因式分解成 query 侧乘 Γ、key 侧除以 Γ 两个逐位置操作,整个 C × C C\times C C × C 矩阵 = 一次稠密 matmul + 两次逐元素乘,完全并行。Tril 保留对角线 :delta rule 里 o i o_i o i 读的是写入当前 token 之后 的状态。
用上面的例子核对 :Q ⊙ Γ = ( 0.5 0.25 0.25 0.0625 ) Q\odot\Gamma = \begin{pmatrix}0.5&0.25\\0.25&0.0625\end{pmatrix} Q ⊙ Γ = ( 0.5 0.25 0.25 0.0625 ) ,K / Γ = ( 2 4 4 32 ) K/\Gamma = \begin{pmatrix}2&4\\4&32\end{pmatrix} K /Γ = ( 2 4 4 32 ) ,乘积 Tril 后 A = ( 2 0 0.75 3 ) A = \begin{pmatrix}2&0\\0.75&3\end{pmatrix} A = ( 2 0.75 0 3 ) ,与上面 A q k A^{qk} A q k 一致。注意 K / Γ K/\Gamma K /Γ 里已出现 32 这种被放大的数——1 / Γ 1/\Gamma 1/Γ 的膨胀就藏在这一步 ,数值后果见下界衰减一节。
C C C 与 d k d_k d k 是否存在倍数关系? 不存在。C 是序列轴的切分(一个 chunk 装多少 token),d k d_k d k 是特征轴的宽度,两根轴独立。有整除要求的是:K3 把 chunk 再切成 16 token 二级瓦片,故 C 需是 16 的倍数;d k d_k d k 、d v d_v d v 按 Tensor Core 友好尺寸取是工程约束,不是数学要求。
KDA 的核心优势是:相比传统 softmax attention,它是线性复杂度 的;相比纯粹的线性 RNN,它通过 delta rule 实现了更强大的信息写入和遗忘控制。
4.4 速查表(KDA 与 GDN)
维度
内容
状态
S ∈ R^{d_v×d_k}(论文记法),固定大小,每序列每头一份
更新
S = S(α(I-βkkT)) + βvkT,全局衰减 α 与定点覆写 δ 的组合
读出
o = Sq,后接 RMSNorm + sigmoid 输出门
α
GDN 标量 / KDA 逐通道,数据相关;β:sigmoid
q/k
Linear → ShortConv → SiLU → L2Norm
等价视角
在线最小二乘 SGD(β 为学习率)+ 自适应权重衰减(α)
并行训练
chunkwise:WY 表示(W 无 γ / Ũ 有 γ)+ UT 变换(下三角求逆)+ 衰减箭头
chunk 输出
O = Q⃖S0T + (QKT ⊙ Γ_causal)(Ũ - W⃖S0T)
chunk 状态
S_C = γ_C S0 + (Ũ→ - W→S0T)TK
记忆瓶颈的应对
覆写(δ)解决叠加干扰;全局衰减(α)解决长期累积
后继
KDA:α→逐通道、换作用顺序、加下界 → 进 K3
参考 :
Yang, Kautz & Hatamizadeh, Gated Delta Networks: Improving Mamba2 with Delta Rule , arXiv:2412.06464 (GDN 出处,ICLR 2025)
Kimi K3 Technical Report, arXiv:2607.24653 (KDA 出处)
Gu & Dao, Mamba: Linear-Time Sequence Modeling with Selective State Spaces , arXiv:2312.00752
Gu et al., Efficiently Modeling Long Sequences with Structured State Spaces (S4), arXiv:2111.00396
Katharopoulos et al., Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention , arXiv:2006.16236