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

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

Kimi Delta Attention 的递推形式。把记忆看成一步在线学习,「加、擦、忘」三个动作各对应目标函数里的一项,拼起来就是 KDA 的状态更新公式;沿谱系图走一遍每一步买来了什么,用三个合成任务验证,再解释为什么转移矩阵的累乘就是位置编码,最后对上 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词嵌入永远是一个来源= Embedding每个块的输出是块内所有子层输出之和已完成的块(每块 12 层之和)本块里已经算完的子层之和,相当于块内的普通残差流当前块 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,每生成一个 token 都要全读一遍。

线性注意力能把这个 cache 压成每头一个 128×128 的矩阵,原因只有一步改写:

  • 两种注意力都是 ot=∑i≤twt,i vio_t = \sum_{i\le t} w_{t,i}\, v_i,区别只在权重。softmax 的权重是 exp⁡(qt⋅ki)\exp(q_t\cdot k_i) 再归一化,线性注意力直接用裸点积 wt,i=qt⋅kiw_{t,i} = q_t\cdot k_i。
  • 线性的一项 (qt⋅ki) vi(q_t\cdot k_i)\,v_i 可以改写成 qt⊤(kivi⊤)q_t^\top (k_i v_i^\top)。查询跑到了括号外面,括号里只依赖 token ii,于是求和可以搬进去:ot=qt⊤∑i≤tkivi⊤=qt⊤Sto_t = q_t^\top \sum_{i\le t} k_i v_i^\top = q_t^\top S_t。每个 token 到来时把自己的外积加进 SS,不需要知道以后谁来查。
  • softmax 做不到这一步:exp⁡(qt⋅ki)\exp(q_t\cdot k_i) 拆不成 f(qt)⊤g(ki)f(q_t)^\top g(k_i),分母又是对所有 jj 求和,查询被钉在权重内部,只能等查询来了再对每个 key 算一次。这就是 KV cache 存在的原因。
softmax:查询在权重里面分母:也含,且要遍历全部的指数和分母都钉在上,求和搬不进去只能等来了再对每个算一次,所以要一直存着这就是 KV cache线性:查询在求和外面128 × 128128 × 128128 × 128128 × 128每个 token 到来时把自己的外积加进去,不必知道以后谁来查16,384 个滑动求和本身不存在任何地方只和连一根线,读出精确,每步代价固定
同样是对过去的 value 加权求和。softmax 的权重里有查询,求和只能等查询来了再做,所以每个 key、value 都要留着;线性注意力的一项 可以写成 ,括号里只依赖 token ,可以提前累加。

代价是记忆有损:128 维空间里最多 128 个正交方向,序列一过 128 个 token 键必然重叠,取出来的总是一堆相似键下面的值的混合。

St∈Rdk×dvS_t \in \mathbb{R}^{d_k \times d_v} 就是这一族模型的记忆,读出是 ot=St⊤qto_t = S_t^\top q_t。对一张固定大小的表,能做的操作只有三种:往里加一条新记录,把某个键下的旧记录擦掉再写,让所有记录随时间忘掉。这三个动作各自做不做、做到什么粒度,就是从线性注意力到 KDA 的全部设计空间。下文一直用「加、擦、忘」指代它们。

一步在线学习:公式 (1) 从哪来

SS 是一张查找表:行是键通道,列是值通道,用 qq 去查得到 S⊤qS^\top q。第 tt 个 token 带来一对 (kt,vt)(k_t, v_t),意思是「以后拿 ktk_t 来查,应该查到 vtv_t」。Kimi Linear(Table 7)用一个统一的视角回答「SS 该怎么变」:写一个只看当前 token 的损失 Lt(S)\mathcal{L}_t(S),做一步梯度下降,步长取 1:

St=St−1−∇SLt(St−1).S_t = S_{t-1} - \nabla_S \mathcal{L}_t(S_{t-1}) .

等号右边只有 St−1S_{t-1} 和当前的 (kt,vt)(k_t, v_t),所以这就是一个 RNN,标题里的「递推形式」指的就是它。三个动作各对应损失里的一项,梯度对 SS 都是线性的,所以每一步都能整理成 St=AtSt−1+BtS_t = A_t S_{t-1} + B_t:

动作目标项梯度 ∇SLt\nabla_S \mathcal{L}_t落在递推里的位置
加−⟨S⊤kt,vt⟩-\langle S^\top k_t, v_t\rangle−ktvt⊤-k_t v_t^\topBt=ktvt⊤B_t = k_t v_t^\top
擦βt2∥S⊤kt−vt∥2\tfrac{\beta_t}{2}\lVert S^\top k_t - v_t\rVert^2βtktkt⊤S−βtktvt⊤\beta_t k_t k_t^\top S - \beta_t k_t v_t^\topAtA_t 里减 βtktkt⊤\beta_t k_t k_t^\top,Bt=βtktvt⊤B_t = \beta_t k_t v_t^\top
忘12∥Diag⁡(1−αt) S∥F2\tfrac12\lVert\operatorname{Diag}(\sqrt{1-\alpha_t})\, S\rVert_F^2Diag⁡(1−αt) S\operatorname{Diag}(1-\alpha_t)\, SAtA_t 里乘 Diag⁡(αt)\operatorname{Diag}(\alpha_t)

「加」的目标对 SS 是线性的,没有最小值,只会往上堆。「擦」把它换成回归目标,最小值处 S⊤kt=vtS^\top k_t = v_t 精确成立,梯度里自带写入项,所以「加」和「擦」不会同时出现。「忘」是把 SS 往零拉的正则,αt\alpha_t 取标量就是所有通道同速衰减,取向量就是每个键通道各有自己的保留率。

KDA 取「擦 + 逐通道忘」,并且先衰减、再对衰减后的状态做回归那一步:

St=(I−βtktkt⊤) Diag⁡(αt) St−1+βtktvt⊤,o~t=St⊤qt,(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)。先衰减再擦,意思是擦的量按衰减后剩下的状态算,和一步到位地把两项梯度相加差一项 βt(1−αt)ktkt⊤St−1\beta_t(1-\alpha_t)k_t k_t^\top S_{t-1}。GDN 和 KDA 用的都是这个顺序,下一篇的 chunkwise 形式依赖它。

三项梯度的完整推导

对矩阵求梯度只用一招:把损失展开成对每个元素 SijS_{ij} 的求和,对 SijS_{ij} 求偏导,再按 (i,j)(i,j) 摆回矩阵。反复用到 (S⊤k)j=∑iSijki(S^\top k)_j = \sum_i S_{ij} k_i。

加。

Lt(S)=−⟨S⊤kt,vt⟩=−∑j(∑iSij kt,i)vt,j,∂Lt∂Sij=−kt,ivt,j,∇SLt=−ktvt⊤.\mathcal{L}_t(S) = -\langle S^\top k_t, v_t \rangle = -\sum_{j}\Big(\sum_i S_{ij}\, k_{t,i}\Big) v_{t,j} , \qquad \frac{\partial \mathcal{L}_t}{\partial S_{ij}} = -k_{t,i} v_{t,j} , \qquad \nabla_S \mathcal{L}_t = -k_t v_t^\top .

代回梯度步:St=St−1+ktvt⊤S_t = S_{t-1} + k_t v_t^\top,即 At=IA_t = I,Bt=ktvt⊤B_t = k_t v_t^\top。这是 Hebb 规则,也是最早的线性注意力。

擦。 记残差 e=S⊤kt−vte = S^\top k_t - v_t,ej=∑iSijkt,i−vt,je_j = \sum_i S_{ij} k_{t,i} - v_{t,j}:

Lt(S)=βt2∑jej2,∂Lt∂Sij=βt ej kt,i  ⇒  ∇SLt=βt kte⊤=βt ktkt⊤S−βt ktvt⊤.\mathcal{L}_t(S) = \tfrac{\beta_t}{2} \sum_j e_j^2 , \qquad \frac{\partial \mathcal{L}_t}{\partial S_{ij}} = \beta_t\, e_j\, k_{t,i} \;\Rightarrow\; \nabla_S \mathcal{L}_t = \beta_t\, k_t e^\top = \beta_t\, k_t k_t^\top S - \beta_t\, k_t v_t^\top .

第二个等号是把 e⊤=kt⊤S−vt⊤e^\top = k_t^\top S - v_t^\top 代进去。代回梯度步并把含 St−1S_{t-1} 的项归到一起:

St=(I−βtktkt⊤)⏟At St−1+βtktvt⊤⏟Bt.S_t = \underbrace{(I - \beta_t k_t k_t^\top)}_{A_t}\, S_{t-1} + \underbrace{\beta_t k_t v_t^\top}_{B_t} .

同一个式子的三种读法:只写入残差 βtkt(vt−St−1⊤kt)⊤\beta_t k_t (v_t - S_{t-1}^\top k_t)^\top,记忆已经答对就什么都不写,这是 delta rule 名字的来源;先把 ktk_t 方向上原来存的东西擦掉 βt\beta_t 的比例再写新的,这是「先擦再写」;∥kt∥=1\|k_t\| = 1 时 I−βtktkt⊤I - \beta_t k_t k_t^\top 是沿 ktk_t 的一个收缩,βt=1\beta_t = 1 时完全替换。

忘。 正则按元素独立,逐通道版本第 ii 行有自己的系数:

Lt(S)+=12∑ij(1−αt,i)Sij2,∂Lt∂Sij=(1−αt,i)Sij,∇SLt=Diag⁡(1−αt) S.\mathcal{L}_t(S) \mathrel{+}= \tfrac12 \sum_{ij} (1 - \alpha_{t,i}) S_{ij}^2 , \qquad \frac{\partial \mathcal{L}_t}{\partial S_{ij}} = (1 - \alpha_{t,i}) S_{ij} , \qquad \nabla_S \mathcal{L}_t = \operatorname{Diag}(1 - \alpha_t)\, S .

代回梯度步:S−Diag⁡(1−αt)S=Diag⁡(αt)SS - \operatorname{Diag}(1-\alpha_t) S = \operatorname{Diag}(\alpha_t) S。标量版本把 αt,i\alpha_{t,i} 换成同一个 αt\alpha_t 即可。让 αt\alpha_t 依赖输入,模型就能逐 token 决定「前面的东西还要不要」,这是 Mamba 一系强调的选择性。

拼起来。 「擦 + 忘」放进同一个目标,梯度相加,一步走完:

St=St−1−(βtktkt⊤St−1−βtktvt⊤)−Diag⁡(1−αt)St−1=(Diag⁡(αt)−βtktkt⊤)St−1+βtktvt⊤.S_t = S_{t-1} - \big(\beta_t k_t k_t^\top S_{t-1} - \beta_t k_t v_t^\top\big) - \operatorname{Diag}(1 - \alpha_t) S_{t-1} = \big(\operatorname{Diag}(\alpha_t) - \beta_t k_t k_t^\top\big) S_{t-1} + \beta_t k_t v_t^\top .

GDN 和 KDA 实际用的是串行版本:令 S~t−1=Diag⁡(αt)St−1\tilde S_{t-1} = \operatorname{Diag}(\alpha_t) S_{t-1},目标是 βt2∥S~t−1⊤kt−vt∥2\tfrac{\beta_t}{2}\|\tilde S_{t-1}^\top k_t - v_t\|^2,更新 St=(I−βtktkt⊤) S~t−1+βtktvt⊤S_t = (I - \beta_t k_t k_t^\top)\,\tilde S_{t-1} + \beta_t k_t v_t^\top,展开即公式 (1)。两者差一项 βt(1−αt)ktkt⊤St−1\beta_t (1 - \alpha_t) k_t k_t^\top S_{t-1}。

谱系:每一步加了哪块积木

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

从线性注意力出发有两条路:左边只加「忘」,先标量后逐通道,始终不擦;右边先加「擦」,再补「忘」。两条路在 KDA 汇合。图里已经写了每个节点的递推式,下面每个节点只说它改了什么、还缺什么。

线性注意力

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

RetNet 和 Mamba2

加「忘」,标量粒度:St=αtSt−1+βtktvt⊤S_t = \alpha_t S_{t-1} + \beta_t k_t v_t^\top。RetNet(2023)的 α\alpha 是每头一个固定常数,一个头一个时间尺度。Mamba2(2024)的 αt=exp⁡(−ΔtAh)\alpha_t = \exp(-\Delta_t A_h) 由输入决定,写入项也乘同一个步长,不重要的 token 可以既不忘也不写。两者都没有擦除:kk 方向上原来存的东西只能等着淡出,不能被覆盖。下一节的合成任务里 Mamba2 全部失败,问题就出在这里。

GLA 和 HGRN2

「忘」从标量变成逐通道:St=Diag⁡(αt)St−1+ktvt⊤S_t = \operatorname{Diag}(\alpha_t) S_{t-1} + k_t v_t^\top,αt∈(0,1)dk\alpha_t \in (0,1)^{d_k}。GLA(2023)的 αt\alpha_t 来自低秩投影加 sigmoid,这个参数化被 KDA 直接沿用。HGRN2(2024)把写入门和遗忘门绑在一起,键换成 1−αt1 - \alpha_t,少一组参数。仍然没有擦除。

DeltaNet

Schlag 等 2021 提出,Yang 等 2024 给出并行形式。换目标而不是加项:回归目标,即上一节的「擦」。(I−βtktkt⊤)(I - \beta_t k_t k_t^\top) 是一个广义 Householder 变换,后面 chunkwise 并行全靠它。能做精确的键值覆盖,但没有任何遗忘:除非同一个键再来一次,写进去的东西永远留在 SS 里。

Longhorn

DeltaNet 的旁支(2024)。不做一步梯度下降,而是把带正则的单步问题精确解出来,结果是 βt\beta_t 换成 βt/(1+βtkt⊤kt)\beta_t / (1 + \beta_t k_t^\top k_t),好处是 βt\beta_t 不用再限制在 (0,1)(0,1)。没有衰减。

Gated DeltaNet

Yang 等 2024。DeltaNet 加标量「忘」,顺序是先衰减再擦写。αt\alpha_t 每头一个,沿 Mamba2 的参数化;βt\beta_t 是 sigmoid。这是第一个同时有「擦」和「忘」的模型,Kimi Linear 的所有消融都拿它当基线。剩下的问题是一个头的 128 个键通道以同一个速率衰减。

Comba

GDN 的旁支(2025)。擦除强度乘一个学习到的标量 b∈(0,1)b \in (0,1),让擦得比写得少;读出前把 query 减去 dd 倍的当前键。衰减仍是每头一个标量。Kimi Linear 的 chunkwise 推导借用了它对转移矩阵累乘的写法,下一篇会碰到。

KDA

Kimi Linear,2025。把 GDN 的标量换成对角矩阵,或者说给 GLA 加上 delta rule,即公式 (1)。逐通道衰减的理由来自位置编码:RoPE 的力量在于每个维度有自己的旋转频率,标量衰减没有这种逐维的多样性。「转移矩阵的累乘就是位置编码」一节展开。

RWKV-7

和 KDA 并排的兄弟(2025):同样是逐通道衰减加擦除,但擦除键 κ^t\hat\kappa_t、写入键 k~t\tilde k_t 和逐通道学习率 ata_t 各自独立,转移矩阵是一般的「对角加低秩」D−ab⊤D - a b^\top。KDA 是 aa、bb、写入键全绑在 ktk_t 上的特例,表达能力让了一步,换来内核少算一半东西,下一篇讲。

一句话总结

按「有没有擦」和「忘的粒度」两个坐标摆开:

无衰减标量衰减逐通道(对角)衰减
无 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):

任务例子考的能力依赖的动作
Palindromea b c d e | → e d c b a把每个 token 分别存下、按位置精确取回键之间互不干扰,只叠加和衰减的记忆读出来是加权和,长了就分不开
MQARA 4 · B 7 · C 2 · … · A ? → 4同时保持很多组 k→vk \to v 绑定,同一个键再来要覆盖旧值擦:查出来不等于 vv 就把差值写进去
Stackpush 3 a · push 3 b · pop 3 → b · pop 3 → a同一个键下反复改写,pop 后露出上一层擦:没有擦除,pop 3 的两次回答无法不同

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

公式 (1) 的图解:忘对付年龄,擦对付碰撞

新状态② 擦除① 衰减旧状态③ 写入=··+从右往左读:先把旧状态每一行按各自衰减,再把落在方向上的内容擦掉的比例,最后写入倍的新关联。读出:。行 = 键通道(,各自的衰减率),列 = 值通道()。

顺序是先衰减、再擦除、再写入。展开后是

St=(Diag⁡(αt)−βtktkt⊤Diag⁡(αt))St−1+β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)转移 D−atbt⊤D - a_t b_t^\top,D=Diag⁡(αt)D = \operatorname{Diag}(\alpha_t),at=βtkta_t = \beta_t k_t,bt=kt⊙αtb_t = k_t \odot \alpha_t。两个低秩向量都绑在 ktk_t 上,这是 KDA 内核比通用 DPLR 内核快约两倍的原因,下一篇细说。

既然有擦,为什么还要忘?两个减法项删的是不同的东西。

忘对付年龄。 SS 固定大小,信息必须随时间淡出才能腾位置;有些信息也会自然失效(变量重新赋值、子任务结束、话题切换)。门控让「淡多少」由当前 token 决定。副产品是衰减编码了新近度,越早写入的东西被乘的衰减越多,这是下一节的来源。

擦对付碰撞。 所有 token 写进同一张表,一个和早先某个键相似的键,查出来是两个值的混合。kt⊤Sk_t^\top S 是「状态现在对这个键的回答」,减去 βtkt(kt⊤S)\beta_t k_t (k_t^\top S) 就是写入之前把旧回答拿掉,ktk_t 不指向的行保持原样。忘做不到这件事:忘按年龄均匀地淡,不知道哪条记录和新键撞了。

衰减在键轴。 Diag⁡(αt)\operatorname{Diag}(\alpha_t) 左乘,衰减按行做,第 jj 个键通道存的所有值一起以 αt,j\alpha_{t,j} 淡出。读出 S⊤qS^\top q 沿键轴乘进去,查询触及的和键写入的是同一组分量,淡化某一行就是削弱「方向接近这个键分量的所有键」下面存的东西,这才是遗忘。若改成淡化某一列,缩的是所有存储值的同一个坐标,那是改输出尺度。内核的张量形状印证了这一点:SS 是 [B,H,K,V][B, H, K, V],门 gg 是 [B,T,H,K][B, T, H, K],没有 VV 轴。

按行衰减:列:值通道,的第个分量行:键通道的第个分量沿键轴进入按列衰减:列:值通道,的第个分量行:键通道的第个分量沿键轴进入第行是「方向接近第个键分量的键」写进来的全部值淡这一行 = 忘掉这一类键下面存的东西这是遗忘。门的形状:每个键分量一个门第列是所有存储值的第个坐标,不管是哪个键写的淡这一列 = 每个读出向量的第维一起缩小这不是遗忘,是改输出尺度。没有轴的门
的两个轴各由键和值的 128 个分量索引,读出 沿键轴乘进去,所以查询触及的和键写入的是同一组分量。只有沿键轴衰减才对应「忘掉某一类键下面的记录」。

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

K3 的 24 层 MLA 全是 NoPE,整个模型里没有 RoPE。不带位置信息的注意力层分不清「A 在 B 前面」和「B 在 A 前面」,那顺序从哪来?从 KDA 的递推里来。

记每一步的转移矩阵 Tj=(I−βjkjkj⊤)Diag⁡(αj)T_j = (I - \beta_j k_j k_j^\top)\operatorname{Diag}(\alpha_j),从 S0=0S_0 = 0 展开公式 (1) 再用 qtq_t 读出:

St=∑i=1t(TtTt−1⋯Ti+1) βikivi⊤,o~t=St⊤qt=∑i=1t(qt⊤(Tt⋯Ti+1) ki)⏟st,i βivi.S_t = \sum_{i=1}^{t} \big( T_t T_{t-1} \cdots T_{i+1} \big)\, \beta_i k_i v_i^\top , \qquad \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_t 和 kik_i 中间夹着一串矩阵,隔得越远夹得越多。RoPE 做的是同一件事:q~t=Rtqt\tilde q_t = R^t q_t、k~i=Riki\tilde k_i = R^i k_i,打分时

st,i=qt⊤(Rt)⊤Ri ki=qt⊤(∏j=i+1tR−1)ki,s_{t,i} = q_t^\top (R^t)^\top R^i\, k_i = q_t^\top \Big( \prod_{j=i+1}^{t} R^{-1} \Big) k_i ,

也是中间夹一串矩阵,只不过每一个都是同一个 R−1R^{-1}。

KDA数据相关、非正交、逐通道衰减⋯RoPE固定、正交、逐维频率⋯:每一格由第个 token 决定,颜色深浅示意该步保留了多少。:同一个分块旋转矩阵累乘次,只和距离有关,和内容无关。

两条链并排看,差别有三处:

  • RoPE 的格子是固定的;KDA 的格子由第 jj 个 token 的 αj\alpha_j、βj\beta_j、kjk_j 决定,可学习、数据相关。
  • RoPE 的 RR 是正交矩阵,分数随距离震荡但不衰减;KDA 的 Diag⁡(αj)\operatorname{Diag}(\alpha_j) 是收缩的,分数随距离衰减,天然偏向近期。
  • RoPE 每个维度有自己的频率,低频看长程、高频看短程;GDN 的标量 αj\alpha_j 相当于所有维度共用一个频率,KDA 的逐通道 αj\alpha_j 才是逐维频率的对应物。Kimi Linear 就是从这里出发把 GDN 的门换成了逐通道。

由此 K3 的两个决定就顺了:每层 MLA 前面都有三层 KDA 往残差流里写位置敏感的内容,所以 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 的工程收益。

参数化:从 xt​ 到 q,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αh∈Rdk\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}
7168Linear7168 → 12288ShortConv(4) + SwishL2Norm逐头96 × 128Linear7168 → 12288ShortConv(4) + SwishL2Norm逐头96 × 128Linear7168 → 12288ShortConv(4) + Swish96 × 128Linear7168 → 96Sigmoid96 个标量Linear7168 → 128Linear128 → 12288有下界,96 × 128Linear(满秩)7168 → 12288SigmoidKDA 递推96 头 × (128 × 128)训练 / prefill:chunkwisedecode:逐 token 递推逐头 RMSNorm128 维,每头独立⊙输出门 96 × 128,Linear12288 → 71687168
一层 KDA。粉色的那一行是 K3 的两个改动所在:衰减映射换成有下界的 scaled sigmoid,输出门从低秩改成满秩。

维度按 K3 的 d=7168d = 7168、96 头、dk=dv=128d_k = d_v = 128:

  • ShortConv。 深度可分的因果卷积,核大小 4(short_conv_kernel_size),给每个 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−βtktkt⊤I - \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_proj、f_b_proj),加一个 12288 维偏置(dt_bias),每个头每个键通道一个 logit。从 logit 到 (0,1)(0,1) 的映射,Kimi Linear 用 GDN 的负 softplus,K3 换成有下界的 scaled sigmoid,另有每头一个的可学习 log 尺度(A_log)。这是 K3 的改动之一,动机在 chunkwise 形式里,放到下一篇。

输出:逐头 RMSNorm、门、投影

递推读出 o~t=St⊤qt\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 的门是 7168 → 128 → 12288 的低秩投影,论文自己说了原因:为了和基线做公平的参数量对比,并且「和满秩门性能相当」。K3 不需要这个约束,直接用 7168 → 12288 的满秩 WgW_g(use_full_rank_gate = true)。MLA 的输出门也一起改成满秩,第 3 篇会看到同样的公式。

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 层 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 的学习率,Kimi Linear Table 7 把它放进目标函数、步长取 1。两种写法得到同一个递推,本文跟 Table 7。
  • 头数和头维。 KDA 和 MLA 共用 num_attention_heads = 96、head_dim = 128,但 MLA 的 query 和 key 实际是 192 维(128 加 64),第 3 篇解释。
  • 行向量还是列向量。 论文按 S∈Rdk×dvS \in \mathbb{R}^{d_k \times d_v} 写,读出是 S⊤qS^\top q;RWKV-7 和 Comba 原文用行向量约定,转移矩阵写在 St−1S_{t-1} 右边,本文全部转成了列向量约定。FLA 内核里 transpose_state_layout=True 存的是转置,只是内存布局的事。
  • 打分式里因子的顺序。 Kimi Linear 公式 (12) 把每一步的转移写成 Aj(I−βjkjkj⊤)A_j (I - \beta_j k_j k_j^\top),先擦再衰减;本文按公式 (1) 写成 (I−βjkjkj⊤)Diag⁡(αj)(I - \beta_j k_j k_j^\top)\operatorname{Diag}(\alpha_j),先衰减再擦。两个矩阵不相等,真正实现的是公式 (1) 的顺序。

下一篇

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

评论