KDA 的来龙去脉:从线性注意力到 Kimi Delta Attention

0. 出发点:固定大小的记忆

KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。它要解决的问题是把随序列线性膨胀的 KV cache 换成一个固定大小的状态矩阵O(nd)O(n \cdot d)O(d2)O(d^2),上下文翻倍时状态大小不变。

代价是这个 dk×dvd_k \times d_v 的矩阵要承载任意长度的序列,记忆必须可写、可改、可遗忘。演进路线就是一步步补齐这三件事:

  • 线性注意力St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,只能累加(同一个 key 写两次读出的是叠加,即记忆碰撞);
  • Mamba-2:加标量衰减门 St=αtSt1+ktvtS_t = \alpha_t S_{t-1} + k_tv_t^\top,会遗忘但衰减全局统一;
  • DeltaNet:改为差值写入 St=(Iβtktkt)St1+βtktvtS_t = (I - \beta_tk_tk_t^\top)S_{t-1} + \beta_tk_tv_t^\top,支持定点改写;
  • GDN:两者结合 St=αt(Iβtktkt)St1+βtktvtS_t = \alpha_t(I - \beta_tk_tk_t^\top)S_{t-1} + \beta_tk_tv_t^\top
  • KDA:把标量门拆成逐通道门 Diag(αt)\operatorname{Diag}(\bm\alpha_t),各维度独立决定衰减速度,并加数值下界。

线性注意力与 SSM 两条路线的推导见前置篇《线性注意力与 SSM:两条技术路线的完整推导》,本文直接用其结论。

记号约定:两套写法互为转置

文献里状态矩阵有两种摆法,内容等价、互为转置,混用会让推导看起来「中途换了个式子」:

主约定(与 KDA 论文一致) 转置约定(chunkwise 推导常用)
状态形状 StRdk×dvS_t \in \mathbb{R}^{d_k \times d_v} S^tRdv×dk\hat S_t \in \mathbb{R}^{d_v \times d_k}
读出 ot=Stqto_t = S_t^\top q_t ot=S^tqto_t = \hat S_t q_t
写入项 ktvtk_t v_t^\top vtktv_t k_t^\top
擦除算子 左乘 (Iβtktkt)St1(I - \beta_t k_tk_t^\top)S_{t-1} 右乘 S^t1(Iβtktkt)\hat S_{t-1}(I - \beta_tk_tk_t^\top)

关系就是一次转置 S^t=St\hat S_t = S_t^\top。由于擦除算子 Ht=IβtktktH_t = I - \beta_tk_tk_t^\top 对称,转置可以直接穿过它。后文若看到 kvkv^\topvkvk^\top 互换、擦除算子从左跑到右,那是切换了约定,不是等式变了。


1. DeltaNet:从加法到差值写入

线性注意力的 SS 只能叠加,修正这一点正是 DeltaNet 的动机。写入差值,不写全值——写入前先用当前 key 把状态里已存的内容读一遍:

vold=St1kt,ut=βt(vtvold),St=St1+ktutv_{\mathrm{old}} = S_{t-1}^\top k_t, \qquad u_t = \beta_t (v_t - v_{\mathrm{old}}), \qquad S_t = S_{t-1} + k_t u_t^\top

其中 ktk_t 经 L2 归一化(kt2=1\|k_t\|_2 = 1),βt(0,1)\beta_t \in (0,1) 是写入强度。

1.1 差值写入 = 先删后写

上式看不出「擦除」在哪,代入 utu_t 展开即可,每一步只用外积的结合律:

St=St1+kt[βt(vtvold)]=St1+βtktvtβtkt(St1kt) ⁣=St1+βtktvtβtktktSt1=(Iβtktkt)St1删除:沿 kt 方向擦除+βtktvt新值\begin{aligned} S_t &= S_{t-1} + k_t\big[\beta_t(v_t - v_{\mathrm{old}})\big]^\top \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t \big(S_{t-1}^\top k_t\big)^{\!\top} \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t k_t^\top S_{t-1} \\[2pt] &= \underbrace{\big(I - \beta_t k_t k_t^\top\big) S_{t-1}}_{\text{删除:沿 } k_t \text{ 方向擦除}} + \underbrace{\beta_t k_t v_t^\top}_{\text{新值}} \end{aligned}

关键是倒数第二步:voldv_{\mathrm{old}} 自己就是由 St1S_{t-1} 算出来的,所以「减去旧值」必然能写成一个作用在 St1S_{t-1} 上的线性算子,而不是额外的加项。提公因式后 βtktkt-\beta_tk_tk_t^\top 并入单位阵,变成擦除算子。预测误差 vtvoldv_t - v_{\mathrm{old}} 就是 delta,Delta Rule 由此得名

差值不是选的,是推出来的。把记忆看成一个在线回归问题:希望 SktvtS^\top k_t \approx v_t,写成平方损失 L(S)=12Sktvt2\mathcal{L}(S) = \tfrac12\|S^\top k_t - v_t\|^2。它的梯度是

SL=kt(Sktvt) ⁣\nabla_S \mathcal{L} = k_t\,\big(S^\top k_t - v_t\big)^{\!\top}

即「钥匙 ⊗ 残差」——和一维情形 12(wxy)2(wxy)x\tfrac12(wx-y)^2 \Rightarrow (wx-y)x(误差乘输入)一模一样,只是乘法变外积。以 βt\beta_t 为步长走一步 SGD:

St=St1βtkt(St1ktvt) ⁣=(Iβtktkt)St1+βtktvtS_t = S_{t-1} - \beta_t k_t\big(S_{t-1}^\top k_t - v_t\big)^{\!\top} = \big(I - \beta_t k_tk_t^\top\big)S_{t-1} + \beta_t k_tv_t^\top

与前面逐字相同。于是 βt\beta_t 的角色也明确了:它就是学习率。写完立即读一次可以验证:

Stkt=(1βt)St1kt+βtvtS_t^\top k_t = (1-\beta_t)\,S_{t-1}^\top k_t + \beta_t v_t

读出是旧值与新值的凸组合,βt=1\beta_t = 1 时完全覆写。这也解释了 L2Norm 为何是前提:只有 kt=1\|k_t\| = 1 时系数才是干净的插值,否则变成 1βtkt21-\beta_t\|k_t\|^2,可能跌出 [0,1][0,1]

一句话:记忆 = 在线回归 → 平方损失 → 梯度 = 残差 ⊗ 钥匙 → 一步 SGD = 先删后写

1.2 擦除算子 HtH_t:秩 1 与特征值

Ht=IβtktktH_t = I - \beta_t k_tk_t^\top 是全文频率最高的矩阵。kkkk^\top 的第 jj 列是 kjkk_j \cdot k——换 jj 只换倍率、方向永远是 kk,所以它秩为 1。作为变换,

(kk)x=k(kx)=(kx)k(k k^\top)\, x = k\,(k^\top x) = (k^\top x)\cdot k

不管输入什么,输出永远落在 kk 张成的直线上。

放回 HtH_tII 让所有方向原样保留,减去 βtktkt\beta_tk_tk_t^\top 只在 ktk_t 一个方向上动刀,其余 d1d-1 个方向碰都不碰。特征值分两种:

Htkt=(1βtkt2)kt,Htx=x(xkt)H_t k_t = \big(1 - \beta_t\|k_t\|^2\big)k_t, \qquad H_t x = x \quad (\forall\, x \perp k_t)

配合 L2Norm 就是沿 ktk_t 缩放 1βt1-\beta_t、正交补方向为 1,即特征值落在 [1βt,1](0,1][1-\beta_t, 1] \subset (0,1]——不放大、不翻转、不发散,这就是数值稳定性的特征值表述。反例很直白:不归一化时取 k=(1,2)k = (1,2)^\topβ=0.6\beta = 0.6,则 1βk2=21-\beta\|k\|^2 = -2,反复作用必然发散。

特征值之所以到处出现,是因为它回答了迭代系统最关心的问题:一个变换反复作用很多次之后会怎样——在特征向量方向上就是 λn\lambda^nλ>1|\lambda|>1 爆炸、λ<1|\lambda|<1 衰减。这条线上的约束几乎都在围着它转(SSM 的 Aˉ\bar A 模长 1\le 1、GDN/KDA 的 α(0,1)\alpha \in (0,1)HtH_t[1βt,1][1-\beta_t,1]、以及后面 1/Γ1/\Gamma 的溢出)。

HtH_tHouseholder 变换 I2uuI - 2uu^\top 同型:取 βk2=2\beta\|k\|^2 = 2 时两者相同,β(0,1)\beta \in (0,1) 的 delta rule 相当于「没照到底的半面镜子」,代价是不再正交,好处是擦除强度成了可学习的连续量。Householder 及其 WY 表示的线性代数细节见《线性注意力的线性代数前置知识》

1.3 转移矩阵形式与并行的可能性

定义 Ht=IβtktktH_t = I - \beta_tk_tk_t^\top,更新写成

St=HtSt1+βtktvtS_t = H_t\, S_{t-1} + \beta_t k_t v_t^\top

这暴露了 DeltaNet 与 SSM 的同构:HtH_t 是随输入变化的状态转移矩阵(SSM 里是 Aˉ\bar A),βtktvt\beta_tk_tv_t^\top 是写入项(SSM 里是 Bˉxt\bar Bx_t)。记 Bt=βtktvtB_t = \beta_tk_tv_t^\top 逐层代入,递归可以彻底消掉:

St=(i=t1Hi)S0+it(j=ti+1Hj)BiS_t = \Big(\prod_{i=t}^{1} H_i\Big) S_0 + \sum_{i \le t} \Big(\prod_{j=t}^{i+1} H_j\Big) B_i

只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化之一——换成二叉树式两两合并,就能在 O(logn)O(\log n) 深度内并行完成。这是 parallel scan 与 chunkwise 并行化共同的理论根基。


2. GDN:把遗忘门与定点改写拼在一起

Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。核心洞察是两者互补:gating 是板擦(大面积擦除,但无法定点修改),delta rule 是铅笔(定点覆写,但无法快速清空)。

St=αt(Iβtktkt)St1+βtktvtS_t = \alpha_t\big(I - \beta_tk_tk_t^\top\big)S_{t-1} + \beta_t k_t v_t^\top

其中 αt(0,1)\alpha_t \in (0,1) 是数据相关的标量门(α=exp(Softplus(Linear(xt)))\alpha = \exp(-\mathrm{Softplus}(\mathrm{Linear}(x_t))),在 log 空间算以保数值稳定)。三种极限:αt1\alpha_t \to 1 退回 DeltaNet;βt1\beta_t \to 1kk 与已有记忆正交时退回 Mamba-2;αt0\alpha_t \to 0 是整表清零再写入——两者单独都做不到的新能力。

几何上,(Iβkk)(I - \beta kk^\top) 沿 kk 方向压缩状态(定向),α\alpha 把整个状态矩阵均匀缩小(全局),作用于不同自由度,所以可以叠加。

注意作用顺序α\alpha 乘的是整个 (Iβkk)(I-\beta kk^\top)。把 α\alpha 只乘到擦除项上,chunkwise 形式会与递归形式对不上——这是实现时最容易踩的坑。

2.1 统一视角:四个模型是同一个在线优化问题

至此四个模型都出现了,它们是同一个在线优化问题的闭式解,差别只在目标函数。把每步看成一次在线学习:已有 St1S_{t-1},新到样本 (kt,vt)(k_t, v_t),求 StS_t

L(St)=StAtF2正则:记忆保留2Stkt, ut拟合:关联学习\mathcal{L}(S_t) = \underbrace{\|S_t - A_t\|_F^2}_{\text{正则:记忆保留}} - \underbrace{2\langle S_tk_t,\ u_t\rangle}_{\text{拟合:关联学习}}

正则项惩罚状态的改动量,锚点 AtA_tSt1S_{t-1} 表示尽量不动、取 αtSt1\alpha_tS_{t-1} 表示容忍遗忘;拟合项要求用 ktk_t 检索的结果朝写入目标 utu_t 对齐。目标对 StS_t 是二次的,=2(StAt)2utkt=0\nabla = 2(S_t - A_t) - 2u_tk_t^\top = 0,闭式解统一为 St=At+utktS_t = A_t + u_tk_t^\top。于是差别完全归结为两个选择:锚点决定怎么遗忘,写入目标决定怎么写入(下表用转置约定)。

模型 锚点 AtA_t 写入目标 utu_t 闭式解
Linear Attn St1S_{t-1} vtv_t St1+vtktS_{t-1} + v_tk_t^\top
Mamba-2 αtSt1\alpha_tS_{t-1} vtv_t αtSt1+vtkt\alpha_tS_{t-1} + v_tk_t^\top
DeltaNet St1S_{t-1} βt(vtSt1kt)\beta_t(v_t - S_{t-1}k_t) St1(Iβtktkt)+βtvtktS_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top
GDN αtSt1\alpha_tS_{t-1} βt(vtαtSt1kt)\beta_t(v_t - \alpha_tS_{t-1}k_t) St1(αt(Iβtktkt))+βtvtktS_{t-1}\big(\alpha_t(I-\beta_tk_tk_t^\top)\big) + \beta_tv_tk_t^\top
KDA St1Diag(αt)S_{t-1}\mathrm{Diag}(\bm\alpha_t) βt(vtSt1Diag(αt)kt)\beta_t(v_t - S_{t-1}\mathrm{Diag}(\bm\alpha_t)k_t) St1Diag(αt)(Iβtktkt)+βtvtktS_{t-1}\mathrm{Diag}(\bm\alpha_t)(I-\beta_tk_tk_t^\top) + \beta_tv_tk_t^\top

说到底,这就是把「kvk \to v 这条记忆是否还对得上」写成损失函数:对不上就修(拟合项),但别为了修这一条把整张表推翻(正则项)。GDN 在两处都取强化版本,KDA 再把锚点的标量收缩换成逐通道的 Diag(αt)\mathrm{Diag}(\bm\alpha_t)

2.2 Chunkwise 并行:WY 表示与 UT 变换

推理用递归(O(1)O(1)/token),但训练和 prefill 必须并行。设 chunk 大小 CC、入口状态 S0S_0,部分展开递归(转置约定):

Sr=γrS0Pr+i=1rγrγiu~iki,γr=jrαj,Pr=ir(Iβikiki)S_r = \gamma_r S_0 P_r + \sum_{i=1}^r\frac{\gamma_r}{\gamma_i}\,\tilde u_i k_i^\top, \qquad \gamma_r = \prod_{j\le r}\alpha_j,\quad P_r = \prod_{i\le r}(I - \beta_ik_ik_i^\top)

PrP_rCCdk×dkd_k\times d_k 矩阵逐个相乘,O(Cdk3)O(C\,d_k^3) 且严格顺序——比原递推还贵,必须处理。

关键观察:这种乘积永远不膨胀。每个因子都是「单位阵减秩 1」,连乘结果仍是「单位阵减一个低秩矩阵」:

Pr:=i=1r(Iβikiki)=IirwikiP_r := \prod_{i=1}^{r}(I - \beta_i k_i k_i^\top) = I - \sum_{i\le r} w_i k_i^\top

证明(对 rr 归纳)P0=IP_0 = I 成立。设 Pr1=Ii<rwikiP_{r-1} = I - \sum_{i<r} w_ik_i^\top,右乘第 rr 个因子:

Pr=(Ii<rwiki)βr(Ii<rwiki)krkrP_r = \Big(I - \sum_{i<r} w_i k_i^\top\Big) - \beta_r\Big(I - \sum_{i<r} w_i k_i^\top\Big)k_r k_r^\top

展开后第三、四项都以 krk_r^\top 结尾,合并同类项:

Pr=Ii<rwikiβr(kri<rwi(kikr))=wrkrP_r = I - \sum_{i<r} w_i k_i^\top - \underbrace{\beta_r\Big(k_r - \sum_{i<r} w_i (k_i^\top k_r)\Big)}_{=\,w_r} k_r^\top

归纳完成。wrw_r 不是发明的技巧,而是「乘积保持低秩」逼出来的——要让结果保持 IiwikiI - \sum_i w_ik_i^\top 的形式,括号里那一坨只能是 wrw_r。它的直觉是修正后的擦除向量:第 rr 步本想擦除 krk_r 方向,但若 krk_r 与之前的 kik_i 有重叠,连乘展开时前面的擦除已经顺带擦过这部分,wrw_r 把已擦的量减掉以避免重复擦除,权重恰是重叠度 kikrk_i^\top k_r。value 侧同理,只多一个衰减比:

u~r=βr(vri<ru~iγrγi(kikr))\tilde u_r = \beta_r\Big(v_r - \sum_{i<r}\tilde u_i\,\tfrac{\gamma_r}{\gamma_i}(k_i^\top k_r)\Big)

γ\gamma 只出现在 u~\tilde u 里而不在 ww 里,因为 GDN 的衰减是标量、与一切矩阵可交换,擦除连乘里的衰减可整体外提成 γr\gamma_r;但 value 侧每次写入的「存活时长」不同,这个相对衰减无法外提。对照 KDA:衰减变成向量后与擦除不可交换,γ\gamma 再也提不出去,只能渗进内积本身。

UT 变换:递归变成一次下三角求解wrw_r 只依赖 wi<rw_{i<r},是严格下三角依赖。把递归移项 wr+i<rβr(kikr)wi=βrkrw_r + \sum_{i<r}\beta_r(k_i^\top k_r)w_i = \beta_rk_r,令 WW 的第 rr 行为 wrw_r^\topL=strictLower(diag(β)KK)L = \mathrm{strictLower}(\mathrm{diag}(\beta)KK^\top),则 CC 个方程一次写成 (I+L)W=diag(β)K(I+L)W = \mathrm{diag}(\beta)K

W=(I+L)1diag(β)KW = (I+L)^{-1}\,\mathrm{diag}(\beta)\,K

求这个逆很便宜,因为严格下三角矩阵幂零LC=0L^C = 0),Neumann 级数有限项精确截断:(I+L)1=IL+L2(I+L)^{-1} = I - L + L^2 - \cdots,一次前代法即可,O(C2)O(C^2)(I+L)(I+L)单位下三角(Unit Triangular,对角为 1 因为 wrw_r 完整依赖自己),这就是「UT」的来源。所以 UT 不是额外发明的东西,它就是这两条递归的矩阵形态

乘法链 → 加法链,这才是并行的真正来源。本质是把擦除矩阵的连乘 r(Iβrkrkr)\prod_r(I - \beta_rk_rk_r^\top) 换成求和 IrwrkrI - \sum_r w_rk_r^\top,代价是求和项带修正。求和为什么就是胜利:加法可交换、可结合,因而可任意分组——树形归约、分块、稠密 matmul,GPU 的全部并行性都在奖励求和结构;而连乘必须一步一步来。同样的手法在这条路线上反复出现:αs=exp(gs)\prod\alpha_s = \exp(\sum g_s)(累积衰减变 cumsum)、SSM 的卷积形式、Mamba-2 的 SSD 半可分矩阵,以及本节的 WY-UT。

辨析:因果掩码 ≠ UT 变换。两者都是下三角,但管的是两件事:因果掩码把 score 的严格上三角置零,管「不许看未来」;UT 变换处理块内历史写入之间的相互影响——delta rule 每次写入都「先读再改」,块内第 2 次写入读到了第 1 次的结果,UT 负责解耦。都是下三角不是巧合,是同一个原因:因果性使干扰系数矩阵天然下三角,对角线天然为 1。因果性决定了它是三角的,但做它的目的是解耦,不是掩码。


3. 从 GDN 到 KDA

维度 GDN KDA(→ K3) 动机
更新式 St1(αt(Iβkk))+βvkS_{t-1}(\alpha_t(I-\beta kk^\top)) + \beta vk^\top (Iβkk)Diag(αt)St1+βvk(I-\beta kk^\top)\mathrm{Diag}(\bm\alpha_t)S_{t-1} + \beta vk^\top α 标量 → 逐通道向量
门粒度 标量(整头同一衰减率) 向量 αt(0,1)dk\bm\alpha_t\in(0,1)^{d_k} 长期记忆通道与短期工作区分离
作用顺序 α 在外乘整个更新 Diag(α) 在内侧先衰减 S 与逐通道参数化配套
α 参数化 负 Softplus,(,0)(-\infty,0) 缩放 sigmoid,下界 gmin=5g_{\min}=-5 1/Γ1/\Gamma 有界 → 全走 Tensor Core
chunkwise Γ 是标量比 Γ 是向量累积比 表达力↑ 数值难度↑

3.1 KDA 的递推公式

回到主约定(StRdk×dv\mathbf{S}_t \in \mathbb{R}^{d_k\times d_v},读出 Stqt\mathbf{S}_t^\top\bm q_t),与论文写法一致:

St=(Iβtktkt)Diag(αt)St1+βtktvt\mathbf{S}_t = \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) \operatorname{Diag}(\bm{\alpha}_t)\, \mathbf{S}_{t-1} + \beta_t \bm{k}_t \bm{v}_t^\top

按从右往左三步:通道衰减 Diag(αt)St1\operatorname{Diag}(\bm\alpha_t)\mathbf{S}_{t-1} 对每一行分别乘对应通道的 α\alpha方向擦除 (Iβtktkt)(\mathbf{I} - \beta_t\bm k_t\bm k_t^\top) 做定点擦除;写入 βtktvt\beta_t\bm k_t\bm v_t^\top

标量门的瓶颈在于所有 key 通道以同一比例衰减——要么一起记住,要么一起忘记。但有的通道在存「长期主题」(希望 α1\alpha \approx 1),有的在存「临时指针」(希望快速腾出容量)。Diag(αt)\mathrm{Diag}(\bm\alpha_t) 相当于给每一列配一个独立的遗忘速度,GDN 是 Diag(αt)=αtI\mathrm{Diag}(\bm\alpha_t) = \alpha_t I 的特例。在线学习视角下,这就是把 §2.1 的正则项换成逐通道版:第 jj 列的信任区域宽度正比于 αt(j)\alpha_t^{(j)}α\alpha 小的通道旧记忆先被压缩、新写入几乎没有阻力。KDA = 逐通道权重衰减 + delta rule。

作用顺序不可颠倒Dt=Diag(αt)D_t = \mathrm{Diag}(\bm\alpha_t)(Iβkk)(I - \beta kk^\top) 不可交换,写成 St1(Iβkk)DtS_{t-1}(I-\beta kk^\top)D_t 或只衰减单位阵部分都是错的。擦除时检索用的也是已衰减的状态。

3.2 逐通道衰减下的 chunkwise

核心记号是累积 log 衰减 γi=sigsRdk\gamma_i = \sum_{s\le i} g_s \in \mathbb{R}^{d_k}。带衰减的 key-key 内积:

Mci=(kceγcγi)ki,1i<cCM_{ci} = \big(k_c \odot e^{\gamma_c - \gamma_i}\big)^\top k_i, \qquad 1 \le i < c \le C

含义是 kik_i 写入的记忆衰减到第 cc 步时与 kck_c 的重叠程度。衰减「长」在 M 内部,无法外提——这是向量门控带来的结构性变化,也是 KDA chunkwise 推导的核心难点。其余步骤与 GDN 同构:构造 Lci=βiMciL_{ci} = \beta_i M_{ci},求 T=(I+L)1T = (I+L)^{-1},得 A^=Tdiag(β)\hat A = T\operatorname{diag}(\beta),然后

W=A^(eγiki)i=1..C,U=A^V,V~=UWS[0]W = \hat{A}\,\big(e^{\gamma_i} \odot k_i\big)_{i=1..C}, \qquad U = \hat{A}\,V, \qquad \tilde{V} = U - W\,S_{[0]}^\top

V~\tilde V伪值:它不是真实的 vv,而是扣除了「历史记忆已有部分 + 块内其他位置已写部分」之后可直接累加、无需再修正的净增量。这与单步 delta rule 的 vtSt1ktv_t - S_{t-1}^\top k_t 一致,V~\tilde V 是它的 chunk 并行版。输出与出口状态:

Acjqk=(qceγcγj)kj,oc=S[0](qceγc)+jcAcjqkv~jA^{qk}_{cj} = \big(q_c \odot e^{\gamma_c - \gamma_j}\big)^\top k_j, \qquad o_c = S_{[0]}\,(q_c \odot e^{\gamma_c}) + \sum_{j \le c} A^{qk}_{cj}\,\tilde v_j

S[C]=S[0]ΓC0+c=1Cv~c(kceγCγc)S_{[C]} = S_{[0]}\,\Gamma_{C\leftarrow 0} + \sum_{c=1}^{C} \tilde v_c \big(k_c \odot e^{\gamma_C - \gamma_c}\big)^\top

所有指数都是 0\le 0 的差,整条链路没有一个 1\ge 1 的因子——这是刻意保持的,原因见 §3.4。

3.3 数值验算(dk=dv=2d_k = d_v = 2C=2C = 2

零初始状态,每步恒定衰减 α=(0.5,0.25)\bm\alpha = (0.5, 0.25)β=1\beta = 1

k1=(11), v1=(20);k2=(12), v2=(02);q1=q2=(11)k_1 = \binom{1}{1},\ v_1 = \binom{2}{0};\quad k_2 = \binom{1}{2},\ v_2 = \binom{0}{2};\quad q_1 = q_2 = \binom{1}{1}

递推式。第 1 步(S0=0S_0 = 0):S1=v1k1=(2200)S_1 = v_1k_1^\top = \begin{pmatrix}2&2\\0&0\end{pmatrix}o1=(4,0)o_1 = (4,0)。第 2 步先衰减 S1D=(10.500)S_1D = \begin{pmatrix}1&0.5\\0&0\end{pmatrix}(第 1 列 ×0.5、第 2 列 ×0.25——逐通道在动),再擦除写入:

S1D(Ik2k2)=(13.500),S2=(13.524),o2=(4.5, 6)S_1D(I - k_2k_2^\top) = \begin{pmatrix}-1&-3.5\\0&0\end{pmatrix}, \qquad S_2 = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix}, \qquad o_2 = (-4.5,\ 6)

Chunkwiseeγ1=(0.5,0.25)e^{\gamma_1} = (0.5, 0.25)eγ2γ1=(0.5,0.25)e^{\gamma_2-\gamma_1} = (0.5, 0.25),衰减 KKT M21=(0.5,0.5)(1,1)=1M_{21} = (0.5, 0.5)\cdot(1,1) = 1。UT:L=(0010)L = \begin{pmatrix}0&0\\1&0\end{pmatrix}T=ILT = I - LA^=T\hat A = T。WY:

W=A^(0.50.250.250.125)=(0.50.250.250.125),U=A^(2002)=(2022)W = \hat A\begin{pmatrix}0.5&0.25\\0.25&0.125\end{pmatrix} = \begin{pmatrix}0.5&0.25\\-0.25&-0.125\end{pmatrix}, \quad U = \hat A\begin{pmatrix}2&0\\0&2\end{pmatrix} = \begin{pmatrix}2&0\\-2&2\end{pmatrix}

伪值(S[0]=0V~=US_{[0]}=0 \Rightarrow \tilde V = U):v~1=(2,0)\tilde v_1 = (2,0)v~2=(2,2)\tilde v_2 = (-2,2)注意 v~2v2\tilde v_2 \ne v_2:因为 k2k_2 与衰减后的 k1k_1 写入重叠(M21=1M_{21} = 1),WY 把 v2v_2 修正为扣除重叠后真正的新增。衰减注意力与输出:

Aqk=(200.753),o1=2v~1=(4,0) ,o2=0.75v~1+3v~2=(4.5, 6) A^{qk} = \begin{pmatrix}2&0\\0.75&3\end{pmatrix}, \qquad o_1 = 2\tilde v_1 = (4,0)\ \checkmark, \qquad o_2 = 0.75\,\tilde v_1 + 3\,\tilde v_2 = (-4.5,\ 6)\ \checkmark

块末状态 S[2]=v~1(k1eγ2γ1)+v~2k2=(13.524)S_{[2]} = \tilde v_1(k_1 \odot e^{\gamma_2-\gamma_1})^\top + \tilde v_2 k_2^\top = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix} \checkmark,与递推式逐项一致。

3.4 下界衰减:1/Γ1/\Gamma 溢出

GDN 的标量衰减在 chunkwise 里只以比值出现(s=j+1iαs1\prod_{s=j+1}^i \alpha_s \le 1),天然安全。KDA 把 Γ\Gamma 变成向量累积后,某些等价写法(如论文公式 (4) 的 K/ΓK/\Gamma 因式分解)就得真的算出

Γt1=exp(stgs)(逐通道)\Gamma_t^{-1} = \exp\Big(-\sum_{s \le t} g_s\Big) \quad \text{(逐通道)}

衰减越快的通道,1/Γ1/\Gamma 呈指数膨胀:恒定 α=0.5\alpha = 0.5t=500t=500 已达 1015010^{150}t=2000t=2000 溢出 float64;BF16 训练下数十步即溢出。标量推广为向量后衰减率的动态范围被放大,溢出由个别情形变为普遍现象。

Kimi Linear 的规避方案是在对数空间算相对衰减、并把 chunk 再切成 16 token 二级瓦片:瓦片之间可交给 Tensor Core,但瓦片内部仍需按位置对显式计算,成为块内瓶颈。K3 的解决方式是不改公式、改参数化,用有界缩放 sigmoid 使溢出在数学上不可能发生:

gth=gminSigmoid(eAhzth)(gmin,0),gmin=5g_t^h = g_{\min} \cdot \text{Sigmoid}(e^{A_h} z_t^h) \in (g_{\min}, 0), \qquad g_{\min} = -5

于是 α>e56.7×103\alpha > e^{-5} \approx 6.7\times10^{-3},16 token 块的累积对数衰减落在 (80,0)(-80, 0),重缩放因子小于 e80e^{80},在 BF16 动态范围内。表达力不受损:快通道在瓦片内仍可衰减到 e801035e^{-80} \approx 10^{-35},长期遗忘靠多步连乘即可,不需要单步清零的能力。以一个衰减下界换取全链路 Tensor Core 化,是有利的取舍。


4. 网络里的其他部件

q,k,v,α,βq, k, v, \alpha, \beta 本身怎么来、o~\tilde o 怎么变成层输出:GDN/KDA 采用 Llama 式宏架构,把 attention 换成 gated delta token mixer。

1
2
3
4
5
6
x ─┬─ W_q/W_k ─ ShortConv ─ SiLU ─ L2Norm ──> q, k
├─ W_v ──── ShortConv ─ SiLU ───────────> v
├─ W_α(线性投影, log 空间)──────────────> α
├─ W_β(线性投影 + Sigmoid)─────────────> β
│ 递归/分块计算 gated delta rule ──────> õ
└─ y = W_o( Sigmoid(W_g x) ⊙ RMSNorm(õ) ), 再接 SwiGLU MLP + 残差

三个部件各负责一部分数值稳定性:

  • ShortConv:窗口 4 的因果 depthwise 卷积,ShortConv(x)t=j<Wwjxtj\mathrm{ShortConv}(x)_t = \sum_{j<W} w_j \odot x_{t-j}。补足线性 RNN 缺失的局部建模能力——递归状态只携带压缩后的全局历史,最近几个 token 的精细模式由这层负责。depthwise 使通道间不混,与逐通道衰减语义配套;
  • SiLU(即 β=1\beta=1 的 Swish):xσ(x)x \cdot \sigma(x),自门控、免参数,平滑非单调、梯度处处非零;
  • L2Norm:只用在 Q/K。kt=1\|k_t\|=1 是 delta rule 擦除稳定的前提(§1.2);Q 归一化使 QKQK^\top 的 score 限制在 [1,1][-1,1]。V 不做,因为写入内容的幅值本身携带信息。

K3 另将输出门升级为输入相关的满秩 sigmoid 门 yt=Wo[Sigmoid(Wgxt)RMSNorm(o~t)]y_t = W_o[\mathrm{Sigmoid}(W_gx_t) \odot \mathrm{RMSNorm}(\tilde o_t)],允许每个 token 独立调制从循环状态读取的通道。

整体骨架自始至终没变:

状态×(衰减/删除算子)+(写入项)\text{状态} \times (\text{衰减/删除算子}) + (\text{写入项})


参考