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}

比上一篇多出的是 βt\beta_t–一个删除力度的系数:它决定在 kt\bm{k}_t 这个方向上,旧内容被擦掉多少。

上两篇见《Chunked 线性注意力》与《标量衰减》,本文沿用其符号约定与四层参考验证框架。本篇会用到 Householder 变换、WY 表示与 UT 变换这些线性代数工具,背景推导见《线性注意力的线性代数前置知识》。


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} 定向替换 + 整体遗忘

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。

值得一提的是,这个视角和 SSM 那边的推导是一致的。SSM 从连续状态方程出发,经 ZOH 离散化得到 Aˉ=eΔA\bar{A} = e^{\Delta A},在 Mamba-2 的标量情形下就是 αt=eΔta\alpha_t = e^{-\Delta_t a}–那个“离散步长” Δt\Delta_t 扮演的角色,和这里学习率与 weight decay 的角色是同一个。换句话说,两条路线都把递推写成“状态 ×(衰减/删除算子)+(写入项)”:SSM 从微分方程推出衰减项,delta rule 从优化目标推出删除项,最后落到同一个递推骨架上。详细对照见《线性注意力与 SSM:两条技术路线的完整推导》。


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)

G\mathbf{G} 同样可以抽出标量。第 ii 项尾巴上的衰减是 j=i+1rαj=γr/γi\prod_{j=i+1}^{r}\alpha^j = \gamma^r/\gamma^i,剩下的 Householder 尾巴记作 Pi+1:r\mathbf{P}^{i+1:r}(约定 Pr+1:r=I\mathbf{P}^{r+1:r} = \mathbf{I}):

G[t]r=i=1rγrγiβivikiPi+1:r,Pi+1:r=j=i+1r(Iβjkjkj)\mathbf{G}^r_{[t]} = \sum_{i=1}^{r}\frac{\gamma^r}{\gamma^i}\,\beta^i\bm{v}^i\bm{k}^{i\intercal}\,\mathbf{P}^{i+1:r}, \qquad \mathbf{P}^{i+1:r} = \prod_{j=i+1}^{r}\big(\mathbf{I} - \beta^j\bm{k}^j\bm{k}^{j\intercal}\big)

两者形式上很像:F\mathbf{F} 是一条从头到尾的 Householder 连乘,G\mathbf{G}CC 条长度不同的连乘尾巴的加权和(第 ii 项的尾巴从 i+1i+1 开始)。写成这个样子后衰减已经全部变成标量系数,剩下的 P\mathbf{P}Pi+1:r\mathbf{P}^{i+1:r} 是纯 DeltaNet 的部分,靠 WY 表示处理–这正是论文 §2.2 的内容,下一节完整展开(到那里 α1\alpha\equiv1G\mathbf{G} 就退化成 H\mathbf{H})。

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

换个角度看Γ\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 门控版的结果

把 §2.2.6 那张表的替换代进去,就得到门控版的 chunkwise 算法。先把递推重写成加法形式 Sr=αrSr1+drkr\mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^{\intercal},其中修正量 dr=β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–删除项作用在已经衰减过的状态上,漏掉不会报错、只会算错)。dr\bm{d}_r 依赖所有 di (i<r)\bm{d}_i\ (i<r),构成一个单位下三角系统:

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

ΔRC×dv\Delta \in \mathbb{R}^{C\times d_v} 的第 rr 行就是 dr\bm{d}_r衰减感知掩码 Γ\Gamma 在这里第二次出现,这次是嵌在三角系统的系数矩阵里。解出 Δ\Delta 用 UT 变换 T=[I+strictLower(A)]1diag(β)\mathbf{T} = [\mathbf{I}+\operatorname{strictLower}(\mathbf{A})]^{-1}\operatorname{diag}(\beta) 即可,于是(论文式 11、12):

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。

系数矩阵 I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A})单位下三角,行列式恒为 1、永不奇异、不需要 pivoting,求逆退化成一次前代(forward substitution),β(0,1)\beta \in (0,1) 时 fp16 够用;完整推导与计算量分析见《线性注意力的线性代数前置知识》§4.4。对 kernel 来说只需要记住一件事:T\mathbf{T} 不必物化,把前代直接作用在右端项 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 与前两篇最本质的区别,下面 §4 会看到它如何限制并行度。


3. 四层参考实现

沿用前两篇的框架:

参考 实现方式 验证目标
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.2 的 WY/UT 表示、§2.3 的三角系统与式 (11)(12)
C kernel 控制流镜像(切 DV、S\mathbf{S}^\intercal 布局、前向替换) kernel 结构
D 退化检验(α1\alpha \equiv 1 / β0\beta \to 0 与前两篇的一致性

3.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

3.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

3.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衰减权重的三处落点完全没变

3.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 里的做法一致(不物化对角矩阵)。

3.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},后者省一次对角矩阵构造。


4. TileLang kernel

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

4.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。

4.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。

4.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 库实际采用的做法,实现复杂度显著上升。

还有第三条路:干脆把 T\mathbf{T} 物化出来。本篇不物化是因为 T\mathbf{T} 只用一次;而论文 §3.3 的标准形式要用同一个 T\mathbf{T}W[t]\mathbf{W}_{[t]}U~[t]\widetilde{\mathbf{U}}_{[t]} 两个量,不物化就得做两次前代,于是值得单起一个 kernel 求逆。分块递归可以把 C=32C = 32 的 31 步依赖链压到 15 步,代价是算术量 1.88 倍–完整实现、bank conflict 处理与数值验证见《TileLang 实战:单位下三角矩阵求逆》。

4.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 Δ

4.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 在循环内

4.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 写入

5. 三篇对照

线性注意力 标量衰减 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 论文那套箭头记号的价值:它把「衰减」这件事隔离成了一个可以独立理解的层面。


6. 总结

  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 的推导只考虑首块(S0=0\mathbf{S}_0=\mathbf{0}),实际 kernel 里每块入口状态非零,修正量要补上 αrβrS[t]-\alpha_r\beta_r\mathbf{S}_{[t]} 那一项(即右端项里的 diag(γ)KS[t]-\operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}_{[t]}^\intercal)——照抄会出现「首块全对、第二块起错」的症状。
  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,求逆只是一次前代,β(0,1)\beta \in (0,1) 时用 fp16 也够;而且 T\mathbf{T} 不必物化,直接把前代作用在右端项的行上即可,省一个 C×CC \times C 中间量。
  9. 代价是并行度。前向替换的 CC 步依赖链无法打破,每步只有 block_DV 个元素并行(32 < 128 线程)。加上累积积的串行前缀和,本篇有两段 CC 步串行–这是 gated delta rule 相比前两篇的固有开销。
  10. CC 被限制在 64。三张 C2C^2 的 f32 fragment(Γ\GammaQK\mathbf{Q}\mathbf{K}^\intercalAmat)在 C=128C = 128 时合计 384 reg/thread,超过 255 上限。
  11. K_s 污染两次是最容易错的地方:构造 Amat 用原始 K\mathbf{K},算右端项时乘 diag(γ)\operatorname{diag}(\gamma),块内项要重载原始 K\mathbf{K},状态更新再乘 γC/γr\gamma^C/\gamma^r。上一篇只污染一次。
  12. 两个退化检验都必须做α1\alpha \equiv 1 回到纯 DeltaNet(实测 8.88×10168.88\times10^{-16})、β0\beta \to 0 回到纯衰减(1.62×10271.62\times10^{-27})。前者让 Γ\Gamma 的指数错误隐身,后者让整个 UT 部分隐身,缺一不可。
  13. 衰减的三处落点三篇未变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)给出门控版的归纳法证明,原文只考虑首块(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