对应论文 §2.1 的开头和 §2.1.2,公式 (7)。前两篇讲的 KDA 用固定大小的状态换掉了 KV cache,代价是记忆有损。K3 每四层留一层 softmax 注意力做无损的全局检索,这一层是 DeepSeek-V2 的 MLA,加上两处改动:不带位置编码,输出加满秩 sigmoid 门。
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
你现在在这里:每个重复块的第四层,以及主干末尾额外的第 93 层。
排布:3
,末尾再补一层
24 层 Gated MLA 的位置:第 4、8、…、92 层,加第 93 层。
每个块是三层 KDA 接一层 Gated MLA,重复 23 次到第 92 层,然后第 93 层再放一层 Gated MLA,论文的解释是保证最后一层一定做全局注意力。这是层级(layerwise)混合,整层整层地换,而不是在一层里混合两种头。Kimi Linear 说选层级混合是为了基础设施简单和训练稳定。
MLA 复习:只缓存 latent
普通多头注意力每 token 每层要缓存所有头的 K 和 V。K3 的规格是 96 头、头维 128,如果直接缓存就是 96 × 128 × 2 = 24576 个数。MLA(DeepSeek-V2)的做法是把 K 和 V 压进一个低维的 latent 向量:
ct=Wkv↓xt∈R512,[ktnope; vt]=Wkv↑RMSNorm(ct),
只缓存 ct,用的时候再由 Wkv↑ 重建每个头的 K 和 V。推理时还可以把 Wkv↑ 吸收进 query 和输出投影,直接在 512 维的 latent 上做注意力,所有头共享同一份 key 和 value,也就是变成 MQA。
一层 Gated MLA。中间那条竖的橙色细框是推理时真正要缓存的东西:512 维 latent 加 64 维头间共享的键,共 576 个数。q 和 k 都是 192 维,其中 64 维在 DeepSeek 原版里用来放 RoPE,K3 保留了这个槽位但不做旋转。
按 K3 的 config 走一遍(q_lora_rank = 1536,kv_lora_rank = 512,qk_nope_head_dim = 128,qk_rope_head_dim = 64,v_head_dim = 128):
query。 xt(7168)→ 1536 维低秩 → RMSNorm → 96 × 192。每个头的 query 是 192 维,由 128 维的内容部分和 64 维的「rope 槽位」拼成。
key 和 value。 xt → 576 维,拆成 512 维 latent 和 64 维共享分量。latent 过 RMSNorm,再由 512 → 96 × 256 的上投影展开成每头 128 维的 knope 和 128 维的 v。64 维共享分量直接广播给所有 96 个头,拼在每个头的 key 后面,凑成 192 维。
注意力。 96 个头各自在 192 维上算 qk⊤/192,因果 softmax,加权 128 维的 v。输出 96 × 128 = 12288 维。
输出门与投影。 乘一个满秩 sigmoid 门,再 12288 → 7168。
KV cache 每 token 每层 512 + 64 = 576 个数。
NoPE:那 64 维不旋转
DeepSeek 原版 MLA 里,那 64 维是专门留给 RoPE 的:内容部分不旋转、低秩压缩;位置部分旋转、所有头共享。这是为了让低秩吸收和 RoPE 兼容。
K3 的 mla_use_nope = true。代码里 rotary_emb = None,64 维的 krot 直接 expand 到所有头、和 knope 拼起来,不做任何旋转。q 那边同样。于是这 64 维退化成一个所有头共享的、每 token 一个的额外键分量。论文没有说为什么保留这个槽位而不是干脆去掉,这是从代码读出来的。
为什么敢不要位置编码。第 1 篇推过:KDA 的读出 qt⊤(∏jDiag(αj)(I−βjkjkj⊤))ki 和 RoPE 的 qt⊤(∏jRj)ki 形状一样,只是 Rj 换成了数据相关、可学习的对角转移。KDA 层已经在做位置敏感、偏向近期的序列混合,MLA 层只管无约束的全局内容交互。论文的原话是「这个分工让扩上下文时不需要改任何位置编码参数,比如重调 RoPE 的频率基或者用 YaRN」。§3.4 说 K3 预训练从 8K 到 64K,冷却阶段从 256K 到 1M,四个阶段,全程不动位置编码,直接外推到 1M。
Kimi Linear 还给了实验证据:带 RoPE 的 MLA 变体在短上下文持平,128K 长上下文上更差(RULER 78.8 对 84.3),他们的解释是全局层里的 RoPE 过分强调短程顺序,让中期训练里的上下文扩展变得不灵活。
顺带的工程好处:NoPE 的 MLA 在推理时可以完全转成 MQA,没有那个必须单独处理的旋转分量。
满秩输出门
论文公式 (7):
yt=Wo[Sigmoid(Wgxt)⊙o~t].
o~t 是没加门的 MLA 输出(12288 维),Wg 是 7168 → 12288 的满秩矩阵,和 K3 的 KDA 输出门同一种参数化。config 里 mla_use_output_gate = true,代码里 g_proj。
# KimiMLAAttention.forward 的收尾
attn_output = attn_output.reshape(batch_size, seq_length, -1) # 96 × 128 = 12288
if self.use_output_gate:
g = self.g_proj(hidden_states).sigmoid() # 7168 -> 12288
attn_output = attn_output * g
attn_output = self.o_proj(attn_output) # 12288 -> 7168
论文引的是 Qwen 团队的 gated attention 工作:在 softmax 注意力输出上加一个逐通道的 sigmoid 门,带来非线性和稀疏性,同时消除 attention sink。K2 的 MLA 没有这个门。和 KDA 那边不同的一点:KDA 是先逐头 RMSNorm 再乘门,MLA 这里直接乘在注意力输出上,没有归一化。
训练时的一个精度细节
论文提到 flash attention 里存在有偏的舍入误差,K3 采用了 Qiu 和 Yao 的方法,训练时把注意力输出保持在 FP32。这让输出 tile 的片上占用翻倍,于是他们重新设计了训练内核,让它和 KV 的暂存缓冲重叠而不是和 query tile 重叠,腾出共享内存给更深的 KV 流水线。这是训练细节,推理不受影响,一句话带过。
3
是怎么选出来的
Kimi Linear 在一个 653M 激活、16 层 16 头的模型上扫了混合比(验证集 PPL,越低越好):
| KDA : MLA | 训练 PPL | 验证 PPL |
|---|
| 0(纯 MLA) | 9.45 | 5.77 |
| 1 | 9.29 | 5.66 |
| 3 | 9.23 | 5.65 |
| 7 | 9.23 | 5.70 |
| 15 | 9.34 | 5.82 |
读法:7
训练损失一样但泛化差(验证集是分布外的高质量数据);1
验证一样但推理贵;纯全局注意力反而最差。3
是质量和吞吐的平衡点。同一篇论文也验证了这个比例下的 48B 模型在短上下文、128K 长上下文和 RL 上都优于同配方的纯 MLA 模型。
1M 上下文的账
把前几篇的数字放到一起。BF16,每个数 2 字节:
| 每 token 每层 | 层数 | 1M 上下文总量 |
|---|
| Gated MLA 的 KV cache | 576 个数 | 24 | 13,824 个数/token → 27.6 GB |
| 假如 K2 那样 61 层全 MLA | 576 个数 | 61 | 35,136 个数/token → 70 GB |
| 假如不用 MLA,96 头 MHA | 24,576 个数 | 24 | 590K 个数/token → 1.18 TB |
| KDA 的状态(S 加卷积窗口) | 与 token 数无关 | 69 | 约 232 MB,固定 |
三件事从表里能看出来:
- MLA 本身把每层 cache 压了 43 倍(24576 → 576)。
- 3 混合又把需要增长的层从 93 降到 24。Kimi Linear 说的「KV cache 减少 75%」就是这个比例。
- KDA 的全部状态只相当于约 8K 个 token 的 MLA cache。上下文越长,这部分越可以忽略。
decode 是访存瓶颈,每生成一个 token 要把 cache 全读一遍,所以 decode 速度的上限由 3
决定。Kimi Linear 的数字:batch 为 1 时 1M 上下文的 TPOT 约为纯 MLA 的 2.2 到 2.3 倍;省下的显存换成更大的 batch 之后才是宣传里的 6.3 倍。prefill 在 1M 时约 2.9 倍。这些是 Kimi Linear 48B 模型的数字,K3 没有单独给。
术语坑
num_key_value_heads = 96 在 config 里,但对 MLA 没有意义,MLA 的 KV 是 latent 共享的。
- q 和 k 是 192 维,v 是 128 维。 用 flash attention 时代码把 v 补零到 192 再切回来。缩放因子是 192−1/2。
- HF 代码里的 cache 是展开后的 K、V(每头 192 维 key 加 128 维 value),不是 576 维的 latent。那是参考实现的偷懒写法,正式推理引擎才做吸收。文章里的 576 是按可压缩的口径算的。
- 「rope 维」这个名字在 K3 里已经名不副实,它只是 64 维的头间共享键分量。
下一篇
序列维度讲完了。第 4 篇换方向,看深度:Attention Residuals,K3 把残差连接换成了深度上的 softmax 注意力。