Kimi Delta Attention:把遗忘门做到每个通道

长上下文最贵的不是算力,是那条 KV cache。序列每长一个 token,cache 就多一行,显存和访存一起涨——这是 full attention 收不住的税。线性注意力想免掉这笔税:把注意力写成一个固定大小的状态,推理时只更新状态、不再堆缓存。问题是它一直打不过 full attention,尤其在需要精确长程检索的地方。

Kimi Delta Attention(KDA)是 Kimi Linear 里那条线性注意力分支,也是 Kimi K3 的主力注意力。它没换掉线性注意力的框架,只做了一件事:把遗忘门从"一个标量"细化到"每个通道一个"。就这一步,配上 3 的混合排布,把线性注意力拉到了能替 full attention 挑大梁的位置。

先把线性注意力写成 RNN

标准 softmax 注意力在第 tt 步要回看所有历史:

ot=itexp(qtki)jtexp(qtkj)vio_t = \sum_{i \le t} \frac{\exp(q_t^\top k_i)}{\sum_{j\le t}\exp(q_t^\top k_j)}\, v_i

分母那个 exp\exp 求和让它没法写成递推,只能把 ki,vik_i, v_i 全存下来——这就是 KV cache 的由来。

线性注意力把 exp\exp 换成一个特征映射(或者干脆去掉),求和就能拆开,注意力退化成一个状态矩阵的读写:

St=St1+ktvt,ot=qtStS_t = S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t

StRdk×dvS_t \in \mathbb{R}^{d_k \times d_v} 是把所有历史键值压进去的一块固定大小内存。推理时不再存 cache,只维护这个 SS——状态大小和序列长度无关。代价也在这:所有历史被加法叠进同一块内存,会互相覆盖、糊在一起,精确检索能力就是这么丢的。

Delta rule:写之前先擦掉旧的

纯加法的毛病是:同一个键 ktk_t 再来一次,新值直接摞在旧值上,谁也没删。DeltaNet 借了个很老的思路——delta 学习规则:写新值之前,先按误差把旧的擦掉。

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

括号里的 (Iβtktkt)(I - \beta_t k_t k_t^\top) 沿 ktk_t 方向把旧内容按比例 βt\beta_t 抹去,再写入新的 ktvtk_t v_t^\top。这一步等价于对 Sktvt2\lVert S k_t - v_t\rVert^2 做一次在线梯度下降——它在"记住键值对"这件事上是有目标的,而不是无脑累加。

Gated DeltaNet:给状态加一个会衰减的门

再往前一步:历史不该永远同等重要,久远的信息该淡出。Gated DeltaNet 在递推前面乘一个标量遗忘门 αt(0,1)\alpha_t \in (0,1)

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

αt\alpha_t 把整块状态统一往下压。有用,但太粗:它假设状态里所有维度以同一个速率遗忘。

KDA 的那一刀:把标量门换成对角门

KDA 的改动就在这里——遗忘门不再是一个标量 αt\alpha_t,而是一个逐通道的向量 at(0,1)dka_t \in (0,1)^{d_k},以对角阵形式作用:

St=Diag(at)St1(Iβtktkt)+βtktvtS_t = \operatorname{Diag}(a_t)\, S_{t-1}\big(I - \beta_t\, k_t k_t^\top\big) + \beta_t\, k_t v_t^\top

含义变了:状态的每一个键通道有自己的遗忘速率。有的通道存的是需要长期挂着的信息(慢遗忘),有的通道是转瞬即逝的局部信号(快遗忘),现在它们可以各走各的。有限大小的 RNN 内存本来就紧张,这种细粒度门控让它把每一位都用在刀刃上——这正是报告里说的"更充分地利用有限状态内存"。

代价几乎为零:ata_t 和原来的 αt\alpha_t 一样由输入投影出来,只是从标量变成向量,多的是一点点投影参数,递推形式没变。

进状态之前,q/k 还要过一道预处理

KDA 的 query、key 在进入上面的递推之前,先走一小段固定流程:

  • 短卷积(short conv):一个小窗口的因果卷积,让每个位置先混入紧邻的几个 token,补上纯 token-wise 投影缺的局部性;
  • Swish 激活;
  • L2 归一化:把 q,kq, k 投到单位球面上,稳住上面那个 delta 更新的数值范围(ktktk_t k_t^\top 的谱不至于炸)。

这几步不新鲜,但它们是让 delta 递推在深层大模型里训得稳的关键配件。

递推怎么并行:分块

写成 RNN 最怕训练时退化成串行。KDA 用分块(chunkwise)来救:把序列切成定长的块,块内用矩阵乘法并行算完,块间只传递那个 dk×dvd_k \times d_v 的状态。于是它对外是线性复杂度的递推,对内又能吃满矩阵乘的并行度——这是这一类现代线性注意力能上规模的前提。

为什么还要留 1/4 的 full attention

线性注意力再好,把历史压进固定状态就必然有损,精确的"大海捞针"式检索是它的软肋。KDA 的解法不是硬扛,而是混合:每 3 层 KDA 配 1 层 full attention(K3 里那层是 MLA,多头潜在注意力)。

  • 3/4 的层是 KDA:常数大小状态,扛住长序列的吞吐和显存;
  • 1/4 的层是 full attention:周期性地恢复一次精确全局回看,把线性层丢掉的检索能力补回来。

据 Kimi Linear 的报告,这个 3 的混合在短上下文、长上下文和 RL scaling 几种设定下都能不输给全 full attention 的基线,同时省掉大部分 KV cache。具体数字以官方报告为准,这里不复述。

从推理成本看它省在哪

full attention 解码时,每生成一个 token 都要把 query 和整条 KV cache 做一遍,访存量随上下文线性增长——长上下文下解码是被 KV cache 的带宽卡住的。KDA 的层没有这条 cache:每步只更新那块固定的 SS,访存量和上下文长度无关。

对纯 CPU、访存受限的部署,这个差别方向上是有利的:把"读一条越来越长的 cache"换成"读写一块定长状态"。但能省多少要看具体形状——状态矩阵 dk×dvd_k \times d_v 本身不小,短上下文时它未必比一小段 cache 便宜。真要上线,按自己的序列长度和硬件实测,别照搬别人的加速比。

踩坑

  • 不是全线性:KDA 单独不构成模型,那 1/4 的 full attention 层不能省,否则长程检索会塌。部署时两种层的 kernel、KV cache 策略都得分别处理。
  • 状态不是免费的dk×dvd_k \times d_v 的状态在每个 KDA 层、每个 batch 都有一份。batch 大、层多时,这块常数内存加起来也可观,只是它不随序列长度涨。
  • L2 归一化别漏:去掉 q,kq,k 的归一化,delta 更新里 kkk k^\top 的尺度会失控,深层容易训崩。这是配件,不是可选项。

小结

KDA 的聪明在于克制:不推翻线性注意力,只把 Gated DeltaNet 的标量遗忘门换成逐通道的对角门,让固定大小的状态内存被更精细地分配。配上 q/k 的短卷积+L2 归一化让它训得稳,分块让它算得快,再用 3 混合 full attention 补上检索短板。结果是一条能替 full attention 挑大梁、又免掉大部分 KV cache 的注意力——这也是 Kimi K3 敢把上下文往百万 token 推的底气之一。

参考:Kimi Linear 技术报告(arXiv.26692)、Kimi K3 技术报告(arXiv.24653)。文中的 delta / gated-delta 递推是这一脉络的通用形式,KDA 的具体贡献是逐通道遗忘门与 q/k 预处理;精确的实现细节与实验数字以官方报告为准。