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

Kimi K3 模型结构(3):序列维度,Gated MLA 与 3:1 混合

每四层一层的全局注意力。MLA 的低秩 KV 压缩按 K3 的真实维度走一遍,解释为什么 K3 敢让所有 MLA 层完全不带位置编码,满秩输出门是什么,3:1 这个比例怎么选出来的,以及 1M 上下文时 KV cache 和 KDA 状态各占多少。

目录9 节

对应论文 §2.1 的开头和 §2.1.2,公式 (7)。前两篇讲的 KDA 用固定大小的状态换掉了 KV cache,代价是记忆有损。K3 每四层留一层 softmax 注意力做无损的全局检索,这一层是 DeepSeek-V2 的 MLA,加上两处改动:不带位置编码,输出加满秩 sigmoid 门。

文本 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零件
你现在在这里:每个重复块的第四层,以及主干末尾额外的第 93 层。

排布:3
,末尾再补一层

第 1 层:KDA(后接 dense FFN)第 2 层:KDA(后接 Stable LatentMoE)第 3 层:KDA(后接 Stable LatentMoE)第 4 层:Gated MLA(后接 Stable LatentMoE)第 5 层:KDA(后接 Stable LatentMoE)第 6 层:KDA(后接 Stable LatentMoE)第 7 层:KDA(后接 Stable LatentMoE)第 8 层:Gated MLA(后接 Stable LatentMoE)第 9 层:KDA(后接 Stable LatentMoE)第 10 层:KDA(后接 Stable LatentMoE)第 11 层:KDA(后接 Stable LatentMoE)第 12 层:Gated MLA(后接 Stable LatentMoE)第 13 层:KDA(后接 Stable LatentMoE)第 14 层:KDA(后接 Stable LatentMoE)第 15 层:KDA(后接 Stable LatentMoE)第 16 层:Gated MLA(后接 Stable LatentMoE)第 17 层:KDA(后接 Stable LatentMoE)第 18 层:KDA(后接 Stable LatentMoE)第 19 层:KDA(后接 Stable LatentMoE)第 20 层:Gated MLA(后接 Stable LatentMoE)第 21 层:KDA(后接 Stable LatentMoE)第 22 层:KDA(后接 Stable LatentMoE)第 23 层:KDA(后接 Stable LatentMoE)第 24 层:Gated MLA(后接 Stable LatentMoE)第 25 层:KDA(后接 Stable LatentMoE)第 26 层:KDA(后接 Stable LatentMoE)第 27 层:KDA(后接 Stable LatentMoE)第 28 层:Gated MLA(后接 Stable LatentMoE)第 29 层:KDA(后接 Stable LatentMoE)第 30 层:KDA(后接 Stable LatentMoE)第 31 层:KDA(后接 Stable LatentMoE)第 32 层:Gated MLA(后接 Stable LatentMoE)第 33 层:KDA(后接 Stable LatentMoE)第 34 层:KDA(后接 Stable LatentMoE)第 35 层:KDA(后接 Stable LatentMoE)第 36 层:Gated MLA(后接 Stable LatentMoE)第 37 层:KDA(后接 Stable LatentMoE)第 38 层:KDA(后接 Stable LatentMoE)第 39 层:KDA(后接 Stable LatentMoE)第 40 层:Gated MLA(后接 Stable LatentMoE)第 41 层:KDA(后接 Stable LatentMoE)第 42 层:KDA(后接 Stable LatentMoE)第 43 层:KDA(后接 Stable LatentMoE)第 44 层:Gated MLA(后接 Stable LatentMoE)第 45 层:KDA(后接 Stable LatentMoE)第 46 层:KDA(后接 Stable LatentMoE)第 47 层:KDA(后接 Stable LatentMoE)第 48 层:Gated MLA(后接 Stable LatentMoE)第 49 层:KDA(后接 Stable LatentMoE)第 50 层:KDA(后接 Stable LatentMoE)第 51 层:KDA(后接 Stable LatentMoE)第 52 层:Gated MLA(后接 Stable LatentMoE)第 53 层:KDA(后接 Stable LatentMoE)第 54 层:KDA(后接 Stable LatentMoE)第 55 层:KDA(后接 Stable LatentMoE)第 56 层:Gated MLA(后接 Stable LatentMoE)第 57 层:KDA(后接 Stable LatentMoE)第 58 层:KDA(后接 Stable LatentMoE)第 59 层:KDA(后接 Stable LatentMoE)第 60 层:Gated MLA(后接 Stable LatentMoE)第 61 层:KDA(后接 Stable LatentMoE)第 62 层:KDA(后接 Stable LatentMoE)第 63 层:KDA(后接 Stable LatentMoE)第 64 层:Gated MLA(后接 Stable LatentMoE)第 65 层:KDA(后接 Stable LatentMoE)第 66 层:KDA(后接 Stable LatentMoE)第 67 层:KDA(后接 Stable LatentMoE)第 68 层:Gated MLA(后接 Stable LatentMoE)第 69 层:KDA(后接 Stable LatentMoE)第 70 层:KDA(后接 Stable LatentMoE)第 71 层:KDA(后接 Stable LatentMoE)第 72 层:Gated MLA(后接 Stable LatentMoE)第 73 层:KDA(后接 Stable LatentMoE)第 74 层:KDA(后接 Stable LatentMoE)第 75 层:KDA(后接 Stable LatentMoE)第 76 层:Gated MLA(后接 Stable LatentMoE)第 77 层:KDA(后接 Stable LatentMoE)第 78 层:KDA(后接 Stable LatentMoE)第 79 层:KDA(后接 Stable LatentMoE)第 80 层:Gated MLA(后接 Stable LatentMoE)第 81 层:KDA(后接 Stable LatentMoE)第 82 层:KDA(后接 Stable LatentMoE)第 83 层:KDA(后接 Stable LatentMoE)第 84 层:Gated MLA(后接 Stable LatentMoE)第 85 层:KDA(后接 Stable LatentMoE)第 86 层:KDA(后接 Stable LatentMoE)第 87 层:KDA(后接 Stable LatentMoE)第 88 层:Gated MLA(后接 Stable LatentMoE)第 89 层:KDA(后接 Stable LatentMoE)第 90 层:KDA(后接 Stable LatentMoE)第 91 层:KDA(后接 Stable LatentMoE)第 92 层:Gated MLA(后接 Stable LatentMoE)第 93 层:Gated MLA(后接 Stable LatentMoE)11224364860728493■ KDA ×69■ Gated MLA ×24(每第 4 层 + 第 93 层)虚线框 = 第 1 层,FFN 是 dense 而不是 MoE
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=WkvxtR512,[ktnope; vt]=WkvRMSNorm(ct),c_t = W^{\downarrow}_{kv}\, x_t \in \mathbb{R}^{512}, \qquad [k^{\text{nope}}_t;\ v_t] = W^{\uparrow}_{kv}\, \operatorname{RMSNorm}(c_t),

只缓存 ctc_t,用的时候再由 WkvW^{\uparrow}_{kv} 重建每个头的 K 和 V。推理时还可以把 WkvW^{\uparrow}_{kv} 吸收进 query 和输出投影,直接在 512 维的 latent 上做注意力,所有头共享同一份 key 和 value,也就是变成 MQA。

q_t:96 头 × (128 + 64)k_nope 96 × 128v 96 × 128k_shared 64,拼到每个头的 k 后面c_t 512共享分量 64KV cache:每 token 576 个数= 512 维 latent c_t + 64 维共享键õ_t 96 × 128 = 12288门 12288,∈ (0,1),逐通道y_t 7168x_t7168Linear W_q↓7168 → 1536RMSNorm1536Linear W_q↑1536 → 96 × 192Linear W_kv↓7168 → 512 + 64RMSNorm512Linear W_kv↑512 → 96 × (128 + 128)广播到 96 个头64 维,不做旋转(NoPE)softmax(q kᵀ / √192) v因果,96 头,无位置编码Linear W_g(满秩)7168 → 12288SigmoidLinear W_o12288 → 7168
一层 Gated MLA。中间那条竖的橙色细框是推理时真正要缓存的东西:512 维 latent 加 64 维头间共享的键,共 576 个数。q 和 k 都是 192 维,其中 64 维在 DeepSeek 原版里用来放 RoPE,K3 保留了这个槽位但不做旋转。

按 K3 的 config 走一遍(q_lora_rank = 1536kv_lora_rank = 512qk_nope_head_dim = 128qk_rope_head_dim = 64v_head_dim = 128):

query。 xtx_t(7168)→ 1536 维低秩 → RMSNorm → 96 × 192。每个头的 query 是 192 维,由 128 维的内容部分和 64 维的「rope 槽位」拼成。

key 和 value。 xtx_t → 576 维,拆成 512 维 latent 和 64 维共享分量。latent 过 RMSNorm,再由 512 → 96 × 256 的上投影展开成每头 128 维的 knopek^{\text{nope}} 和 128 维的 vv。64 维共享分量直接广播给所有 96 个头,拼在每个头的 key 后面,凑成 192 维。

注意力。 96 个头各自在 192 维上算 qk/192q k^\top / \sqrt{192},因果 softmax,加权 128 维的 vv。输出 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 维的 krotk^{\text{rot}} 直接 expand 到所有头、和 knopek^{\text{nope}} 拼起来,不做任何旋转。q 那边同样。于是这 64 维退化成一个所有头共享的、每 token 一个的额外键分量。论文没有说为什么保留这个槽位而不是干脆去掉,这是从代码读出来的。

为什么敢不要位置编码。第 1 篇推过:KDA 的读出 qt(jDiag(αj)(Iβjkjkj))kiq_t^\top \big(\prod_{j} \operatorname{Diag}(\alpha_j)(I - \beta_j k_j k_j^\top)\big) k_i 和 RoPE 的 qt(jRj)kiq_t^\top \big(\prod_j R_j\big) k_i 形状一样,只是 RjR_j 换成了数据相关、可学习的对角转移。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].y_t = W_o\big[\operatorname{Sigmoid}(W_g x_t) \odot \tilde o_t\big].

o~t\tilde o_t 是没加门的 MLA 输出(12288 维),WgW_g 是 7168 → 12288 的满秩矩阵,和 K3 的 KDA 输出门同一种参数化。config 里 mla_use_output_gate = true,代码里 g_proj

python
# 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.455.77
1
9.295.66
3
9.235.65
7
9.235.70
15
9.345.82

读法:7

训练损失一样但泛化差(验证集是分布外的高质量数据);1
验证一样但推理贵;纯全局注意力反而最差。3
是质量和吞吐的平衡点。同一篇论文也验证了这个比例下的 48B 模型在短上下文、128K 长上下文和 RL 上都优于同配方的纯 MLA 模型。

1M 上下文的账

把前几篇的数字放到一起。BF16,每个数 2 字节:

每 token 每层层数1M 上下文总量
Gated MLA 的 KV cache576 个数2413,824 个数/token → 27.6 GB
假如 K2 那样 61 层全 MLA576 个数6135,136 个数/token → 70 GB
假如不用 MLA,96 头 MHA24,576 个数24590K 个数/token → 1.18 TB
KDA 的状态(SS 加卷积窗口)与 token 数无关69约 232 MB,固定

三件事从表里能看出来:

  1. MLA 本身把每层 cache 压了 43 倍(24576 → 576)。
  2. 3
    混合
    又把需要增长的层从 93 降到 24。Kimi Linear 说的「KV cache 减少 75%」就是这个比例。
  3. 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 再切回来。缩放因子是 1921/2192^{-1/2}
  • HF 代码里的 cache 是展开后的 K、V(每头 192 维 key 加 128 维 value),不是 576 维的 latent。那是参考实现的偷懒写法,正式推理引擎才做吸收。文章里的 576 是按可压缩的口径算的。
  • 「rope 维」这个名字在 K3 里已经名不副实,它只是 64 维的头间共享键分量。

下一篇

序列维度讲完了。第 4 篇换方向,看深度:Attention Residuals,K3 把残差连接换成了深度上的 softmax 注意力。