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

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

KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。本文不铺开 K3 的全部架构,只讲一条演化线:KDA 之前的几个模型(线性注意力、Mamba、DeltaNet/GDN)分别解决了什么问题,KDA 又在它们的基础上改了什么

故事线很直白:线性注意力先把「固定大小的记忆」造出来,但这个记忆只会往里堆、不会清理;Mamba 教会它遗忘,但遗忘只能整片擦;DeltaNet 让它能精准改写某条记忆;GDN 把「擦除」和「改写」两样拼在一起;KDA 最后一步,让每个维度各自决定忘得多快。每一步都有公式推导、手动验算和工程上的理由。

0. 先算一笔账:为什么需要「固定大小的 KV 记忆」

标准 softmax 注意力在生成文本时,必须把每个历史 token 的 Key 和 Value 都存下来(KV cache)–因为每生成一个新 token,都要和前面所有 token 重新算一遍注意力。序列越长,这笔账越贵:

开销 softmax 注意力(KV cache) 固定状态(线性注意力一族)
解码显存 O(nd)O(n \cdot d)(全部历史的 KV 都得缓存) O(d2)O(d^2)(一个固定矩阵 S,与 n 无关)
单步解码 O(nd)O(n \cdot d)(扫全部历史) O(d2)O(d^2)(更新 S + 读出,两次矩阵乘)
预填充 O(n2d)O(n^2 d) O(nd2)O(n d^2)
上下文长度翻倍 显存、延迟一起翻倍 完全不变

(单头视角,d 为特征维度。)

这笔账具体有多痛:几十万 token 的上下文就能把 KV cache 撑到几十 GB。长上下文场景(长文档、agent 记忆、仓库级代码理解)里,第一瓶颈往往不是模型权重,而是 KV cache 的显存和带宽。K3 把「序列维度缩放」列为头号工程目标,93 层里 69 层换成 KDA,买的就是右列这张表。

但固定大小不是白来的。一个 dv×dkd_v \times d_k 的矩阵要装下任意长的序列,记忆就必须会写、会改、会忘–只进不出的仓库,序列一长就糊成一团。怎么设计这块「会自我管理的固定记忆」,就是本文主线,按出场顺序:

  • 线性注意力:先把固定状态本身造出来(用外积累积出一个矩阵 S),但只会做加法,只进不出;
  • Mamba/SSM:加上一个全局衰减系数 α–会忘了,但只能整片擦;
  • DeltaNet:加上定点改写(delta rule)–会改了,写入的是新旧值之差;
  • GDN:把「板擦」(α)和「铅笔」(delta)拼在一起,先删后写;
  • KDA:把一个总的衰减系数拆成每个维度一个,各维度自己决定记多久,再加固数值下界,进 K3。

线性注意力和 SSM 这两条技术路线的完整推导–softmax 为什么挡路、核化怎么搬、积分因子法、ZOH 离散化、卷积形式、S4 到 Mamba-2 的演化–单独成了一篇:《线性注意力与 SSM:两条技术路线的完整推导》。本文只要求读者记得结论,主线是:每个模型解决前一个的什么问题、留下什么新问题,讲完之后再回头看总表。


1. KDA 是什么:先看递推式

KDA 是 K3 序列维度缩放的核心,基于 delta rule 循环,并加上了通道维度的遗忘门

路线一:线性注意力–固定状态是怎么来的(概述)

线性注意力的完整推导(softmax 为什么挡路、核函数怎么把非线性搬走、结合律怎么把 O(n2)O(n^2) 降到 O(n)O(n))在前置篇《线性注意力与 SSM:两条技术路线的完整推导》里一步步展开,这里只留骨架:

  • 去掉 softmax,(QK)VQ(KV)(QK^\top)V \to Q(K^\top V):换一种括号方式,先算 KVK^\top V,复杂度 O(dn2)O(nd2)O(dn^2) \to O(nd^2);
  • 核化 sim(q,k)=ϕ(q)ϕ(k)\mathrm{sim}(q,k) = \phi(q)^\top\phi(k):非线性提前到 Q/K 各自身上,n×nn\times n 的注意力矩阵就消失了,换成一个固定大小的状态 SRd×dS \in \mathbb{R}^{d\times d};
  • 状态可以逐步累积:St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,每来一个 token 更新一次、读一次(fast weights 视角:S 是被逐 token 编程的快速权重,查询是读出)。

代价也在这三行里:S 只会加法。写入只有「往上堆」,没有「删掉」–序列长度远超状态容量时,不同 kvk \to v 的关联互相干扰,旧信息永远赖在状态里,新信息覆盖不掉。前置篇里有个数值演示:同一个 key 写两次,线性注意力读出来的是两个 value 的叠加,而 delta rule 读出来的是新的那个。这个「记忆碰撞」就是线性注意力最大的问题,修复它是 DeltaNet 的全部动机。

DeltaNet:从加法到差值写入

线性注意力留下的 S 是个只会「贴便利贴」的状态,修正它就是 DeltaNet 的全部动机。解法:写入差值,不写全值。写入前,先用当前 key 把状态里已存的内容读一遍:

vold=St1kt(当前 key 指向的旧内容),ut=βt(vtvold)v_{\mathrm{old}} = S_{t-1}^\top k_t \quad (\text{当前 key 指向的旧内容}), \qquad u_t = \beta_t (v_t - v_{\mathrm{old}})

St=St1+ktut=(Iβtktkt)St1+βtktvtS_t = S_{t-1} + k_t u_t^\top = (I - \beta_t k_t k_t^\top)\, S_{t-1} + \beta_t k_t v_t^\top

符号 形状 含义
voldv_{\mathrm{old}} [dv][d_v] 当前 key 从状态里读出的已有内容
βt\beta_t 标量 (0,1)\in (0,1) 内容替换强度
utu_t [dv][d_v] 实际写入的差值

展开后两项分解:删除项 (Iβtktkt)(I - \beta_t k_t k_t^\top) 沿 ktk_t 方向抹掉旧内容,再加回新值项 βtktvt\beta_t k_t v_t^\top–先删旧、再写新,一次外积同时完成两件事。预测误差 vtvoldv_t - v_{\mathrm{old}} 就是 delta,Delta Rule 由此得名:每次写入的从来不是新值本身,而是新值与旧值之差。

为什么恰好是差值?因为它就是对回归损失做一步 SGD 的结果。注意 DeltaNet 没有动线性注意力的两个前置条件–ϕ\phi、外积状态、递推形式全部保留(实现里常取 $\phi = $ L2Norm,即下文单位球约束),它改的只是写入规则。而且这个规则不是拍脑袋的设计,它可以从回归损失严格推出来:

L(S)=12Sktvt2,SL=(Sktvt)kt\mathcal{L}(S) = \tfrac{1}{2}\|S k_t - v_t\|^2, \qquad \nabla_S\mathcal{L} = (Sk_t - v_t)k_t^\top

St=St1βtSL=St1(Iβtktkt)先删:擦掉 St1kt 方向的旧值+βtvtkt后写S_t = S_{t-1} - \beta_t\nabla_S\mathcal{L} = \underbrace{S_{t-1}(I - \beta_tk_tk_t^\top)}_{\text{先删:擦掉 } S_{t-1}k_t \text{ 方向的旧值}} + \underbrace{\beta_tv_tk_t^\top}_{\text{后写}}

βt\beta_t 就是学习率。论文原版写法还提供了一个插值视角:St=St1vtoldkt+vtnewktS_t = S_{t-1} - v_t^{\mathrm{old}}k_t^\top + v_t^{\mathrm{new}}k_t^\top,其中 vtnew=βtvt+(1βt)St1ktv_t^{\mathrm{new}} = \beta_tv_t + (1-\beta_t)S_{t-1}k_t–δ-rule 是软覆写,β 控制新旧比例(β=1 且 k 为单位向量时完全覆写)。q/k 做 L2 归一化(钉在单位球上)保证 kkkk^\top 特征值 1\le 1,更新稳定。

直观类比:state 是白板,key 是指针。纯加性写入是往白板上不停贴便利贴,贴满了就糊;DeltaNet 是先擦掉指针 ktk_t 指向的那块,再贴新的。βt=0\beta_t = 0:完全不写入,状态不动;βt=1\beta_t = 1:完全替换,指针指向的内容整个换掉。这种「先擦后写」的机制让状态能覆盖错误记忆,这是 GDN/DeltaNet 一线的核心改进。

关键整理:写成转移矩阵形式。定义 Ht=IβtktktH_t = I - \beta_t k_t k_t^\top,状态更新变成

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

这个形式之所以非常关键,是因为它暴露了 DeltaNet 和 SSM 的同构:HtH_t 正是随输入变化的状态转移矩阵(SSM 里是 Aˉ=eΔA\bar A = e^{\Delta A}),βtktvt\beta_t k_t v_t^\top 正是写入项(SSM 里是 Bˉxt\bar B x_t)。

把递归展开,消掉时间依赖。记 Bt=βtktvtB_t = \beta_t k_t v_t^\top,递推 St=St1Ht+BtS_t = S_{t-1}H_t + B_t 逐层代入:

S1=S0H1+B1S_1 = S_0 H_1 + B_1

S2=S0H1H2+B1H2+B2S_2 = S_0 H_1 H_2 + B_1 H_2 + B_2

S3=S0H1H2H3+B1H2H3+B2H3+B3S_3 = S_0 H_1 H_2 H_3 + B_1 H_2 H_3 + B_2 H_3 + B_3

S4=S0H1H2H3H4+B1H2H3H4+B2H3H4+B3H4+B4S_4 = S_0 H_1 H_2 H_3 H_4 + B_1 H_2 H_3 H_4 + B_2 H_3 H_4 + B_3 H_4 + B_4

规律一目了然:St=S0itHi+itBii<jtHjS_t = S_0 \prod_{i \le t} H_i + \sum_{i \le t} B_i \prod_{i < j \le t} H_j。展开后递归被彻底消掉了:每个 StS_t 都是初始状态、各步写入项与转移矩阵连乘的线性组合,只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化方案之一–换个括号方式(比如二叉树式两两合并),HH 的连乘与写入项的累积可以在 O(logn)O(\log n) 深度内并行完成。这就是 parallel scan / associative scan(并行扫描/结合扫描) 类方法的核心思想,也是 chunkwise 并行化和 SSM 训练并行(如 Mamba 的 selective scan)共同的理论根基。

但问题还没解决:HtH_t 是什么? 回看定义 Ht=IβtktktH_t = I - \beta_t k_t k_t^\top。设 ktRdk_t \in \mathbb{R}^d,则 ktktk_t k_t^\topd×dd \times d秩 1(rank-1)矩阵–所有列都是 ktk_t 的标量倍,只张出一维。所以 HtH_t 本质是单位矩阵减去一个秩 1 矩阵:ktk_t 方向被缩放 1βtkt21 - \beta_t\|k_t\|^2,正交补方向原封不动。这类「I 减秩 1」的结构与**豪斯霍尔德变换(Householder transformation)**同族–QR 分解里用它做反射/消元。区别在于:Householder 取 β=2/k2\beta = 2/\|k\|^2 且要求 k=1\|k\|=1 时是精确反射(范数保长),DeltaNet 的 βt(0,1)\beta_t \in (0,1) 是「部分反射」,软化为可学习的写入强度。秩 1 结构也解释了删除项为什么便宜:抹掉一个方向不需要满秩运算,一次外积就够。

到这里两条路线可以拼在一起了:DeltaNet 的 St=HtSt1+βtktvtS_t = H_t S_{t-1} + \beta_t k_t v_t^\top 与 SSM 的 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t 结构完全同构,差的只有一件事–HtH_t 的遗忘是「沿 ktk_t 方向删一块」,没有 SSM 那种全通道的指数衰减。

拼在一起:Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。论文的核心洞察:gating 和 delta rule 是互补的两种记忆管理机制:

机制 比喻 能力 缺陷
Gating(Mamba-2) 板擦 一键大面积擦除(全局衰减) 无法定点修改
Delta rule(DeltaNet) 铅笔 定点覆写某个 key 的关联 无法快速清空

序列长度超过状态容量时,记忆碰撞必然发生;GDN 用板擦 + 铅笔的组合管理这块固定大小的白板:

St=St1(αt(Iβtktkt))+βtvtktS_t = S_{t-1}\big(\alpha_t(I - \beta_tk_tk_t^\top)\big) + \beta_tv_tk_t^\top

其中 αt(0,1)\alpha_t \in (0,1) 是数据相关的标量门(Mamba-2 的参数化:α=exp(\alpha = \exp(负 Softplus(Linear(xt)))(x_t))),在 log 空间计算以保证数值稳定)。三个极限情形读懂这个式子:

极限 行为 对应模型
αt1\alpha_t \to 1 纯 delta rule,只定点改写 DeltaNet
βt1\beta_t \to 1,k ⊥ 已有记忆 退化为 St=αtSt1+vtktS_t = \alpha_tS_{t-1} + v_tk_t^\top Mamba-2
αt0\alpha_t \to 0 整表清零再写入(硬重置) 新能力:两者都做不到

几何直觉:(Iβkk)(I - \beta kk^\top) 是广义 Householder 反射–把状态往 k 方向「压扁擦除」;再乘标量 α\alpha 把整张表均匀缩小。一个是定向操作,一个是全局操作,作用在不同维度上,所以可以叠加。

注意作用顺序:擦除量是 βt(St1kt)\beta_t(S_{t-1}k_t),用的是未衰减的旧读出;而 delta 对照的是衰减后的旧值(vtαtSt1ktv_t - \alpha_tS_{t-1}k_t)。α\alpha 乘的是整个 (Iβkk)(I-\beta kk^\top)–这一点写代码时极易弄错(把 α 只乘到擦除项上,chunkwise 形式会与递归形式对不上)。读取侧同样随时间累积衰减:token 在时间步 xx 写入,在 x+tx+t 读取时已经被 αxαx+1αx+t\alpha_x\alpha_{x+1}\dots\alpha_{x+t} 衰减过。实现中通过 γr/γi\gamma^r/\gamma^i 项修正–分子分母都是 α\alpha 连乘,相除就是区间衰减,本质是乘法形式的前缀和(prefix-sum),与 SSM 的 Aˉ\bar A 连乘完全同源。

回头看:一张图 + 一张表把全家统一(论文 Table 1)。到这里四个模型都出场了,回头看,它们其实是同一个在线优化问题的闭式解,差别只在目标函数:

Linear Attn(纯加性)Mamba-2(全局遗忘门)DeltaNet(定点覆写)    GDN(板擦+铅笔)KDA(逐通道门+下界)\text{Linear Attn}(纯加性) \to \text{Mamba-2}(全局遗忘门) \searrow \\ \text{DeltaNet}(定点覆写) \nearrow \;\; \text{GDN}(板擦+铅笔) \to \text{KDA}(逐通道门+下界)

模型 在线学习目标 闭式解(状态更新)
Linear Attn StSt1F22Stkt,vt|S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle St=St1+vtktS_t = S_{t-1} + v_tk_t^\top
Mamba-2 StαtSt1F22Stkt,vt|S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle St=αtSt1+vtktS_t = \alpha_tS_{t-1} + v_tk_t^\top
DeltaNet StSt1F22Stkt,βt(vtSt1kt)|S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - S_{t-1}k_t)\rangle St=St1(Iβtktkt)+βtvtktS_t = S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top
GDN StαtSt1F22Stkt,βt(vtαtSt1kt)|S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - \alpha_tS_{t-1}k_t)\rangle St=St1(αt(Iβtktkt))+βtvtktS_t = S_{t-1}(\alpha_t(I-\beta_tk_tk_t^\top)) + \beta_tv_tk_t^\top
KDA (GDN 的逐通道化) (Iβtktkt)Diag(αt)St1+βtktvt(I-\beta_tk_tk_t^\top)\,\mathrm{Diag}(\alpha_t)S_{t-1} + \beta_tk_tv_t^\top

读法:第一项是正则(别离上一刻状态太远 = 记忆保留),第二项是拟合(SkvSk \approx v = 关联学习)。Mamba-2 与 DeltaNet 是两条独立的进化路线–前者把正则锚点换成可收缩的 αtSt1\alpha_tS_{t-1}(选择性遗忘),后者把拟合项换成回归目标(修正误差);GDN 把两条路线拼在了一起,KDA 再把拼完之后的标量门拆成逐通道门。沿着两个轴看这张表的演化:

  • 拟合轴:线性注意力/Mamba-2 的拟合项 2Sk,v-2\langle Sk, v\rangle 只要求内积大,太弱,所以只能「越叠越厚」;DeltaNet/GDN 换成回归目标 Skv2\|Sk - v\|^2 类的构造,才获得「修正误差」的闭式解–delta 写入是回归目标的直接推论;
  • 正则轴:Mamba-2/GDN 的正则锚点是 αtSt1\alpha_tS_{t-1}(可收缩的锚):正则放松 \Rightarrow 允许状态偏离旧值 \Rightarrow 选择性遗忘

所以这棵「进化树」可以反过来读:不是「谁加了什么模块」,而是目标函数从弱拟合 + 硬正则,逐级加强到强拟合 + 软正则,每一级都是上一级的严格泛化(GDN 令 βt\beta_t 退化为 0/边界值就退回全部下游模型)。再往前推一步就是**测试时训练(Test-Time Training, TTT)**一族:delta rule 是对 L(S)=12Sktvt2\mathcal{L}(S)=\frac12\|Sk_t-v_t\|^2 逐步做 SGD(β=学习率),GDN 等价于 SGD + 自适应权重衰减(α)–把「序列处理」看成「在线优化」,线性 RNN 全家都是这个视角的特例。

KDA 要做的最后一步就清楚了。Kimi Linear(后演化为 KDA)在 Gated DeltaNet 基础上的核心改进是细粒度门控:不再是每个注意力头一个标量 α\alpha,而是每个通道一个独立衰减值:

αt(0,1)dk(channel-wise 遗忘)\alpha_t \in (0,1)^{d_k} \quad (\text{channel-wise 遗忘})

作用是模型可以对不同维度做不同程度的记忆衰减:部分通道 α\alpha 接近 1,保留长期信息;部分通道 α\alpha 接近 0,快速遗忘。打个比方,Gated DeltaNet 的标量门是「全屋一个总闸」,KDA 的逐通道门是「每盏灯一个旋钮」–同一个状态里,慢通道记长程依赖,快通道只管局部上下文,记忆容量按维度重新分配。再加上衰减下界(log\log-decay 限制在 (gmin,0)(g_{\min}, 0))保证数值稳定,就得到 KDA 的完整递推式–正是下一小节的内容。

下面先把 GDN(Gated DeltaNet)的机制吃透:手算一遍、推导 chunkwise 并行形式,最后对照 KDA 看它继承了什么、改了什么。

GDN 数值验算:手算一遍

设定 dk=dv=2d_k = d_v = 2,C=3C = 3 个 token,S0=0S_0 = 0。精心设计的数据–k1=k2=e1k_1 = k_2 = e_1 制造 key 碰撞:

K=[101001], V=[102031], α=[0.8,0.5,0.9], β=[1,1,0.6]K = \begin{bmatrix}1&0\\1&0\\0&1\end{bmatrix},\ V = \begin{bmatrix}1&0\\2&0\\3&1\end{bmatrix},\ \alpha = [0.8,\,0.5,\,0.9],\ \beta = [1,\,1,\,0.6]

顺序递归(ground truth)。累积衰减 γj=ijαi=[0.8,0.4,0.36]\gamma_j = \prod_{i\le j}\alpha_i = [0.8, 0.4, 0.36]

t=1(k=e1k=e_1, v=[1,0]v=[1,0], α=0.8\alpha=0.8, β=1\beta=1):S0k=0S_0k = 0,无旧记忆可删:

S1=0.80+1([1,0]0)e1=[1000],o1=S1q1=[1,0]S_1 = 0.8\cdot 0 + 1\cdot([1,0]-0)\,e_1^\top = \begin{bmatrix}1&0\\0&0\end{bmatrix}, \qquad o_1 = S_1q_1 = [1,0]

t=2(k=e1k=e_1, v=[2,0]v=[2,0], α=0.5\alpha=0.5, β=1\beta=1):旧读出 S1k2=[1,0]=v1S_1k_2 = [1,0] = v_1碰撞发生。擦除 1[1,0]e1-1\cdot[1,0]e_1^\topv1v_1 完全擦掉;写入 delta =[2,0]0.5[1,0]=[1.5,0]= [2,0] - 0.5\cdot[1,0] = [1.5, 0]:

S2=0.5S1[1,0]e1+[1.5,0]e1=[2000]S_2 = 0.5S_1 - [1,0]e_1^\top + [1.5,0]e_1^\top = \begin{bmatrix}2&0\\0&0\end{bmatrix}

现在查 e1e_1[2,0] = v2v_2,v1v_1 已被覆写(对照线性注意力:纯加性会得到 [3,0] 的污染叠加)。

t=3(k=e2k=e_2, v=[3,1]v=[3,1], α=0.9\alpha=0.9, β=0.6\beta=0.6):旧读出 S2k3=0S_2k_3 = 0,正交无碰撞。写入 0.6[3,1]e20.6\cdot[3,1]e_2^\top,同时全表再乘 α\alpha(此处也乘了已写入的 e1e_1 行):

S3=[1.81.800.6],o3=S3q3=[1.8,0.6]S_3 = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}, \qquad o_3 = S_3q_3 = [1.8,\,0.6]

检查点:为什么 S3[0,0]=1.8S_3[0,0] = 1.8 不是 2.0? t=3 的 α3=0.9\alpha_3=0.9 作用在整张表上:第 1 行 2.0×0.9=1.82.0\times0.9 = 1.8。这就是 gating 与 delta 的交互:即使 token 3 的 key 与 e1e_1 正交,它的遗忘门仍然衰减了 e1e_1 通道上的记忆。逐通道门(KDA)与下界衰减都是围绕这一约束做文章。

GDN 的 Chunkwise 并行形式:递归怎么塞进 GPU

推理用递归(O(1)O(1)/token),但训练/prefill 必须并行。GDN 的贡献是把 gating 并入 DeltaNet 的 WY 表示 chunkwise 框架。设 chunk 大小 C,chunk 入口状态 S0S_0,目标:一次矩阵乘算出整个 chunk 的 O 和 chunk 出口状态 SCS_C

第一步:部分展开递归:

Sr=γrS0PrFr+i=1rγrγiu~ikiGrS_r = \underbrace{\gamma_r S_0 P_r}_{F_r} + \underbrace{\sum_{i=1}^r\frac{\gamma_r}{\gamma_i}\,\tilde u_i k_i^\top}_{G_r}

  • γr=jrαj\gamma_r = \prod_{j\le r}\alpha_j(α\alpha 是标量,可提到矩阵连乘外面);
  • Pr=ir(Iβikiki)P_r = \prod_{i\le r}(I - \beta_ik_ik_i^\top):纯 Householder 连乘,与 gating 无关;
  • u~i\tilde u_i:吸收了 β\beta 和衰减修正的「伪 value」。

第二步:WY 表示–秩 1 连乘压缩成两个小矩阵。关键观察:Pr=IirwikiP_r = I - \sum_{i\le r}w_ik_i^\top(乘积不膨胀,始终是 I 减一个秩 r\le r 的项)。wrw_ru~r\tilde u_r 的递归(W 不带 γ\gammaU~\tilde Uγ\gamma,这是 gating 并入的位置):

wr=βr(kri<rwi(kikr)),u~r=βr(vri<ru~iγrγi(kikr))w_r = \beta_r\Big(k_r - \sum_{i<r}w_i\,(k_i^\top k_r)\Big), \qquad \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)

直觉:wrw_r 是「擦除向量」–前面的写入若与 krk_r 重叠,先从自己的擦除量里扣掉(避免重复擦除);u~r\tilde u_r 是「修正量 value 版」。两者形式相同,差别只在 u~\tilde u 的内积多了一个衰减比 γr/γi\gamma_r/\gamma_i:i 越久远、衰减越多,它对当前擦除的干扰越小

第三步:UT 变换–递归变成一次下三角方程求解。w 和 u~\tilde u 的递归是「自回归」的(依赖前面的 w/u~\tilde u),但系数矩阵是单位下三角,可整体解方程。

辨析:因果掩码 ≠ UT 变换。两者容易混,因为碰巧都是下三角,但要管的完全是两件事:

  • 因果掩码(tril):把 score 矩阵的严格上三角直接置零,一行掩码操作,没有任何「变换」可言。它管的是「第 i 个输出不许看未来的 token」;
  • UT 变换(求 (I+L)1(I+L)^{-1}):管的是已经看过的历史之间怎么算账。delta rule 每次写入是「先读、再改」,块内第 2 次写入读到了第 1 次的结果,第 3 次读到前两次–历史写入之间互相纠缠,UT 把这团纠缠解开。

两者都是下三角不是巧合,是同一个原因:因果性。位置 i 的写入只能影响 i 之后的位置,所以「干扰系数矩阵」天然下三角;对角线天然是 1(自己的写入自己完整可见),于是要逆的矩阵恰好是单位下三角–这就是「UT」(Unit Triangular)名字的由来。一句话:因果性决定了它是三角的,但做它的目的是解耦,不是掩码

为什么必须做(最短版本):并行化要求把 C 次顺序写入折成一锤子外积累加 iv~iki\sum_i \tilde v_i k_i^\top。如果直接用原始 viv_i 累加,重叠部分会被重复计算–第 1 次写入的内容会透过后续写入的「先读」环节被间接再写一遍。UT 变换算出每个 viv_i 该扣除多少,使等式精确成立。不做的代价:要么结果错,要么退回逐 token 循环。

旁注:这类下三角变换是个大家族。「顺序递推 ↔ 三角矩阵求逆」是个通用模式,KDA 的 UT 只是其中一员:

家族成员 三角结构 与 UT 的关系
三角方程组求解(前代/回代) LU、Cholesky 分解之后的三角系统 UT 变换的计算过程就是一次前代法–同一个算法
Householder QR 的紧凑 WY 表示 反射连乘 (Iβkk)=I+YTY\prod(I - \beta kk^\top) = I + YTY^\top,TT 三角 「WY」「UT」两个名字的学术出处(Schreiber-Van Loan 1989);DeltaNet 把同样的打包思想借到 delta rule
因果卷积的逆(去卷积) 因果线性系统 = 下三角 Toeplitz,其逆也是下三角 Toeplitz 信号处理经典:「用三角逆矩阵解顺序依赖」
Mamba-2/SSD 的半可分矩阵 块内注意力 (CB)L(CB^\top)\odot L,LL 下三角衰减 同一枚硬币另一面:顺序 SSM 递推等价于带结构下三角矩阵,KDA 的 AqkA^{qk} 衰减注意力项完全是这个结构
幂零矩阵 Neumann 级数 (I+L)1=IL+L2(I+L)^{-1} = I - L + L^2 - \cdots 一般矩阵是无穷级数,但严格下三角矩阵幂零(LC=0L^C = 0),有限项精确截断–UT 能精确、便宜算出的数学原因

放进大图景里记:凡是「顺序执行的因果更新」,分块并行化时都会变成一个下三角矩阵的求逆/求解问题–SSM、delta rule、因果卷积、QR 分解,全是这个模式的化身。KDA chunkwise 里的 WY 表示(第二步)与 UT 变换(第三步),分别是这个模式在「写入打包」与「依赖解耦」上的具体化。

回到公式。定义衰减感知掩码 Γij=γi/γj (i>j)\Gamma_{ij} = \gamma_i/\gamma_j\ (i>j),则:

Tplain=[I+strictLower(diag(β)KK)]1diag(β),W=TplainKT_{\text{plain}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)\,KK^\top)\big]^{-1}\mathrm{diag}(\beta), \qquad W = T_{\text{plain}}K

Tgated=[I+strictLower(diag(β)(ΓKK))]1diag(β),U~=TgatedVT_{\text{gated}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)\,(\Gamma\odot KK^\top))\big]^{-1}\mathrm{diag}(\beta), \qquad \tilde U = T_{\text{gated}}V

W 和 U~\tilde U 只差内积处的一个 γ\gamma 比。下三角求逆(C×CC\times C,实践中 C=64)用前代法(forward substitution)完成。

第四步:chunk 输出与出口状态。记号(论文的箭头约定):()r=γr()r\overleftarrow{(\cdot)}_r = \gamma_r(\cdot)_r(衰减到 chunk 首端),()r=γCγr()r\overrightarrow{(\cdot)}_r = \frac{\gamma_C}{\gamma_r}(\cdot)_r(衰减到 chunk 末端)。

输出(两项:读 chunk 外旧状态 + chunk 内交互):

O=QS0+(QKΓcausal)(U~WS0)O = \overleftarrow{Q}\,S_0^\top + \big(QK^\top\odot\Gamma_{\text{causal}}\big)\big(\tilde U - \overleftarrow W S_0^\top\big)

出口状态:

SC=γCS0+(U~WS0)K(U~r=γCγru~r)S_C = \gamma_C S_0 + \big(\overrightarrow{\tilde U} - \overrightarrow W S_0^\top\big)^\top K \qquad(\overrightarrow{\tilde U}_r = \tfrac{\gamma_C}{\gamma_r}\tilde u_r)

结构读法:

  • 输出的第二项是「chunk 内小注意力」:QKΓcausalQK^\top\odot\Gamma_{\text{causal}} 就是带衰减的因果注意力矩阵,attend 的对象不是 V 而是修正后的伪 value U~WS0\tilde U - \overleftarrow WS_0^\top(后者是「旧状态在这个 chunk 里该被擦掉的部分」);
  • 出口状态 = 旧状态整体衰减 γC\gamma_C + 修正量加权写入(权重 γC/γi\gamma_C/\gamma_i:越早写入衰减越多);
  • 全部计算都是 C×CC\times CC×dkC\times d_kC×dvC\times d_v 的稠密矩阵乘–Tensor Core 友好;chunk 间只传一个 dv×dkd_v\times d_k 矩阵。

数值验算:拿顺序递归的例子走一遍 chunkwise

Step 1:KK=[110110001]KK^\top = \begin{bmatrix}1&1&0\\1&1&0\\0&0&1\end{bmatrix},ΓstrictKK\Gamma_{\text{strict}}\odot KK^\top 只有 (2,1) 处非零 =γ2/γ1×1=0.5= \gamma_2/\gamma_1\times1 = 0.5

Step 2:解两个下三角方程(β=[1,1,0.6]\beta = [1,1,0.6]):

Tplain=[100110000.6],Tgated=[1000.510000.6]T_{\text{plain}} = \begin{bmatrix}1&0&0\\-1&1&0\\0&0&0.6\end{bmatrix}, \qquad T_{\text{gated}} = \begin{bmatrix}1&0&0\\-0.5&1&0\\0&0&0.6\end{bmatrix}

对比唯一差别 (2,1):-1 -> -0.5,正是 γ2/γ1=0.5\gamma_2/\gamma_1 = 0.5 的衰减–t=2 时 v1v_1 已被衰减一半,擦除它的需求也减半。

Step 3:W=TplainK=[100000.6]W = T_{\text{plain}}K = \begin{bmatrix}1&0\\0&0\\0&0.6\end{bmatrix},U~=TgatedV=[101.501.80.6]\tilde U = T_{\text{gated}}V = \begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix}

u~\tilde u 的第二行:$v_2 - 0.5\cdot\tilde u_1(k_1^\top k_2) = [2,0] - 0.5[1,0] = $ [1.5, 0]–与顺序递归里手算的 delta 完全一致 ✓(W2=0W_2 = 0 则因为 w1w_1 已把 e1e_1 方向擦干净,β=1\beta=1 时无需再擦)。

Step 4(S0=0S_0=0,修正项消失):U~=γCγiu~i\overrightarrow{\tilde U} = \frac{\gamma_C}{\gamma_i}\tilde u_i 逐行乘 [0.36/0.8,0.36/0.4,1]=[0.45,0.9,1][0.36/0.8, 0.36/0.4, 1] = [0.45, 0.9, 1]:

U~=[0.4501.3501.80.6],SC=U~K=[1.81.800.6] \overrightarrow{\tilde U} = \begin{bmatrix}0.45&0\\1.35&0\\1.8&0.6\end{bmatrix}, \qquad S_C = \overrightarrow{\tilde U}^\top K = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}\ \checkmark

与顺序递归的 S3S_3 完全一致

Step 5:Γcausal=[1000.5100.450.91]\Gamma_{\text{causal}} = \begin{bmatrix}1&0&0\\0.5&1&0\\0.45&0.9&1\end{bmatrix},O=(QKΓcausal)U~O = (QK^\top\odot\Gamma_{\text{causal}})\tilde U:

O=[1000.510001][101.501.80.6]=[10201.80.6] O = \begin{bmatrix}1&0&0\\0.5&1&0\\0&0&1\end{bmatrix}\begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix} = \begin{bmatrix}1&0\\2&0\\1.8&0.6\end{bmatrix}\ \checkmark

与顺序递归的 o1=[1,0]o_1=[1,0]o2=[2,0]o_2=[2,0]o3=[1.8,0.6]o_3=[1.8,0.6] 逐步精确一致(含非零 S0S_0 的一般情形同样对拍通过)。

完整 block:GDN 在网络里长什么样

Llama 式宏架构,把 attention 换成 gated delta token mixer:

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

要点:q/k 走 ShortConv+SiLU+L2Norm(局部上下文 + 单位球稳定擦除几何),α/β 只走线性投影(它们是标量,无需卷积);输出门与 Mamba 的 SiLU 门一脉相承(K3 把它升级为全秩)。

旁注:ShortConv 的数学定义

上面流程图里的 ShortConv(Short Convolution,短因果深度卷积)出自 Kimi Linear(K3 技术报告引用 [64]),位于 Q/K/V 线性投影之后、进入 KDA 递归之前–顺序是 xWq/k/vShortConvx \to W_{q/k/v} \to \mathrm{ShortConv} \to \dots,即卷积是「KDA 前置」而非「投影前置」。它只捕捉当前 token 前面少量局部上下文,无未来 token 泄露;且在 Kimi Linear / FLA 实现里卷积不是裸用的,后面紧跟 SiLU:silu(conv(x))\mathrm{silu}(\mathrm{conv}(x))(见下一个旁注)。

输入单头向量 xtRdx_t \in \mathbb{R}^d,卷积核窗口大小 W(标准取 4),深度可分离(depthwise)且因果:

ShortConv(x)t=j=0W1wjxtj\mathrm{ShortConv}(x)_t = \sum_{j=0}^{W-1} w_j \odot x_{t-j}

记号 含义
wjRdw_j \in \mathbb{R}^d 每通道独立卷积权重–depthwise,每个特征通道一套独立卷积核
\odot 逐元素相乘
因果约束 tj0t - j \ge 0;t<jt < j 时零填充,看不到未来 token

三个性质决定了它为什么放在这个位置:

  1. 时序维度滑动,只混合最近 W 个历史 token,开销远小于全局注意力;
  2. depthwise:wkw_kxtkx_{t-k} 逐通道相乘,通道间不混–局部上下文的注入不破坏各通道独立的衰减语义(与 Diag(αt)\mathrm{Diag}(\alpha_t) 逐通道门配套);
  3. 因果:与上一节因果时序卷积同理,LLM 自回归必备。

直觉上,它给 KDA 补上了线性 RNN 天生缺失的「局部听觉」–递归状态只携带压缩后的全局历史,最近几个 token 的精细局部模式(词内字符、短程搭配)由这层小卷积负责。

旁注:Swish(SiLU,自门控激活函数)

流程图里 ShortConv 之后、L2Norm 之前的 SiLU,即 Swish–严格说 Swish 带可学参数时是 xσ(βx)x \cdot \sigma(\beta x),论文和实现里固定 β=1\beta = 1,两个名字就此等价:

Swish(x)=xσ(x),σ(x)=11+ex\mathrm{Swish}(x) = x \cdot \sigma(x), \qquad \sigma(x) = \frac{1}{1 + e^{-x}}

分层计算只有两步:1对输入的每个元素算 σ(x)=1/(1+exp(x))\sigma(x) = 1/(1+\exp(-x));2原输入 x 与 sigmoid 结果逐元素相乘得到输出。Swish = 输入 × 输入自己的 sigmoid 门–自门控(self-gated):门控信号不是外部来的,是输入自身,免参数。

为什么这层用它:1平滑非单调(负区间有个小下凹,xx \to -\infty 时输出趋于 0 而非 ReLU 的硬截断),梯度处处非零,深网训练稳;2门控形式与整个 block 的「信息流过 vs 被抑制」语义一致–Q/K/V 进递归状态前先过一道自门控,相当于对局部卷积混合后的特征做一次软筛选,再交给 L2Norm 钉上单位球。K3 把输出侧的这门升级为全秩输入相关门(满秩门控一节),源头就是这里的 SiLU。

旁注:L2Norm(逐头 L2 归一化,Q/K 专用)

流程图最后一框。对单头向量 zRdk\bm{z} \in \mathbb{R}^{d_k},逐通道 L2 标准化到单位范数:

L2Norm(z)=zz2+ϵ,z2=i=1dkzi2\mathrm{L2Norm}(\bm{z}) = \frac{\bm{z}}{\|\bm{z}\|_2 + \epsilon}, \qquad \|\bm{z}\|_2 = \sqrt{\sum_{i=1}^{d_k} z_i^2}

记号 含义
ϵ\epsilon 极小防除零常数(10610^{-6} 左右)
操作维度 每个注意力头独立归一化,跨头不共享统计

为什么只有 Q/K 用、V 不用?回到 DeltaNet 一节埋下的伏笔:归一化后 kt2=1\|k_t\|_2 = 1,于是 delta rule 的擦除矩阵 IβkkI - \beta k k^\top 的特征值被夹在 [1β, 1][1-\beta,\ 1]–写入步长天然稳定,β=1 时是精确保长反射。Q 同样归一化则是让读出 o=Sqo = S^\top q 的尺度可控,查询和 key 在同一单位球上内积才有可比性,chunkwise 公式里 QKQK^\top 的 score 变成有余弦界的小量。V 不归一化:写入内容的幅值本身携带信息(value 的模长是特征,不是 bug),而且 βtktvt\beta_t k_t v_t^\top 的稳定性由 k 侧保证,与 v 的尺度无关。这与 softmax 注意力里的 QK-norm(防 logit 爆炸)动机同源,但线性 RNN 里它 additionally 承担擦除几何的稳定性–归一化不是锦上添花,是 delta rule 能正常工作的前提。

至此三个旁注凑齐流程图的 Q/K 支路:Linear 投影 -> ShortConv(局部混合) -> SiLU(软筛选) -> L2Norm(单位球),三步全部因果、逐通道、免跨头统计,正好匹配 Diag(αt)\mathrm{Diag}(\alpha_t) 的逐通道语义。更重要的是三者各管一段数值安全:ShortConv 管局部混合、L2Norm 管幅度稳定、逐通道衰减差 γiγj0\gamma_i - \gamma_j \le 0 管指数安全(下界衰减一节)–全链路数值安全正是 KDA 能在 BF16 上跑起来的主线。

混合架构(论文的另一贡献):GDN + 滑窗注意力(H1)或 Mamba-2 + GDN + SWA(H2)交错堆叠,取长补短–与 K3 的「3 KDA + 1 MLA」同款思路,GDN 论文可视为这个混合范式的先声。

从 GDN 到 KDA:K3 继承了什么、改了什么
维度 GDN KDA(Kimi Linear -> 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}(\alpha_t)S_{t-1} + \beta vk^\top α 从标量 -> 逐通道向量
遗忘门粒度 标量(整头同一衰减率) 向量 αt(0,1)dk\alpha_t\in(0,1)^{d_k} 长期记忆通道与短期工作区分离
门作用顺序 α 在外乘整个更新 Diag(α) 在内侧先衰减 S,再删写 与逐通道参数化配套(KCP 推导需要)
α 参数化 负 Softplus,值域 (,0)(-\infty,0) K3:缩放 sigmoid,下界 gmin=5g_{\min}=-5 1/Γ1/\Gamma 有界 -> 对角 tile 全走 Tensor Core
chunkwise WY + UT + 衰减箭头 同框架,Γ 从标量比变成向量累积比 表达力↑ 数值难度↑(K3 用下界解决)
输出门 低秩/简单门 K3:输入相关全秩 sigmoid 门 逐通道调节读出

两处值得注意的「暗改」:

  1. 作用顺序变了:GDN 是 S(α(Iβkk))S(\alpha(I-\beta kk^\top))(α 吸进 Householder 连乘),KDA 是 (Iβkk)Diag(α)S(I-\beta kk^\top)\mathrm{Diag}(\alpha)S(α 先作用、再删写)。KDA 的逐通道门是矩阵 Diag(α)\mathrm{Diag}(\alpha),与 Householder 不交换,顺序成为实质设计选择–它使 KCP 的「段转移分解」成为可能。
  2. 衰减在 Γ 处的复杂化:GDN 的 γ 是标量连乘,Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 只是数;KDA 的 Γ 是向量逐元素累积,chunk 公式里 K/Γ、Q⊙Γ 等运算随之复杂化,数值范围问题(1/Γ1/\Gamma 爆炸)由此而生–K3 的下界衰减正是对 GDN->KDA 这一步引入的新问题的修复。

SSM 部分速览,详细推导见前置篇

KDA 递推式的 SSM 部分来自一条完整的推导链:连续状态方程 h(t)=Ah(t)+Bx(t)h'(t) = Ah(t) + Bx(t),经积分因子法得通解,ZOH 离散化得递推 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,其中 Aˉ=eΔA\bar A = e^{\Delta A}Bˉ=A1(eΔAI)B\bar B = A^{-1}(e^{\Delta A} - I)B;递归展开等价于卷积(S4 双形式);Mamba 让 Δt\Delta_t 依赖输入,Aˉt=eΔta(0,1)\bar A_t = e^{-\Delta_t a} \in (0,1) 变成看内容的遗忘门;Mamba-2 把状态写成外积形式 St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top,和线性注意力在数学上就是一回事了(SSD 对偶)。每一步的严格推导、数值验算(ZOH 精确性 1.2642 = 1.2642)和 S4/Mamba/Mamba-2 三站的演化,见前置篇《线性注意力与 SSM:两条技术路线的完整推导》

这里只需要记住结论:SSM 这条路线最终给出的,是乘在旧状态上的标量衰减系数 αt\alpha_t–它回答了线性注意力一族「状态该怎么遗忘」的问题,但答案是整片擦除、不分通道。它和 DeltaNet(定点覆写)在 GDN 处碰头;GDN 又被 KDA 逐通道化。骨架始终是同一条:

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

递推公式(论文 2.1.1 节,核心式 1)

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

o~t=Stqt\tilde{\bm{o}}_t = \mathbf{S}_t^\top \bm{q}_t

符号逐个过一遍(论文原定义):

符号 含义
αtRdk\bm{\alpha}_t \in \mathbb{R}^{d_k} 逐通道一维保留因子向量(channel-wise one-step retention factor)
Diag(αt)\operatorname{Diag}(\bm{\alpha}_t) 向量转对角矩阵算子,把通道级衰减系数变成矩阵乘
I\mathbf{I} 同维度单位矩阵
βt(0,1)\beta_t \in (0,1) delta rule 写入强度

读这个式子的三层动作(从右往左):

  1. 通道衰减:Diag(αt)St1\operatorname{Diag}(\bm{\alpha}_t)\mathbf{S}_{t-1} 等价于对 St1\mathbf{S}_{t-1} 的每一行分别乘以对应通道的 α\alpha 系数–逐通道缩放历史记忆。普通线性注意力/SSM 大多全局统一衰减(αt\alpha_t 是标量),KDA 给每一个 key 通道分配独立衰减权重:不同语义通道可以选择更快/更慢遗忘 St1\mathbf{S}_{t-1};
  2. 方向擦除:(Iβtktkt)\left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) 再对衰减后的状态做当前 token 的定点擦除(秩 1 Householder 型,详见 DeltaNet 一节);
  3. 新信息写入:叠加 βtktvt\beta_t \bm{k}_t \bm{v}_t^\top

三步合起来:先通道衰减,再方向擦除,最后写入–带通道精细遗忘的 delta 递推。作用顺序是设计选择:Diag 与 Householder 不交换,这个顺序使 KCP 的段转移分解成为可能(见下界衰减与基础设施篇)。

记号约定:论文里大写 Diag()\operatorname{Diag}(\cdot) 是「向量 -> 对角矩阵」算子;小写 diag()\operatorname{diag}(\cdot) 有时指反向操作(输入矩阵、提取对角线为向量)。本文全程大写表示向量转对角矩阵。

为什么是逐通道门:动机与在线学习视角

标量门的表达力瓶颈。GDN 的 αt(0,1)\alpha_t \in (0,1) 是标量:每一步遗忘时,所有 key 通道以同一个比例衰减–要么一起记住,要么一起忘记。但不同通道承担的角色不同:

  • 有的通道在存「长期主题」(希望 α1\alpha \approx 1,几乎不遗忘);
  • 有的通道在存「临时指针」(希望快速衰减,腾出容量)。

KDA 的核心改动只有一处:把标量换成向量 αt(0,1)dk\bm{\alpha}_t \in (0,1)^{d_k}(βt\beta_t 仍是标量,这是 KDA 的选择而非必须)。直觉:S 的第 j 列对应 key 空间的第 j 个通道,Diag(αt)\mathrm{Diag}(\alpha_t) 作用上去就是给每一列配一个独立的遗忘速度。GDN 是 KDA 在 Diag(αt)=αtI\mathrm{Diag}(\alpha_t) = \alpha_t I 时的特例。

在线学习视角:逐通道权重衰减的 delta rule。与 GDN 一节的在线学习表同构,只需把正则项换成逐通道版。每步给定新样本 (kt,vt)(k_t, v_t),希望新状态 S:1拟合新样本 SktvtSk_t \approx v_t;2不偏离「衰减后的旧记忆」太远 SSt1Diag(αt)S \approx S_{t-1}\mathrm{Diag}(\alpha_t)(以下记 Dt=Diag(αt)D_t = \mathrm{Diag}(\alpha_t),αt=exp(gt)\alpha_t = \exp(g_t),gtR<0dkg_t \in \mathbb{R}_{<0}^{d_k}):

Lt(S)=12Sktvt2拟合+12SSt1DtF2逐通道正则\mathcal{L}_t(S) = \underbrace{\tfrac12 \|S k_t - v_t\|^2}_{\text{拟合}} + \underbrace{\tfrac12 \|S - S_{t-1} D_t\|_F^2}_{\text{逐通道正则}}

St1DtS_{t-1}D_t 出发对第一项做一步梯度下降(步长 βt\beta_t):

St=St1Dtβt(St1Dtktvt)kt=St1Dt(Iβtktkt)+βtvtktS_t = S_{t-1}D_t - \beta_t\,(S_{t-1}D_t\, k_t - v_t)\,k_t^\top = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top

正是递推式(转置约定下)。逐通道正则的含义:第 j 列的「信任区域」宽度正比于 αt(j)\alpha_t^{(j)}α\alpha 小的通道,旧记忆先被压缩,新写入覆盖几乎没有阻力(快遗忘);α1\alpha \approx 1 的通道,旧记忆原样进入下一步,delta rule 只做精细增量(慢遗忘)。这就是「细粒度记忆控制」的优化论表述:KDA = 逐通道权重衰减 + delta rule。

实现注意:作用顺序是坑DtD_t 作用于整个 St1S_{t-1}、先于擦除项–擦除时检索用的也是已衰减的状态:

St=St1Dt先衰减(Iβtktkt)再擦除/写入S_t = \underbrace{S_{t-1} D_t}_{\text{先衰减}}\underbrace{(I - \beta_t k_t k_t^\top)}_{\text{再擦除/写入}}

写成 St1(Iβkk)DtS_{t-1}(I-\beta kk^\top)D_t(顺序颠倒)或只衰减单位阵部分都是错的–DtD_t(Iβkk)(I - \beta kk^\top) 不可交换,顺序错了结果就错(与 GDN 一节「α 乘整个括号」的警告同源,但逐通道化之后错误更隐蔽)。FLA 参考实现里对应 S = S * g.exp() 之后立刻用衰减后的 S 做检索 v - k^T S

下界衰减(Lower-bounded decay)

KDA 的衰减参数化是一个关键改进。Kimi Linear 使用无界的负 Softplus 映射 g=eASoftplus(z)(,0)g = -e^A \text{Softplus}(z) \in (-\infty, 0),而 K3 改用有界缩放 sigmoid:

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

其中 gmin=5g_{\min} = -5 固定,AhA_h 是可学习的每头对数尺度。这意味着每个保留因子满足 α>e56.7×103\alpha > e^{-5} \approx 6.7 \times 10^{-3},16 token 块的累积对数衰减落在 (80,0)(-80, 0) 内,对应的重缩放因子小于 e80e^{80},在 BF16 动态范围内。

计算收益:有限范围使得因果对角块和离对角块都能用密集 Tensor Core 矩阵乘法,消除了 Kimi Linear 中需要的 position-pair 对角计算路径。

问题从哪来:1/Γ1/\Gamma 爆炸是向量门控自带的代价

为什么这个修复在 GDN 上不必要、在 KDA 上变成刚需?GDN 的标量衰减在 chunkwise 里只以比值出现(s=j+1iαs1\prod_{s=j+1}^{i}\alpha_s \le 1,iji \ge j),天然安全。KDA 把 Γ\Gamma 变成向量累积 Γt=exp(stgs)\Gamma_t = \exp(\sum_{s\le t} g_s),逐通道独立:第 4 步式的推导可以刻意只用 iji \ge j 的差(指数 0\le 0,安全);但只要换一种等价写法–把状态「反归一化」回 chunk 起点、或把衰减从 key 上整体外提(论文公式 (4) 的 K/ΓK/\Gamma 因式分解正是这种写法)–就得真的算出

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

衰减越快的通道,γt-\gamma_t 越大,1/Γ1/\Gamma 指数膨胀。量级感受(实测):

恒定 α\alpha t=100t=100 t=500t=500 t=1000t=1000 t=2000t=2000
0.9 3.8×1043.8\times10^{4} 7.6×10227.6\times10^{22} 5.7×10455.7\times10^{45} 3.3×10913.3\times10^{91}
0.5 1.3×10301.3\times10^{30} 3.3×101503.3\times10^{150} 1.1×103011.1\times10^{301} 溢出 float64
0.1 1010010^{100} 溢出 溢出 溢出

float64 上界约 e7091.8×10308e^{709} \approx 1.8\times10^{308};BF16 训练下几十步就出事。标量 -> 向量的推广放大了衰减率的动态范围,爆炸从例外变成常态–所以必须给衰减率本身夹界。

Kimi Linear 的绕行方案(补丁,不是修复):在对数空间算相对衰减(减法代替除法,不溢出),并把每个 chunk 再切成 16 token 的二级瓦片。效果:瓦片之间(非对角块)可以安全交给 Tensor Core 稠密矩阵乘;但瓦片内部(对角块)衰减可能极端,仍需逐位置对显式计算–position-pair 路径无法组织成大矩阵乘,吃不满 Tensor Core,成为块内主要瓶颈。

K3 的根治:不改公式,改参数化,让爆炸在数学上不可能发生(负 Softplus 允许 α\alpha 任意接近 0,即「一步清零」;缩放 sigmoid 不允许)。

表达力为什么不受损:快通道在瓦片内照样能把记忆衰减到 e801035e^{-80} \approx 10^{-35}(约等于零,只是不真的为零);长期遗忘靠多步连乘(每步 ×0.0067\times 0.0067,几十步后照样忘干净),不需要单步清零的能力。用一个下界换来全 Tensor Core 化,划算的买卖。

满秩门控(Full-rank gate)

K3 将 KDA 的输出门从低秩参数化改为输入相关的满秩投影。在递推输出经过 head-wise RMSNorm 后,应用数据相关的输出门控:

yt=Wo[Sigmoid(Wgxt)RMSNorm(o~t)]y_t = W_o [\text{Sigmoid}(W_g x_t) \odot \text{RMSNorm}(\tilde{o}_t)]

满秩门控允许每个 token 独立调制从循环状态读取的通道。

Chunkwise 并行形式

递推形式在推理(decode)时是优势–O(1)O(1) 状态更新;但在训练/prefill 时是灾难:每个 token 的状态依赖前一个 token,顺序循环让 GPU 的数千个核心集体围观一个 for 循环。Chunkwise 并行化把序列切成长度 C 的块:块内矩阵运算并行,块间只传状态。先看通用的复杂度框架,再看 KDA 逐通道门带来的推导细节。

通用框架Chunk 内(intra-chunk):token 交互用带衰减的因果注意力直接算,O(C2d)O(C^2 d)–C 是常数(64/128),对序列长度 N 线性。Chunk 间(inter-chunk):每个 chunk 对状态做一次递推更新,chunk 内外积先归约成固定大小的 state 增量,块间只传 d×dd \times d 的 state。总计算量:

O(NCd)chunk 内注意力+O(N/Cd2)chunk 间状态更新  =  2Nd2(固定项)+2NCd\underbrace{O\big(N \cdot C \cdot d\big)}_{\text{chunk 内注意力}} + \underbrace{O\big(N/C \cdot d^2\big)}_{\text{chunk 间状态更新}} \;=\; 2Nd^2(\text{固定项}) + 2NCd

C 形态 复杂度
C=1C = 1 每个 token 一个 chunk,chunk 内注意力消失 纯线性注意力递推,FLOPs 最少但不一定最快(GPU 对小矩阵乘利用率低)
C=NC = N 整个序列一个 chunk,递推消失 标准 O(N2)O(N^2) 注意力

实践中 C 取 64/128:小到让 chunk 间项不贵,大到让 C×CC\times C 注意力矩阵铺满 Tensor Core 的 tile。这与 S4 时代「训练用卷积、推理用递归」双形式一脉相承:选择性打破 LTI 后,chunkwise 就是卷积的继任者。

KDA 版推导:逐通道衰减下的 WY / UT

记 chunk 起点传入状态 S[0]S_{[0]}(以下用局部下标 i=1..Ci = 1..C;乘积按时间倒序)。核心记号是累积 log 衰减:

γi=s=1igsRdk,Γij=diag(exp(γiγj))(ij)\gamma_i = \sum_{s=1}^{i} g_s \in \mathbb{R}^{d_k}, \qquad \Gamma_{i \leftarrow j} = \operatorname{diag}\big(\exp(\gamma_i - \gamma_j)\big) \quad (i \ge j)

Γij\Gamma_{i\leftarrow j} 是「从第 j 步衰减到第 i 步」的逐通道算子。两条性质:iji \ge jγiγj0\gamma_i - \gamma_j \le 0 逐分量成立(gs<0g_s < 0),元素都在 (0,1](0,1],数值安全;反向 Γji1\Gamma_{j\leftarrow i}^{-1} 的元素 1\ge 1 且随 chunk 长度指数增长–这是 1/Γ1/\Gamma 爆炸的根源,先记住这个观察。

第一步:衰减 KKT 矩阵 M。类比 GDN chunkwise 里的普通 kikjk_i^\top k_j,逐通道版需要带衰减的 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 写入的记忆衰减到第 c 步时,与 kck_c 的重叠程度。Hadamard 积 \odot 作用在 key 通道维–正是逐通道衰减出现的位置。

第二步:UT 变换。构造严格下三角 LL:Lci=βiMci (c>i)L_{ci} = \beta_i M_{ci}\ (c > i),求 T=(I+L)1T = (I+L)^{-1}(幂零,有限项截断;实践中不显式求逆,前代法逐行解),再每列乘 βj\beta_jA^=Tdiag(β)\hat{A} = T\,\operatorname{diag}(\beta)。FLA 实现里那个三行循环,数学上就是解 (I+L)T=I(I+L)T = I

与 GDN 对照:GDN 这一步的 M 是普通 kckik_c^\top k_i(标量衰减被拆成比值吸收进别的项);KDA 里衰减「长」在 M 内部,无法外提–这是向量门控带来的结构性变化,KDA chunkwise 推导的核心难点。

第三步:WY 表示:

W=A^(eγiki)i=1..CRC×dk,U=A^VRC×dv,V~=UWS[0]W = \hat{A}\,\big(e^{\gamma_i} \odot k_i\big)_{i=1..C} \in \mathbb{R}^{C \times d_k}, \qquad U = \hat{A}\,V \in \mathbb{R}^{C \times d_v}, \qquad \tilde{V} = U - W\,S_{[0]}^\top

伪值 V~\tilde V 到底是什么:一本账。UT 变换产出 U 和 W,用它们定义 V~:=UWS\tilde V := U - WS。三个符号各司其职:

符号 形状 角色
S[0]S_{[0]} dk×dvd_k \times d_v 旧账本:chunk 之前所有 token 写入的记忆,块内计算时固定不变
WW C×dkC \times d_k 翻账本的钥匙:每行是某位置的衰减 key 修正组合;WSWS 即「这个位置能从旧记忆里读到什么」
UU C×dvC \times d_v 想记的新账(已和同伴对过账):A^\hat A 作用在 V 上,块内写入重叠已扣除
V~\tilde V C×dvC \times d_v 真正上账的净额 = 新账 - (旧记忆已有的 + 同伴已写的)

为什么叫「伪」值:它们不是真实的 v(手算例子里 v~2=(2,2)v2\tilde v_2 = (-2,2) \ne v_2),而是预扣完所有重叠后可以不做任何修正直接累加的修正值。这正是 delta rule 的灵魂的单步版:单步写入 vtSt1ktv_t - S_{t-1}^\top k_t = 目标值 - 旧记忆读数;V~\tilde V 是它的 chunk 并行版,每行回答「这个位置真正新增了多少」。账在 V~\tilde V 里算清了,后面的块内注意力 AV~A\tilde V 和状态更新才敢放心一锤子矩阵乘。

第四步:输出与跨 chunk 状态:

Acjqk=(qceγcγj)kj (jc),oc=S[0](qceγc)块间:query 打折查旧状态+jcAcjqkv~j块内:衰减注意力A^{qk}_{cj} = \big(q_c \odot e^{\gamma_c - \gamma_j}\big)^\top k_j \ (j \le c), \qquad o_c = \underbrace{S_{[0]}\,(q_c \odot e^{\gamma_c})}_{\text{块间:query 打折查旧状态}} + \underbrace{\sum_{j \le c} A^{qk}_{cj}\,\tilde v_j}_{\text{块内:衰减注意力}}

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

出口状态读法:旧状态整体按整 chunk 累积衰减缩小;每个伪值以「写入时刻衰减到块末」的 key 为地址写入。所有指数都是 0\le 0 的差,整条链路没有一个 1\ge 1 的因子–刻意保持,原因见下界衰减一节。

手算一遍(dk=dv=2d_k = d_v = 2,C = 2)

零初始状态,每步恒定衰减 α=(0.5,0.25)\alpha = (0.5, 0.25)(即 g=(ln0.5,ln0.25)g = (\ln 0.5, \ln 0.25)),β 取 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 步,D=diag(0.5,0.25)D = \mathrm{diag}(0.5, 0.25),先衰减 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=(0.25,0.0625)e^{\gamma_2} = (0.25, 0.0625),eγ2γ1=(0.5,0.25)e^{\gamma_2-\gamma_1} = (0.5, 0.25)。衰减 KKT:M21=(k2eγ2γ1)k1=(0.5,0.5)(1,1)=1M_{21} = (k_2 \odot e^{\gamma_2-\gamma_1})^\top k_1 = (0.5, 0.5)\cdot(1,1) = 1。UT:L=(0010)L = \begin{pmatrix}0&0\\1&0\end{pmatrix},T=IL=(1011)T = I - L = \begin{pmatrix}1&0\\-1&1\end{pmatrix},A^=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=10M_{21} = 1 \ne 0),WY 把 v2v_2 修正为扣除重叠后真正的新增–单步 delta rule 的 chunk 版样子。衰减注意力与输出:

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 与递推式逐项一致。(随机对拍:非零初始状态、随机门控下递推 vs chunkwise 最大误差 10910^{-9} 量级。)

与论文公式 (4) 的对照:K/Γ 因式分解

K3 报告(沿用 Kimi Linear)用乘积记号:γij=r=ijαr=exp(r=ijgr)\gamma_{i\to j} = \prod_{r=i}^{j}\alpha_r = \exp(\sum_{r=i}^j g_r);ΓRC×dk\Gamma \in \mathbb{R}^{C\times d_k} 把各步 γ\gamma 按行堆叠(注意 Γ\Gamma 本身不是对角阵,每个位置 rrdiag(γr)\mathrm{diag}(\gamma_r) 才是遗忘对角阵,Γ\Gamma 是 C 个对角阵的打包)。论文状态约定 SRdk×dvS \in \mathbb{R}^{d_k \times d_v}(本文的转置),其公式 (4):

A=Tril[(QΓ)(K/Γ)],O=(ΓQ)S块间+AV~块内A = \mathrm{Tril}\big[(Q\odot\Gamma)(K/\Gamma)^\top\big], \qquad O = \underbrace{(\Gamma\odot Q)S}_{\text{块间}} + \underbrace{A\,\tilde V}_{\text{块内}}

为什么 A 能这样拆:看 (i,j)(i,j) 元素 Aij=dqi,dγi,dkj,d/γj,d=(qiγji)kjA_{ij} = \sum_d q_{i,d}\gamma_{i,d}\cdot k_{j,d}/\gamma_{j,d} = (q_i \odot \gamma_{j\to i})^\top k_j,与 AqkA^{qk} 逐元素相同–逐对位置的衰减比值被因式分解成 query 侧乘 Γ、key 侧除以 Γ 两个逐位置操作,整个 C×CC\times C 矩阵 = 一次稠密 matmul + 两次逐元素乘,完全并行。Tril 保留对角线:delta rule 里 oio_i 读的是写入当前 token 之后的状态。

用手算例子核对:QΓ=(0.50.250.250.0625)Q\odot\Gamma = \begin{pmatrix}0.5&0.25\\0.25&0.0625\end{pmatrix},K/Γ=(24432)K/\Gamma = \begin{pmatrix}2&4\\4&32\end{pmatrix},乘积 Tril 后 A=(200.753)A = \begin{pmatrix}2&0\\0.75&3\end{pmatrix},与上面 AqkA^{qk} 一致。注意 K/ΓK/\Gamma 里已出现 32 这种被放大的数–1/Γ1/\Gamma 的膨胀就藏在这一步,数值后果见下界衰减一节。

常见疑问:CCdkd_k 有倍数关系吗? 没有。C 是序列轴的切分(一个 chunk 装多少 token),dkd_k 是特征轴的宽度,两根轴独立。有整除要求的是:K3 把 chunk 再切成 16 token 二级瓦片,故 C 需是 16 的倍数;dkd_kdvd_v 按 Tensor Core 友好尺寸取是工程约束,不是数学要求。

KDA 的核心优势是:相比传统 softmax attention,它是线性复杂度的;相比纯粹的线性 RNN,它通过 delta rule 实现了更强大的信息写入和遗忘控制。

一页纸速查(KDA/GDN 全家)
维度 内容
状态 S ∈ R^{d_v×d_k}(论文记法),固定大小,每序列每头一份
更新 S = S(α(I-βkkT)) + βvkT:板擦 α + 铅笔 δ
读出 o = Sq,后接 RMSNorm + sigmoid 输出门
α GDN 标量 / KDA 逐通道,数据相关;β:sigmoid
q/k Linear -> ShortConv -> SiLU -> L2Norm
等价视角 在线最小二乘 SGD(β=学习率)+ 自适应权重衰减(α)
并行训练 chunkwise:WY 表示(W 无 γ / Ũ 有 γ)+ UT 变换(下三角求逆)+ 衰减箭头
chunk 输出 O = Q⃖S0T + (QKT⊙Γ因果)(Ũ - W⃖S0T)
chunk 状态 S_C = γ_C S0 + (Ũ-> - W->S0T)TK
记忆瓶颈修复 覆写(δ)治叠加污染;清场(α)治长期堆积
后继 KDA:α->逐通道、换作用顺序、加下界 -> 进 K3

参考: