TileLang 实战:KDA 从零到一–Gated DeltaNet
前两篇分别实现了无遗忘的 chunked 线性注意力,以及带标量衰减的版本。两者的状态更新都只做加法 :新的键值对累加进状态,旧信息靠衰减系数被动遗忘。本篇补上最后一块–删除 。
递推式变成:
S t = S t − 1 ( α t ( I − β t k t k t ⊺ ) ) + β t v t k t ⊺ \mathbf{S}_t = \mathbf{S}_{t-1}\Big(\alpha_t\big(\mathbf{I} - \beta_t \bm{k}_t\bm{k}_t^{\intercal}\big)\Big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}
S t = S t − 1 ( α t ( I − β t k t k t ⊺ ) ) + β t v t k t ⊺
比上一篇多出的是 I − β t k t k t ⊺ \mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} I − β t k t k t ⊺ 这个广义 Householder 变换 。它带来的不是又一个逐元素权重,而是一个矩阵乘在状态右侧–这一改把整个 chunkwise 算法的结构改掉了:块内不再能靠一张下三角权重表解决,需要解一个 C × C C \times C C × C 的三角系统。
上两篇见《Chunked 线性注意力 》与《标量衰减 》,本文沿用其符号约定与四层参考验证框架。
0. 符号约定
沿用 GDN(arXiv:2412.06464v3)§2.2、§3.1、§3.3 与附录 A 的记号,本篇新增两个:
符号
含义
说明
β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 )
写入强度 (writing strength)
也是 delta rule 视角下的学习率
T [ t ] ∈ R C × C \mathbf{T}_{[t]} \in \mathbb{R}^{C \times C} T [ t ] ∈ R C × C
UT 变换矩阵
下三角系统的逆,本篇的核心开销
U ~ [ t ] ∈ R C × d v \widetilde{\mathbf{U}}_{[t]} \in \mathbb{R}^{C \times d_v} U [ t ] ∈ R C × d v
修正后的 value
T \mathbf{T} T 作用在 diag ( β ) V \operatorname{diag}(\beta)\mathbf{V} diag ( β ) V 上
W [ t ] ∈ R C × d k \mathbf{W}_{[t]} \in \mathbb{R}^{C \times d_k} W [ t ] ∈ R C × d k
修正后的 key
T \mathbf{T} T 作用在 diag ( β ) K \operatorname{diag}(\beta)\mathbf{K} diag ( β ) K 上
沿用的:α t \alpha_t α t 单步衰减、γ [ t ] r = ∏ i = 1 r α i \gamma^r_{[t]} = \prod_{i=1}^{r}\alpha_i γ [ t ] r = ∏ i = 1 r α i 累积衰减积(按 chunk 重置)、Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j 衰减感知因果掩码、C C C 块长、[ t ] [t] [ t ] 块序号、r ∈ [ 1 , C ] r \in [1,C] r ∈ [ 1 , C ] 块内位置。
一个前提 :论文对 q , k \bm{q}, \bm{k} q , k 做 L2 归一化,即 ∥ k t ∥ = 1 \|\bm{k}_t\| = 1 ∥ k t ∥ = 1 。这不只是训练稳定性的考虑–下面 §1.2 会看到它直接决定了 Householder 变换会不会把状态越推越大。
1. 从加法到删除:delta rule 在做什么
1.1 三种更新方式的对照
把三篇的递推式并排放,差异一目了然:
形态
递推式
状态如何变化
线性注意力
S t = S t − 1 + v t k t ⊺ \mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S t = S t − 1 + v t k t ⊺
只增不减
标量衰减
S t = α t S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S t = α t S t − 1 + v t k t ⊺
整体等比遗忘
gated delta rule
S t = S t − 1 ( α t ( I − β t k t k t ⊺ ) ) + β t v t k t ⊺ \mathbf{S}_t = \mathbf{S}_{t-1}\big(\alpha_t(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal})\big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal} S t = S t − 1 ( α t ( I − β t k t k t ⊺ ) ) + β t v t k t ⊺
定向替换 + 整体遗忘
标量衰减的问题在于它不分对象 :α t \alpha_t α t 一乘,所有历史信息按同一比例衰减。要腾出空间写入新内容,只能把无关的旧信息一起冲淡。GDN 论文的说法是,门控擅长「快速擦除」,delta rule 擅长「定向修改」,两者互补。
delta rule 的定向性来自哪里?把它拆开看:
S t = S t − 1 − ( S t − 1 k t ) ⏟ v t old k t ⊺ + ( β t v t + ( 1 − β t ) S t − 1 k t ) ⏟ v t new k t ⊺ \mathbf{S}_t = \mathbf{S}_{t-1} - \underbrace{(\mathbf{S}_{t-1}\bm{k}_t)}_{\bm{v}_t^{\text{old}}}\bm{k}_t^{\intercal} + \underbrace{\big(\beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t\big)}_{\bm{v}_t^{\text{new}}}\bm{k}_t^{\intercal}
S t = S t − 1 − v t old ( S t − 1 k t ) k t ⊺ + v t new ( β t v t + ( 1 − β t ) S t − 1 k t ) k t ⊺
读法是:先把 k t \bm{k}_t k t 这个键上原有的值 v t old = S t − 1 k t \bm{v}^{\text{old}}_t = \mathbf{S}_{t-1}\bm{k}_t v t old = S t − 1 k t 减掉,再写入新值 v t new \bm{v}^{\text{new}}_t v t new ,而新值是旧值与目标值的凸组合,β t \beta_t β t 控制替换的彻底程度。β t → 1 \beta_t \to 1 β t → 1 是完全覆盖,β t → 0 \beta_t \to 0 β t → 0 是不动。
关键在于这个减法只作用在 k t \bm{k}_t k t 方向上 ,与 k t \bm{k}_t k t 正交的记忆分毫不动。这就是「定向」–相比标量衰减的一刀切,delta rule 只擦掉要覆盖的那一条。
1.2 为什么是 Householder,以及 L2 归一化的作用
I − β t k t k t ⊺ \mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} I − β t k t k t ⊺ 是广义 Householder 变换。∥ k t ∥ = 1 \|\bm{k}_t\| = 1 ∥ k t ∥ = 1 时,它的特征值只有两种取值,结构一目了然:
沿 k t \bm{k}_t k t 方向:特征值 1 − β t 1 - \beta_t 1 − β t
与 k t \bm{k}_t k t 正交的 d k − 1 d_k - 1 d k − 1 个方向:特征值 1 1 1
于是 β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 ) 时全部特征值的绝对值都不超过 1,状态每步只会被压缩、不会被放大,递推因此稳定。这正是论文对 k \bm{k} k 做 L2 归一化的深层原因–若 ∥ k t ∥ ≠ 1 \|\bm{k}_t\| \ne 1 ∥ k t ∥ = 1 ,沿 k t \bm{k}_t k t 的特征值变成 1 − β t ∥ k t ∥ 2 1 - \beta_t\|\bm{k}_t\|^2 1 − β t ∥ k t ∥ 2 ,β t ∥ k t ∥ 2 > 2 \beta_t\|\bm{k}_t\|^2 > 2 β t ∥ k t ∥ 2 > 2 时就会翻到 − 1 -1 − 1 以下,递推放大。
顺带一提,论文脚注提到可以放开到 β t ∈ ( 0 , 2 ) \beta_t \in (0,2) β t ∈ ( 0 , 2 ) 以允许负特征值,那是为了解锁状态跟踪能力(state tracking)。本文按 ( 0 , 1 ) (0,1) ( 0 , 1 ) 处理。
1.3 test-time SGD 视角
论文给了一个很有启发的解释:把状态 S \mathbf{S} S 看成一个快速权重矩阵,delta rule 就是在做在线回归的一步梯度下降。目标是 L ( S t ) = 1 2 ∥ S t k t − v t ∥ 2 \mathcal{L}(\mathbf{S}_t) = \frac{1}{2}\|\mathbf{S}_t\bm{k}_t - \bm{v}_t\|^2 L ( S t ) = 2 1 ∥ S t k t − v t ∥ 2 ,那么:
S t − β t ∇ L ( S t ) = S t − β t ( S t k t − v t ) k t ⊺ = S t ( I − β t k t k t ⊺ ) + β t v t k t ⊺ \mathbf{S}_t - \beta_t\nabla\mathcal{L}(\mathbf{S}_t)
= \mathbf{S}_t - \beta_t(\mathbf{S}_t\bm{k}_t - \bm{v}_t)\bm{k}_t^{\intercal}
= \mathbf{S}_t(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}
S t − β t ∇ L ( S t ) = S t − β t ( S t k t − v t ) k t ⊺ = S t ( I − β t k t k t ⊺ ) + β t v t k t ⊺
β t \beta_t β t 就是学习率,α t \alpha_t α t 就是 weight decay。 这个视角下 gated delta rule 没有任何神秘之处–它是带权重衰减的 test-time SGD。
2. Chunkwise 形式:为什么需要解三角系统
2.1 展开递推:转移矩阵不再是标量
按块展开 r r r 步(论文式 10):
S [ t ] r = S [ t ] ∏ i = 1 r α [ t ] i ( I − β [ t ] i k [ t ] i k [ t ] i ⊺ ) ⏟ F [ t ] r + ∑ i = 1 r β [ t ] i v [ t ] i k [ t ] i ⊺ ∏ j = i + 1 r α [ t ] j ( I − β [ t ] j k [ t ] j k [ t ] j ⊺ ) ⏟ G [ t ] r \mathbf{S}_{[t]}^{r} = \mathbf{S}_{[t]}\underbrace{\prod_{i=1}^{r}\alpha^i_{[t]}\big(\mathbf{I} - \beta^i_{[t]}\bm{k}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\big)}_{\mathbf{F}^r_{[t]}} + \underbrace{\sum_{i=1}^{r}\beta^i_{[t]}\bm{v}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\prod_{j=i+1}^{r}\alpha^j_{[t]}\big(\mathbf{I} - \beta^j_{[t]}\bm{k}^j_{[t]}\bm{k}^{j\intercal}_{[t]}\big)}_{\mathbf{G}^r_{[t]}}
S [ t ] r = S [ t ] F [ t ] r i = 1 ∏ r α [ t ] i ( I − β [ t ] i k [ t ] i k [ t ] i ⊺ ) + G [ t ] r i = 1 ∑ r β [ t ] i v [ t ] i k [ t ] i ⊺ j = i + 1 ∏ r α [ t ] j ( I − β [ t ] j k [ t ] j k [ t ] j ⊺ )
对比上一篇:那里的转移量是标量 α r − i \alpha^{r-i} α r − i ,可以直接查表。这里是矩阵连乘 ,F \mathbf{F} F 与 G \mathbf{G} G 都不能靠逐元素权重表达。
衰减部分可以先提出来–∏ α i = γ r \prod\alpha_i = \gamma^r ∏ α i = γ r 是标量,与 Householder 部分可交换:
F [ t ] r = γ [ t ] r P [ t ] r , P [ t ] r = ∏ i = 1 r ( I − β i k i k i ⊺ ) \mathbf{F}^r_{[t]} = \gamma^r_{[t]}\,\mathbf{P}^r_{[t]},
\qquad \mathbf{P}^r_{[t]} = \prod_{i=1}^{r}\big(\mathbf{I} - \beta^i\bm{k}^i\bm{k}^{i\intercal}\big)
F [ t ] r = γ [ t ] r P [ t ] r , P [ t ] r = i = 1 ∏ r ( I − β i k i k i ⊺ )
剩下的 P \mathbf{P} P 是纯 DeltaNet 的部分,靠 WY 表示处理–这正是论文 §2.2 的内容,下一节完整展开。
2.2 论文 §2.2:无门控 DeltaNet 的 WY 表示
在处理带门控的版本之前,先把论文 §2.2 那套无门控 DeltaNet 的 chunkwise 推导完整走一遍。GDN 的做法本质上是在这套框架上打补丁,先看清基线,后面的改动才有参照。
本节 α t ≡ 1 \alpha_t \equiv 1 α t ≡ 1 (无遗忘门),递推退化为 S t = S t − 1 ( I − β t k t k t ⊺ ) + β t v t k t ⊺ \mathbf{S}_t = \mathbf{S}_{t-1}(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}) + \beta_t\bm{v}_t\bm{k}_t^{\intercal} S t = S t − 1 ( I − β t k t k t ⊺ ) + β t v t k t ⊺ 。
2.2.1 部分展开:两个连乘(论文式 3)
按块部分展开递推:
S [ t ] r = S [ t ] ( ∏ i = 1 r ( I − β [ t ] i k [ t ] i k [ t ] i ⊺ ) ) ⏟ : = P [ t ] r + ∑ i = 1 r β [ t ] i v [ t ] i k [ t ] i ⊺ ∏ j = i + 1 r ( I − β [ t ] j k [ t ] j k [ t ] j ⊺ ) ⏟ : = H [ t ] r \mathbf{S}^r_{[t]} = \mathbf{S}_{[t]}\underbrace{\Big(\prod_{i=1}^{r}\big(\mathbf{I} - \beta^i_{[t]}\bm{k}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\big)\Big)}_{\textstyle :=\mathbf{P}^r_{[t]}} + \underbrace{\sum_{i=1}^{r}\beta^i_{[t]}\bm{v}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\prod_{j=i+1}^{r}\big(\mathbf{I} - \beta^j_{[t]}\bm{k}^j_{[t]}\bm{k}^{j\intercal}_{[t]}\big)}_{\textstyle :=\mathbf{H}^r_{[t]}}
S [ t ] r = S [ t ] := P [ t ] r ( i = 1 ∏ r ( I − β [ t ] i k [ t ] i k [ t ] i ⊺ ) ) + := H [ t ] r i = 1 ∑ r β [ t ] i v [ t ] i k [ t ] i ⊺ j = i + 1 ∏ r ( I − β [ t ] j k [ t ] j k [ t ] j ⊺ )
两个部分的角色不同,值得分清:
形状
含义
结构
P [ t ] r \mathbf{P}^r_{[t]} P [ t ] r
d k × d k d_k \times d_k d k × d k
历史状态的遗忘算子 :入口状态经过 r r r 步 Householder 后剩下什么
Householder 的纯连乘
H [ t ] r \mathbf{H}^r_{[t]} H [ t ] r
d v × d k d_v \times d_k d v × d k
块内新写入的累积 :前 r r r 个 token 写进来的内容(互相已扣除重叠)
连乘的加权和
P \mathbf{P} P 是纯连乘,H \mathbf{H} H 的每一项后面还挂着一截连乘尾巴–两者都不能直接算,C = 64 C = 64 C = 64 时 P \mathbf{P} P 要 64 个 d k × d k d_k\times d_k d k × d k 矩阵相乘。WY 表示的作用就是把这两个连乘各自压成一次求和。
2.2.2 经典 WY:把 Householder 连乘压成秩-C C C 更新(论文式 4)
这是 Bischof & Van Loan (1985) 的经典结果。核心事实:C C C 个 Householder 矩阵的乘积可以写成单位矩阵减去一个秩至多 C C C 的修正 :
P [ t ] r = I − ∑ i = 1 r w [ t ] i k [ t ] i ⊺ ∈ R d k × d k , w [ t ] r = β [ t ] r ( k [ t ] r − ∑ i = 1 r − 1 w [ t ] i ( k [ t ] i ⊺ k [ t ] r ) ) ∈ R d k \mathbf{P}^r_{[t]} = \mathbf{I} - \sum_{i=1}^{r}\bm{w}^i_{[t]}\bm{k}^{i\intercal}_{[t]} \in \mathbb{R}^{d_k\times d_k},
\qquad
\bm{w}^r_{[t]} = \beta^r_{[t]}\Big(\bm{k}^r_{[t]} - \sum_{i=1}^{r-1}\bm{w}^i_{[t]}\big(\bm{k}^{i\intercal}_{[t]}\bm{k}^r_{[t]}\big)\Big) \in \mathbb{R}^{d_k}
P [ t ] r = I − i = 1 ∑ r w [ t ] i k [ t ] i ⊺ ∈ R d k × d k , w [ t ] r = β [ t ] r ( k [ t ] r − i = 1 ∑ r − 1 w [ t ] i ( k [ t ] i ⊺ k [ t ] r ) ) ∈ R d k
为什么成立,看一步归纳就够:
P r = P r − 1 ( I − β r k r k r ⊺ ) = P r − 1 − β r P r − 1 k r ⏟ w r k r ⊺ \mathbf{P}^{r} = \mathbf{P}^{r-1}\big(\mathbf{I} - \beta_r\bm{k}_r\bm{k}_r^{\intercal}\big)
= \mathbf{P}^{r-1} - \underbrace{\beta_r\mathbf{P}^{r-1}\bm{k}_r}_{\textstyle \bm{w}_r}\bm{k}_r^{\intercal}
P r = P r − 1 ( I − β r k r k r ⊺ ) = P r − 1 − w r β r P r − 1 k r k r ⊺
于是 w r = β r P r − 1 k r \bm{w}_r = \beta_r\mathbf{P}^{r-1}\bm{k}_r w r = β r P r − 1 k r ,把 P r − 1 = I − ∑ i < r w i k i ⊺ \mathbf{P}^{r-1} = \mathbf{I} - \sum_{i<r}\bm{w}_i\bm{k}_i^{\intercal} P r − 1 = I − ∑ i < r w i k i ⊺ 代入就得到上面那个递推。每多一个 Householder,秩只增加 1 –这就是"连乘变求和"的全部内容。
2.2.3 同一个模具:H \mathbf{H} H 的 WY 表示(论文式 5)
H \mathbf{H} H 的推导结构与 P \mathbf{P} P 完全平行 :
H [ t ] r = ∑ i = 1 r u [ t ] i k [ t ] i ⊺ ∈ R d v × d k , u [ t ] r = β [ t ] r ( v [ t ] r − ∑ i = 1 r − 1 u [ t ] i ( k [ t ] i ⊺ k [ t ] r ) ) ∈ R d v \mathbf{H}^r_{[t]} = \sum_{i=1}^{r}\bm{u}^i_{[t]}\bm{k}^{i\intercal}_{[t]} \in \mathbb{R}^{d_v\times d_k},
\qquad
\bm{u}^r_{[t]} = \beta^r_{[t]}\Big(\bm{v}^r_{[t]} - \sum_{i=1}^{r-1}\bm{u}^i_{[t]}\big(\bm{k}^{i\intercal}_{[t]}\bm{k}^r_{[t]}\big)\Big) \in \mathbb{R}^{d_v}
H [ t ] r = i = 1 ∑ r u [ t ] i k [ t ] i ⊺ ∈ R d v × d k , u [ t ] r = β [ t ] r ( v [ t ] r − i = 1 ∑ r − 1 u [ t ] i ( k [ t ] i ⊺ k [ t ] r ) ) ∈ R d v
把 w \bm{w} w 的递推和 u \bm{u} u 的递推并排看,会发现它们是同一个式子 :
w r = β r ( k r ‾ − ∑ i < r w i ( k i ⊺ k r ) ) , u r = β r ( v r ‾ − ∑ i < r u i ( k i ⊺ k r ) ) \bm{w}_r = \beta_r\Big(\underline{\bm{k}_r} - \sum_{i<r}\bm{w}_i(\bm{k}_i^{\intercal}\bm{k}_r)\Big),
\qquad
\bm{u}_r = \beta_r\Big(\underline{\bm{v}_r} - \sum_{i<r}\bm{u}_i(\bm{k}_i^{\intercal}\bm{k}_r)\Big)
w r = β r ( k r − i < r ∑ w i ( k i ⊺ k r ) ) , u r = β r ( v r − i < r ∑ u i ( k i ⊺ k r ) )
只有下划线处不同:一个是 k r \bm{k}_r k r 、一个是 v r \bm{v}_r v r 。系数完全一样 –这个观察是后面式 (6)(7) 能共用一个 T \mathbf{T} T 的全部原因,也是 kernel 里 W \mathbf{W} W 与 U \mathbf{U} U 能拼进一次前向替换的依据(第四篇 KDA 用的就是这个技巧)。
写成矩阵形式,两个连乘都消失了:
P [ t ] = I − W [ t ] ⊺ K [ t ] ∈ R d k × d k , H [ t ] = U [ t ] ⊺ K [ t ] ∈ R d v × d k \mathbf{P}_{[t]} = \mathbf{I} - \mathbf{W}_{[t]}^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_k\times d_k},
\qquad
\mathbf{H}_{[t]} = \mathbf{U}_{[t]}^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_v\times d_k}
P [ t ] = I − W [ t ] ⊺ K [ t ] ∈ R d k × d k , H [ t ] = U [ t ] ⊺ K [ t ] ∈ R d v × d k
2.2.4 UT 变换:递推也不必串行(论文式 6、7)
式 (4)(5) 虽然把连乘变成了求和,但 w r \bm{w}_r w r / u r \bm{u}_r u r 自身还是串行递推 。Joffrain et al. (2006) 的 UT 变换把它变成一次矩阵求逆:
T [ t ] = [ I + strictLower ( diag ( β [ t ] ) K [ t ] K [ t ] ⊺ ) ] − 1 diag ( β [ t ] ) ∈ R C × C \mathbf{T}_{[t]} = \Big[\mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta_{[t]})\mathbf{K}_{[t]}\mathbf{K}_{[t]}^{\intercal}\big)\Big]^{-1}\operatorname{diag}(\beta_{[t]}) \in \mathbb{R}^{C\times C}
T [ t ] = [ I + strictLower ( diag ( β [ t ] ) K [ t ] K [ t ] ⊺ ) ] − 1 diag ( β [ t ] ) ∈ R C × C
W [ t ] = T [ t ] K [ t ] ∈ R C × d k , U [ t ] = T [ t ] V [ t ] ∈ R C × d v \mathbf{W}_{[t]} = \mathbf{T}_{[t]}\mathbf{K}_{[t]} \in \mathbb{R}^{C\times d_k},
\qquad
\mathbf{U}_{[t]} = \mathbf{T}_{[t]}\mathbf{V}_{[t]} \in \mathbb{R}^{C\times d_v}
W [ t ] = T [ t ] K [ t ] ∈ R C × d k , U [ t ] = T [ t ] V [ t ] ∈ R C × d v
W \mathbf{W} W 与 U \mathbf{U} U 共用同一个 T \mathbf{T} T ,正是因为 §2.2.3 那两个递推的系数相同。这一步的意义在于:T \mathbf{T} T 只跟 K \mathbf{K} K 和 β \beta β 有关,与 V \mathbf{V} V 无关。
2.2.5 代回式 3:可用 Tensor Core 的算法(论文式 8、9)
S [ t + 1 ] = S [ t ] P [ t ] + H [ t ] = S [ t ] + ( U [ t ] − W [ t ] S [ t ] ⊺ ) ⊺ K [ t ] ∈ R d v × d k O [ t ] = Q [ t ] S [ t ] ⊺ + ( Q [ t ] K [ t ] ⊺ ⊙ M ) ( U [ t ] − W [ t ] S [ t ] ⊺ ) ∈ R C × d v \begin{aligned}
\mathbf{S}_{[t+1]} &= \mathbf{S}_{[t]}\mathbf{P}_{[t]} + \mathbf{H}_{[t]}
= \mathbf{S}_{[t]} + \big(\mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]}^{\intercal}\big)^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_v\times d_k} \\[4pt]
\mathbf{O}_{[t]} &= \mathbf{Q}_{[t]}\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal}\odot\mathbf{M}\big)\big(\mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]}^{\intercal}\big) \in \mathbb{R}^{C\times d_v}
\end{aligned}
S [ t + 1 ] O [ t ] = S [ t ] P [ t ] + H [ t ] = S [ t ] + ( U [ t ] − W [ t ] S [ t ] ⊺ ) ⊺ K [ t ] ∈ R d v × d k = Q [ t ] S [ t ] ⊺ + ( Q [ t ] K [ t ] ⊺ ⊙ M ) ( U [ t ] − W [ t ] S [ t ] ⊺ ) ∈ R C × d v
式 (8) 的化简值得看一眼–S P = S ( I − W ⊺ K ) = S − ( W S ⊺ ) ⊺ K \mathbf{S}\mathbf{P} = \mathbf{S}(\mathbf{I}-\mathbf{W}^\intercal\mathbf{K}) = \mathbf{S} - (\mathbf{W}\mathbf{S}^\intercal)^\intercal\mathbf{K} SP = S ( I − W ⊺ K ) = S − ( W S ⊺ ) ⊺ K ,与 H = U ⊺ K \mathbf{H} = \mathbf{U}^\intercal\mathbf{K} H = U ⊺ K 合并后 K \mathbf{K} K 被提到右侧,括号里剩下 U − W S ⊺ \mathbf{U} - \mathbf{W}\mathbf{S}^\intercal U − W S ⊺ 。这个量在式 (8) 和式 (9) 里是同一个 ,算一次用两回。
M \mathbf{M} M 是下三角全 1 的因果掩码。注意此处没有任何衰减 –这正是 GDN 要改的地方。
2.2.6 数值验证:九个等式逐一核对
d k = 6 , d v = 5 , C = 7 d_k=6, d_v=5, C=7 d k = 6 , d v = 5 , C = 7 ,β ∼ U ( 0.1 , 0.9 ) \beta \sim \mathcal{U}(0.1,0.9) β ∼ U ( 0.1 , 0.9 ) ,k \bm{k} k 已 L2 归一化,入口状态 S 0 ≠ 0 \mathbf{S}_0 \ne \mathbf{0} S 0 = 0 (比论文附录 A 的首块假设更严格),fp64:
论文式
内容
max abs 误差
(3)
S r = S 0 P r + H r \mathbf{S}^r = \mathbf{S}_0\mathbf{P}^r + \mathbf{H}^r S r = S 0 P r + H r
4.44 × 10 − 16 4.44\times10^{-16} 4.44 × 1 0 − 16
(4)
P r = I − ∑ i w i k i ⊺ \mathbf{P}^r = \mathbf{I} - \sum_i\bm{w}_i\bm{k}_i^\intercal P r = I − ∑ i w i k i ⊺
4.44 × 10 − 16 4.44\times10^{-16} 4.44 × 1 0 − 16
(5)
H r = ∑ i u i k i ⊺ \mathbf{H}^r = \sum_i\bm{u}_i\bm{k}_i^\intercal H r = ∑ i u i k i ⊺
3.33 × 10 − 16 3.33\times10^{-16} 3.33 × 1 0 − 16
矩阵形式
P = I − W ⊺ K \mathbf{P} = \mathbf{I} - \mathbf{W}^\intercal\mathbf{K} P = I − W ⊺ K
2.22 × 10 − 16 2.22\times10^{-16} 2.22 × 1 0 − 16
矩阵形式
H = U ⊺ K \mathbf{H} = \mathbf{U}^\intercal\mathbf{K} H = U ⊺ K
2.78 × 10 − 16 2.78\times10^{-16} 2.78 × 1 0 − 16
(6)(7)
W = T K \mathbf{W} = \mathbf{T}\mathbf{K} W = TK (UT 变换 vs 式 4 递推)
5.55 × 10 − 17 5.55\times10^{-17} 5.55 × 1 0 − 17
(6)(7)
U = T V \mathbf{U} = \mathbf{T}\mathbf{V} U = TV
2.22 × 10 − 16 2.22\times10^{-16} 2.22 × 1 0 − 16
(8)
S [ t + 1 ] \mathbf{S}_{[t+1]} S [ t + 1 ] 完整式
4.44 × 10 − 16 4.44\times10^{-16} 4.44 × 1 0 − 16
(9)
O [ t ] \mathbf{O}_{[t]} O [ t ] 完整式
8.88 × 10 − 16 8.88\times10^{-16} 8.88 × 1 0 − 16
另外验证了 P = I − W ⊺ K \mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K} P = I − W ⊺ K 的特征值绝对值分别为 0.9918 , 0.9082 , 0.5805 , 0.3131 , 0.1951 , 0.0676 0.9918, 0.9082, 0.5805, 0.3131, 0.1951, 0.0676 0.9918 , 0.9082 , 0.5805 , 0.3131 , 0.1951 , 0.0676 –全部 ≤ 1 \le 1 ≤ 1 ,与 §1.2 说的"每步只压缩不放大"一致。
2.2.7 GDN 改了哪两处
把门控加回来(α t ≠ 1 \alpha_t \ne 1 α t = 1 ),论文 §3.3 的做法只动了式 (6)(7) 两个地方 :
无门控(式 6、7)
GDN(门控版)
T \mathbf{T} T 里的 KK 矩阵
K K ⊺ \mathbf{K}\mathbf{K}^{\intercal} K K ⊺
Γ ⊙ K K ⊺ \Gamma \odot \mathbf{K}\mathbf{K}^{\intercal} Γ ⊙ K K ⊺
W \mathbf{W} W 的输入
K \mathbf{K} K
diag ( γ ) K = K ← \operatorname{diag}(\gamma)\mathbf{K} = \overleftarrow{\mathbf{K}} diag ( γ ) K = K
U \mathbf{U} U 的输入
V \mathbf{V} V
V \mathbf{V} V (不变)
式 (8) 的旧状态项
S [ t ] \mathbf{S}_{[t]} S [ t ]
γ C S [ t ] = S → \gamma^C\mathbf{S}_{[t]} = \overrightarrow{\mathbf{S}} γ C S [ t ] = S
式 (8) 的 K \mathbf{K} K
K \mathbf{K} K
γ C γ r K = K → \frac{\gamma^C}{\gamma^r}\mathbf{K} = \overrightarrow{\mathbf{K}} γ r γ C K = K
式 (9) 的 Q \mathbf{Q} Q 与掩码
Q \mathbf{Q} Q ,M \mathbf{M} M
γ r Q = Q ← \gamma^r\mathbf{Q} = \overleftarrow{\mathbf{Q}} γ r Q = Q ,Γ \Gamma Γ
实测这个替换的正确性:门控版 S [ t + 1 ] \mathbf{S}_{[t+1]} S [ t + 1 ] 与逐 token 递归差 3.33 × 10 − 16 3.33\times10^{-16} 3.33 × 1 0 − 16 ;令 α ≡ 1 \alpha\equiv1 α ≡ 1 时,Γ \Gamma Γ 与下三角全 1 掩码 M \mathbf{M} M 完全相等(差 0.0) 、T \mathbf{T} T 与无门控版完全相等(差 0.0) 、W \mathbf{W} W 也完全相等。三个 0.0 说明门控版是无门控版的严格推广,没有引入任何额外近似。
换个角度看 :Γ \Gamma Γ 就是把式 (9) 那个 0/1 因果掩码 M \mathbf{M} M 升级成了"带衰减的因果掩码"。M i j ∈ { 0 , 1 } \mathbf{M}_{ij} \in \{0,1\} M ij ∈ { 0 , 1 } 只回答"j j j 能否影响 i i i ",Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j 还回答"影响衰减了多少"。前两篇反复出现的那个 Γ \Gamma Γ ,在这个视角下就是 M \mathbf{M} M 的自然推广。
2.3 手推一遍:Householder 连乘怎么变成三角系统
直接推 C = 4 C = 4 C = 4 的情形最清楚。定义每步的修正量 d r \bm{d}_r d r ,使得递推写成加法形式:
S r = α r S r − 1 + d r k r ⊺ , d r = β r v r − α r β r S r − 1 k r \mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^{\intercal},
\qquad
\bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r
S r = α r S r − 1 + d r k r ⊺ , d r = β r v r − α r β r S r − 1 k r
这一步只是把 S r − 1 ( α r ( I − β r k k ⊺ ) ) + β r v k ⊺ \mathbf{S}^{r-1}(\alpha_r(\mathbf{I}-\beta_r\bm{k}\bm{k}^\intercal)) + \beta_r\bm{v}\bm{k}^\intercal S r − 1 ( α r ( I − β r k k ⊺ )) + β r v k ⊺ 重新分组,恒等变形。注意 d r \bm{d}_r d r 里带 α r \alpha_r α r –删除项作用在已经衰减过的状态上,这个细节写错不会报错、只会算错。
有了加法形式,就能像上一篇那样倒代换(每项系数是「距块尾的步数」):
S [ t + 1 ] = γ C S [ t ] + ∑ r = 1 C γ C γ r d r k r ⊺ \mathbf{S}_{[t+1]} = \gamma^C\mathbf{S}_{[t]} + \sum_{r=1}^{C}\frac{\gamma^C}{\gamma^r}\,\bm{d}_r\bm{k}_r^{\intercal}
S [ t + 1 ] = γ C S [ t ] + r = 1 ∑ C γ r γ C d r k r ⊺
问题在于 d r \bm{d}_r d r 依赖 S r − 1 \mathbf{S}^{r-1} S r − 1 ,而 S r − 1 \mathbf{S}^{r-1} S r − 1 又依赖 d 1 , … , d r − 1 \bm{d}_1, \ldots, \bm{d}_{r-1} d 1 , … , d r − 1 –串行依赖,无法并行。把 S r − 1 \mathbf{S}^{r-1} S r − 1 也倒代换开(利用 α r γ r − 1 = γ r \alpha_r\gamma^{r-1} = \gamma^r α r γ r − 1 = γ r ):
α r S r − 1 = γ r S [ t ] + ∑ i < r γ r γ i d i k i ⊺ \alpha_r\mathbf{S}^{r-1} = \gamma^r\mathbf{S}_{[t]} + \sum_{i<r}\frac{\gamma^r}{\gamma^i}\bm{d}_i\bm{k}_i^{\intercal}
α r S r − 1 = γ r S [ t ] + i < r ∑ γ i γ r d i k i ⊺
代回 d r \bm{d}_r d r 的定义:
d r = β r ( v r − γ r S [ t ] k r − ∑ i < r γ r γ i ( k i ⊺ k r ) d i ) \bm{d}_r = \beta_r\Big(\bm{v}_r - \gamma^r\mathbf{S}_{[t]}\bm{k}_r - \sum_{i<r}\frac{\gamma^r}{\gamma^i}\big(\bm{k}_i^{\intercal}\bm{k}_r\big)\bm{d}_i\Big)
d r = β r ( v r − γ r S [ t ] k r − i < r ∑ γ i γ r ( k i ⊺ k r ) d i )
这是一个下三角线性系统 :d r \bm{d}_r d r 只依赖 d i ( i < r ) \bm{d}_i\ (i < r) d i ( i < r ) ,系数是 β r γ r γ i ( k i ⊺ k r ) \beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r) β r γ i γ r ( k i ⊺ k r ) 。写成矩阵形式,令 A r i = β r γ r γ i ( k i ⊺ k r ) \mathbf{A}_{ri} = \beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r) A r i = β r γ i γ r ( k i ⊺ k r ) 的严格下三角部分:
( I + strictLower ( A ) ) Δ = diag ( β ) ( V − diag ( γ ) K S [ t ] ⊺ ) (\mathbf{I} + \operatorname{strictLower}(\mathbf{A}))\,\Delta = \operatorname{diag}(\beta)\big(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}_{[t]}^{\intercal}\big)
( I + strictLower ( A )) Δ = diag ( β ) ( V − diag ( γ ) K S [ t ] ⊺ )
其中 Δ ∈ R C × d v \Delta \in \mathbb{R}^{C \times d_v} Δ ∈ R C × d v 的第 r r r 行是 d r \bm{d}_r d r 。注意 A = diag ( β ) ( Γ ⊙ K K ⊺ ) \mathbf{A} = \operatorname{diag}(\beta)\big(\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}\big) A = diag ( β ) ( Γ ⊙ K K ⊺ ) –衰减感知掩码 Γ \Gamma Γ 在这里第二次出现 ,这次是嵌在三角系统的系数矩阵里。
2.4 论文附录 A:扩展 WY 表示的归纳证明
上面 §2.3 是"把递推硬拆开"的推法。论文附录 A 给了一个更漂亮的等价路线:先猜出闭式,再用数学归纳法证明 。这一节按论文原文复述(GDN 论文附录 A,为减少符号负担,论文同样只考虑首块,即 S 0 = 0 \mathbf{S}_0 = \mathbf{0} S 0 = 0 )。
命题(扩展 WY 表示). 对 S t \mathbf{S}_t S t 有
S t = ∑ i = 1 t γ t γ i u i k i ⊺ , u t = β t ( v t − ∑ i = 1 t − 1 γ t γ i u i k i ⊺ k t ) \mathbf{S}_t = \sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal},
\qquad
\bm{u}_t = \beta_t\Big(\bm{v}_t - \sum_{i=1}^{t-1}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_t\Big)
S t = i = 1 ∑ t γ i γ t u i k i ⊺ , u t = β t ( v t − i = 1 ∑ t − 1 γ i γ t u i k i ⊺ k t )
证明. 对 t t t 作归纳。
S t + 1 = S t ( α t + 1 ( I − β t + 1 k t + 1 k t + 1 ⊺ ) ) + β t + 1 v t + 1 k t + 1 ⊺ = α t + 1 ( ∑ i = 1 t γ t γ i u i k i ⊺ ) − α t + 1 β t + 1 ( ∑ i = 1 t γ t γ i u i k i ⊺ k t + 1 k t + 1 ⊺ ) + β t + 1 v t + 1 k t + 1 ⊺ = ∑ i = 1 t γ t + 1 γ i u i k i ⊺ + β t + 1 ( v t + 1 − ∑ i = 1 t γ t + 1 γ i u i k i ⊺ k t + 1 ) ⏟ u t + 1 k t + 1 ⊺ = ∑ i = 1 t γ t + 1 γ i u i k i ⊺ + γ t + 1 γ t + 1 ⏟ 1 u t + 1 k t + 1 ⊺ = ∑ i = 1 t + 1 γ t + 1 γ i u i k i ⊺ \begin{aligned}
\mathbf{S}_{t+1}
&= \mathbf{S}_t\Big(\alpha_{t+1}\big(\mathbf{I} - \beta_{t+1}\bm{k}_{t+1}\bm{k}_{t+1}^{\intercal}\big)\Big) + \beta_{t+1}\bm{v}_{t+1}\bm{k}_{t+1}^{\intercal} \\[4pt]
&= \alpha_{t+1}\Big(\sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\Big)
- \alpha_{t+1}\beta_{t+1}\Big(\sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_{t+1}\bm{k}_{t+1}^{\intercal}\Big)
+ \beta_{t+1}\bm{v}_{t+1}\bm{k}_{t+1}^{\intercal} \\[4pt]
&= \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}
+ \underbrace{\beta_{t+1}\Big(\bm{v}_{t+1} - \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_{t+1}\Big)}_{\textstyle \bm{u}_{t+1}}\bm{k}_{t+1}^{\intercal} \\[4pt]
&= \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}
+ \underbrace{\frac{\gamma_{t+1}}{\gamma_{t+1}}}_{\textstyle 1}\bm{u}_{t+1}\bm{k}_{t+1}^{\intercal}
\;=\; \sum_{i=1}^{t+1}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}
\end{aligned}
S t + 1 = S t ( α t + 1 ( I − β t + 1 k t + 1 k t + 1 ⊺ ) ) + β t + 1 v t + 1 k t + 1 ⊺ = α t + 1 ( i = 1 ∑ t γ i γ t u i k i ⊺ ) − α t + 1 β t + 1 ( i = 1 ∑ t γ i γ t u i k i ⊺ k t + 1 k t + 1 ⊺ ) + β t + 1 v t + 1 k t + 1 ⊺ = i = 1 ∑ t γ i γ t + 1 u i k i ⊺ + u t + 1 β t + 1 ( v t + 1 − i = 1 ∑ t γ i γ t + 1 u i k i ⊺ k t + 1 ) k t + 1 ⊺ = i = 1 ∑ t γ i γ t + 1 u i k i ⊺ + 1 γ t + 1 γ t + 1 u t + 1 k t + 1 ⊺ = i = 1 ∑ t + 1 γ i γ t + 1 u i k i ⊺
□ \square □
第二个等号到第三个等号是全部关键 ,值得拆开看:
α t + 1 ⋅ γ t γ i = γ t + 1 γ i \alpha_{t+1}\cdot\frac{\gamma_t}{\gamma_i} = \frac{\gamma_{t+1}}{\gamma_i} α t + 1 ⋅ γ i γ t = γ i γ t + 1 –衰减系数被吸收进累积积的比值,这是 γ \gamma γ 定义为连乘 才有的性质,也是四篇一直在用的那条恒等式。
第二项里 α t + 1 \alpha_{t+1} α t + 1 同样被吸收进 γ t + 1 γ i \frac{\gamma_{t+1}}{\gamma_i} γ i γ t + 1 ,然后整项与第三项合并、共同提出右侧的 k t + 1 ⊺ \bm{k}_{t+1}^{\intercal} k t + 1 ⊺ –括号里剩下的就是 u t + 1 \bm{u}_{t+1} u t + 1 的定义。
最后一步只是注意 γ t + 1 / γ t + 1 = 1 \gamma_{t+1}/\gamma_{t+1} = 1 γ t + 1 / γ t + 1 = 1 ,于是新项能并入求和,下标从 t t t 推到 t + 1 t+1 t + 1 ,归纳闭合。
这个证明独立印证了 §2.3 那个坑。 注意第二项的系数是 α t + 1 β t + 1 \alpha_{t+1}\beta_{t+1} α t + 1 β t + 1 –α \alpha α 和 β \beta β 都在 ,因为删除项 I − β k k ⊺ \mathbf{I}-\beta\bm{k}\bm{k}^\intercal I − β k k ⊺ 作用在已经乘过 α t + 1 \alpha_{t+1} α t + 1 的状态上。我在 §2.3 用倒代换推 d r \bm{d}_r d r 时漏掉这个 α r \alpha_r α r ,refB 就与逐 token 递归差了 2.66 × 10 − 1 2.66\times10^{-1} 2.66 × 1 0 − 1 。两条路线在同一个位置要求同一个因子,可以互为校验。
实测验证这个命题(d k = 5 , d v = 4 , C = 6 d_k=5, d_v=4, C=6 d k = 5 , d v = 4 , C = 6 ,fp64,逐步比对 S t \mathbf{S}_t S t ):
检验
结果
WY 闭式 vs 逐 token 递推(t = 1..6 t=1..6 t = 1..6 逐步)
最大 1.67 × 10 − 16 1.67\times10^{-16} 1.67 × 1 0 − 16
u t \bm{u}_t u t 是否等于 §2.3 的修正量 d t \bm{d}_t d t
1.11 × 10 − 16 1.11\times10^{-16} 1.11 × 1 0 − 16
若 d t \bm{d}_t d t 漏掉 α t \alpha_t α t
偏差 2.99 × 10 − 2 2.99\times10^{-2} 2.99 × 1 0 − 2
第二行说明论文的 u t \bm{u}_t u t 与我 §2.3 手推的 d t \bm{d}_t d t 是同一个量 ,只是推导路径不同:论文归纳法从闭式出发验证,§2.3 从递推倒代换构造。两者都给出同一个下三角系统,下面的 UT 变换对二者通用。
关于跨块 :论文只推首块(S 0 = 0 \mathbf{S}_0=\mathbf{0} S 0 = 0 )。实际 kernel 里每块的入口状态非零,u t \bm{u}_t u t 的定义要补上 − β t γ t S 0 ⊺ k t -\beta_t\gamma_t\mathbf{S}_0^{\intercal}\bm{k}_t − β t γ t S 0 ⊺ k t 这一项,即 §2.3 里那个 − γ r S [ t ] k r -\gamma^r\mathbf{S}_{[t]}\bm{k}_r − γ r S [ t ] k r 。照抄附录 A 而漏掉这项,是我最初 refB 失配(max abs 2.66 × 10 − 1 2.66\times10^{-1} 2.66 × 1 0 − 1 、rel L2 6.01 × 10 − 2 6.01\times10^{-2} 6.01 × 1 0 − 2 )的另一个来源 –首块测试全对、第二块开始错,这种症状基本可以直接定位到跨块项。
2.5 UT 变换:把解系统变成矩阵乘
定义(论文 §2.2 与 §3.3):
T [ t ] = [ I + strictLower ( diag ( β [ t ] ) ( Γ [ t ] ⊙ K [ t ] K [ t ] ⊺ ) ) ] − 1 diag ( β [ t ] ) \mathbf{T}_{[t]} = \Big[\mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta_{[t]})(\Gamma_{[t]} \odot \mathbf{K}_{[t]}\mathbf{K}_{[t]}^{\intercal})\big)\Big]^{-1}\operatorname{diag}(\beta_{[t]})
T [ t ] = [ I + strictLower ( diag ( β [ t ] ) ( Γ [ t ] ⊙ K [ t ] K [ t ] ⊺ ) ) ] − 1 diag ( β [ t ] )
于是修正量一次算出:
Δ = T V ⏟ U ~ − T diag ( γ ) K ⏟ W ← S [ t ] ⊺ \Delta = \underbrace{\mathbf{T}\mathbf{V}}_{\widetilde{\mathbf{U}}} - \underbrace{\mathbf{T}\operatorname{diag}(\gamma)\mathbf{K}}_{\overleftarrow{\mathbf{W}}}\mathbf{S}_{[t]}^{\intercal}
Δ = U TV − W T diag ( γ ) K S [ t ] ⊺
这正是论文式 (11)(12) 里的 ( U ~ [ t ] − W [ t ] ← S [ t ] ⊺ ) \big(\widetilde{\mathbf{U}}_{[t]} - \overleftarrow{\mathbf{W}_{[t]}}\mathbf{S}_{[t]}^{\intercal}\big) ( U [ t ] − W [ t ] S [ t ] ⊺ ) 。完整的 chunkwise 算法:
S [ t + 1 ] = S [ t ] → + Δ ⊺ K [ t ] → O [ t ] = Q [ t ] ← S [ t ] ⊺ + ( Q [ t ] K [ t ] ⊺ ⊙ Γ [ t ] ) Δ \begin{aligned}
\mathbf{S}_{[t+1]} &= \overrightarrow{\mathbf{S}_{[t]}} + \Delta^{\intercal}\overrightarrow{\mathbf{K}_{[t]}} \\
\mathbf{O}_{[t]} &= \overleftarrow{\mathbf{Q}_{[t]}}\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big)\Delta
\end{aligned}
S [ t + 1 ] O [ t ] = S [ t ] + Δ ⊺ K [ t ] = Q [ t ] S [ t ] ⊺ + ( Q [ t ] K [ t ] ⊺ ⊙ Γ [ t ] ) Δ
三个箭头量与前两篇完全一致:q r ← = γ r q r \overleftarrow{\bm{q}^r} = \gamma^r\bm{q}^r q r = γ r q r 、k r → = γ C γ r k r \overrightarrow{\bm{k}^r} = \frac{\gamma^C}{\gamma^r}\bm{k}^r k r = γ r γ C k r 、S → = γ C S \overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S} S = γ C S 。
对比上一篇,结构上只有一处变化 :块内那一项的右乘对象从 V [ t ] \mathbf{V}_{[t]} V [ t ] 变成了 Δ \Delta Δ ,而 Δ \Delta Δ 需要解一个三角系统才能得到。β → 0 \beta \to 0 β → 0 时 T → 0 \mathbf{T} \to \mathbf{0} T → 0 、Δ → 0 \Delta \to \mathbf{0} Δ → 0 ,退化为纯衰减;α ≡ 1 \alpha \equiv 1 α ≡ 1 时 Γ → M \Gamma \to \mathbf{M} Γ → M ,退化为纯 DeltaNet。
2.6 数值验证
B = 2 , H = 2 , N = 12 , d k = d v = 4 , C = 4 B=2, H=2, N=12, d_k=d_v=4, C=4 B = 2 , H = 2 , N = 12 , d k = d v = 4 , C = 4 ,α t ∼ U ( 0.90 , 0.999 ) \alpha_t \sim \mathcal{U}(0.90, 0.999) α t ∼ U ( 0.90 , 0.999 ) 、β t ∼ U ( 0.10 , 0.90 ) \beta_t \sim \mathcal{U}(0.10, 0.90) β t ∼ U ( 0.10 , 0.90 ) ,k \bm{k} k 已 L2 归一化,fp64:
比较
max abs 误差
相对 L2
chunkwise + UT vs 逐 token 递归
8.88 × 10 − 16 8.88 \times 10^{-16} 8.88 × 1 0 − 16
2.49 × 10 − 16 2.49 \times 10^{-16} 2.49 × 1 0 − 16
α ≡ 1 \alpha \equiv 1 α ≡ 1 (退化为纯 DeltaNet)
8.88 × 10 − 16 8.88 \times 10^{-16} 8.88 × 1 0 − 16
–
β → 0 \beta \to 0 β → 0 (退化为纯衰减)
1.62 × 10 − 27 1.62 \times 10^{-27} 1.62 × 1 0 − 27
–
两个退化检验都必须做,理由与上一篇同构:α ≡ 1 \alpha \equiv 1 α ≡ 1 时 Γ \Gamma Γ 退化成 0/1 掩码,Γ \Gamma Γ 里任何指数写错都看不出来;β → 0 \beta \to 0 β → 0 时整个三角系统消失,UT 部分的错误全部隐身。
3. 三角系统的数值性质
T \mathbf{T} T 要求一个 C × C C \times C C × C 矩阵的逆,这是本篇最值得担心的地方。但实际上它比看起来温和得多。
3.1 单位下三角,条件数可控
I + strictLower ( A ) \mathbf{I} + \operatorname{strictLower}(\mathbf{A}) I + strictLower ( A ) 是单位下三角矩阵 –对角线恒为 1,严格下三角才是 A \mathbf{A} A 。这有两个直接后果:
行列式恒为 1,永不奇异。 不存在需要 pivoting 的情形。
求逆可以用前向替换 ,不需要通用矩阵求逆。
实测条件数(C = 64 C = 64 C = 64 ,k \bm{k} k L2 归一化,α t ∼ U ( 0.9 , 0.999 ) \alpha_t \sim \mathcal{U}(0.9, 0.999) α t ∼ U ( 0.9 , 0.999 ) ,20 组随机采样):
β t \beta_t β t 采样区间
cond 中位数
cond 最大
[ 0.05 , 0.2 ] [0.05,\ 0.2] [ 0.05 , 0.2 ]
1.76
1.91
[ 0.1 , 0.9 ] [0.1,\ 0.9] [ 0.1 , 0.9 ]
5.67
6.86
[ 0.5 , 0.99 ] [0.5,\ 0.99] [ 0.5 , 0.99 ]
7.88
9.05
[ 0.9 , 0.999 ] [0.9,\ 0.999] [ 0.9 , 0.999 ]
10.38
11.99
[ 1.0 , 1.99 ] [1.0,\ 1.99] [ 1.0 , 1.99 ]
34.27
40.85
β ∈ ( 0 , 1 ) \beta \in (0,1) β ∈ ( 0 , 1 ) 时条件数不超过 12,fp16 完全够用 。这与上一篇那个「因式分解会溢出」的结论形成有意思的对比:那里是代数变形引入了 1 / γ 1/\gamma 1/ γ 这种指数增长量,而这里虽然要求逆,但矩阵结构本身保证了良态。
放开到 β ∈ ( 0 , 2 ) \beta \in (0,2) β ∈ ( 0 , 2 ) (论文脚注提到的负特征值情形)条件数跳到 34,仍可接受,但已需要留意。
3.2 前向替换代替求逆
T \mathbf{T} T 从不需要显式求逆。逐行前向替换:
T [ r , : ] = e r − ∑ j < r A [ r , j ] T [ j , : ] \mathbf{T}[r,:] = \bm{e}_r - \sum_{j<r}\mathbf{A}[r,j]\,\mathbf{T}[j,:]
T [ r , : ] = e r − j < r ∑ A [ r , j ] T [ j , : ]
实测与 np.linalg.inv 的差异在 10 − 17 10^{-17} 1 0 − 17 量级(三组随机种子分别 7.81 , 5.55 , 8.33 × 10 − 17 7.81, 5.55, 8.33 \times 10^{-17} 7.81 , 5.55 , 8.33 × 1 0 − 17 ),即机器精度。
更进一步,T \mathbf{T} T 本身也不必物化。 需要的只是 Δ = T R \Delta = \mathbf{T}\mathbf{R} Δ = TR (其中 R = diag ( β ) ( V − diag ( γ ) K S ⊺ ) \mathbf{R} = \operatorname{diag}(\beta)(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}^\intercal) R = diag ( β ) ( V − diag ( γ ) K S ⊺ ) ),直接对 R \mathbf{R} R 做前向替换:
Δ [ r , : ] = R [ r , : ] − ∑ j < r A [ r , j ] Δ [ j , : ] \Delta[r,:] = \mathbf{R}[r,:] - \sum_{j<r}\mathbf{A}[r,j]\,\Delta[j,:]
Δ [ r , : ] = R [ r , : ] − j < r ∑ A [ r , j ] Δ [ j , : ]
这省掉一个 C × C C \times C C × C 的中间量。代价是引入了 C C C 步串行–这是本篇 kernel 与前两篇最本质的区别,下面 §5 会看到它如何限制并行度。
4. 四层参考实现
沿用前两篇的框架:
参考
实现方式
验证目标
A
逐 token 递归
递推式定义 S t = S t − 1 ( α t ( I − β t k k ⊺ ) ) + β t v k ⊺ \mathbf{S}_t = \mathbf{S}_{t-1}(\alpha_t(\mathbf{I}-\beta_t\bm{k}\bm{k}^\intercal)) + \beta_t\bm{v}\bm{k}^\intercal S t = S t − 1 ( α t ( I − β t k k ⊺ )) + β t v k ⊺
B
chunkwise + UT 变换
§2.3 的三角系统推导、§2.4 的 WY 表示与式 (11)(12)
C
kernel 控制流镜像(切 DV、S ⊺ \mathbf{S}^\intercal S ⊺ 布局、前向替换)
kernel 结构
D
退化检验(α ≡ 1 \alpha \equiv 1 α ≡ 1 / β → 0 \beta \to 0 β → 0 )
与前两篇的一致性
4.0 状态方向
与前两篇一致:论文的状态是 S ∈ R d v × d k \mathbf{S} \in \mathbb{R}^{d_v \times d_k} S ∈ R d v × d k ,kernel 存转置 S ⊺ ∈ R d k × d v \mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v} S ⊺ ∈ R d k × d v 。本篇多一层要注意的是 Δ ∈ R C × d v \Delta \in \mathbb{R}^{C \times d_v} Δ ∈ R C × d v –它的布局与 V \mathbf{V} V 相同,所以 Δ ⊺ K → \Delta^\intercal\overrightarrow{\mathbf{K}} Δ ⊺ K 在转置布局下写作 K → ⊺ Δ \overrightarrow{\mathbf{K}}^\intercal\Delta K ⊺ Δ 。
4.1 参考 A:逐 token 递归
1 2 3 4 5 6 7 8 9 10 11 12 13 14 def ref_A (Q, K, V, alpha, beta ): B, N, H, DK = Q.shape DV = V.shape[-1 ] O = np.zeros((B, N, H, DV)) for b in range (B): for h in range (H): S = np.zeros((DV, DK)) for t in range (N): k, v, q = K[b, t, h], V[b, t, h], Q[b, t, h] a, be = alpha[b, t, h], beta[b, t, h] S = S @ (a * (np.eye(DK) - be * np.outer(k, k))) + be * np.outer(v, k) O[b, t, h] = S @ q return O
4.2 参考 B:chunkwise + UT 变换
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 def ref_B (Q, K, V, alpha, beta, C ): B, N, H, DK = Q.shape DV = V.shape[-1 ] O = np.zeros((B, N, H, DV)) for b in range (B): for h in range (H): S = np.zeros((DV, DK)) for c in range (N // C): sl = slice (c * C, (c + 1 ) * C) Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h] a, be = alpha[b, sl, h], beta[b, sl, h] g = np.cumprod(a) gC = g[-1 ] i = np.arange(C) Gam = np.where(i[:, None ] >= i[None , :], g[i][:, None ] / g[i][None , :], 0.0 ) A = np.diag(be) @ (Gam * (Kc @ Kc.T)) T = np.linalg.inv(np.eye(C) + np.tril(A, -1 )) @ np.diag(be) Delta = T @ Vc - T @ (np.diag(g) @ (Kc @ S.T)) O[b, sl, h] = (Qc * g[:, None ]) @ S.T + ((Qc @ Kc.T) * Gam) @ Delta S = gC * S + Delta.T @ (Kc * (gC / g)[:, None ]) return O
与上一篇参考 B 的差别只有两处:多出 T 的构造,以及块内项的右乘对象从 Vc 变成 Delta。衰减权重的三处落点完全没变 。
4.3 参考 C:kernel 控制流镜像
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 def ref_C (Q, K, V, alpha, beta, C, block_DV ): B, N, H, DK = Q.shape DV = V.shape[-1 ] O = np.zeros((B, N, H, DV)) for bv in range (DV // block_DV): dv = slice (bv * block_DV, (bv + 1 ) * block_DV) for bbh in range (B * H): b, h = bbh // H, bbh % H S = np.zeros((DK, block_DV)) for c in range (N // C): sl = slice (c * C, (c + 1 ) * C) Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h, dv] a, be = alpha[b, sl, h], beta[b, sl, h] g = np.cumprod(a); gC = g[-1 ] i = np.arange(C) Gam = np.where(i[:, None ] >= i[None , :], g[i][:, None ] / g[i][None , :], 0.0 ) A = (Kc * be[:, None ]) @ Kc.T * Gam Tm = np.linalg.inv(np.eye(C) + np.tril(A, -1 )) R = be[:, None ] * Vc - be[:, None ] * ((g[:, None ] * Kc) @ S) Delta = Tm @ R O[b, sl, h, dv] = (Qc * g[:, None ]) @ S + ((Qc @ Kc.T) * Gam) @ Delta S = gC * S + (Kc * (gC / g)[:, None ]).T @ Delta return O
注意 R 的写法:diag(β) 被展开成逐行乘 be[:, None],这与 kernel 里的做法一致(不物化对角矩阵)。
4.4 一致性验证
B = 2 , H = 2 , N = 12 , d k = d v = 4 , C = 4 B=2, H=2, N=12, d_k=d_v=4, C=4 B = 2 , H = 2 , N = 12 , d k = d v = 4 , C = 4 ,block D V = 2 \text{block}_{DV}=2 block D V = 2 ,fp64:
比较
max abs 误差
相对 L2
B chunkwise + UT vs A 逐 token
8.88 × 10 − 16 8.88 \times 10^{-16} 8.88 × 1 0 − 16
2.49 × 10 − 16 2.49 \times 10^{-16} 2.49 × 1 0 − 16
C kernel 镜像 vs A 逐 token
8.88 × 10 − 16 8.88 \times 10^{-16} 8.88 × 1 0 − 16
2.61 × 10 − 16 2.61 \times 10^{-16} 2.61 × 1 0 − 16
α ≡ 1 \alpha \equiv 1 α ≡ 1 (纯 DeltaNet)B vs A
8.88 × 10 − 16 8.88 \times 10^{-16} 8.88 × 1 0 − 16
–
β → 10 − 12 \beta \to 10^{-12} β → 1 0 − 12 (纯衰减)B vs A
1.62 × 10 − 27 1.62 \times 10^{-27} 1.62 × 1 0 − 27
–
另验证了 T \mathbf{T} T 的两种等价写法(diag(β) @ (Γ * KKᵀ) 与 (K * β) @ Kᵀ * Γ)差异 5.55 × 10 − 17 5.55 \times 10^{-17} 5.55 × 1 0 − 17 ,后者省一次对角矩阵构造。
5. TileLang kernel
grid 划分沿用前两篇:切 ( b v , b ⋅ h ) (bv,\ b \cdot h) ( b v , b ⋅ h ) ,序列轴走 kernel 内的 T.Pipelined。但本篇多了一个 C × C C \times C C × C 的 A 矩阵和 C C C 步串行的前向替换,寄存器与并行度都要重新算。
5.0 寄存器账
d k = d v = 128 d_k = d_v = 128 d k = d v = 128 、C = 64 C = 64 C = 64 、128 线程:
fragment
上一篇
本篇
说明
S_f [ d k , bDV ] [d_k, \text{bDV}] [ d k , bDV ]
32 reg
32 reg
不变
acc_o [ C , bDV ] [C, \text{bDV}] [ C , bDV ]
16 reg
16 reg
不变
Dtri / Gam [ C , C ] [C,C] [ C , C ]
32 reg
32 reg
衰减掩码
A [ C , C ] [C,C] [ C , C ]
32 reg
32 reg
Q K ⊺ \mathbf{Q}\mathbf{K}^\intercal Q K ⊺ 复用
Amat [ C , C ] [C,C] [ C , C ]
–
32 reg
新增:三角系统系数
Delta [ C , bDV ] [C, \text{bDV}] [ C , bDV ]
–
16 reg
新增:修正量
合计
≈ 112 reg
≈ 160 reg
上限 255
C = 64 C = 64 C = 64 时仍有余量。但 C = 128 C = 128 C = 128 时 Gam + A + Amat 三张 C 2 C^2 C 2 表就是 384 reg,直接超限 –本篇的 C C C 实际上被限制在 64。
5.1 声明与权重预计算
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 @tilelang.jit(out_idx=[5 ] ) def gdn_chunk (B, H, S, DK, DV, blk=64 , block_DV=32 , num_stages=2 , threads=128 , dtype=T.float16, accum_dtype=T.float32 ): C = blk NS = T.ceildiv(S, C) @T.prim_func def main ( Q: T.Tensor([B, S, H, DK], dtype ), K: T.Tensor([B, S, H, DK], dtype ), V: T.Tensor([B, S, H, DV], dtype ), Alpha: T.Tensor([B, S, H], accum_dtype ), Beta: T.Tensor([B, S, H], accum_dtype ), O: T.Tensor([B, S, H, DV], dtype ), ): with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh): bb, bh = bbh // H, bbh % H dv0 = bv * block_DV Q_s = T.alloc_shared([C, DK], dtype) K_s = T.alloc_shared([C, DK], dtype) V_s = T.alloc_shared([C, block_DV], dtype) S_s = T.alloc_shared([DK, block_DV], dtype) O_s = T.alloc_shared([C, block_DV], dtype) D_s = T.alloc_shared([C, block_DV], dtype) S_f = T.alloc_fragment([DK, block_DV], accum_dtype) acc_o = T.alloc_fragment([C, block_DV], accum_dtype) A = T.alloc_fragment([C, C], accum_dtype) A_cast = T.alloc_fragment([C, C], dtype) Amat = T.alloc_fragment([C, C], accum_dtype) Delta = T.alloc_fragment([C, block_DV], accum_dtype) R = T.alloc_fragment([C, block_DV], accum_dtype) Gam = T.alloc_fragment([C, C], accum_dtype) lg = T.alloc_fragment([C], accum_dtype) g_r = T.alloc_fragment([C], accum_dtype) w_decay = T.alloc_fragment([C], accum_dtype) be_f = T.alloc_fragment([C], accum_dtype)
α t \alpha_t α t 现在是数据依赖的,累积积必须在 kernel 内算。用 log 域 cumsum 而非 cumprod –上一篇实测 fp16 直接连乘在 C = 128 C = 128 C = 128 时完全下溢:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 for i_s in T.Pipelined(NS, num_stages=num_stages): s0 = i_s * C T.copy(Q[bb, s0:s0+C, bh, :], Q_s) T.copy(K[bb, s0:s0+C, bh, :], K_s) T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s) for r in T.serial(C): lg[r] = T.log2(Alpha[bb, s0 + r, bh]) for r in T.serial(1 , C): lg[r] += lg[r - 1 ] for r in T.Parallel(C): g_r[r] = T.exp2(lg[r]) w_decay[r] = T.exp2(lg[C - 1 ] - lg[r]) be_f[r] = Beta[bb, s0 + r, bh] for i, j in T.Parallel(C, C): Gam[i, j] = T.if_then_else( j <= i, T.exp2(lg[i] - lg[j]), 0.0 )
Γ 的构造沿用上一篇的教训:log 域相减再 exp2,不能先各自取指数再相除 ,否则 1 / γ j 1/\gamma_j 1/ γ j 会溢出。条件求值也不能省成「先全算再掩」,j > i j > i j > i 处指数为正会先溢出成 inf。
5.2 三角系统:构造系数并前向替换
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 T.gemm(K_s, K_s, Amat, transpose_B=True , clear_accum=True ) for i, j in T.Parallel(C, C): Amat[i, j] = T.if_then_else( j < i, be_f[i] * Gam[i, j] * Amat[i, j], 0.0 ) T.copy(S_f, S_s) for r, d in T.Parallel(C, DK): K_s[r, d] *= g_r[r] T.gemm(K_s, S_s, R, clear_accum=True ) for r, d in T.Parallel(C, block_DV): R[r, d] = be_f[r] * (V_s[r, d] - R[r, d]) for r in T.serial(C): for d in T.Parallel(block_DV): Delta[r, d] = R[r, d] for j in T.serial(r): for d in T.Parallel(block_DV): Delta[r, d] -= Amat[r, j] * Delta[j, d]
前向替换是本篇唯一的串行段。C = 64 C = 64 C = 64 时是 64 步,每步内部 block_DV = 32 个元素并行–并行度只有 32,远低于 128 线程 。这是 gated delta rule 相比前两篇的固有代价:三角系统的依赖链无法打破。
有两个缓解方向:一是增大 block_DV(但吃寄存器),二是把 C C C 切成更小的子块做分块前向替换(blocked forward substitution),用矩阵乘处理子块间的耦合。后者是 fla 库实际采用的做法,但实现复杂度显著上升。
5.3 两项输出与状态更新
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 for r, d in T.Parallel(C, DK): Q_s[r, d] *= g_r[r] T.gemm(Q_s, S_s, acc_o, clear_accum=True ) T.copy(Q[bb, s0:s0+C, bh, :], Q_s) T.copy(K[bb, s0:s0+C, bh, :], K_s) T.gemm(Q_s, K_s, A, transpose_B=True , clear_accum=True ) for i, j in T.Parallel(C, C): A[i, j] *= Gam[i, j] T.copy(A, A_cast) T.copy(Delta, D_s) T.gemm(A_cast, D_s, acc_o) T.copy(acc_o, O_s) T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV]) for i, j in T.Parallel(DK, block_DV): S_f[i, j] *= T.exp2(lg[C - 1 ]) for r, d in T.Parallel(C, DK): K_s[r, d] *= w_decay[r] T.gemm(K_s, D_s, S_f, transpose_A=True )
5.4 五处次序约束
比上一篇多一处,全部写错都不报错、只算错:
状态更新排在输出写回之后 。S f S_f S f 在步骤②③被读时代表进块状态 S [ t ] \mathbf{S}_{[t]} S [ t ] 。
S_f *= γ^C 在 T.gemm 之前 。先衰减旧状态再累加新贡献。
K_s 被复用三次、污染两次 。构造 Amat 用原始 K \mathbf{K} K ,算 R 时被乘上 diag ( γ ) \operatorname{diag}(\gamma) diag ( γ ) ,块内项要用原始 K \mathbf{K} K (须重载),状态更新时又要乘 γ C / γ r \gamma^C/\gamma^r γ C / γ r 。这是本篇最容易错的地方 –上一篇 K_s 只污染一次。
Amat 必须只取严格下三角 (j < i j < i j < i ,不含对角)。含对角就变成了求解 ( I + A full ) (\mathbf{I} + \mathbf{A}_{\text{full}}) ( I + A full ) ,对角上的 β r ∥ k r ∥ 2 = β r \beta_r\|\bm{k}_r\|^2 = \beta_r β r ∥ k r ∥ 2 = β r 会被重复计入。
T.clear(S_f) 在流水线循环外,clear_accum=True 在循环内 。
5.5 与上一篇的改动汇总
位置
上一篇
本篇
新增开销
输入
Q, K, V
+ α t \alpha_t α t , β t \beta_t β t 两个张量
2 次 HBM 读 / chunk
累积积
编译期常量
kernel 内 log 域 cumsum
C C C 步串行前缀和
Γ \Gamma Γ
编译期可算
运行时 exp2(lg[i]-lg[j])
C 2 C^2 C 2 次 exp2
三角系统
无
Amat 构造 + 前向替换
C 2 C^2 C 2 reg + C C C 步串行
块内右乘
V \mathbf{V} V
Δ \Delta Δ
一次 C × bDV C \times \text{bDV} C × bDV f16 转换
K 复用
污染 1 次
污染 2 次,重载 1 次
一次 shared 写入
6. 三篇对照
线性注意力
标量衰减
Gated DeltaNet
递推
S + v k ⊺ \mathbf{S} + \bm{v}\bm{k}^\intercal S + v k ⊺
α S + v k ⊺ \alpha\mathbf{S} + \bm{v}\bm{k}^\intercal α S + v k ⊺
S ( α ( I − β k k ⊺ ) ) + β v k ⊺ \mathbf{S}(\alpha(\mathbf{I}-\beta\bm{k}\bm{k}^\intercal)) + \beta\bm{v}\bm{k}^\intercal S ( α ( I − β k k ⊺ )) + β v k ⊺
块内掩码
M \mathbf{M} M (0/1)
Γ \Gamma Γ (衰减感知)
Γ \Gamma Γ + 三角系统
块内右乘
V \mathbf{V} V
V \mathbf{V} V
Δ \Delta Δ
串行段
无
无
前缀和 + 前向替换(各 C C C 步)
C 2 C^2 C 2 fragment
1(A \mathbf{A} A )
2(+ Γ \Gamma Γ )
3(+ Amat)
可用 C C C
128
128(切 DV 后)
64
对应架构
–
RetNet / Lightning-Attn
Gated DeltaNet / KDA
三篇的衰减权重落点完全一致 –q ← \overleftarrow{\bm{q}} q 、k → \overrightarrow{\bm{k}} k 、S → \overrightarrow{\mathbf{S}} S 三处,从第一篇到第三篇没有变过。delta rule 加进来的是块内那一项的内容(V → Δ \mathbf{V} \to \Delta V → Δ ),而不是衰减的结构。这是 GDN 论文那套箭头记号的价值:它把「衰减」这件事隔离成了一个可以独立理解的层面。
7. 总结
delta rule 的本质是定向替换 。I − β t k t k t ⊺ \mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^\intercal I − β t k t k t ⊺ 先减掉 k t \bm{k}_t k t 键上的旧值 S t − 1 k t \mathbf{S}_{t-1}\bm{k}_t S t − 1 k t ,再写入新值 β t v t + ( 1 − β t ) S t − 1 k t \beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t β t v t + ( 1 − β t ) S t − 1 k t ;与 k t \bm{k}_t k t 正交的记忆完全不受影响。这与标量衰减的一刀切互补:门控负责快速擦除,delta rule 负责精确修改。
L2 归一化 k \bm{k} k 不只是训练技巧 。∥ k t ∥ = 1 \|\bm{k}_t\| = 1 ∥ k t ∥ = 1 时 Householder 变换的特征值是 { 1 − β t } ∪ { 1 } d k − 1 \{1-\beta_t\} \cup \{1\}^{d_k-1} { 1 − β t } ∪ { 1 } d k − 1 ,β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 ) 时它们的绝对值都不超过 1,状态不会被越推越大;若不归一化,β t ∥ k t ∥ 2 > 2 \beta_t\|\bm{k}_t\|^2 > 2 β t ∥ k t ∥ 2 > 2 会让特征值翻到 − 1 -1 − 1 以下导致发散。
test-time SGD 视角 :L = 1 2 ∥ S k − v ∥ 2 \mathcal{L} = \frac{1}{2}\|\mathbf{S}\bm{k}-\bm{v}\|^2 L = 2 1 ∥ S k − v ∥ 2 的一步梯度下降就是 delta rule,β t \beta_t β t 是学习率、α t \alpha_t α t 是 weight decay。
WY 表示的核心事实:C C C 个 Householder 的乘积只是秩至多 C C C 的修正 。P r = P r − 1 ( I − β r k r k r ⊺ ) = P r − 1 − β r P r − 1 k r ⏟ w r k r ⊺ \mathbf{P}^r = \mathbf{P}^{r-1}(\mathbf{I}-\beta_r\bm{k}_r\bm{k}_r^\intercal) = \mathbf{P}^{r-1} - \underbrace{\beta_r\mathbf{P}^{r-1}\bm{k}_r}_{\bm{w}_r}\bm{k}_r^\intercal P r = P r − 1 ( I − β r k r k r ⊺ ) = P r − 1 − w r β r P r − 1 k r k r ⊺ ,每多一个 Householder 秩只增 1,于是连乘变求和:P = I − W ⊺ K \mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K} P = I − W ⊺ K 。H = U ⊺ K \mathbf{H} = \mathbf{U}^\intercal\mathbf{K} H = U ⊺ K 同理。
矩阵值转移矩阵迫使块内计算变成解三角系统 。前两篇的转移量是标量,可以查表;Householder 连乘不行。把递推重写成 S r = α r S r − 1 + d r k r ⊺ \mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^\intercal S r = α r S r − 1 + d r k r ⊺ 后倒代换,d r \bm{d}_r d r 的依赖关系构成单位下三角系统,系数矩阵恰是 diag ( β ) ( Γ ⊙ K K ⊺ ) \operatorname{diag}(\beta)(\Gamma \odot \mathbf{K}\mathbf{K}^\intercal) diag ( β ) ( Γ ⊙ K K ⊺ ) –Γ \Gamma Γ 在这里第二次出现 。
两条推导路线互为校验 。论文附录 A 的归纳法从闭式 S t = ∑ i γ t γ i u i k i ⊺ \mathbf{S}_t = \sum_i\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^\intercal S t = ∑ i γ i γ t u i k i ⊺ 出发验证,§2.3 的倒代换从递推构造,两者给出同一个量(实测 u t = d t \bm{u}_t = \bm{d}_t u t = d t ,差 1.11 × 10 − 16 1.11\times10^{-16} 1.11 × 1 0 − 16 )。归纳法第二项的系数 α t + 1 β t + 1 \alpha_{t+1}\beta_{t+1} α t + 1 β t + 1 独立确认了下一条那个必须带的 α \alpha α 。另注意附录 A 只推首块,跨块要补 − β t γ t S 0 ⊺ k t -\beta_t\gamma_t\mathbf{S}_0^\intercal\bm{k}_t − β t γ t S 0 ⊺ k t ——照抄会出现「首块全对、第二块起错」的症状。
d r \bm{d}_r d r 的定义里带 α r \alpha_r α r :d r = β r v r − α r β r S r − 1 k r \bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r d r = β r v r − α r β r S r − 1 k r 。删除项作用在已衰减的状态上,漏掉 α r \alpha_r α r 不会报错,α ≡ 1 \alpha \equiv 1 α ≡ 1 时也看不出来,只在两者都非退化时表现为数值不符。
三角系统是良态的,这是好消息 。I + strictLower ( A ) \mathbf{I} + \operatorname{strictLower}(\mathbf{A}) I + strictLower ( A ) 是单位下三角,行列式恒为 1、永不奇异、不需 pivoting。实测 C = 64 C = 64 C = 64 、β ∈ ( 0 , 1 ) \beta \in (0,1) β ∈ ( 0 , 1 ) 时条件数不超过 12,fp16 够用。放开到 β ∈ ( 0 , 2 ) \beta \in (0,2) β ∈ ( 0 , 2 ) 升到 34。
T \mathbf{T} T 与其逆都不必物化 ,直接对右端项做前向替换即可,实测与 inv 差异在 10 − 17 10^{-17} 1 0 − 17 量级。省一个 C × C C \times C C × C 中间量。
代价是并行度 。前向替换的 C C C 步依赖链无法打破,每步只有 block_DV 个元素并行(32 < 128 线程)。加上累积积的串行前缀和,本篇有两段 C C C 步串行–这是 gated delta rule 相比前两篇的固有开销。
C C C 被限制在 64 。三张 C 2 C^2 C 2 的 f32 fragment(Γ \Gamma Γ 、Q K ⊺ \mathbf{Q}\mathbf{K}^\intercal Q K ⊺ 、Amat)在 C = 128 C = 128 C = 128 时合计 384 reg/thread,超过 255 上限。
K_s 污染两次是最容易错的地方 :构造 Amat 用原始 K \mathbf{K} K ,算右端项时乘 diag ( γ ) \operatorname{diag}(\gamma) diag ( γ ) ,块内项要重载原始 K \mathbf{K} K ,状态更新再乘 γ C / γ r \gamma^C/\gamma^r γ C / γ r 。上一篇只污染一次。
两个退化检验都必须做 :α ≡ 1 \alpha \equiv 1 α ≡ 1 回到纯 DeltaNet(实测 8.88 × 10 − 16 8.88\times10^{-16} 8.88 × 1 0 − 16 )、β → 0 \beta \to 0 β → 0 回到纯衰减(1.62 × 10 − 27 1.62\times10^{-27} 1.62 × 1 0 − 27 )。前者让 Γ \Gamma Γ 的指数错误隐身,后者让整个 UT 部分隐身,缺一不可。
衰减的三处落点三篇未变 。q ← = γ r q \overleftarrow{\bm{q}} = \gamma^r\bm{q} q = γ r q 、k → = γ C γ r k \overrightarrow{\bm{k}} = \frac{\gamma^C}{\gamma^r}\bm{k} k = γ r γ C k 、S → = γ C S \overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S} S = γ C S 从第一篇到第三篇完全一致,delta rule 只改变块内项右乘的内容。
可迁移的启示 :引入一个"看起来只是多一项"的机制,实际代价往往不在 FLOPs 而在依赖结构 。delta rule 的算术开销并不大(一个 C × C C\times C C × C 的三角系统),但它把块内计算从"纯矩阵乘"变成了"带串行依赖的求解",并行度从 128 掉到 32。评估一个改动时,先问它引入了什么依赖,再问它加了多少乘法。
参考
Gated Delta Networks(GDN):Yang, Kautz & Hatamizadeh, Gated Delta Networks: Improving Mamba2 with Delta Rule , arXiv:2412.06464,ICLR 2025。§3.1 给出 gated delta rule、§3.3 给出 chunkwise 算法与 UT 变换;§2.2 给出本文 §2.2 复述的无门控 WY/UT 推导(式 3–9)、附录 A(Extended WY Representation for Gated Delta Rule)给出本文 §2.4 复述的归纳法证明 ,原文只考虑首块(S 0 = 0 \mathbf{S}_0=\mathbf{0} S 0 = 0 )
DeltaNet 的硬件高效 chunkwise 算法:Yang et al., 2024b(arXiv:2406.06484)
WY 表示:Bischof & Van Loan, 1985;UT 变换:Joffrain et al., 2006
delta rule 溯源:Widrow & Hoff, 1960;用于线性 Transformer:Schlag et al., 2021a