内容提要
本文介绍Kimi K3模型中KDA(Kimi Delta Attention)的递推形式,从线性注意力、DeltaNet到Gated DeltaNet演进而来。KDA通过逐通道衰减和擦除写入机制更新状态,每头128×128矩阵,不依赖上下文长度。它采用逐头RMSNorm和满秩sigmoid输出门,在合成任务中表现优于基线,并天然具备位置编码能力。
延伸解读
线性注意力谱系:从“加”到“逐通道忘”
文章梳理了线性注意力家族的演进:线性注意力只加不擦,状态无限累积;DeltaNet引入“先擦再写”的delta rule;Gated DeltaNet加上标量遗忘门,但所有通道衰减速率相同;KDA则用对角矩阵实现逐通道衰减。每一步都针对前一步的缺陷,最终让模型能更精细地控制记忆的写入、擦除和遗忘。
KDA状态的内存优势
KDA每层状态固定为96头×128×128矩阵,约3.4MB,69层共约232MB,与上下文长度无关。相比之下,MLA的KV cache随上下文线性增长,1M上下文时约27.6GB。KDA全部状态仅相当于约8K token的MLA缓存,这使得超长上下文推理的内存占用大幅降低。
递推层自带位置编码能力
KDA的递推形式中,旧信息的衰减由数据相关的对角矩阵控制,与RoPE的固定旋转矩阵在形式上相似,但可学习且非正交。因此KDA层本身就能编码位置信息,这解释了K3为何所有MLA层都设为NoPE,且无需额外位置编码即可外推到1M上下文。
Q&A
KDA(Kimi Delta Attention)的状态更新公式是什么?
KDA的状态更新公式为:S_t = (I - β_t k_t k_t^T) Diag(α_t) S_{t-1} + β_t k_t v_t^T,其中α_t是逐通道的遗忘门,β_t是写入强度,k_t和v_t分别是键和值。读出为o_t = S_t^T q_t。
KDA与线性注意力、DeltaNet、Gated DeltaNet有何区别?
线性注意力只加不擦,状态不断累积;DeltaNet引入delta rule,先擦除再写入;Gated DeltaNet增加标量遗忘门,但所有通道共享衰减率;KDA将标量遗忘门替换为对角矩阵,实现逐通道衰减,每个键通道有自己的遗忘率。
KDA的逐通道衰减有什么作用?
逐通道衰减允许每个键通道以不同的速率遗忘,类似于RoPE中每个维度有不同的旋转频率。这增强了模型的位置编码能力,使KDA层本身可学习位置编码,从而支持K3在无位置编码的情况下外推到长上下文。
KDA的输出门与Kimi Linear相比有何改动?
K3将Kimi Linear中的低秩输出门(先降维再升维)改为满秩输出门,即直接使用一个7168到12288的线性层W_g,然后经过sigmoid激活。这一改动去掉了参数量约束,且性能相当。
KDA在解码时每层需要多少状态?
KDA每层解码时携带约1.68M个数值,包括96头×128×128的状态矩阵(1,572,864个)和三路ShortConv的窗口(110,592个),BF16下约3.4 MB。69层总计约232 MB,与上下文长度无关。
KDA在合成任务上的表现如何?
在Palindrome、MQAR和Stack三个合成任务上,KDA在所有长度上准确率最高,且收敛速度明显快于Gated DeltaNet。Mamba2(只有衰减、没有delta rule)在三个任务上全部失败。这表明先擦再写带来精确回忆和状态跟踪,逐通道衰减带来收敛速度。
KDA如何实现位置编码?
KDA的递推形式中,当前token读取历史信息时,中间隔着数据相关的转移矩阵乘积,形式与RoPE相似,但转移矩阵是可学习、非正交的。逐通道的α对应RoPE逐维的频率,因此KDA层本身可学习位置编码,无需额外位置编码。