专栏Kimi K3 模型结构·序列2 / 8
20 min学习

Kimi K3 模型结构(1):序列维度(上),从线性注意力到 KDA

Kimi Delta Attention 的递推形式。先把记忆看成在线学习,从目标函数推出「加、擦、忘」三个动作,再沿两条路径走完线性注意力到 KDA 的谱系(RetNet、Mamba2、GLA、HGRN2、DeltaNet、GDN、RWKV-7 等),用三个合成任务看每一步买来了什么,把状态更新公式的每一项画出来,解释为什么转移矩阵的累乘就是位置编码,最后对上 K3 的真实维度和 K3 在输出门上的第一个改动。

目录21 节

对应论文 §2.1.1 的前半,公式 (1)、(2) 和 (6)。K3 的 93 层里有 69 层是 KDA,它来自 Kimi 团队 2025 年 10 月的 Kimi Linear。这一篇只讲递推形式,也就是「一个 token 进来,状态怎么变」;chunkwise 并行和 K3 的下界衰减放到下一篇。

文本 token图像 / 视频视觉 token 与文本 token 交错后进入同一条主干残差流(prefix sum)重复 23 次3 KDA : 1 Gated MLA第 2 – 92 层每 12 层是一个 AttnRes 块AttnRes 来源(最多 9 个)每个 α 前都有这一组来源α = softmax(wₗ · RMSNorm(来源))wₗ 是每个子层各一个的可学习向量logits输出前再聚合一次全部块最终隐藏状态 + 下一个 token 的 embeddingAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和α词表 163840,隐藏维 7168Token Embedding163840 × 7168从零训练的视觉编码器,hidden 1024,12 头MoonViT-V227 层 · 0.4B · patch 14merge_type = sd2_tpool2×2 Pixel-shuffle+ 时间池化 · token ÷4PatchMergerMLPV2,GELU,RMSNormProjector4096 → 4096 → 7168词嵌入永远是一个来源b₀ = Embedding每个块的输出是块内所有子层输出之和b₁ … bₙ₋₁已完成的块(每块 12 层之和)本块里已经算完的子层之和,相当于块内的普通残差流bₙ⁽ⁱ⁾ 当前块 partial sumKimi Delta Attention:96 头 × 128 维,状态 128×128/头KDA第 1 层第一层不用 MoE,用一个稠密 SiTU-GLU FFNDense FFN仅第 1 层 · 中间维 33792三层 KDA,每层后接一个 Stable LatentMoEKDA×37168 → 3584 latent → 16 个 routed expert → RMSNorm → 7168Stable LatentMoE896 选 16 + 2 sharedDeepSeek 式 MLA,无位置编码,满秩 sigmoid 输出门Gated MLA×1 · NoPE同上Stable LatentMoE896 选 16 + 2 shared主干末尾额外放一层全局注意力Gated MLA第 93 层 · 收尾同上Stable LatentMoE第 93 层最终归一化RMSNorm不与 embedding 共享权重LM Head→ 163840预训练时 1 层多 token 预测;post-training 微调成投机解码的 draftMTP 层 ×1镜像主干 block · 部署时做 EAGLE-3 draft
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
你现在在这里:每个重复块里的三层 KDA,以及第 1 层。

为什么要线性注意力

softmax 注意力的代价是 KV cache 随上下文线性增长。1M 上下文下,即便用 MLA 把每 token 每层压到 576 个数,24 层也要 27 GB 的 BF16 cache,decode 时每生成一个 token 都要把它全读一遍。

线性注意力换了一种记忆方式:不存 token,存一个固定大小的矩阵 StRdk×dvS_t \in \mathbb{R}^{d_k \times d_v}。每个头 128×128,不管上下文多长都是这么大。读出是 ot=Stqto_t = S_t^\top q_t。代价是记忆有损,怎么写、怎么忘、怎么擦,就成了这一族模型全部的设计空间。

这个设计空间看起来很散:每篇论文给一个递推式,符号还不一样。Kimi Linear 论文(Table 7)用一个统一的视角把它们收拢了:SS 看成一组快速权重,每个 token 到来时对它做一步在线学习。不同的模型只是在优化不同的目标函数。先把这个视角讲清楚,谱系就不用背了。

记忆就是在线学习:目标函数从哪来

SS 是一张查找表:行是键通道,列是值通道,用 qq 去查得到 SqS^\top q。第 tt 个 token 带来一对 (kt,vt)(k_t, v_t),意思是「以后拿 ktk_t 来查,应该查到 vtv_t」。问题是:SS 该怎么变?

在线学习的回答:定义一个只看这一个 token 的损失 Lt(S)\mathcal{L}_t(S),衡量当前的 SS 对这一对 (kt,vt)(k_t, v_t) 服务得有多差,然后做一步梯度下降,步长取 1:

St=St1SLt(St1).S_t = S_{t-1} - \nabla_S \mathcal{L}_t(S_{t-1}) .

Kimi Linear Table 7 里所有模型都写成这个形式,差别全在 Lt\mathcal{L}_t 里。Lt\mathcal{L}_t 一共只用到三种项,分别对应三个动作。

加:内积目标。 最朴素的要求是「用 ktk_t 查出来的东西要和 vtv_t 对齐」,写成内积:

Lt(S)=Skt,vt=ktSvt,SLt=ktvt,St=St1+ktvt.\mathcal{L}_t(S) = -\langle S^\top k_t, v_t \rangle = -k_t^\top S v_t , \qquad \nabla_S \mathcal{L}_t = -k_t v_t^\top , \qquad S_t = S_{t-1} + k_t v_t^\top .

这就是 Hebb 规则,也是最早的线性注意力。它的问题在目标函数本身:Lt\mathcal{L}_tSS 是线性的,没有最小值。往 ktvtk_t v_t^\top 方向走多远都能让损失更低,所以每一步只会往上堆,永远不会有「够了」或者「原来存错了,去掉」。查的时候 Stq=i(kiq)viS_t^\top q = \sum_i (k_i^\top q)\, v_i,所有历史值按相似度叠在一起,谁也清不掉。

擦:回归目标。 把「对齐」换成「查出来正好等于」:

Lt(S)=βt2Sktvt2.\mathcal{L}_t(S) = \tfrac{\beta_t}{2}\, \| S^\top k_t - v_t \|^2 .

et=Sktvte_t = S^\top k_t - v_t,这是「记忆现在对 ktk_t 的回答」减去「应该的回答」,也就是残差。梯度是

SLt=βtktet=βtktktSβtktvt,\nabla_S \mathcal{L}_t = \beta_t\, k_t\, e_t^\top = \beta_t\, k_t k_t^\top S - \beta_t\, k_t v_t^\top ,

代进去:

St=St1βtktktSt1+βtktvt=(Iβtktkt)St1+βtktvt.S_t = S_{t-1} - \beta_t k_t k_t^\top S_{t-1} + \beta_t k_t v_t^\top = (I - \beta_t k_t k_t^\top)\, S_{t-1} + \beta_t k_t v_t^\top .

同一个式子有三种读法,都对。第一,它只写入残差 βtkt(vtSt1kt)\beta_t k_t (v_t - S_{t-1}^\top k_t)^\top,如果记忆已经答对了,什么都不写,这是 delta rule 这个名字的来源。第二,它先把 ktk_t 方向上原来存的东西擦掉 βt\beta_t 的比例(IβtktktI - \beta_t k_t k_t^\top,当 kt=1\|k_t\| = 1 时是沿 ktk_t 的一个收缩),再写入新的,这是「先擦再写」。第三,这个目标有最小值,最小值处 Skt=vtS^\top k_t = v_t 精确成立,所以它瞄准的是精确回忆βt(0,1)\beta_t \in (0,1) 是这个样本的权重,也叫写入强度:βt=1\beta_t = 1kt=1\|k_t\| = 1 时,ktk_t 方向上的旧内容被完全替换。

忘:正则项。 内积目标无界,回归目标也只管当前这一个键,SS 的 128 行终归是有限的容量,旧的绑定占着方向不走,新的进来就互相干扰。标准做法是加一个把 SS 往零拉的正则:

Lt(S)+=121αt  SF2=1αt2SF2,S=(1αt)S,\mathcal{L}_t(S) \mathrel{+}= \tfrac12\, \big\| \sqrt{1 - \alpha_t}\; S \big\|_F^2 = \tfrac{1 - \alpha_t}{2}\, \|S\|_F^2 , \qquad \nabla_S = (1 - \alpha_t)\, S ,

这一项在梯度步里贡献 S(1αt)S=αtSS - (1 - \alpha_t) S = \alpha_t S,就是权重衰减。αt(0,1)\alpha_t \in (0,1) 越小忘得越多。让 αt\alpha_t 依赖输入,模型就能逐 token 决定「前面的东西还要不要」,这是 Mamba 一系强调的选择性。再把标量换成逐通道的向量,正则写成 12Diag(1αt)SF2\tfrac12 \|\operatorname{Diag}(\sqrt{1 - \alpha_t})\, S\|_F^2,梯度步变成 Diag(αt)S\operatorname{Diag}(\alpha_t) S,第 jj 行有自己的正则强度。

一个顺序上的细节。 把三项放进同一个目标里一步走完,得到的是 αtSβtktktS+βtktvt\alpha_t S - \beta_t k_t k_t^\top S + \beta_t k_t v_t^\top。GDN 和 KDA 实际用的不是这个,而是先衰减、再对衰减后的状态做回归那一步:令 S~t1=αtSt1\tilde S_{t-1} = \alpha_t S_{t-1}(或 Diag(αt)St1\operatorname{Diag}(\alpha_t) S_{t-1}),目标是 βt2S~t1ktvt2\tfrac{\beta_t}{2}\|\tilde S_{t-1}^\top k_t - v_t\|^2,更新是 St=S~t1S~LtS_t = \tilde S_{t-1} - \nabla_{\tilde S}\mathcal{L}_t。两者差一项 βt(1αt)ktktSt1\beta_t (1 - \alpha_t) k_t k_t^\top S_{t-1}:串行版本是从衰减之后剩下的东西里擦,擦的量和当前状态自洽。下一篇的 chunkwise 形式依赖这个顺序。

有了这三块积木,谱系就只剩两个问题:用了哪几块,衰减的粒度是标量还是逐通道

谱系:加、擦、忘、逐通道忘

「忘」路径:加衰减,不擦「擦」路径:delta rule,不忘旁支+ 标量衰减 α(忘)内积目标 → 回归目标(擦)标量 → 逐通道+ 标量衰减 α(忘)+ delta rule(擦)标量 → 逐通道隐式解擦除 × b放开绑定线性注意力 2020:S = S + k vᵀ线性注意力 2020S = S + k vᵀ只加不擦,目标无下界RetNet 2023 · Mamba2 2024:S = α S + k vᵀRetNet 2023 · Mamba2 2024S = α S + k vᵀα 每头一个标量Mamba2:三个合成任务全挂DeltaNet 2021:S = (I − β k kᵀ) S + β k vᵀDeltaNet 2021S = (I − β k kᵀ) S + β k vᵀ在 k 方向先擦再写Longhorn 2024:β → β / (1 + β kᵀk)Longhorn 2024β → β / (1 + β kᵀk)一步 SGD 换成闭式解GLA 2023 · HGRN2 2024:S = Diag(α) S + k vᵀGLA 2023 · HGRN2 2024S = Diag(α) S + k vᵀα 每个键通道一个Gated DeltaNet 2024:S = (I − β k kᵀ) α S + β k vᵀGated DeltaNet 2024S = (I − β k kᵀ) α S + β k vᵀ先忘,再擦,再写收敛慢于 KDAComba 2025:S = (α − b β k kᵀ) S + β k vᵀComba 2025S = (α − b β k kᵀ) S + β k vᵀ擦得比写得少(b < 1),读出用 q − d kKDA 2025:S = (I − β k kᵀ) Diag(α) S + β k vᵀKDA 2025S = (I − β k kᵀ) Diag(α) S + β k vᵀ先逐通道忘,再擦,再写RWKV-7 2025:S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀRWKV-7 2025S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀ一般 DPLR:擦除键 κ̂、写入键 k̃ 各不相同竖边 = 主线,每条边标注「加了什么」;横边 = 旁支。虚线:KDA 是 RWKV-7 这类一般 DPLR 转移的特例。点击节点跳到对应小节。「忘」路径:加衰减,不擦「擦」路径:delta rule,不忘+ 标量衰减 α(忘)→ 回归目标(擦)标量 → 逐通道+ 标量衰减 α(忘)+ delta rule(擦)标量 → 逐通道隐式解擦除 × b放开绑定线性注意力 2020:S = S + k vᵀ线性注意力 2020S = S + k vᵀ只加不擦,目标无下界RetNet 2023 · Mamba2 2024:S = α S + k vᵀRetNet 2023 · Mamba2 2024S = α S + k vᵀα 每头一个标量Mamba2:三个合成任务全挂DeltaNet 2021:S = (I − β k kᵀ) S + β k vᵀDeltaNet 2021S = (I − β k kᵀ) S + β k vᵀ在 k 方向先擦再写Longhorn 2024:β → β / (1 + β kᵀk)Longhorn 2024β → β / (1 + β kᵀk)一步 SGD 换成闭式解GLA 2023 · HGRN2 2024:S = Diag(α) S + k vᵀGLA 2023 · HGRN2 2024S = Diag(α) S + k vᵀα 每个键通道一个Gated DeltaNet 2024:S = (I − β k kᵀ) α S + β k vᵀGated DeltaNet 2024S = (I − β k kᵀ) α S + β k vᵀ先忘,再擦,再写收敛慢于 KDAComba 2025:S = (α − b β k kᵀ) S + β k vᵀComba 2025S = (α − b β k kᵀ) S + β k vᵀ擦得比写得少(b < 1),读出用 q − d kKDA 2025:S = (I − β k kᵀ) Diag(α) S + β k vᵀKDA 2025S = (I − β k kᵀ) Diag(α) S + β k vᵀ先逐通道忘,再擦,再写RWKV-7 2025:S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀRWKV-7 2025S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀ一般 DPLR:擦除键 κ̂、写入键 k̃ 各不相同实线边标注「加了什么」;虚线框是旁支,挂在它改动的那个节点下面。点击节点跳到对应小节。

从线性注意力出发有两条路。左边一路只加「忘」,先标量后逐通道,始终不擦;右边一路先加「擦」,再补「忘」。两条路在 KDA 汇合:从 GLA 看是加了 delta rule,从 GDN 看是把标量衰减换成了逐通道。右侧三个旁支各自在某个节点上改了一个零件,和 KDA 没有直接的继承关系,但 RWKV-7 值得并排放着看。下面按节点走,主线节点给目标函数和递推,旁支只给递推和一句话。

线性注意力

Katharopoulos 等,2020。上一节的「加」:

Lt=Skt,vt,St=St1+ktvt.\mathcal{L}_t = -\langle S^\top k_t, v_t \rangle , \qquad S_t = S_{t-1} + k_t v_t^\top .

只加不擦不忘。它证明了 softmax 注意力可以换成固定大小的状态加一次递推,但长上下文下旧关联堆积、互相干扰,效果远逊于 softmax。后面每一个模型都是在补这个目标函数的缺陷。

RetNet 和 Mamba2

加「忘」,标量粒度:

Lt=βtSkt,vt+121αt  SF2,St=αtSt1+βtktvt.\mathcal{L}_t = -\beta_t \langle S^\top k_t, v_t \rangle + \tfrac12 \big\| \sqrt{1 - \alpha_t}\; S \big\|_F^2 , \qquad S_t = \alpha_t S_{t-1} + \beta_t k_t v_t^\top .

两者的区别在 α\alpha 从哪来。RetNet(2023)的 α\alpha 是每头一个固定常数,不看输入,βt=1\beta_t = 1,所以它的记忆是纯粹的指数衰减,一个头一个时间尺度。Mamba2(2024)的 αt=exp(ΔtAh)\alpha_t = \exp(-\Delta_t A_h) 由输入决定,写入项也乘了同一个步长 Δt\Delta_t(即上式的 βt\beta_t),模型可以对不重要的 token 选 Δt0\Delta_t \approx 0,既不忘也不写。这一路把「状态该保留多久」交给了数据,但目标函数里仍然是内积项,没有擦除:kk 方向上原来存的东西只能等着淡出,不能被覆盖。下一节的合成任务里 Mamba2 全部失败,问题就出在这里。

GLA 和 HGRN2

「忘」从标量变成逐通道:

Lt=Skt,vt+12Diag(1αt)  SF2,St=Diag(αt)St1+ktvt,αt(0,1)dk.\mathcal{L}_t = -\langle S^\top k_t, v_t \rangle + \tfrac12 \big\| \operatorname{Diag}(\sqrt{1 - \alpha_t})\; S \big\|_F^2 , \qquad S_t = \operatorname{Diag}(\alpha_t)\, S_{t-1} + k_t v_t^\top , \qquad \alpha_t \in (0,1)^{d_k} .

GLA(2023)的 αt\alpha_t 来自一个低秩投影加 sigmoid,每个键通道一个遗忘率,这个参数化被 KDA 直接沿用。HGRN2(2024)是从逐元素门控 RNN 那边走过来的,把一维的状态扩成外积,它的特点是写入门和遗忘门绑在一起:键被换成 1αt1 - \alpha_t,即 St=Diag(αt)St1+(1αt)vtS_t = \operatorname{Diag}(\alpha_t) S_{t-1} + (1 - \alpha_t) v_t^\top,忘得多的通道写得也多,少一组参数。两者都还是内积目标,仍然没有擦除。

DeltaNet

回到起点,换目标而不是加项。Schlag 等 2021 提出,Yang 等 2024 给出并行形式:

Lt=βt2Sktvt2,St=(Iβtktkt)St1+βtktvt.\mathcal{L}_t = \tfrac{\beta_t}{2}\, \| S^\top k_t - v_t \|^2 , \qquad S_t = (I - \beta_t k_t k_t^\top)\, S_{t-1} + \beta_t k_t v_t^\top .

这是上一节推过的「擦」。(Iβtktkt)(I - \beta_t k_t k_t^\top) 是一个广义 Householder 变换,后面 chunkwise 并行能做出来全靠它。它能做精确的键值覆盖,但没有任何遗忘:除非同一个键再来一次把它擦掉,写进去的东西永远留在 SS 里。

Longhorn

DeltaNet 的旁支(2024)。目标一样是回归,但不做一步梯度下降,而是把带正则的单步问题精确解出来(隐式在线学习),结果是把 βt\beta_t 换成 βt/(1+βtktkt)\beta_t / (1 + \beta_t k_t^\top k_t)

St=(Iβt1+βtktktktkt)St1+βtktvt.S_t = \Big( I - \tfrac{\beta_t}{1 + \beta_t k_t^\top k_t}\, k_t k_t^\top \Big) S_{t-1} + \beta_t k_t v_t^\top .

好处是当 kt=1\|k_t\| = 1 时擦除系数 βt/(1+βt)<1\beta_t / (1 + \beta_t) < 1 自动成立,βt\beta_t 不用再限制在 (0,1)(0,1)。没有衰减。

Gated DeltaNet

Yang 等 2024。在 DeltaNet 上加「忘」,标量粒度,并且按上一节末尾说的顺序,先衰减再擦写:

Lt=βt2S~t1ktvt2,S~t1=αtSt1,St=(Iβtktkt)αtSt1+βtktvt.\mathcal{L}_t = \tfrac{\beta_t}{2}\, \| \tilde S_{t-1}^\top k_t - v_t \|^2 ,\quad \tilde S_{t-1} = \alpha_t S_{t-1} , \qquad S_t = (I - \beta_t k_t k_t^\top)\, \alpha_t S_{t-1} + \beta_t k_t v_t^\top .

αt=exp(Ahsoftplus())\alpha_t = \exp(-A_h \cdot \operatorname{softplus}(\cdot)) 每头一个,沿 Mamba2 的参数化;βt\beta_t 是 sigmoid。这是第一个同时有「擦」和「忘」的模型,Kimi Linear 的所有消融都拿它当基线。剩下的问题是一个头的 128 个键通道以同一个速率衰减。

Comba

GDN 的旁支(2025)。两个改动:擦除强度乘一个学习到的标量 b(0,1)b \in (0,1),让擦得比写得少;读出前先把 query 减去 dd 倍的当前键:

St=(αtbβtktkt)St1+βtktvt,ot=St(qtdkt).S_t = (\alpha_t - b\,\beta_t k_t k_t^\top)\, S_{t-1} + \beta_t k_t v_t^\top , \qquad o_t = S_t^\top (q_t - d\, k_t) .

Kimi Linear Table 7 里写的 k^t\hat k_t 就是 bktb\, k_t。衰减仍是每头一个标量。Kimi Linear 的 chunkwise 推导里借用了 Comba 对转移矩阵累乘的写法,下一篇会碰到。

KDA

Kimi Linear,2025。把 GDN 的标量换成对角矩阵,或者说给 GLA 加上 delta rule:

Lt=βt2S~t1ktvt2,S~t1=Diag(αt)St1,\mathcal{L}_t = \tfrac{\beta_t}{2}\, \| \tilde S_{t-1}^\top k_t - v_t \|^2 ,\quad \tilde S_{t-1} = \operatorname{Diag}(\alpha_t)\, S_{t-1} , St=(Iβtktkt)Diag(αt)St1+βtktvt,o~t=Stqt.(K3 Eq. 1)S_t = \big(I - \beta_t k_t k_t^\top\big)\,\operatorname{Diag}(\alpha_t)\, S_{t-1} + \beta_t k_t v_t^\top , \qquad \tilde o_t = S_t^\top q_t . \tag{K3 Eq. 1}

αt(0,1)dk\alpha_t \in (0,1)^{d_k}βt(0,1)\beta_t \in (0,1)。逐通道衰减的动机不只是「更细的粒度」,Kimi Linear 给的理由来自位置编码:RoPE 的力量在于每个维度有自己的旋转频率,标量衰减没有这种逐维的多样性,如果要让递推层承担位置信息,逐通道的门是自然的对应物。这个类比在「转移矩阵的累乘就是位置编码」一节里展开。

RWKV-7

和 KDA 并排的兄弟(2025):同样是逐通道衰减加擦除,但擦除和写入用的不是同一个键。按本文的列向量约定写:

St=(Diag(wt)(atκ^t)κ^t)St1+k~tvt.S_t = \big( \operatorname{Diag}(w_t) - (a_t \odot \hat\kappa_t)\, \hat\kappa_t^\top \big)\, S_{t-1} + \tilde k_t v_t^\top .

κ^t\hat\kappa_t 是把 ktk_t 逐通道乘一个学习到的尺度再做 L2 归一化得到的擦除键at(0,1)dka_t \in (0,1)^{d_k} 是逐通道的「上下文学习率」,替代了 KDA 里每头一个的标量 βt\beta_tk~t=ktlerp(1,at,)\tilde k_t = k_t \odot \operatorname{lerp}(1, a_t, \cdot) 是另一个写入键wt(0.55,1)w_t \in (0.55, 1) 逐通道衰减。它的转移矩阵是一般的「对角加低秩」DabD - a b^\topaabb 各自独立。KDA 是 aabb 都绑在 ktk_t 上、写入键也等于 ktk_t 的特例。表达能力上 KDA 让了一步,换来的是内核少算一半东西,下一篇讲。

一句话总结

线性注意力是「加」,RetNet / Mamba2 是「加,然后所有东西以同一个速率淡出」,GLA / HGRN2 是「加,然后每个通道以自己的速率淡出」,DeltaNet 是「在这个键下先擦再写」,GDN 是「先擦再写,然后所有东西以同一个速率淡出」,KDA 是「先擦再写,然后每个特征通道以自己的速率淡出」。按两个坐标摆开:

无衰减标量衰减逐通道(对角)衰减
无 delta rule线性注意力RetNet、Mamba2GLA、HGRN2
有 delta ruleDeltaNet、LonghornGated DeltaNet、CombaKDA、RWKV-7

递推形式能做什么:三个合成任务

谱系走完,自然要问:「擦」和「逐通道忘」这两步各买来了什么?语言模型的 PPL 差异只有零点几,看不出机制。Kimi Linear §5.1 用三个合成任务分别考察状态的三种能力,模型只有 2 层 2 头、头维 128,序列长度 256 到 2048,比较 KDA、GDN 和 Mamba2。三个任务分别对应状态必须做到的三件事。

Palindrome,考精确复制。 输入一串随机 token,输出它的逆序:

text
a b c d e | → e d c b a

状态要把每个 token 分别存下来,还要能按位置精确取回,一个都不能糊。只做叠加和衰减的记忆读出来的是所有 token 的加权和,长度一上去就分不开。

MQAR,多查询关联回忆。 先给一串键值对,后面用键去查:

text
A 4 · B 7 · C 2 · … · A ? → 4 · C ? → 2

状态要同时保持很多组 kvk \to v 绑定互不干扰,而且如果同一个键后来又出现了新值,旧值必须被覆盖。这正是回归目标做的事:查出来不等于 vv 就把差值写进去。这个任务被认为和语言模型质量强相关。

Stack,考状态跟踪。 64 个独立的栈,push 和 pop 交错:

text
push 3 a · push 3 b · push 7 x · pop 3 → b · pop 3 → a

同一个栈 id 反复被写,pop 之后要露出上一层。这是最纯粹的「同一个键下擦掉再写」,没有擦除就没有办法让 pop 3 的两次回答不同。

结果:KDA 在三个任务、所有长度上准确率最高,在 Palindrome 和 MQAR 上收敛明显快于 GDN。Mamba2 三个任务全部失败。论文没有逐任务分析机制,但对着更新规则看是清楚的:Mamba2 和 KDA 的差别是有没有「擦」,它全挂,说明擦除是精确回忆和状态跟踪的前提;GDN 和 KDA 的差别只在衰减的粒度,它慢,说明逐通道衰减买来的是收敛速度。这两句话就是流程图上两条汇入 KDA 的边的实验注脚。

公式 (1) 的图解

S_t新状态k_tI − β_t k_t k_tᵀ② 擦除Diag(α_t)① 衰减S_{t−1}旧状态k_tv_tᵀβ_t k_t v_tᵀ③ 写入=··+从右往左读:先把旧状态每一行按 α_t 各自衰减,再把落在 k_t 方向上的内容擦掉 β_t 的比例,最后写入 β_t 倍的新关联 k_t → v_t。读出:õ_t = S_tᵀ q_t。行 = 键通道(d_k = 128,各自的衰减率),列 = 值通道(d_v = 128)。

三个操作的顺序有讲究:先衰减,再擦除,再写入。展开后是

St=(Diag(αt)βtktktDiag(αt))St1+βtktvt,S_t = \big(\operatorname{Diag}(\alpha_t) - \beta_t k_t k_t^\top \operatorname{Diag}(\alpha_t)\big) S_{t-1} + \beta_t k_t v_t^\top ,

即一般的「对角加低秩」(DPLR)转移 DatbtD - a_t b_t^\top,其中 D=Diag(αt)D = \operatorname{Diag}(\alpha_t)at=βtkta_t = \beta_t k_tbt=ktαtb_t = k_t \odot \alpha_t。两个低秩向量都绑在 ktk_t 上,不像 RWKV-7 那样各自独立。这个绑定看起来是表达能力上的让步,但它让内核少算一半东西,是 Kimi Linear 的 KDA 内核比通用 DPLR 内核快约两倍的原因。下一篇细说。

对状态矩阵的直觉:行是键通道,列是值通道Diag(αt)\operatorname{Diag}(\alpha_t) 左乘,所以衰减是按行做的,第 jj 个键通道存的所有值一起以 αt,j\alpha_{t,j} 的速率淡出。擦除沿 ktk_t 这一个方向。写入是一个秩 1 的外积。

转移矩阵的累乘就是位置编码

K3 的 24 层 MLA 全是 NoPE,整个模型里没有 RoPE。一个不带位置信息的注意力层对 token 的排列是等变的,它分不清「A 在 B 前面」和「B 在 A 前面」。那顺序从哪来?答案是:从 KDA 的递推里来。把公式 (1) 展开一遍就能看到。

记每一步的转移矩阵为 Tj=(Iβjkjkj)Diag(αj)T_j = (I - \beta_j k_j k_j^\top)\operatorname{Diag}(\alpha_j),公式 (1) 就是 St=TtSt1+βtktvtS_t = T_t S_{t-1} + \beta_t k_t v_t^\top。从 S0=0S_0 = 0 开始往前套:

S1=β1k1v1,S2=T2β1k1v1+β2k2v2,St=i=1t(TtTt1Ti+1)βikivi.S_1 = \beta_1 k_1 v_1^\top ,\qquad S_2 = T_2\, \beta_1 k_1 v_1^\top + \beta_2 k_2 v_2^\top ,\qquad S_t = \sum_{i=1}^{t} \big( T_t T_{t-1} \cdots T_{i+1} \big)\, \beta_i k_i v_i^\top .

ii 个 token 写进去的 kivik_i v_i^\top,到第 tt 步时已经被从 i+1i+1tt 的每一个转移矩阵各乘了一遍。用 qtq_t 读出:

o~t=Stqt=i=1t(qt(TtTi+1)ki)st,iβivi.\tilde o_t = S_t^\top q_t = \sum_{i=1}^{t} \underbrace{\Big( q_t^\top \big( T_t \cdots T_{i+1} \big)\, k_i \Big)}_{s_{t,i}}\, \beta_i v_i .

这和注意力的形状完全一样:对每个历史位置 ii 算一个分数 st,is_{t,i},用它给 viv_i 加权求和。区别在于 qtq_tkik_i 中间夹着一串矩阵,iitt 隔得越远,夹的越多。

RoPE 那边做的其实是同一件事。RoPE 把位置 tt 的 query 和位置 ii 的 key 各旋转一次,q~t=Rtqt\tilde q_t = R^t q_tk~i=Riki\tilde k_i = R^i k_i,其中 RR 是一个固定的分块对角旋转矩阵,第 dd 个二维块转 θd\theta_d 角。打分时

st,i=q~tk~i=qt(Rt)Riki=qt(R1)tiki=qt(j=i+1tR1)ki.s_{t,i} = \tilde q_t^\top \tilde k_i = q_t^\top (R^t)^\top R^i\, k_i = q_t^\top \big( R^{-1} \big)^{t-i} k_i = q_t^\top \Big( \prod_{j=i+1}^{t} R^{-1} \Big) k_i .

也是 qtq_tkik_i 中间夹一串矩阵,只不过每一个都是同一个 R1R^{-1},夹几个只取决于距离 tit - i

KDA数据相关、非正交、逐通道衰减k_iTi+1Ti+2Tt−1Ttq_tRoPE固定、正交、逐维频率k_iRRRRq_tT_j = (I − β_j k_j k_jᵀ) Diag(α_j):每一格由第 j 个 token 决定,颜色深浅示意该步保留了多少。R:同一个分块旋转矩阵累乘 t − i 次,只和距离有关,和内容无关。

两条链并排看,差别有三处。第一,RoPE 的格子是固定的,KDA 的格子由第 jj 个 token 的 αj\alpha_jβj\beta_jkjk_j 决定,是可学习、数据相关的。第二,RoPE 的 RR 是正交矩阵,只旋转不缩放,分数随距离震荡但不衰减;KDA 的 Diag(αj)\operatorname{Diag}(\alpha_j) 元素在 (0,1)(0,1) 里,是收缩的,分数随距离衰减,天然偏向近期,再加一个 RoPE 里没有对应物的擦除因子。第三,也是 KDA 比 GDN 多出来的那一步的理由:RoPE 每个维度有自己的频率 θd\theta_d,低频维度看长程、高频维度看短程,这种逐维的多样性是它有效的关键;GDN 的标量 αj\alpha_j 相当于所有维度共用一个频率,KDA 的逐通道 αj\alpha_j 才是逐维频率的对应物。Kimi Linear 论文就是从这里出发,把 GDN 的门换成了逐通道。

有了这个对应,K3 的两个决定就顺了。第一,每层 MLA 前面都有三层 KDA 在往残差流里写位置敏感的内容,MLA 自己不需要位置编码,所以 24 层 MLA 全部 NoPE。第二,不存在需要为长上下文重调的 RoPE 频率基或 YaRN 之类的插值,K3 §3.4 从 8K 一路扩到 1M,不改任何位置编码参数。Kimi Linear Table 5 给了反面证据:把 MLA 层加回 RoPE 的变体在短上下文上打平,128K 长上下文的 RULER 从 84.3 掉到 78.8。第 3 篇讲 MLA 时再展开 NoPE 带来的工程收益。

参数化:从 xtq,k,v,α,β

论文公式 (2),每个头 hh

qth,kth=L2Norm(Swish(ShortConv(Wq/khxt)))Rdkvth=Swish(ShortConv(Wvhxt))Rdvβth=Sigmoid(Wβhxt)(0,1)zth=WαWαxt+bαhRdk\begin{aligned} q_t^h, k_t^h &= \operatorname{L_2Norm}\big(\operatorname{Swish}(\operatorname{ShortConv}(W_{q/k}^h x_t))\big) \in \mathbb{R}^{d_k} \\ v_t^h &= \operatorname{Swish}(\operatorname{ShortConv}(W_v^h x_t)) \in \mathbb{R}^{d_v} \\ \beta_t^h &= \operatorname{Sigmoid}(W_\beta^h x_t) \in (0,1) \\ z_t^h &= W_\alpha^{\uparrow} W_\alpha^{\downarrow} x_t + b_\alpha^h \in \mathbb{R}^{d_k} \end{aligned}
x_t7168Linear W_q7168 → 12288ShortConv(4) + SwishL2Norm逐头q_t 96 × 128Linear W_k7168 → 12288ShortConv(4) + SwishL2Norm逐头k_t 96 × 128Linear W_v7168 → 12288ShortConv(4) + Swishv_t 96 × 128Linear W_β7168 → 96Sigmoidβ_t 96 个标量Linear W_α↓7168 → 128Linear W_α↑ + b_α128 → 12288g = g_min · σ(e^A z)有下界,g_min = −5α_t = e^{g} 96 × 128Linear W_g(满秩)7168 → 12288SigmoidKDA 递推S_t = (I − β k kᵀ) Diag(α) S_{t−1}+ β k vᵀõ_t = S_tᵀ q_tS96 头 × (128 × 128)训练 / prefill:chunkwisedecode:逐 token 递推逐头 RMSNorm128 维,每头独立输出门 96 × 128,∈ (0,1)Linear W_o12288 → 7168y_t 7168
一层 KDA。粉色的那一行是 K3 的两个改动所在:衰减映射换成有下界的 scaled sigmoid,输出门从低秩改成满秩。

逐项说明,维度按 K3 的 d=7168d = 7168、96 头、dk=dv=128d_k = d_v = 128

  • ShortConv。 深度可分的因果卷积,核大小 4(short_conv_kernel_size),后接 Swish,沿用 GDN 的做法。作用是给每个 token 的 q、k、v 掺进前三个 token 的信息,Kimi Linear 的消融显示去掉它 PPL 从 5.65 变 5.70。decode 时每层要为 q、k、v 各缓存最近 3 个 token 的投影,这是 KDA 除了 SS 之外唯一的状态。
  • q、k 的 L2 归一化。kt=1\|k_t\| = 1,于是 IβtktktI - \beta_t k_t k_t^\top 的特征值落在 [1βt,1][1-\beta_t, 1],递推不会爆。
  • β\beta 7168 → 96 的线性层加 sigmoid,每头一个标量。代码里叫 b_proj
  • α\alpha 的 logit zz 低秩投影 7168 → 128 → 12288(f_a_projf_b_proj),加一个 12288 维的偏置 bαb_\alphadt_bias),得到每个头每个键通道一个 logit。从 logit 到 (0,1)(0,1) 里的 α\alpha 还要过一个映射,Kimi Linear 用 GDN 的负 softplus,K3 换成了有下界的 scaled sigmoid,还有一个每头一个的可学习 log 尺度 AhA_hA_log)。这是 K3 的改动之一,放到下一篇讲,因为它的动机在 chunkwise 形式里。

输出:逐头 RMSNorm、门、投影

递推读出 o~t=Stqt\tilde o_t = S_t^\top q_t 之后,论文公式 (6):

yt=Wo[Sigmoid(Wgxt)RMSNorm(o~t)].y_t = W_o \big[\operatorname{Sigmoid}(W_g x_t) \odot \operatorname{RMSNorm}(\tilde o_t)\big].

RMSNorm 是逐头的,每头独立在 128 维上归一化。门是 sigmoid,作用在归一化之后的每个通道上,然后才是 WoW_o(12288 → 7168)。代码里归一化和门融在一个 FusedRMSNormGated(activation='sigmoid') 里。

K3 的改动:门从低秩改满秩。 Kimi Linear 的门是 Sigmoid(WgWgxt)\operatorname{Sigmoid}(W_g^{\uparrow} W_g^{\downarrow} x_t),7168 → 128 → 12288,和遗忘门一样走低秩。Kimi Linear 论文自己说了原因:为了和基线做公平的参数量对比,并且「和满秩门性能相当」。K3 不需要这个约束,直接用 7168 → 12288 的满秩 WgW_guse_full_rank_gate = true,代码里 g_proj)。MLA 那边的输出门也一起改成满秩,下下篇会看到同样的公式。

python
# KimiDeltaAttention.forward 的收尾(HF modeling_kimi_linear.py)
if self.use_full_rank_gate:
    g = self.g_proj(hidden_states)                  # 7168 -> 12288
else:
    g = self.g_b_proj(self.g_a_proj(hidden_states)) # 7168 -> 128 -> 12288(Kimi Linear)
g = rearrange(g, '... (h d) -> ... h d', d=self.head_dim)
o = self.o_norm(o, g)                               # 逐头 RMSNorm,再乘 sigmoid(g)
o = self.o_proj(rearrange(o, 'b t h d -> b t (h d)'))

为什么是 sigmoid 门。Kimi Linear 的消融:sigmoid 门 PPL 5.65,无门 5.67,GDN 默认的 Swish 门 5.81。sigmoid 输出在 (0,1)(0,1),能把某些通道彻底关掉,引用的是 Qwen 团队关于 gated attention 的工作:门带来非线性和稀疏性,并缓解 attention sink。

每层的状态有多大

一层 KDA 在 decode 时携带的状态:

状态形状数量
SS96 头 × 128 × 1281,572,864
三路 ShortConv 的窗口3 × 12288 × 3 个 token110,592

约 1.68M 个数,BF16 下约 3.4 MB。69 层加起来约 232 MB,和上下文长度无关。对照 MLA:每 token 每层 576 个数,24 层 13,824 个数,1M 上下文约 27.6 GB。KDA 全部 69 层的状态相当于大约 8K 个 token 的 MLA cache。完整的账在第 3 篇。

术语坑

  • α\alpha 在两篇论文里是同一个东西,但 Kimi Linear 写 αt=f()\alpha_t = f(\cdot) 没有给出 ff,K3 把映射写全了(公式 (5)),并且改了它。
  • β\beta 是学习率还是样本权重。 DeltaNet 原文把 βt\beta_t 写成 SGD 的学习率,St=St1βtLS_t = S_{t-1} - \beta_t \nabla \mathcal{L};Kimi Linear Table 7 把它放进目标函数、步长取 1。两种写法得到同一个递推,本文跟 Table 7。
  • 头数和头维。 KDA 和 MLA 共用 num_attention_heads = 96head_dim = 128,但 MLA 的 query 和 key 实际是 192 维(128 加 64),下下篇解释。
  • 状态是 dk×dvd_k \times d_v 还是转置。 论文按 SRdk×dvS \in \mathbb{R}^{d_k \times d_v} 写,读出是 SqS^\top q;FLA 内核里有 transpose_state_layout=True,存的是转置,只是内存布局的事。RWKV-7 和 Comba 的原文用行向量约定,转移矩阵写在 St1S_{t-1} 右边,本文全部转成了列向量约定。
  • 打分式里因子的顺序。 Kimi Linear 公式 (12) 把每一步的转移写成 Aj(Iβjkjkj)A_j (I - \beta_j k_j k_j^\top),作用到右边的 kik_i 上是先擦再衰减;本文按公式 (1) 写成 (Iβjkjkj)Diag(αj)(I - \beta_j k_j k_j^\top)\operatorname{Diag}(\alpha_j),先衰减再擦。两个矩阵不相等,真正实现的是公式 (1) 的顺序,公式 (12) 只是在说明「累乘」这个形状。

下一篇

递推形式一个 token 一步,训练时没法用 Tensor Core。下一篇讲 Kimi Linear 的 chunkwise 并行形式(WY 表示、UT 变换、块间递推加块内并行),然后就能看懂 K3 为什么要给衰减加一个 5-5 的下界,以及这个改动怎么把对角 tile 也搬上 Tensor Core。