TileLang 实战:KDA 从零到一--Gated DeltaNet

TileLang 实战:KDA 从零到一–Gated DeltaNet

前两篇分别实现了无遗忘的 chunked 线性注意力,以及带标量衰减的版本。两者的状态更新都只做加法:新的键值对累加进状态,旧信息靠衰减系数被动遗忘。本篇补上最后一块–删除

递推式变成:

St=St1(αt(Iβtktkt))+βtvtkt\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}

比上一篇多出的是 Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} 这个广义 Householder 变换。它带来的不是又一个逐元素权重,而是一个矩阵乘在状态右侧–这一改把整个 chunkwise 算法的结构改掉了:块内不再能靠一张下三角权重表解决,需要解一个 C×CC \times C 的三角系统。

上两篇见《Chunked 线性注意力》与《标量衰减》,本文沿用其符号约定与四层参考验证框架。


0. 符号约定

沿用 GDN(arXiv:2412.06464v3)§2.2、§3.1、§3.3 与附录 A 的记号,本篇新增两个:

符号 含义 说明
βt(0,1)\beta_t \in (0,1) 写入强度(writing strength) 也是 delta rule 视角下的学习率
T[t]RC×C\mathbf{T}_{[t]} \in \mathbb{R}^{C \times C} UT 变换矩阵 下三角系统的逆,本篇的核心开销
U~[t]RC×dv\widetilde{\mathbf{U}}_{[t]} \in \mathbb{R}^{C \times d_v} 修正后的 value T\mathbf{T} 作用在 diag(β)V\operatorname{diag}(\beta)\mathbf{V}
W[t]RC×dk\mathbf{W}_{[t]} \in \mathbb{R}^{C \times d_k} 修正后的 key T\mathbf{T} 作用在 diag(β)K\operatorname{diag}(\beta)\mathbf{K}

沿用的:αt\alpha_t 单步衰减、γ[t]r=i=1rαi\gamma^r_{[t]} = \prod_{i=1}^{r}\alpha_i 累积衰减积(按 chunk 重置)、Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 衰减感知因果掩码、CC 块长、[t][t] 块序号、r[1,C]r \in [1,C] 块内位置。

一个前提:论文对 q,k\bm{q}, \bm{k} 做 L2 归一化,即 kt=1\|\bm{k}_t\| = 1。这不只是训练稳定性的考虑–下面 §1.2 会看到它直接决定了 Householder 变换会不会把状态越推越大。


1. 从加法到删除:delta rule 在做什么

1.1 三种更新方式的对照

把三篇的递推式并排放,差异一目了然:

形态 递推式 状态如何变化
线性注意力 St=St1+vtkt\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} 只增不减
标量衰减 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} 整体等比遗忘
gated delta rule St=St1(αt(Iβtktkt))+βtvtkt\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} 定向替换 + 整体遗忘

标量衰减的问题在于它不分对象αt\alpha_t 一乘,所有历史信息按同一比例衰减。要腾出空间写入新内容,只能把无关的旧信息一起冲淡。GDN 论文的说法是,门控擅长「快速擦除」,delta rule 擅长「定向修改」,两者互补。

delta rule 的定向性来自哪里?把它拆开看:

St=St1(St1kt)vtoldkt+(βtvt+(1βt)St1kt)vtnewkt\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}

读法是:先把 kt\bm{k}_t 这个键上原有的值 vtold=St1kt\bm{v}^{\text{old}}_t = \mathbf{S}_{t-1}\bm{k}_t 减掉,再写入新值 vtnew\bm{v}^{\text{new}}_t,而新值是旧值与目标值的凸组合,βt\beta_t 控制替换的彻底程度。βt1\beta_t \to 1 是完全覆盖,βt0\beta_t \to 0 是不动。

关键在于这个减法只作用在 kt\bm{k}_t 方向上,与 kt\bm{k}_t 正交的记忆分毫不动。这就是「定向」–相比标量衰减的一刀切,delta rule 只擦掉要覆盖的那一条。

1.2 为什么是 Householder,以及 L2 归一化的作用

Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} 是广义 Householder 变换。kt=1\|\bm{k}_t\| = 1 时,它的特征值只有两种取值,结构一目了然:

  • 沿 kt\bm{k}_t 方向:特征值 1βt1 - \beta_t
  • kt\bm{k}_t 正交的 dk1d_k - 1 个方向:特征值 11

于是 βt(0,1)\beta_t \in (0,1) 时全部特征值的绝对值都不超过 1,状态每步只会被压缩、不会被放大,递推因此稳定。这正是论文对 k\bm{k} 做 L2 归一化的深层原因–若 kt1\|\bm{k}_t\| \ne 1,沿 kt\bm{k}_t 的特征值变成 1βtkt21 - \beta_t\|\bm{k}_t\|^2βtkt2>2\beta_t\|\bm{k}_t\|^2 > 2 时就会翻到 1-1 以下,递推放大。

顺带一提,论文脚注提到可以放开到 βt(0,2)\beta_t \in (0,2) 以允许负特征值,那是为了解锁状态跟踪能力(state tracking)。本文按 (0,1)(0,1) 处理。

1.3 test-time SGD 视角

论文给了一个很有启发的解释:把状态 S\mathbf{S} 看成一个快速权重矩阵,delta rule 就是在做在线回归的一步梯度下降。目标是 L(St)=12Stktvt2\mathcal{L}(\mathbf{S}_t) = \frac{1}{2}\|\mathbf{S}_t\bm{k}_t - \bm{v}_t\|^2,那么:

StβtL(St)=Stβt(Stktvt)kt=St(Iβtktkt)+βtvtkt\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}

βt\beta_t 就是学习率,αt\alpha_t 就是 weight decay。 这个视角下 gated delta rule 没有任何神秘之处–它是带权重衰减的 test-time SGD。


2. Chunkwise 形式:为什么需要解三角系统

2.1 展开递推:转移矩阵不再是标量

按块展开 rr 步(论文式 10):

S[t]r=S[t]i=1rα[t]i(Iβ[t]ik[t]ik[t]i)F[t]r+i=1rβ[t]iv[t]ik[t]ij=i+1rα[t]j(Iβ[t]jk[t]jk[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]}}

对比上一篇:那里的转移量是标量 αri\alpha^{r-i},可以直接查表。这里是矩阵连乘F\mathbf{F}G\mathbf{G} 都不能靠逐元素权重表达。

衰减部分可以先提出来–αi=γr\prod\alpha_i = \gamma^r 是标量,与 Householder 部分可交换:

F[t]r=γ[t]rP[t]r,P[t]r=i=1r(Iβikiki)\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)

剩下的 P\mathbf{P} 是纯 DeltaNet 的部分,靠 WY 表示处理–这正是论文 §2.2 的内容,下一节完整展开。

2.2 论文 §2.2:无门控 DeltaNet 的 WY 表示

在处理带门控的版本之前,先把论文 §2.2 那套无门控 DeltaNet 的 chunkwise 推导完整走一遍。GDN 的做法本质上是在这套框架上打补丁,先看清基线,后面的改动才有参照。

本节 αt1\alpha_t \equiv 1(无遗忘门),递推退化为 St=St1(Iβtktkt)+βtvtkt\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}

2.2.1 部分展开:两个连乘(论文式 3)

按块部分展开递推:

S[t]r=S[t](i=1r(Iβ[t]ik[t]ik[t]i)):=P[t]r+i=1rβ[t]iv[t]ik[t]ij=i+1r(Iβ[t]jk[t]jk[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]}}

两个部分的角色不同,值得分清:

形状 含义 结构
P[t]r\mathbf{P}^r_{[t]} dk×dkd_k \times d_k 历史状态的遗忘算子:入口状态经过 rr 步 Householder 后剩下什么 Householder 的纯连乘
H[t]r\mathbf{H}^r_{[t]} dv×dkd_v \times d_k 块内新写入的累积:前 rr 个 token 写进来的内容(互相已扣除重叠) 连乘的加权和

P\mathbf{P} 是纯连乘,H\mathbf{H} 的每一项后面还挂着一截连乘尾巴–两者都不能直接算,C=64C = 64P\mathbf{P} 要 64 个 dk×dkd_k\times d_k 矩阵相乘。WY 表示的作用就是把这两个连乘各自压成一次求和。

2.2.2 经典 WY:把 Householder 连乘压成秩-CC 更新(论文式 4)

这是 Bischof & Van Loan (1985) 的经典结果。核心事实:CC 个 Householder 矩阵的乘积可以写成单位矩阵减去一个秩至多 CC 的修正

P[t]r=Ii=1rw[t]ik[t]iRdk×dk,w[t]r=β[t]r(k[t]ri=1r1w[t]i(k[t]ik[t]r))Rdk\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}

为什么成立,看一步归纳就够:

Pr=Pr1(Iβrkrkr)=Pr1βrPr1krwrkr\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}

于是 wr=βrPr1kr\bm{w}_r = \beta_r\mathbf{P}^{r-1}\bm{k}_r,把 Pr1=Ii<rwiki\mathbf{P}^{r-1} = \mathbf{I} - \sum_{i<r}\bm{w}_i\bm{k}_i^{\intercal} 代入就得到上面那个递推。每多一个 Householder,秩只增加 1–这就是"连乘变求和"的全部内容。

2.2.3 同一个模具:H\mathbf{H} 的 WY 表示(论文式 5)

H\mathbf{H} 的推导结构与 P\mathbf{P} 完全平行

H[t]r=i=1ru[t]ik[t]iRdv×dk,u[t]r=β[t]r(v[t]ri=1r1u[t]i(k[t]ik[t]r))Rdv\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}

w\bm{w} 的递推和 u\bm{u} 的递推并排看,会发现它们是同一个式子

wr=βr(kri<rwi(kikr)),ur=βr(vri<rui(kikr))\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)

只有下划线处不同:一个是 kr\bm{k}_r、一个是 vr\bm{v}_r系数完全一样–这个观察是后面式 (6)(7) 能共用一个 T\mathbf{T} 的全部原因,也是 kernel 里 W\mathbf{W}U\mathbf{U} 能拼进一次前向替换的依据(第四篇 KDA 用的就是这个技巧)。

写成矩阵形式,两个连乘都消失了:

P[t]=IW[t]K[t]Rdk×dk,H[t]=U[t]K[t]Rdv×dk\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}

2.2.4 UT 变换:递推也不必串行(论文式 6、7)

式 (4)(5) 虽然把连乘变成了求和,但 wr\bm{w}_r / ur\bm{u}_r 自身还是串行递推。Joffrain et al. (2006) 的 UT 变换把它变成一次矩阵求逆:

T[t]=[I+strictLower(diag(β[t])K[t]K[t])]1diag(β[t])RC×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}

W[t]=T[t]K[t]RC×dk,U[t]=T[t]V[t]RC×dv\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\mathbf{W}U\mathbf{U} 共用同一个 T\mathbf{T},正是因为 §2.2.3 那两个递推的系数相同。这一步的意义在于:T\mathbf{T} 只跟 K\mathbf{K}β\beta 有关,与 V\mathbf{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]Rdv×dkO[t]=Q[t]S[t]+(Q[t]K[t]M)(U[t]W[t]S[t])RC×dv\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}

式 (8) 的化简值得看一眼–SP=S(IWK)=S(WS)K\mathbf{S}\mathbf{P} = \mathbf{S}(\mathbf{I}-\mathbf{W}^\intercal\mathbf{K}) = \mathbf{S} - (\mathbf{W}\mathbf{S}^\intercal)^\intercal\mathbf{K},与 H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 合并后 K\mathbf{K} 被提到右侧,括号里剩下 UWS\mathbf{U} - \mathbf{W}\mathbf{S}^\intercal这个量在式 (8) 和式 (9) 里是同一个,算一次用两回。

M\mathbf{M} 是下三角全 1 的因果掩码。注意此处没有任何衰减–这正是 GDN 要改的地方。

2.2.6 数值验证:九个等式逐一核对

dk=6,dv=5,C=7d_k=6, d_v=5, C=7βU(0.1,0.9)\beta \sim \mathcal{U}(0.1,0.9)k\bm{k} 已 L2 归一化,入口状态 S00\mathbf{S}_0 \ne \mathbf{0}(比论文附录 A 的首块假设更严格),fp64:

论文式 内容 max abs 误差
(3) Sr=S0Pr+Hr\mathbf{S}^r = \mathbf{S}_0\mathbf{P}^r + \mathbf{H}^r 4.44×10164.44\times10^{-16}
(4) Pr=Iiwiki\mathbf{P}^r = \mathbf{I} - \sum_i\bm{w}_i\bm{k}_i^\intercal 4.44×10164.44\times10^{-16}
(5) Hr=iuiki\mathbf{H}^r = \sum_i\bm{u}_i\bm{k}_i^\intercal 3.33×10163.33\times10^{-16}
矩阵形式 P=IWK\mathbf{P} = \mathbf{I} - \mathbf{W}^\intercal\mathbf{K} 2.22×10162.22\times10^{-16}
矩阵形式 H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 2.78×10162.78\times10^{-16}
(6)(7) W=TK\mathbf{W} = \mathbf{T}\mathbf{K}(UT 变换 vs 式 4 递推) 5.55×10175.55\times10^{-17}
(6)(7) U=TV\mathbf{U} = \mathbf{T}\mathbf{V} 2.22×10162.22\times10^{-16}
(8) S[t+1]\mathbf{S}_{[t+1]} 完整式 4.44×10164.44\times10^{-16}
(9) O[t]\mathbf{O}_{[t]} 完整式 8.88×10168.88\times10^{-16}

另外验证了 P=IWK\mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K} 的特征值绝对值分别为 0.9918,0.9082,0.5805,0.3131,0.1951,0.06760.9918, 0.9082, 0.5805, 0.3131, 0.1951, 0.0676全部 1\le 1,与 §1.2 说的"每步只压缩不放大"一致。

2.2.7 GDN 改了哪两处

把门控加回来(αt1\alpha_t \ne 1),论文 §3.3 的做法只动了式 (6)(7) 两个地方

无门控(式 6、7) GDN(门控版)
T\mathbf{T} 里的 KK 矩阵 KK\mathbf{K}\mathbf{K}^{\intercal} ΓKK\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}
W\mathbf{W} 的输入 K\mathbf{K} diag(γ)K=K\operatorname{diag}(\gamma)\mathbf{K} = \overleftarrow{\mathbf{K}}
U\mathbf{U} 的输入 V\mathbf{V} V\mathbf{V}(不变)
式 (8) 的旧状态项 S[t]\mathbf{S}_{[t]} γCS[t]=S\gamma^C\mathbf{S}_{[t]} = \overrightarrow{\mathbf{S}}
式 (8) 的 K\mathbf{K} K\mathbf{K} γCγrK=K\frac{\gamma^C}{\gamma^r}\mathbf{K} = \overrightarrow{\mathbf{K}}
式 (9) 的 Q\mathbf{Q} 与掩码 Q\mathbf{Q}M\mathbf{M} γrQ=Q\gamma^r\mathbf{Q} = \overleftarrow{\mathbf{Q}}Γ\Gamma

实测这个替换的正确性:门控版 S[t+1]\mathbf{S}_{[t+1]} 与逐 token 递归差 3.33×10163.33\times10^{-16};令 α1\alpha\equiv1 时,Γ\Gamma 与下三角全 1 掩码 M\mathbf{M} 完全相等(差 0.0)T\mathbf{T} 与无门控版完全相等(差 0.0)W\mathbf{W} 也完全相等。三个 0.0 说明门控版是无门控版的严格推广,没有引入任何额外近似。

换个角度看Γ\Gamma 就是把式 (9) 那个 0/1 因果掩码 M\mathbf{M} 升级成了"带衰减的因果掩码"。Mij{0,1}\mathbf{M}_{ij} \in \{0,1\} 只回答"jj 能否影响 ii",Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 还回答"影响衰减了多少"。前两篇反复出现的那个 Γ\Gamma,在这个视角下就是 M\mathbf{M} 的自然推广。


2.3 手推一遍:Householder 连乘怎么变成三角系统

直接推 C=4C = 4 的情形最清楚。定义每步的修正量 dr\bm{d}_r,使得递推写成加法形式:

Sr=αrSr1+drkr,dr=βrvrαrβrSr1kr\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

这一步只是把 Sr1(αr(Iβrkk))+βrvk\mathbf{S}^{r-1}(\alpha_r(\mathbf{I}-\beta_r\bm{k}\bm{k}^\intercal)) + \beta_r\bm{v}\bm{k}^\intercal 重新分组,恒等变形。注意 dr\bm{d}_r 里带 αr\alpha_r–删除项作用在已经衰减过的状态上,这个细节写错不会报错、只会算错。

有了加法形式,就能像上一篇那样倒代换(每项系数是「距块尾的步数」):

S[t+1]=γCS[t]+r=1CγCγrdrkr\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}

问题在于 dr\bm{d}_r 依赖 Sr1\mathbf{S}^{r-1},而 Sr1\mathbf{S}^{r-1} 又依赖 d1,,dr1\bm{d}_1, \ldots, \bm{d}_{r-1}–串行依赖,无法并行。把 Sr1\mathbf{S}^{r-1} 也倒代换开(利用 αrγr1=γr\alpha_r\gamma^{r-1} = \gamma^r):

αrSr1=γrS[t]+i<rγrγidiki\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}

代回 dr\bm{d}_r 的定义:

dr=βr(vrγrS[t]kri<rγrγi(kikr)di)\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)

这是一个下三角线性系统dr\bm{d}_r 只依赖 di (i<r)\bm{d}_i\ (i < r),系数是 βrγrγi(kikr)\beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r)。写成矩阵形式,令 Ari=βrγrγi(kikr)\mathbf{A}_{ri} = \beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r) 的严格下三角部分:

(I+strictLower(A))Δ=diag(β)(Vdiag(γ)KS[t])(\mathbf{I} + \operatorname{strictLower}(\mathbf{A}))\,\Delta = \operatorname{diag}(\beta)\big(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}_{[t]}^{\intercal}\big)

其中 ΔRC×dv\Delta \in \mathbb{R}^{C \times d_v} 的第 rr 行是 dr\bm{d}_r。注意 A=diag(β)(ΓKK)\mathbf{A} = \operatorname{diag}(\beta)\big(\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}\big)衰减感知掩码 Γ\Gamma 在这里第二次出现,这次是嵌在三角系统的系数矩阵里。

2.4 论文附录 A:扩展 WY 表示的归纳证明

上面 §2.3 是"把递推硬拆开"的推法。论文附录 A 给了一个更漂亮的等价路线:先猜出闭式,再用数学归纳法证明。这一节按论文原文复述(GDN 论文附录 A,为减少符号负担,论文同样只考虑首块,即 S0=0\mathbf{S}_0 = \mathbf{0})。

命题(扩展 WY 表示).St\mathbf{S}_t

St=i=1tγtγiuiki,ut=βt(vti=1t1γtγiuikikt)\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)

证明.tt 作归纳。

St+1=St(αt+1(Iβt+1kt+1kt+1))+βt+1vt+1kt+1=αt+1(i=1tγtγiuiki)αt+1βt+1(i=1tγtγiuikikt+1kt+1)+βt+1vt+1kt+1=i=1tγt+1γiuiki+βt+1(vt+1i=1tγt+1γiuikikt+1)ut+1kt+1=i=1tγt+1γiuiki+γt+1γt+11ut+1kt+1  =  i=1t+1γt+1γiuiki\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}

\square

第二个等号到第三个等号是全部关键,值得拆开看:

  • αt+1γtγi=γt+1γi\alpha_{t+1}\cdot\frac{\gamma_t}{\gamma_i} = \frac{\gamma_{t+1}}{\gamma_i}–衰减系数被吸收进累积积的比值,这是 γ\gamma 定义为连乘才有的性质,也是四篇一直在用的那条恒等式。
  • 第二项里 αt+1\alpha_{t+1} 同样被吸收进 γt+1γi\frac{\gamma_{t+1}}{\gamma_i},然后整项与第三项合并、共同提出右侧的 kt+1\bm{k}_{t+1}^{\intercal}–括号里剩下的就是 ut+1\bm{u}_{t+1} 的定义。
  • 最后一步只是注意 γt+1/γt+1=1\gamma_{t+1}/\gamma_{t+1} = 1,于是新项能并入求和,下标从 tt 推到 t+1t+1,归纳闭合。

这个证明独立印证了 §2.3 那个坑。 注意第二项的系数是 αt+1βt+1\alpha_{t+1}\beta_{t+1}α\alphaβ\beta 都在,因为删除项 Iβkk\mathbf{I}-\beta\bm{k}\bm{k}^\intercal 作用在已经乘过 αt+1\alpha_{t+1} 的状态上。我在 §2.3 用倒代换推 dr\bm{d}_r 时漏掉这个 αr\alpha_r,refB 就与逐 token 递归差了 2.66×1012.66\times10^{-1}。两条路线在同一个位置要求同一个因子,可以互为校验。

实测验证这个命题(dk=5,dv=4,C=6d_k=5, d_v=4, C=6,fp64,逐步比对 St\mathbf{S}_t):

检验 结果
WY 闭式 vs 逐 token 递推(t=1..6t=1..6 逐步) 最大 1.67×10161.67\times10^{-16}
ut\bm{u}_t 是否等于 §2.3 的修正量 dt\bm{d}_t 1.11×10161.11\times10^{-16}
dt\bm{d}_t 漏掉 αt\alpha_t 偏差 2.99×1022.99\times10^{-2}

第二行说明论文的 ut\bm{u}_t 与我 §2.3 手推的 dt\bm{d}_t 是同一个量,只是推导路径不同:论文归纳法从闭式出发验证,§2.3 从递推倒代换构造。两者都给出同一个下三角系统,下面的 UT 变换对二者通用。

关于跨块:论文只推首块(S0=0\mathbf{S}_0=\mathbf{0})。实际 kernel 里每块的入口状态非零,ut\bm{u}_t 的定义要补上 βtγtS0kt-\beta_t\gamma_t\mathbf{S}_0^{\intercal}\bm{k}_t 这一项,即 §2.3 里那个 γrS[t]kr-\gamma^r\mathbf{S}_{[t]}\bm{k}_r照抄附录 A 而漏掉这项,是我最初 refB 失配(max abs 2.66×1012.66\times10^{-1}、rel L2 6.01×1026.01\times10^{-2})的另一个来源–首块测试全对、第二块开始错,这种症状基本可以直接定位到跨块项。

2.5 UT 变换:把解系统变成矩阵乘

定义(论文 §2.2 与 §3.3):

T[t]=[I+strictLower(diag(β[t])(Γ[t]K[t]K[t]))]1diag(β[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]})

于是修正量一次算出:

Δ=TVU~Tdiag(γ)KWS[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}

这正是论文式 (11)(12) 里的 (U~[t]W[t]S[t])\big(\widetilde{\mathbf{U}}_{[t]} - \overleftarrow{\mathbf{W}_{[t]}}\mathbf{S}_{[t]}^{\intercal}\big)。完整的 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}

三个箭头量与前两篇完全一致:qr=γrqr\overleftarrow{\bm{q}^r} = \gamma^r\bm{q}^rkr=γCγrkr\overrightarrow{\bm{k}^r} = \frac{\gamma^C}{\gamma^r}\bm{k}^rS=γCS\overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S}

对比上一篇,结构上只有一处变化:块内那一项的右乘对象从 V[t]\mathbf{V}_{[t]} 变成了 Δ\Delta,而 Δ\Delta 需要解一个三角系统才能得到。β0\beta \to 0T0\mathbf{T} \to \mathbf{0}Δ0\Delta \to \mathbf{0},退化为纯衰减;α1\alpha \equiv 1ΓM\Gamma \to \mathbf{M},退化为纯 DeltaNet。

2.6 数值验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4αtU(0.90,0.999)\alpha_t \sim \mathcal{U}(0.90, 0.999)βtU(0.10,0.90)\beta_t \sim \mathcal{U}(0.10, 0.90)k\bm{k} 已 L2 归一化,fp64:

比较 max abs 误差 相对 L2
chunkwise + UT vs 逐 token 递归 8.88×10168.88 \times 10^{-16} 2.49×10162.49 \times 10^{-16}
α1\alpha \equiv 1(退化为纯 DeltaNet) 8.88×10168.88 \times 10^{-16}
β0\beta \to 0(退化为纯衰减) 1.62×10271.62 \times 10^{-27}

两个退化检验都必须做,理由与上一篇同构:α1\alpha \equiv 1Γ\Gamma 退化成 0/1 掩码,Γ\Gamma 里任何指数写错都看不出来;β0\beta \to 0 时整个三角系统消失,UT 部分的错误全部隐身。


3. 三角系统的数值性质

T\mathbf{T} 要求一个 C×CC \times C 矩阵的逆,这是本篇最值得担心的地方。但实际上它比看起来温和得多。

3.1 单位下三角,条件数可控

I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A})单位下三角矩阵–对角线恒为 1,严格下三角才是 A\mathbf{A}。这有两个直接后果:

  1. 行列式恒为 1,永不奇异。 不存在需要 pivoting 的情形。
  2. 求逆可以用前向替换,不需要通用矩阵求逆。

实测条件数(C=64C = 64k\bm{k} L2 归一化,αtU(0.9,0.999)\alpha_t \sim \mathcal{U}(0.9, 0.999),20 组随机采样):

βt\beta_t 采样区间 cond 中位数 cond 最大
[0.05, 0.2][0.05,\ 0.2] 1.76 1.91
[0.1, 0.9][0.1,\ 0.9] 5.67 6.86
[0.5, 0.99][0.5,\ 0.99] 7.88 9.05
[0.9, 0.999][0.9,\ 0.999] 10.38 11.99
[1.0, 1.99][1.0,\ 1.99] 34.27 40.85

β(0,1)\beta \in (0,1) 时条件数不超过 12,fp16 完全够用。这与上一篇那个「因式分解会溢出」的结论形成有意思的对比:那里是代数变形引入了 1/γ1/\gamma 这种指数增长量,而这里虽然要求逆,但矩阵结构本身保证了良态。

放开到 β(0,2)\beta \in (0,2)(论文脚注提到的负特征值情形)条件数跳到 34,仍可接受,但已需要留意。

3.2 前向替换代替求逆

T\mathbf{T} 从不需要显式求逆。逐行前向替换:

T[r,:]=erj<rA[r,j]T[j,:]\mathbf{T}[r,:] = \bm{e}_r - \sum_{j<r}\mathbf{A}[r,j]\,\mathbf{T}[j,:]

实测与 np.linalg.inv 的差异在 101710^{-17} 量级(三组随机种子分别 7.81,5.55,8.33×10177.81, 5.55, 8.33 \times 10^{-17}),即机器精度。

更进一步,T\mathbf{T} 本身也不必物化。 需要的只是 Δ=TR\Delta = \mathbf{T}\mathbf{R}(其中 R=diag(β)(Vdiag(γ)KS)\mathbf{R} = \operatorname{diag}(\beta)(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}^\intercal)),直接对 R\mathbf{R} 做前向替换:

Δ[r,:]=R[r,:]j<rA[r,j]Δ[j,:]\Delta[r,:] = \mathbf{R}[r,:] - \sum_{j<r}\mathbf{A}[r,j]\,\Delta[j,:]

这省掉一个 C×CC \times C 的中间量。代价是引入了 CC 步串行–这是本篇 kernel 与前两篇最本质的区别,下面 §5 会看到它如何限制并行度。


4. 四层参考实现

沿用前两篇的框架:

参考 实现方式 验证目标
A 逐 token 递归 递推式定义 St=St1(αt(Iβtkk))+βtvk\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
B chunkwise + UT 变换 §2.3 的三角系统推导、§2.4 的 WY 表示与式 (11)(12)
C kernel 控制流镜像(切 DV、S\mathbf{S}^\intercal 布局、前向替换) kernel 结构
D 退化检验(α1\alpha \equiv 1 / β0\beta \to 0 与前两篇的一致性

4.0 状态方向

与前两篇一致:论文的状态是 SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k},kernel 存转置 SRdk×dv\mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v}。本篇多一层要注意的是 ΔRC×dv\Delta \in \mathbb{R}^{C \times d_v}–它的布局与 V\mathbf{V} 相同,所以 ΔK\Delta^\intercal\overrightarrow{\mathbf{K}} 在转置布局下写作 KΔ\overrightarrow{\mathbf{K}}^\intercal\Delta

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]
# 广义 Householder:注意 α 乘在整个括号上
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) # γ^r,chunk 内重置
gC = g[-1]
i = np.arange(C)
Gam = np.where(i[:, None] >= i[None, :],
g[i][:, None] / g[i][None, :], 0.0) # Γ

# UT 变换:T = [I + strictLower(diag(β)(Γ ⊙ K K^T))]^{-1} diag(β)
A = np.diag(be) @ (Gam * (Kc @ Kc.T))
T = np.linalg.inv(np.eye(C) + np.tril(A, -1)) @ np.diag(be)

# Δ = Ũ - \overleftarrow{W} S^T
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): # ← grid.x
dv = slice(bv * block_DV, (bv + 1) * block_DV)
for bbh in range(B * H): # ← grid.y
b, h = bbh // H, bbh % H
S = np.zeros((DK, block_DV)) # 转置布局
for c in range(N // C): # ← T.Pipelined
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 # diag(β)(Γ ⊙ K K^T)
Tm = np.linalg.inv(np.eye(C) + np.tril(A, -1))
R = be[:, None] * Vc - be[:, None] * ((g[:, None] * Kc) @ S)
Delta = Tm @ R # C × block_DV

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,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4blockDV=2\text{block}_{DV}=2,fp64:

比较 max abs 误差 相对 L2
B chunkwise + UT vs A 逐 token 8.88×10168.88 \times 10^{-16} 2.49×10162.49 \times 10^{-16}
C kernel 镜像 vs A 逐 token 8.88×10168.88 \times 10^{-16} 2.61×10162.61 \times 10^{-16}
α1\alpha \equiv 1(纯 DeltaNet)B vs A 8.88×10168.88 \times 10^{-16}
β1012\beta \to 10^{-12}(纯衰减)B vs A 1.62×10271.62 \times 10^{-27}

另验证了 T\mathbf{T} 的两种等价写法(diag(β) @ (Γ * KKᵀ)(K * β) @ Kᵀ * Γ)差异 5.55×10175.55 \times 10^{-17},后者省一次对角矩阵构造。


5. TileLang kernel

grid 划分沿用前两篇:切 (bv, bh)(bv,\ b \cdot h),序列轴走 kernel 内的 T.Pipelined。但本篇多了一个 C×CC \times CA 矩阵和 CC 步串行的前向替换,寄存器与并行度都要重新算。

5.0 寄存器账

dk=dv=128d_k = d_v = 128C=64C = 64、128 线程:

fragment 上一篇 本篇 说明
S_f [dk,bDV][d_k, \text{bDV}] 32 reg 32 reg 不变
acc_o [C,bDV][C, \text{bDV}] 16 reg 16 reg 不变
Dtri / Gam [C,C][C,C] 32 reg 32 reg 衰减掩码
A [C,C][C,C] 32 reg 32 reg QK\mathbf{Q}\mathbf{K}^\intercal 复用
Amat [C,C][C,C] 32 reg 新增:三角系统系数
Delta [C,bDV][C, \text{bDV}] 16 reg 新增:修正量
合计 ≈ 112 reg ≈ 160 reg 上限 255

C=64C = 64 时仍有余量。C=128C = 128Gam + A + Amat 三张 C2C^2 表就是 384 reg,直接超限–本篇的 CC 实际上被限制在 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), # α_t,数据依赖
Beta: T.Tensor([B, S, H], accum_dtype), # β_t,数据依赖
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) # Δ 的 f16 副本

S_f = T.alloc_fragment([DK, block_DV], accum_dtype) # loop-carried
acc_o = T.alloc_fragment([C, block_DV], accum_dtype)
A = T.alloc_fragment([C, C], accum_dtype) # Q K^T
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) # log2 γ^r
g_r = T.alloc_fragment([C], accum_dtype) # γ^r
w_decay = T.alloc_fragment([C], accum_dtype) # γ^C/γ^r
be_f = T.alloc_fragment([C], accum_dtype) # β_r

αt\alpha_t 现在是数据依赖的,累积积必须在 kernel 内算。用 log 域 cumsum 而非 cumprod–上一篇实测 fp16 直接连乘在 C=128C = 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)

# ── 累积衰减积:log 域前缀和 ──
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]) # γ^r
w_decay[r] = T.exp2(lg[C - 1] - lg[r]) # γ^C/γ^r
be_f[r] = Beta[bb, s0 + r, bh]

# ── Γ[i,j] = γ^i/γ^j (j≤i),log 域相减后取指数 ──
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/γj1/\gamma_j 会溢出。条件求值也不能省成「先全算再掩」,j>ij > 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
# ── Amat = diag(β)(Γ ⊙ K K^T) 的严格下三角 ──
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)

# ── 右端项 R = diag(β)(V - diag(γ) K S^T) ──
T.copy(S_f, S_s)
for r, d in T.Parallel(C, DK):
K_s[r, d] *= g_r[r] # diag(γ) K,原地
T.gemm(K_s, S_s, R, clear_accum=True) # (γK) S
for r, d in T.Parallel(C, block_DV):
R[r, d] = be_f[r] * (V_s[r, d] - R[r, d])

# ── 前向替换:Δ[r] = R[r] - Σ_{j<r} Amat[r,j] Δ[j] ──
for r in T.serial(C): # 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=64C = 64 时是 64 步,每步内部 block_DV = 32 个元素并行–并行度只有 32,远低于 128 线程。这是 gated delta rule 相比前两篇的固有代价:三角系统的依赖链无法打破。

有两个缓解方向:一是增大 block_DV(但吃寄存器),二是把 CC 切成更小的子块做分块前向替换(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
# ② 跨块项:\overleftarrow{Q} S^T
for r, d in T.Parallel(C, DK):
Q_s[r, d] *= g_r[r] # 位置三:query 侧
T.gemm(Q_s, S_s, acc_o, clear_accum=True)

# ③ 块内项:(Q K^T ⊙ Γ) Δ —— 注意要用未加权的原始 Q、K
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.copy(K[bb, s0:s0+C, bh, :], K_s) # K_s 被 diag(γ) 污染过
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]) # γ^C
for r, d in T.Parallel(C, DK):
K_s[r, d] *= w_decay[r] # 位置二:γ^C/γ^r
T.gemm(K_s, D_s, S_f, transpose_A=True) # \overrightarrow{K}^T Δ

5.4 五处次序约束

比上一篇多一处,全部写错都不报错、只算错:

  1. 状态更新排在输出写回之后SfS_f 在步骤②③被读时代表进块状态 S[t]\mathbf{S}_{[t]}
  2. S_f *= γ^CT.gemm 之前。先衰减旧状态再累加新贡献。
  3. K_s 被复用三次、污染两次。构造 Amat 用原始 K\mathbf{K},算 R 时被乘上 diag(γ)\operatorname{diag}(\gamma),块内项要用原始 K\mathbf{K}(须重载),状态更新时又要乘 γC/γr\gamma^C/\gamma^r这是本篇最容易错的地方–上一篇 K_s 只污染一次。
  4. Amat 必须只取严格下三角j<ij < i,不含对角)。含对角就变成了求解 (I+Afull)(\mathbf{I} + \mathbf{A}_{\text{full}}),对角上的 βrkr2=βr\beta_r\|\bm{k}_r\|^2 = \beta_r 会被重复计入。
  5. T.clear(S_f) 在流水线循环外,clear_accum=True 在循环内

5.5 与上一篇的改动汇总

位置 上一篇 本篇 新增开销
输入 Q, K, V + αt\alpha_t, βt\beta_t 两个张量 2 次 HBM 读 / chunk
累积积 编译期常量 kernel 内 log 域 cumsum CC 步串行前缀和
Γ\Gamma 编译期可算 运行时 exp2(lg[i]-lg[j]) C2C^2exp2
三角系统 Amat 构造 + 前向替换 C2C^2 reg + CC 步串行
块内右乘 V\mathbf{V} Δ\Delta 一次 C×bDVC \times \text{bDV} f16 转换
K 复用 污染 1 次 污染 2 次,重载 1 次 一次 shared 写入

6. 三篇对照

线性注意力 标量衰减 Gated DeltaNet
递推 S+vk\mathbf{S} + \bm{v}\bm{k}^\intercal αS+vk\alpha\mathbf{S} + \bm{v}\bm{k}^\intercal S(α(Iβkk))+βvk\mathbf{S}(\alpha(\mathbf{I}-\beta\bm{k}\bm{k}^\intercal)) + \beta\bm{v}\bm{k}^\intercal
块内掩码 M\mathbf{M}(0/1) Γ\Gamma(衰减感知) Γ\Gamma + 三角系统
块内右乘 V\mathbf{V} V\mathbf{V} Δ\Delta
串行段 前缀和 + 前向替换(各 CC 步)
C2C^2 fragment 1(A\mathbf{A} 2(+ Γ\Gamma 3(+ Amat
可用 CC 128 128(切 DV 后) 64
对应架构 RetNet / Lightning-Attn Gated DeltaNet / KDA

三篇的衰减权重落点完全一致q\overleftarrow{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 三处,从第一篇到第三篇没有变过。delta rule 加进来的是块内那一项的内容(VΔ\mathbf{V} \to \Delta),而不是衰减的结构。这是 GDN 论文那套箭头记号的价值:它把「衰减」这件事隔离成了一个可以独立理解的层面。


7. 总结

  1. delta rule 的本质是定向替换Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^\intercal 先减掉 kt\bm{k}_t 键上的旧值 St1kt\mathbf{S}_{t-1}\bm{k}_t,再写入新值 βtvt+(1βt)St1kt\beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t;与 kt\bm{k}_t 正交的记忆完全不受影响。这与标量衰减的一刀切互补:门控负责快速擦除,delta rule 负责精确修改。
  2. L2 归一化 k\bm{k} 不只是训练技巧kt=1\|\bm{k}_t\| = 1 时 Householder 变换的特征值是 {1βt}{1}dk1\{1-\beta_t\} \cup \{1\}^{d_k-1}βt(0,1)\beta_t \in (0,1) 时它们的绝对值都不超过 1,状态不会被越推越大;若不归一化,βtkt2>2\beta_t\|\bm{k}_t\|^2 > 2 会让特征值翻到 1-1 以下导致发散。
  3. test-time SGD 视角L=12Skv2\mathcal{L} = \frac{1}{2}\|\mathbf{S}\bm{k}-\bm{v}\|^2 的一步梯度下降就是 delta rule,βt\beta_t 是学习率、αt\alpha_t 是 weight decay。
  4. WY 表示的核心事实:CC 个 Householder 的乘积只是秩至多 CC 的修正Pr=Pr1(Iβrkrkr)=Pr1βrPr1krwrkr\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,每多一个 Householder 秩只增 1,于是连乘变求和:P=IWK\mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K}H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 同理。
  5. 矩阵值转移矩阵迫使块内计算变成解三角系统。前两篇的转移量是标量,可以查表;Householder 连乘不行。把递推重写成 Sr=αrSr1+drkr\mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^\intercal 后倒代换,dr\bm{d}_r 的依赖关系构成单位下三角系统,系数矩阵恰是 diag(β)(ΓKK)\operatorname{diag}(\beta)(\Gamma \odot \mathbf{K}\mathbf{K}^\intercal)Γ\Gamma 在这里第二次出现
  6. 两条推导路线互为校验。论文附录 A 的归纳法从闭式 St=iγtγiuiki\mathbf{S}_t = \sum_i\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^\intercal 出发验证,§2.3 的倒代换从递推构造,两者给出同一个量(实测 ut=dt\bm{u}_t = \bm{d}_t,差 1.11×10161.11\times10^{-16})。归纳法第二项的系数 αt+1βt+1\alpha_{t+1}\beta_{t+1} 独立确认了下一条那个必须带的 α\alpha。另注意附录 A 只推首块,跨块要补 βtγtS0kt-\beta_t\gamma_t\mathbf{S}_0^\intercal\bm{k}_t——照抄会出现「首块全对、第二块起错」的症状。
  7. dr\bm{d}_r 的定义里带 αr\alpha_rdr=βrvrαrβrSr1kr\bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r。删除项作用在已衰减的状态上,漏掉 αr\alpha_r 不会报错,α1\alpha \equiv 1 时也看不出来,只在两者都非退化时表现为数值不符。
  8. 三角系统是良态的,这是好消息I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A}) 是单位下三角,行列式恒为 1、永不奇异、不需 pivoting。实测 C=64C = 64β(0,1)\beta \in (0,1) 时条件数不超过 12,fp16 够用。放开到 β(0,2)\beta \in (0,2) 升到 34。
  9. T\mathbf{T} 与其逆都不必物化,直接对右端项做前向替换即可,实测与 inv 差异在 101710^{-17} 量级。省一个 C×CC \times C 中间量。
  10. 代价是并行度。前向替换的 CC 步依赖链无法打破,每步只有 block_DV 个元素并行(32 < 128 线程)。加上累积积的串行前缀和,本篇有两段 CC 步串行–这是 gated delta rule 相比前两篇的固有开销。
  11. CC 被限制在 64。三张 C2C^2 的 f32 fragment(Γ\GammaQK\mathbf{Q}\mathbf{K}^\intercalAmat)在 C=128C = 128 时合计 384 reg/thread,超过 255 上限。
  12. K_s 污染两次是最容易错的地方:构造 Amat 用原始 K\mathbf{K},算右端项时乘 diag(γ)\operatorname{diag}(\gamma),块内项要重载原始 K\mathbf{K},状态更新再乘 γC/γr\gamma^C/\gamma^r。上一篇只污染一次。
  13. 两个退化检验都必须做α1\alpha \equiv 1 回到纯 DeltaNet(实测 8.88×10168.88\times10^{-16})、β0\beta \to 0 回到纯衰减(1.62×10271.62\times10^{-27})。前者让 Γ\Gamma 的指数错误隐身,后者让整个 UT 部分隐身,缺一不可。
  14. 衰减的三处落点三篇未变q=γrq\overleftarrow{\bm{q}} = \gamma^r\bm{q}k=γCγrk\overrightarrow{\bm{k}} = \frac{\gamma^C}{\gamma^r}\bm{k}S=γCS\overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S} 从第一篇到第三篇完全一致,delta rule 只改变块内项右乘的内容。

可迁移的启示:引入一个"看起来只是多一项"的机制,实际代价往往不在 FLOPs 而在依赖结构。delta rule 的算术开销并不大(一个 C×CC\times 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 复述的归纳法证明,原文只考虑首块(S0=0\mathbf{S}_0=\mathbf{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