对应论文 §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 的 embedding 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 加权求和 α Attention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和 α 词表 163840,隐藏维 7168 Token Embedding 163840 × 7168 从零训练的视觉编码器,hidden 1024,12 头 MoonViT-V2 27 层 · 0.4B · patch 14 merge_type = sd2_tpool 2×2 Pixel-shuffle + 时间池化 · token ÷4 PatchMergerMLPV2,GELU,RMSNorm Projector 4096 → 4096 → 7168 词嵌入永远是一个来源 b₀ = Embedding 每个块的输出是块内所有子层输出之和 b₁ … bₙ₋₁ 已完成的块(每块 12 层之和) 本块里已经算完的子层之和,相当于块内的普通残差流 bₙ⁽ⁱ⁾ 当前块 partial sum Kimi Delta Attention:96 头 × 128 维,状态 128×128/头 KDA 第 1 层 第一层不用 MoE,用一个稠密 SiTU-GLU FFN Dense FFN 仅第 1 层 · 中间维 33792 三层 KDA,每层后接一个 Stable LatentMoE KDA ×3 7168 → 3584 latent → 16 个 routed expert → RMSNorm → 7168 Stable LatentMoE 896 选 16 + 2 shared DeepSeek 式 MLA,无位置编码,满秩 sigmoid 输出门 Gated MLA ×1 · NoPE 同上 Stable LatentMoE 896 选 16 + 2 shared 主干末尾额外放一层全局注意力 Gated MLA 第 93 层 · 收尾 同上 Stable LatentMoE 第 93 层 最终归一化 RMSNorm 不与 embedding 共享权重 LM Head → 163840 预训练时 1 层多 token 预测;post-training 微调成投机解码的 draft MTP 层 ×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,存一个固定大小的矩阵 S t ∈ R d k × d v S_t \in \mathbb{R}^{d_k \times d_v} S t ∈ R d k × d v 。每个头 128×128,不管上下文多长都是这么大。读出是 o t = S t ⊤ q t o_t = S_t^\top q_t o t = S t ⊤ q t 。代价是记忆有损,怎么写、怎么忘、怎么擦,就成了这一族模型全部的设计空间。
这个设计空间看起来很散:每篇论文给一个递推式,符号还不一样。Kimi Linear 论文(Table 7)用一个统一的视角把它们收拢了:把 S S S 看成一组快速权重,每个 token 到来时对它做一步在线学习 。不同的模型只是在优化不同的目标函数。先把这个视角讲清楚,谱系就不用背了。
记忆就是在线学习:目标函数从哪来
S S S 是一张查找表:行是键通道,列是值通道,用 q q q 去查得到 S ⊤ q S^\top q S ⊤ q 。第 t t t 个 token 带来一对 ( k t , v t ) (k_t, v_t) ( k t , v t ) ,意思是「以后拿 k t k_t k t 来查,应该查到 v t v_t v t 」。问题是:S S S 该怎么变?
在线学习的回答:定义一个只看这一个 token 的损失 L t ( S ) \mathcal{L}_t(S) L t ( S ) ,衡量当前的 S S S 对这一对 ( k t , v t ) (k_t, v_t) ( k t , v t ) 服务得有多差,然后做一步梯度下降 ,步长取 1:
S t = S t − 1 − ∇ S L t ( S t − 1 ) . S_t = S_{t-1} - \nabla_S \mathcal{L}_t(S_{t-1}) . S t = S t − 1 − ∇ S L t ( S t − 1 ) .
Kimi Linear Table 7 里所有模型都写成这个形式,差别全在 L t \mathcal{L}_t L t 里。L t \mathcal{L}_t L t 一共只用到三种项,分别对应三个动作。
加:内积目标。 最朴素的要求是「用 k t k_t k t 查出来的东西要和 v t v_t v t 对齐」,写成内积:
L t ( S ) = − ⟨ S ⊤ k t , v t ⟩ = − k t ⊤ S v t , ∇ S L t = − k t v t ⊤ , S t = S t − 1 + k t v t ⊤ . \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 . L t ( S ) = − ⟨ S ⊤ k t , v t ⟩ = − k t ⊤ S v t , ∇ S L t = − k t v t ⊤ , S t = S t − 1 + k t v t ⊤ .
这就是 Hebb 规则,也是最早的线性注意力。它的问题在目标函数本身:L t \mathcal{L}_t L t 对 S S S 是线性的,没有最小值 。往 k t v t ⊤ k_t v_t^\top k t v t ⊤ 方向走多远都能让损失更低,所以每一步只会往上堆,永远不会有「够了」或者「原来存错了,去掉」。查的时候 S t ⊤ q = ∑ i ( k i ⊤ q ) v i S_t^\top q = \sum_i (k_i^\top q)\, v_i S t ⊤ q = ∑ i ( k i ⊤ q ) v i ,所有历史值按相似度叠在一起,谁也清不掉。
擦:回归目标。 把「对齐」换成「查出来正好等于」:
L t ( S ) = β t 2 ∥ S ⊤ k t − v t ∥ 2 . \mathcal{L}_t(S) = \tfrac{\beta_t}{2}\, \| S^\top k_t - v_t \|^2 . L t ( S ) = 2 β t ∥ S ⊤ k t − v t ∥ 2 .
记 e t = S ⊤ k t − v t e_t = S^\top k_t - v_t e t = S ⊤ k t − v t ,这是「记忆现在对 k t k_t k t 的回答」减去「应该的回答」,也就是残差。梯度是
∇ S L t = β t k t e t ⊤ = β t k t k t ⊤ S − β t k t v t ⊤ , \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 , ∇ S L t = β t k t e t ⊤ = β t k t k t ⊤ S − β t k t v t ⊤ ,
代进去:
S t = S t − 1 − β t k t k t ⊤ S t − 1 + β t k t v t ⊤ = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ . 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 . S t = S t − 1 − β t k t k t ⊤ S t − 1 + β t k t v t ⊤ = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ .
同一个式子有三种读法,都对。第一,它只写入残差 β t k t ( v t − S t − 1 ⊤ k t ) ⊤ \beta_t k_t (v_t - S_{t-1}^\top k_t)^\top β t k t ( v t − S t − 1 ⊤ k t ) ⊤ ,如果记忆已经答对了,什么都不写,这是 delta rule 这个名字的来源。第二,它先把 k t k_t k t 方向上原来存的东西擦掉 β t \beta_t β t 的比例(I − β t k t k t ⊤ I - \beta_t k_t k_t^\top I − β t k t k t ⊤ ,当 ∥ k t ∥ = 1 \|k_t\| = 1 ∥ k t ∥ = 1 时是沿 k t k_t k t 的一个收缩),再写入新的,这是「先擦再写」。第三,这个目标有最小值,最小值处 S ⊤ k t = v t S^\top k_t = v_t S ⊤ k t = v t 精确成立,所以它瞄准的是精确回忆 。β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 ) 是这个样本的权重,也叫写入强度:β t = 1 \beta_t = 1 β t = 1 且 ∥ k t ∥ = 1 \|k_t\| = 1 ∥ k t ∥ = 1 时,k t k_t k t 方向上的旧内容被完全替换。
忘:正则项。 内积目标无界,回归目标也只管当前这一个键,S S S 的 128 行终归是有限的容量,旧的绑定占着方向不走,新的进来就互相干扰。标准做法是加一个把 S S S 往零拉的正则:
L t ( S ) + = 1 2 ∥ 1 − α t S ∥ F 2 = 1 − α t 2 ∥ S ∥ F 2 , ∇ 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 , L t ( S ) + = 2 1 1 − α t S F 2 = 2 1 − α t ∥ S ∥ F 2 , ∇ S = ( 1 − α t ) S ,
这一项在梯度步里贡献 S − ( 1 − α t ) S = α t S S - (1 - \alpha_t) S = \alpha_t S S − ( 1 − α t ) S = α t S ,就是权重衰减。α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 ) 越小忘得越多。让 α t \alpha_t α t 依赖输入,模型就能逐 token 决定「前面的东西还要不要」,这是 Mamba 一系强调的选择性。再把标量换成逐通道的向量,正则写成 1 2 ∥ Diag ( 1 − α t ) S ∥ F 2 \tfrac12 \|\operatorname{Diag}(\sqrt{1 - \alpha_t})\, S\|_F^2 2 1 ∥ Diag ( 1 − α t ) S ∥ F 2 ,梯度步变成 Diag ( α t ) S \operatorname{Diag}(\alpha_t) S Diag ( α t ) S ,第 j j j 行有自己的正则强度。
一个顺序上的细节。 把三项放进同一个目标里一步走完,得到的是 α t S − β t k t k t ⊤ S + β t k t v t ⊤ \alpha_t S - \beta_t k_t k_t^\top S + \beta_t k_t v_t^\top α t S − β t k t k t ⊤ S + β t k t v t ⊤ 。GDN 和 KDA 实际用的不是这个,而是先衰减、再对衰减后的状态做回归那一步 :令 S ~ t − 1 = α t S t − 1 \tilde S_{t-1} = \alpha_t S_{t-1} S ~ t − 1 = α t S t − 1 (或 Diag ( α t ) S t − 1 \operatorname{Diag}(\alpha_t) S_{t-1} Diag ( α t ) S t − 1 ),目标是 β t 2 ∥ S ~ t − 1 ⊤ k t − v t ∥ 2 \tfrac{\beta_t}{2}\|\tilde S_{t-1}^\top k_t - v_t\|^2 2 β t ∥ S ~ t − 1 ⊤ k t − v t ∥ 2 ,更新是 S t = S ~ t − 1 − ∇ S ~ L t S_t = \tilde S_{t-1} - \nabla_{\tilde S}\mathcal{L}_t S t = S ~ t − 1 − ∇ S ~ L t 。两者差一项 β t ( 1 − α t ) k t k t ⊤ S t − 1 \beta_t (1 - \alpha_t) k_t k_t^\top S_{t-1} β t ( 1 − α t ) k t k t ⊤ S t − 1 :串行版本是从衰减之后剩下的东西里擦,擦的量和当前状态自洽。下一篇的 chunkwise 形式依赖这个顺序。
有了这三块积木,谱系就只剩两个问题:用了哪几块,衰减的粒度是标量还是逐通道 。
谱系:加、擦、忘、逐通道忘
「忘」路径:加衰减,不擦 「擦」路径:delta rule,不忘 旁支 + 标量衰减 α(忘) 内积目标 → 回归目标(擦) 标量 → 逐通道 + 标量衰减 α(忘) + delta rule(擦) 标量 → 逐通道 隐式解 擦除 × b 放开绑定 线性注意力 2020:S = S + k vᵀ 线性注意力 2020 S = S + k vᵀ 只加不擦,目标无下界 RetNet 2023 · Mamba2 2024:S = α S + k vᵀ RetNet 2023 · Mamba2 2024 S = α S + k vᵀ α 每头一个标量 Mamba2:三个合成任务全挂 DeltaNet 2021:S = (I − β k kᵀ) S + β k vᵀ DeltaNet 2021 S = (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 2024 S = Diag(α) S + k vᵀ α 每个键通道一个 Gated DeltaNet 2024:S = (I − β k kᵀ) α S + β k vᵀ Gated DeltaNet 2024 S = (I − β k kᵀ) α S + β k vᵀ 先忘,再擦,再写 收敛慢于 KDA Comba 2025:S = (α − b β k kᵀ) S + β k vᵀ Comba 2025 S = (α − b β k kᵀ) S + β k vᵀ 擦得比写得少(b < 1),读出用 q − d k KDA 2025:S = (I − β k kᵀ) Diag(α) S + β k vᵀ KDA 2025 S = (I − β k kᵀ) Diag(α) S + β k vᵀ 先逐通道忘,再擦,再写 RWKV-7 2025:S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀ RWKV-7 2025 S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀ 一般 DPLR:擦除键 κ̂、写入键 k̃ 各不相同 「忘」路径:加衰减,不擦 「擦」路径:delta rule,不忘 + 标量衰减 α(忘) → 回归目标(擦) 标量 → 逐通道 + 标量衰减 α(忘) + delta rule(擦) 标量 → 逐通道 隐式解 擦除 × b 放开绑定 线性注意力 2020:S = S + k vᵀ 线性注意力 2020 S = S + k vᵀ 只加不擦,目标无下界 RetNet 2023 · Mamba2 2024:S = α S + k vᵀ RetNet 2023 · Mamba2 2024 S = α S + k vᵀ α 每头一个标量 Mamba2:三个合成任务全挂 DeltaNet 2021:S = (I − β k kᵀ) S + β k vᵀ DeltaNet 2021 S = (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 2024 S = Diag(α) S + k vᵀ α 每个键通道一个 Gated DeltaNet 2024:S = (I − β k kᵀ) α S + β k vᵀ Gated DeltaNet 2024 S = (I − β k kᵀ) α S + β k vᵀ 先忘,再擦,再写 收敛慢于 KDA Comba 2025:S = (α − b β k kᵀ) S + β k vᵀ Comba 2025 S = (α − b β k kᵀ) S + β k vᵀ 擦得比写得少(b < 1),读出用 q − d k KDA 2025:S = (I − β k kᵀ) Diag(α) S + β k vᵀ KDA 2025 S = (I − β k kᵀ) Diag(α) S + β k vᵀ 先逐通道忘,再擦,再写 RWKV-7 2025:S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀ RWKV-7 2025 S = (Diag(α) − (a ⊙ κ̂) κ̂ᵀ) S + k̃ vᵀ 一般 DPLR:擦除键 κ̂、写入键 k̃ 各不相同
从线性注意力出发有两条路。左边一路只加「忘」,先标量后逐通道,始终不擦;右边一路先加「擦」,再补「忘」。两条路在 KDA 汇合:从 GLA 看是加了 delta rule,从 GDN 看是把标量衰减换成了逐通道。右侧三个旁支各自在某个节点上改了一个零件,和 KDA 没有直接的继承关系,但 RWKV-7 值得并排放着看。下面按节点走,主线节点给目标函数和递推,旁支只给递推和一句话。
线性注意力
Katharopoulos 等,2020。上一节的「加」:
L t = − ⟨ S ⊤ k t , v t ⟩ , S t = S t − 1 + k t v t ⊤ . \mathcal{L}_t = -\langle S^\top k_t, v_t \rangle ,
\qquad
S_t = S_{t-1} + k_t v_t^\top . L t = − ⟨ S ⊤ k t , v t ⟩ , S t = S t − 1 + k t v t ⊤ .
只加不擦不忘。它证明了 softmax 注意力可以换成固定大小的状态加一次递推,但长上下文下旧关联堆积、互相干扰,效果远逊于 softmax。后面每一个模型都是在补这个目标函数的缺陷。
RetNet 和 Mamba2
加「忘」,标量粒度:
L t = − β t ⟨ S ⊤ k t , v t ⟩ + 1 2 ∥ 1 − α t S ∥ F 2 , S t = α t S t − 1 + β t k t v t ⊤ . \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 . L t = − β t ⟨ S ⊤ k t , v t ⟩ + 2 1 1 − α t S F 2 , S t = α t S t − 1 + β t k t v t ⊤ .
两者的区别在 α \alpha α 从哪来。RetNet(2023)的 α \alpha α 是每头一个固定常数 ,不看输入,β t = 1 \beta_t = 1 β t = 1 ,所以它的记忆是纯粹的指数衰减,一个头一个时间尺度。Mamba2(2024)的 α t = exp ( − Δ t A h ) \alpha_t = \exp(-\Delta_t A_h) α t = exp ( − Δ t A h ) 由输入决定,写入项也乘了同一个步长 Δ t \Delta_t Δ t (即上式的 β t \beta_t β t ),模型可以对不重要的 token 选 Δ t ≈ 0 \Delta_t \approx 0 Δ t ≈ 0 ,既不忘也不写。这一路把「状态该保留多久」交给了数据,但目标函数里仍然是内积项,没有擦除:k k k 方向上原来存的东西只能等着淡出,不能被覆盖。下一节的合成任务里 Mamba2 全部失败,问题就出在这里。
GLA 和 HGRN2
「忘」从标量变成逐通道:
L t = − ⟨ S ⊤ k t , v t ⟩ + 1 2 ∥ Diag ( 1 − α t ) S ∥ F 2 , S t = Diag ( α t ) S t − 1 + k t v t ⊤ , α t ∈ ( 0 , 1 ) d k . \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} . L t = − ⟨ S ⊤ k t , v t ⟩ + 2 1 Diag ( 1 − α t ) S F 2 , S t = Diag ( α t ) S t − 1 + k t v t ⊤ , α t ∈ ( 0 , 1 ) d k .
GLA(2023)的 α t \alpha_t α t 来自一个低秩投影加 sigmoid,每个键通道一个遗忘率,这个参数化被 KDA 直接沿用。HGRN2(2024)是从逐元素门控 RNN 那边走过来的,把一维的状态扩成外积,它的特点是写入门和遗忘门绑在一起 :键被换成 1 − α t 1 - \alpha_t 1 − α t ,即 S t = Diag ( α t ) S t − 1 + ( 1 − α t ) v t ⊤ S_t = \operatorname{Diag}(\alpha_t) S_{t-1} + (1 - \alpha_t) v_t^\top S t = Diag ( α t ) S t − 1 + ( 1 − α t ) v t ⊤ ,忘得多的通道写得也多,少一组参数。两者都还是内积目标,仍然没有擦除。
DeltaNet
回到起点,换目标而不是加项。Schlag 等 2021 提出,Yang 等 2024 给出并行形式:
L t = β t 2 ∥ S ⊤ k t − v t ∥ 2 , S t = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ . \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 . L t = 2 β t ∥ S ⊤ k t − v t ∥ 2 , S t = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ .
这是上一节推过的「擦」。( I − β t k t k t ⊤ ) (I - \beta_t k_t k_t^\top) ( I − β t k t k t ⊤ ) 是一个广义 Householder 变换,后面 chunkwise 并行能做出来全靠它。它能做精确的键值覆盖,但没有任何遗忘 :除非同一个键再来一次把它擦掉,写进去的东西永远留在 S S S 里。
Longhorn
DeltaNet 的旁支(2024)。目标一样是回归,但不做一步梯度下降,而是把带正则的单步问题精确解出来(隐式在线学习),结果是把 β t \beta_t β t 换成 β t / ( 1 + β t k t ⊤ k t ) \beta_t / (1 + \beta_t k_t^\top k_t) β t / ( 1 + β t k t ⊤ k t ) :
S t = ( I − β t 1 + β t k t ⊤ k t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ . 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 . S t = ( I − 1 + β t k t ⊤ k t β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ .
好处是当 ∥ k t ∥ = 1 \|k_t\| = 1 ∥ k t ∥ = 1 时擦除系数 β t / ( 1 + β t ) < 1 \beta_t / (1 + \beta_t) < 1 β t / ( 1 + β t ) < 1 自动成立,β t \beta_t β t 不用再限制在 ( 0 , 1 ) (0,1) ( 0 , 1 ) 。没有衰减。
Gated DeltaNet
Yang 等 2024。在 DeltaNet 上加「忘」,标量粒度,并且按上一节末尾说的顺序,先衰减再擦写:
L t = β t 2 ∥ S ~ t − 1 ⊤ k t − v t ∥ 2 , S ~ t − 1 = α t S t − 1 , S t = ( I − β t k t k t ⊤ ) α t S t − 1 + β t k t v t ⊤ . \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 . L t = 2 β t ∥ S ~ t − 1 ⊤ k t − v t ∥ 2 , S ~ t − 1 = α t S t − 1 , S t = ( I − β t k t k t ⊤ ) α t S t − 1 + β t k t v t ⊤ .
α t = exp ( − A h ⋅ softplus ( ⋅ ) ) \alpha_t = \exp(-A_h \cdot \operatorname{softplus}(\cdot)) α t = exp ( − A h ⋅ softplus ( ⋅ )) 每头一个,沿 Mamba2 的参数化;β t \beta_t β t 是 sigmoid。这是第一个同时有「擦」和「忘」的模型,Kimi Linear 的所有消融都拿它当基线。剩下的问题是一个头的 128 个键通道以同一个速率衰减。
Comba
GDN 的旁支(2025)。两个改动:擦除强度乘一个学习到的标量 b ∈ ( 0 , 1 ) b \in (0,1) b ∈ ( 0 , 1 ) ,让擦得比写得少;读出前先把 query 减去 d d d 倍的当前键:
S t = ( α t − b β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ , o t = S t ⊤ ( q t − d k t ) . 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) . S t = ( α t − b β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ , o t = S t ⊤ ( q t − d k t ) .
Kimi Linear Table 7 里写的 k ^ t \hat k_t k ^ t 就是 b k t b\, k_t b k t 。衰减仍是每头一个标量。Kimi Linear 的 chunkwise 推导里借用了 Comba 对转移矩阵累乘的写法,下一篇会碰到。
KDA
Kimi Linear,2025。把 GDN 的标量换成对角矩阵,或者说给 GLA 加上 delta rule:
L t = β t 2 ∥ S ~ t − 1 ⊤ k t − v t ∥ 2 , S ~ t − 1 = Diag ( α t ) S t − 1 , \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} , L t = 2 β t ∥ S ~ t − 1 ⊤ k t − v t ∥ 2 , S ~ t − 1 = Diag ( α t ) S t − 1 ,
S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤ , o ~ t = S t ⊤ q t . (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} S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤ , o ~ t = S t ⊤ q t . ( K3 Eq. 1 )
α t ∈ ( 0 , 1 ) d k \alpha_t \in (0,1)^{d_k} α t ∈ ( 0 , 1 ) d k ,β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 ) 。逐通道衰减的动机不只是「更细的粒度」,Kimi Linear 给的理由来自位置编码:RoPE 的力量在于每个维度有自己的旋转频率,标量衰减没有这种逐维的多样性,如果要让递推层承担位置信息,逐通道的门是自然的对应物。这个类比在「转移矩阵的累乘就是位置编码」一节里展开。
RWKV-7
和 KDA 并排的兄弟(2025):同样是逐通道衰减加擦除,但擦除和写入用的不是同一个键。按本文的列向量约定写:
S t = ( Diag ( w t ) − ( a t ⊙ κ ^ t ) κ ^ t ⊤ ) S t − 1 + k ~ t v t ⊤ . 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 . S t = ( Diag ( w t ) − ( a t ⊙ κ ^ t ) κ ^ t ⊤ ) S t − 1 + k ~ t v t ⊤ .
κ ^ t \hat\kappa_t κ ^ t 是把 k t k_t k t 逐通道乘一个学习到的尺度再做 L2 归一化得到的擦除键 ;a t ∈ ( 0 , 1 ) d k a_t \in (0,1)^{d_k} a t ∈ ( 0 , 1 ) d k 是逐通道的「上下文学习率」,替代了 KDA 里每头一个的标量 β t \beta_t β t ;k ~ t = k t ⊙ lerp ( 1 , a t , ⋅ ) \tilde k_t = k_t \odot \operatorname{lerp}(1, a_t, \cdot) k ~ t = k t ⊙ lerp ( 1 , a t , ⋅ ) 是另一个写入键 ;w t ∈ ( 0.55 , 1 ) w_t \in (0.55, 1) w t ∈ ( 0.55 , 1 ) 逐通道衰减。它的转移矩阵是一般的「对角加低秩」D − a b ⊤ D - a b^\top D − a b ⊤ ,a a a 、b b b 各自独立。KDA 是 a a a 、b b b 都绑在 k t k_t k t 上、写入键也等于 k t k_t k t 的特例。表达能力上 KDA 让了一步,换来的是内核少算一半东西,下一篇讲。
一句话总结
线性注意力是「加」,RetNet / Mamba2 是「加,然后所有东西以同一个速率淡出」,GLA / HGRN2 是「加,然后每个通道以自己的速率淡出」,DeltaNet 是「在这个键下先擦再写」,GDN 是「先擦再写,然后所有东西以同一个速率淡出」,KDA 是「先擦再写,然后每个特征通道以自己的速率淡出」。按两个坐标摆开:
无衰减 标量衰减 逐通道(对角)衰减 无 delta rule 线性注意力 RetNet、Mamba2 GLA、HGRN2 有 delta rule DeltaNet、Longhorn Gated DeltaNet、Comba KDA 、RWKV-7
递推形式能做什么:三个合成任务
谱系走完,自然要问:「擦」和「逐通道忘」这两步各买来了什么?语言模型的 PPL 差异只有零点几,看不出机制。Kimi Linear §5.1 用三个合成任务分别考察状态的三种能力,模型只有 2 层 2 头、头维 128,序列长度 256 到 2048,比较 KDA、GDN 和 Mamba2。三个任务分别对应状态必须做到的三件事。
Palindrome,考精确复制。 输入一串随机 token,输出它的逆序:
状态要把每个 token 分别存下来,还要能按位置精确取回,一个都不能糊。只做叠加和衰减的记忆读出来的是所有 token 的加权和,长度一上去就分不开。
MQAR,多查询关联回忆。 先给一串键值对,后面用键去查:
A 4 · B 7 · C 2 · … · A ? → 4 · C ? → 2
状态要同时保持很多组 k → v k \to v k → v 绑定互不干扰,而且如果同一个键后来又出现了新值,旧值必须被覆盖。这正是回归目标做的事:查出来不等于 v v v 就把差值写进去。这个任务被认为和语言模型质量强相关。
Stack,考状态跟踪。 64 个独立的栈,push 和 pop 交错:
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_t I − β_t k_t k_tᵀ ② 擦除 Diag(α_t) ① 衰减 S_{t−1} 旧状态 k_t v_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)。
三个操作的顺序有讲究:先衰减,再擦除,再写入。展开后是
S t = ( Diag ( α t ) − β t k t k t ⊤ Diag ( α t ) ) S t − 1 + β t k t v t ⊤ , 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 , S t = ( Diag ( α t ) − β t k t k t ⊤ Diag ( α t ) ) S t − 1 + β t k t v t ⊤ ,
即一般的「对角加低秩」(DPLR)转移 D − a t b t ⊤ D - a_t b_t^\top D − a t b t ⊤ ,其中 D = Diag ( α t ) D = \operatorname{Diag}(\alpha_t) D = Diag ( α t ) ,a t = β t k t a_t = \beta_t k_t a t = β t k t ,b t = k t ⊙ α t b_t = k_t \odot \alpha_t b t = k t ⊙ α t 。两个低秩向量都绑在 k t k_t k t 上,不像 RWKV-7 那样各自独立。这个绑定看起来是表达能力上的让步,但它让内核少算一半东西,是 Kimi Linear 的 KDA 内核比通用 DPLR 内核快约两倍的原因。下一篇细说。
对状态矩阵的直觉:行是键通道,列是值通道 。Diag ( α t ) \operatorname{Diag}(\alpha_t) Diag ( α t ) 左乘,所以衰减是按行做的,第 j j j 个键通道存的所有值一起以 α t , j \alpha_{t,j} α t , j 的速率淡出。擦除沿 k t k_t k t 这一个方向。写入是一个秩 1 的外积。
转移矩阵的累乘就是位置编码
K3 的 24 层 MLA 全是 NoPE,整个模型里没有 RoPE。一个不带位置信息的注意力层对 token 的排列是等变的,它分不清「A 在 B 前面」和「B 在 A 前面」。那顺序从哪来?答案是:从 KDA 的递推里来 。把公式 (1) 展开一遍就能看到。
记每一步的转移矩阵为 T j = ( I − β j k j k j ⊤ ) Diag ( α j ) T_j = (I - \beta_j k_j k_j^\top)\operatorname{Diag}(\alpha_j) T j = ( I − β j k j k j ⊤ ) Diag ( α j ) ,公式 (1) 就是 S t = T t S t − 1 + β t k t v t ⊤ S_t = T_t S_{t-1} + \beta_t k_t v_t^\top S t = T t S t − 1 + β t k t v t ⊤ 。从 S 0 = 0 S_0 = 0 S 0 = 0 开始往前套:
S 1 = β 1 k 1 v 1 ⊤ , S 2 = T 2 β 1 k 1 v 1 ⊤ + β 2 k 2 v 2 ⊤ , S t = ∑ i = 1 t ( T t T t − 1 ⋯ T i + 1 ) β i k i v i ⊤ . 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 . S 1 = β 1 k 1 v 1 ⊤ , S 2 = T 2 β 1 k 1 v 1 ⊤ + β 2 k 2 v 2 ⊤ , S t = i = 1 ∑ t ( T t T t − 1 ⋯ T i + 1 ) β i k i v i ⊤ .
第 i i i 个 token 写进去的 k i v i ⊤ k_i v_i^\top k i v i ⊤ ,到第 t t t 步时已经被从 i + 1 i+1 i + 1 到 t t t 的每一个转移矩阵各乘了一遍。用 q t q_t q t 读出:
o ~ t = S t ⊤ q t = ∑ i = 1 t ( q t ⊤ ( T t ⋯ T i + 1 ) k i ) ⏟ s t , i β i v i . \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 . o ~ t = S t ⊤ q t = i = 1 ∑ t s t , i ( q t ⊤ ( T t ⋯ T i + 1 ) k i ) β i v i .
这和注意力的形状完全一样:对每个历史位置 i i i 算一个分数 s t , i s_{t,i} s t , i ,用它给 v i v_i v i 加权求和。区别在于 q t q_t q t 和 k i k_i k i 中间夹着一串矩阵,i i i 和 t t t 隔得越远,夹的越多。
RoPE 那边做的其实是同一件事。RoPE 把位置 t t t 的 query 和位置 i i i 的 key 各旋转一次,q ~ t = R t q t \tilde q_t = R^t q_t q ~ t = R t q t 、k ~ i = R i k i \tilde k_i = R^i k_i k ~ i = R i k i ,其中 R R R 是一个固定的分块对角旋转矩阵,第 d d d 个二维块转 θ d \theta_d θ d 角。打分时
s t , i = q ~ t ⊤ k ~ i = q t ⊤ ( R t ) ⊤ R i k i = q t ⊤ ( R − 1 ) t − i k i = q t ⊤ ( ∏ j = i + 1 t R − 1 ) k i . 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 . s t , i = q ~ t ⊤ k ~ i = q t ⊤ ( R t ) ⊤ R i k i = q t ⊤ ( R − 1 ) t − i k i = q t ⊤ ( j = i + 1 ∏ t R − 1 ) k i .
也是 q t q_t q t 和 k i k_i k i 中间夹一串矩阵,只不过每一个都是同一个 R − 1 R^{-1} R − 1 ,夹几个只取决于距离 t − i t - i t − i 。
KDA 数据相关、非正交、逐通道衰减 k_i Ti+1 Ti+2 ⋯ Tt−1 Tt q_t RoPE 固定、正交、逐维频率 k_i R R ⋯ R R q_t T_j = (I − β_j k_j k_jᵀ) Diag(α_j):每一格由第 j 个 token 决定,颜色深浅示意该步保留了多少。 R:同一个分块旋转矩阵累乘 t − i 次,只和距离有关,和内容无关。
两条链并排看,差别有三处。第一,RoPE 的格子是固定 的,KDA 的格子由第 j j j 个 token 的 α j \alpha_j α j 、β j \beta_j β j 、k j k_j k j 决定,是可学习、数据相关的。第二,RoPE 的 R R R 是正交矩阵,只旋转不缩放,分数随距离震荡但不衰减;KDA 的 Diag ( α j ) \operatorname{Diag}(\alpha_j) Diag ( α j ) 元素在 ( 0 , 1 ) (0,1) ( 0 , 1 ) 里,是收缩的,分数随距离衰减,天然偏向近期,再加一个 RoPE 里没有对应物的擦除因子。第三,也是 KDA 比 GDN 多出来的那一步的理由:RoPE 每个维度有自己的频率 θ d \theta_d θ d ,低频维度看长程、高频维度看短程,这种逐维的多样性是它有效的关键;GDN 的标量 α j \alpha_j α j 相当于所有维度共用一个频率,KDA 的逐通道 α j \alpha_j α 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 带来的工程收益。
参数化:从 x t 到 q , k , v , α , β
论文公式 (2),每个头 h h h :
q t h , k t h = L 2 N o r m ( Swish ( ShortConv ( W q / k h x t ) ) ) ∈ R d k v t h = Swish ( ShortConv ( W v h x t ) ) ∈ R d v β t h = Sigmoid ( W β h x t ) ∈ ( 0 , 1 ) z t h = W α ↑ W α ↓ x t + b α h ∈ R d k \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} q t h , k t h v t h β t h z t h = L 2 Norm ( Swish ( ShortConv ( W q / k h x t )) ) ∈ R d k = Swish ( ShortConv ( W v h x t )) ∈ R d v = Sigmoid ( W β h x t ) ∈ ( 0 , 1 ) = W α ↑ W α ↓ x t + b α h ∈ R d k
x_t 7168 Linear W_q 7168 → 12288 ShortConv(4) + Swish L2Norm 逐头 q_t 96 × 128 Linear W_k 7168 → 12288 ShortConv(4) + Swish L2Norm 逐头 k_t 96 × 128 Linear W_v 7168 → 12288 ShortConv(4) + Swish v_t 96 × 128 Linear W_β 7168 → 96 Sigmoid β_t 96 个标量 Linear W_α↓ 7168 → 128 Linear W_α↑ + b_α 128 → 12288 g = g_min · σ(e^A z) 有下界,g_min = −5 α_t = e^{g} 96 × 128 Linear W_g(满秩) 7168 → 12288 Sigmoid KDA 递推 S_t = (I − β k kᵀ) Diag(α) S_{t−1} + β k vᵀ õ_t = S_tᵀ q_t S 96 头 × (128 × 128) 训练 / prefill:chunkwise decode:逐 token 递推 逐头 RMSNorm 128 维,每头独立 ⊙ 输出门 96 × 128,∈ (0,1) Linear W_o 12288 → 7168 y_t 7168 一层 KDA。粉色的那一行是 K3 的两个改动所在:衰减映射换成有下界的 scaled sigmoid,输出门从低秩改成满秩。
逐项说明,维度按 K3 的 d = 7168 d = 7168 d = 7168 、96 头、d k = d v = 128 d_k = d_v = 128 d 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 除了 S S S 之外唯一的状态。
q、k 的 L2 归一化。 让 ∥ k t ∥ = 1 \|k_t\| = 1 ∥ k t ∥ = 1 ,于是 I − β t k t k t ⊤ I - \beta_t k_t k_t^\top I − β t k t k t ⊤ 的特征值落在 [ 1 − β t , 1 ] [1-\beta_t, 1] [ 1 − β t , 1 ] ,递推不会爆。
β \beta β 。 7168 → 96 的线性层加 sigmoid,每头一个标量。代码里叫 b_proj。
α \alpha α 的 logit z z z 。 低秩投影 7168 → 128 → 12288(f_a_proj、f_b_proj),加一个 12288 维的偏置 b α b_\alpha b α (dt_bias),得到每个头每个键通道一个 logit。从 logit 到 ( 0 , 1 ) (0,1) ( 0 , 1 ) 里的 α \alpha α 还要过一个映射,Kimi Linear 用 GDN 的负 softplus,K3 换成了有下界的 scaled sigmoid,还有一个每头一个的可学习 log 尺度 A h A_h A h (A_log)。这是 K3 的改动之一,放到下一篇讲,因为它的动机在 chunkwise 形式里。
输出:逐头 RMSNorm、门、投影
递推读出 o ~ t = S t ⊤ q t \tilde o_t = S_t^\top q_t o ~ t = S t ⊤ q t 之后,论文公式 (6):
y t = W o [ Sigmoid ( W g x t ) ⊙ RMSNorm ( o ~ t ) ] . y_t = W_o \big[\operatorname{Sigmoid}(W_g x_t) \odot \operatorname{RMSNorm}(\tilde o_t)\big]. y t = W o [ Sigmoid ( W g x t ) ⊙ RMSNorm ( o ~ t ) ] .
RMSNorm 是逐头的,每头独立在 128 维上归一化。门是 sigmoid,作用在归一化之后的每个通道上,然后才是 W o W_o W o (12288 → 7168)。代码里归一化和门融在一个 FusedRMSNormGated(activation='sigmoid') 里。
K3 的改动:门从低秩改满秩。 Kimi Linear 的门是 Sigmoid ( W g ↑ W g ↓ x t ) \operatorname{Sigmoid}(W_g^{\uparrow} W_g^{\downarrow} x_t) Sigmoid ( W g ↑ W g ↓ x t ) ,7168 → 128 → 12288,和遗忘门一样走低秩。Kimi Linear 论文自己说了原因:为了和基线做公平的参数量对比,并且「和满秩门性能相当」。K3 不需要这个约束,直接用 7168 → 12288 的满秩 W g W_g W g (use_full_rank_gate = true,代码里 g_proj)。MLA 那边的输出门也一起改成满秩,下下篇会看到同样的公式。
# 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) ( 0 , 1 ) ,能把某些通道彻底关掉,引用的是 Qwen 团队关于 gated attention 的工作:门带来非线性和稀疏性,并缓解 attention sink。
每层的状态有多大
一层 KDA 在 decode 时携带的状态:
状态 形状 数量 S S S 96 头 × 128 × 128 1,572,864 三路 ShortConv 的窗口 3 × 12288 × 3 个 token 110,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) α t = f ( ⋅ ) 没有给出 f f f ,K3 把映射写全了(公式 (5)),并且改了它。
β \beta β 是学习率还是样本权重。 DeltaNet 原文把 β t \beta_t β t 写成 SGD 的学习率,S t = S t − 1 − β t ∇ L S_t = S_{t-1} - \beta_t \nabla \mathcal{L} S t = S t − 1 − β t ∇ L ;Kimi Linear Table 7 把它放进目标函数、步长取 1。两种写法得到同一个递推,本文跟 Table 7。
头数和头维。 KDA 和 MLA 共用 num_attention_heads = 96、head_dim = 128,但 MLA 的 query 和 key 实际是 192 维(128 加 64),下下篇解释。
状态是 d k × d v d_k \times d_v d k × d v 还是转置。 论文按 S ∈ R d k × d v S \in \mathbb{R}^{d_k \times d_v} S ∈ R d k × d v 写,读出是 S ⊤ q S^\top q S ⊤ q ;FLA 内核里有 transpose_state_layout=True,存的是转置,只是内存布局的事。RWKV-7 和 Comba 的原文用行向量约定,转移矩阵写在 S t − 1 S_{t-1} S t − 1 右边,本文全部转成了列向量约定。
打分式里因子的顺序。 Kimi Linear 公式 (12) 把每一步的转移写成 A j ( I − β j k j k j ⊤ ) A_j (I - \beta_j k_j k_j^\top) A j ( I − β j k j k j ⊤ ) ,作用到右边的 k i k_i k i 上是先擦再衰减;本文按公式 (1) 写成 ( I − β j k j k j ⊤ ) Diag ( α j ) (I - \beta_j k_j k_j^\top)\operatorname{Diag}(\alpha_j) ( I − β j k j k j ⊤ ) Diag ( α j ) ,先衰减再擦。两个矩阵不相等,真正实现的是公式 (1) 的顺序,公式 (12) 只是在说明「累乘」这个形状。
下一篇
递推形式一个 token 一步,训练时没法用 Tensor Core。下一篇讲 Kimi Linear 的 chunkwise 并行形式(WY 表示、UT 变换、块间递推加块内并行),然后就能看懂 K3 为什么要给衰减加一个 − 5 -5 − 5 的下界,以及这个改动怎么把对角 tile 也搬上 Tensor Core。