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

Kimi K3 模型结构(2):序列维度(下),chunkwise 并行与下界衰减

KDA 的 chunkwise 并行形式:WY 表示、UT 变换、块间递推加块内矩阵乘。然后是 K3 的改动:给每步的 log-decay 加一个 −5 的下界,让 16-token 对角 tile 也能用 Tensor Core 算。附一个可以拖的曲线对比。

目录11 节

对应论文 §2.1.1 的后半:公式 (3) 到 (5) 和 Figure 3。上一篇的递推形式一个 token 走一步,全是向量和矩阵的逐元素运算,训练时用不上 Tensor Core。这一篇讲怎么把一个块(C=64C = 64 个 token)的计算改写成几次矩阵乘,然后你会看到 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 里。这一篇讲的是同一个模块在训练和 prefill 时的算法形式。

目标:块间递推,块内并行

把序列切成长度 CC 的块。记 X[t]\mathbf{X}_{[t]} 为第 tt 块里的 token 堆成的矩阵(Q,K,VRC×d\mathbf{Q}, \mathbf{K}, \mathbf{V} \in \mathbb{R}^{C \times d}),S[t]\mathbf{S}_{[t]} 为进入第 tt 块时的状态。想要的形式是:

  • 块间:状态只在块边界更新一次,S[t]S[t+1]\mathbf{S}_{[t]} \to \mathbf{S}_{[t+1]}
  • 块内:块里 CC 个输出一起算,用矩阵乘。

先定义逐通道的累积衰减(论文公式 (3)):

γ[t]ij:=r=ijα[t]r,γ[t]r:=γ[t]1r,\gamma_{[t]}^{i \to j} := \prod_{r=i}^{j} \alpha_{[t]}^{r}, \qquad \gamma_{[t]}^{r} := \gamma_{[t]}^{1 \to r},

Γ[t]1CRC×dk\Gamma_{[t]}^{1\to C} \in \mathbb{R}^{C \times d_k}γ1,,γC\gamma^1, \dots, \gamma^C 按行堆起来。γr\gamma^r 是「从块开头到位置 rr 一共衰减了多少」,逐通道。

块内展开:PH

把递推在块内展开 rr 步(Kimi Linear 公式 (2)):

S[t]r=i=1r(Iβikiki)Diag(αi)Pr S[t]0 + i=1r(j=i+1r(Iβjkjkj)Diag(αj))βikiviHr.\mathbf{S}^{r}_{[t]} = \underbrace{\prod_{i=1}^{r} \big(I - \beta^i k^i k^{i\top}\big)\operatorname{Diag}(\alpha^i)}_{\mathbf{P}^r}\ \mathbf{S}^0_{[t]} \ +\ \underbrace{\sum_{i=1}^{r} \Big(\prod_{j=i+1}^{r} \big(I - \beta^j k^j k^{j\top}\big)\operatorname{Diag}(\alpha^j)\Big)\beta^i k^i v^{i\top}}_{\mathbf{H}^r}.

Pr\mathbf{P}^r 是进入块时的状态被怎样变换,Hr\mathbf{H}^r 是块内写入的贡献。两项都是一串带衰减的 Householder 矩阵的乘积,直接算还是顺序的。

WY 表示:一串 Householder 乘积等于「对角减低秩和」

Kimi Linear 附录 B 的两个命题(归纳法证):

Pr=Diag(γr)i=1rDiag(γir)kiwi,Hr=i=1rDiag(γir)kiui,\mathbf{P}^r = \operatorname{Diag}(\gamma^r) - \sum_{i=1}^{r} \operatorname{Diag}(\gamma^{i\to r})\, k^i w^{i\top}, \qquad \mathbf{H}^r = \sum_{i=1}^{r} \operatorname{Diag}(\gamma^{i\to r})\, k^i u^{i\top},

其中辅助向量由两个三角递推给出:

wr=βr(Diag(γr)kri=1r1wi(kiDiag(γir)kr)),ur=βr(vri=1r1ui(kiDiag(γir)kr)).w^r = \beta^r \Big(\operatorname{Diag}(\gamma^r) k^r - \sum_{i=1}^{r-1} w^i \big(k^{i\top}\operatorname{Diag}(\gamma^{i\to r}) k^r\big)\Big), \qquad u^r = \beta^r \Big(v^r - \sum_{i=1}^{r-1} u^i \big(k^{i\top}\operatorname{Diag}(\gamma^{i\to r}) k^r\big)\Big).

怎么理解 uuwwuru^r伪值,等于 vrv^r 减掉块内更早的键已经解释掉的部分,也就是 delta rule 的「先擦再写」被限制在块内的版本。wrw^r 是配套的衰减键,告诉你这个修正要怎么作用到进入块时的状态上。

UT 变换:一次 C×C 三角求逆

两个递推可以对整个块一次解出(Kimi Linear 公式 (6)、(7)):

M=(I+StrictTril(Diag(β)(ΓK)(KΓ) ⁣))1Diag(β),W=M(ΓK),U=MV.\mathbf{M} = \Big(I + \operatorname{StrictTril}\Big(\operatorname{Diag}(\beta)\,(\Gamma \odot \mathbf{K})\Big(\frac{\mathbf{K}}{\Gamma}\Big)^{\!\top}\Big)\Big)^{-1} \operatorname{Diag}(\beta), \qquad \mathbf{W} = \mathbf{M}(\Gamma \odot \mathbf{K}), \qquad \mathbf{U} = \mathbf{M}\mathbf{V}.

单位下三角矩阵的逆用前向替换做,代价 O(C2)O(C^2) 一次,之后全是矩阵乘。这一步把非矩阵乘的顺序工作变成了矩阵乘,Kimi Linear 说它「对硬件利用率至关重要」。

块间递推与块内输出

有了 U\mathbf{U}W\mathbf{W},定义伪值项 V~[t]:=U[t]W[t]S[t]\widetilde{\mathbf{V}}_{[t]} := \mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]},K3 论文公式 (4):

A[t]=Tril ⁣[(Q[t]Γ[t]1C)(K[t]/Γ[t]1C) ⁣],O[t]=(Γ[t]1CQ[t])S[t]块间+A[t]V~[t]块内.\mathbf{A}_{[t]} = \operatorname{Tril}\!\Big[(\mathbf{Q}_{[t]} \odot \Gamma_{[t]}^{1\to C})\big(\mathbf{K}_{[t]} / \Gamma_{[t]}^{1\to C}\big)^{\!\top}\Big], \qquad \mathbf{O}_{[t]} = \underbrace{(\Gamma_{[t]}^{1\to C} \odot \mathbf{Q}_{[t]})\,\mathbf{S}_{[t]}}_{\text{块间}} + \underbrace{\mathbf{A}_{[t]}\,\widetilde{\mathbf{V}}_{[t]}}_{\text{块内}}.

状态更新(Kimi Linear 公式 (8)):

S[t+1]=Diag(γ[t]C)S[t]+(Γ[t]iCK[t]) ⁣V~[t].\mathbf{S}_{[t+1]} = \operatorname{Diag}(\gamma^C_{[t]})\,\mathbf{S}_{[t]} + \big(\Gamma^{i\to C}_{[t]} \odot \mathbf{K}_{[t]}\big)^{\!\top}\,\widetilde{\mathbf{V}}_{[t]}.
块间通道:状态 S 逐块递推(每块一次)S_[t]× Diag(γ^C)整块的逐通道衰减+S_[t+1](Γ^{i→C} ⊙ K)ᵀ Ṽ本块的写入,一次矩阵乘块内通道:本块所有输出并行算(一次 C×C 掩码矩阵乘)(Γ ⊙ Q) S_[t]读进入本块时的状态+A Ṽ,A = Tril[(Q⊙Γ)(K/Γ)ᵀ]块内 token 之间的因果交互O_[t]S_[t] 也喂给块内通道A 的下三角,C = 64,16-token 二级 tile查询 i键 j ≤ i非对角 tile:exp(g_i − g_s)·exp(g_s − g_j),两个因子都 ≤ 1,直接 BF16 矩阵乘对角 tile:exp(g_s − g_j) 可能 ≫ 1Kimi Linear:逐位置对、FP32,是瓶颈K3:16 步累计 log-decay 大于 −80,e^80 在 BF16 范围内,也走矩阵乘

读法:

  • 块间那一行,先把整个状态按这一块的总衰减 γC\gamma^C 逐通道缩一下,再加上这一块的写入。写入是「每个键按它到块尾的剩余距离衰减」乘「修正后的伪值」。
  • 块内那一行,A\mathbf{A} 是一个 C×CC \times C 的下三角矩阵,第 (i,j)(i, j) 项是 qiq_ikjk_j 的内积,每个通道再乘上 γi/γj\gamma^i / \gamma^j,也就是从 jjii 之间的衰减。保留对角线,因为每个输出读的是「当前 token 更新之后」的状态。
  • Kimi Linear 的 FLOPs 公式(每头,C=64C = 64):6Tdh2+3TCdh+TC26Td_h^2 + 3TCd_h + TC^2,对 TT 线性;softmax 注意力是 2T2dh2T^2 d_h

数值问题:1/Γ 会溢出

公式 (4) 里 K/Γ\mathbf{K}/\Gamma 把键除以累积衰减。Γ\Gamma 是一串 (0,1)(0,1) 里的数的乘积,倒数可以任意大,半精度下溢出。

标准解法(GLA 的做法,Kimi Linear 沿用)是进 log 空间:令 gr=logγrg_r = \log \gamma^r(对每步的 log-decay 做前缀和),那么 γi/γj=exp(gigj)\gamma^i/\gamma^j = \exp(g_i - g_j),对 jij \le i 永远不超过 1。问题是矩阵乘要求把 exp(gigj)\exp(g_i - g_j) 拆成「只依赖行」乘「只依赖列」的两个因子,拆开就又回到了 exp(gi)exp(gj)\exp(g_i)\cdot\exp(-g_j),后者会炸。

于是再把 C=64C = 64 的块切成 16-token 的二级 tile,以 tile 起点 ss 为参考:

  • 非对角 tile(查询 ii 在后面的 tile,键 jj 在前面的 tile):exp(gigj)=exp(gigs)exp(gsgj)\exp(g_i - g_j) = \exp(g_i - g_s)\cdot\exp(g_s - g_j)is>ji \ge s > j,两个因子都不超过 1,可以放心拆成两边、用 BF16 矩阵乘。
  • 对角 tileiijj 同在一个 tile):同样拆法里 exp(gsgj)\exp(g_s - g_j)1\ge 1 的,16 步累计能有多大取决于每步的 log-decay 有没有下界。Kimi Linear 的映射没有下界,所以对角 tile 只能逐位置对、用 FP32 算 exp(gigj)\exp(g_i - g_j),不走 Tensor Core。论文说这是块内计算的主要瓶颈。

K3 的改动:下界衰减

从 logit zz 到每步 log-decay gg 的映射,Kimi Linear 沿用 GDN 和 Mamba-2:

gth=eAhSoftplus(zth)(,0)dk.g_t^h = -e^{A_h}\operatorname{Softplus}(z_t^h) \in (-\infty, 0)^{d_k}.

K3 换成一个带尺度的 sigmoid(论文公式 (5)):

gth=gminSigmoid ⁣(eAhzth)(gmin,0)dk,αth=exp(gth)(egmin,1)dk,g_t^h = g_{\min}\,\operatorname{Sigmoid}\!\big(e^{A_h} z_t^h\big) \in (g_{\min}, 0)^{d_k}, \qquad \alpha_t^h = \exp(g_t^h) \in \big(e^{g_{\min}}, 1\big)^{d_k},

gmin=5g_{\min} = -5 固定,AhA_h 是每头一个的可学习 log 尺度,初始化为 0。config 里就是 gate_lower_bound = -5.0

z(衰减 logit)
Kimi Linear:g = −e^A · softplus(z)Kimi K3:g = −5 · σ(e^A z)g_min = −5(α 最小 e^−5 ≈ 0.0067)

两条曲线的差别:负 softplus 在 z+z \to +\infty 时线性往下走,没有底;scaled sigmoid 在 z+z \to +\infty 时贴到 gming_{\min}AhA_h 只改横轴的伸缩。

后果一:数值范围有界。 每步 α>e56.7×103\alpha > e^{-5} \approx 6.7\times10^{-3},16 步累计的 log-decay 在 (80,0)(-80, 0) 里,exp(gsgj)<e805.5×1034\exp(g_s - g_j) < e^{80} \approx 5.5\times10^{34},在 BF16 的动态范围(最大约 3.4×10383.4\times10^{38})之内。于是对角 tile 也能拆成两个因子做 BF16 矩阵乘,Figure 3b 说的「消除逐位置对的对角路径」就是这个意思。gmin=5g_{\min} = -5 和 tile 宽度 16 是配套选的:5×16=805 \times 16 = 80 正好卡在 BF16 指数范围以内。

后果二:记忆不会靠衰减被瞬间清空。 一个通道靠衰减能做到的最快遗忘是每步剩 0.67%。真正需要精确擦除的时候,delta rule 的 (Iβkk)(I - \beta k k^\top) 还在。论文没有讨论这一点对表达力的影响,只提了这种下界门在 RWKV-7、Griffin、HGRN2 里有先例。

后果三:AhA_h 的作用变了。 在负 softplus 里 eAhe^{A_h} 是纵向的尺度,直接决定衰减能多狠;在 scaled sigmoid 里它只能改横向的斜率,衰减的上限被 gming_{\min} 钉死。

为什么 KDA 内核比通用 DPLR 快一倍

上一篇提到 KDA 的转移是 DatbtD - a_t b_t^\top,且 at=βtkta_t = \beta_t k_tbt=ktαtb_t = k_t \odot \alpha_t 都绑在 ktk_t 上。在 chunkwise 形式里这意味着:通用 DPLR 内核需要四张二级 tile 矩阵(Aab,Aak,Aqb,Aqk\mathbf{A}_{ab}, \mathbf{A}_{ak}, \mathbf{A}_{qb}, \mathbf{A}_{qk}),KDA 只需要两张(Aqk,Akk\mathbf{A}_{qk}, \mathbf{A}_{kk});输出阶段 DPLR 的三项输出和两次状态更新,在 KDA 里合成一行输出和一次状态更新。Kimi Linear Figure 2 显示在 2K 到 64K 长度上 KDA 内核约为 DPLR 内核的 2 倍速度。K3 在此基础上做的 FlashKDA 内核放到第 7 篇。

训练、prefill、decode 各用哪种形式

阶段形式代码
训练chunkwisechunk_kda(..., safe_gate=True, lower_bound=-5.0)
prefillchunkwise同上
decode(每次 1 个 token)递推fused_recurrent_kda(..., lower_bound=-5.0)

HF 代码里 mode = 'fused_recurrent' if use_cache and q_len == 1 else 'chunk'。L2 归一化、β\beta 的 sigmoid、α\alpha 的映射都在内核里做(use_qk_l2norm_in_kerneluse_beta_sigmoid_in_kerneluse_gate_in_kernel),Python 侧只算到 logit。

术语坑

  • 两个 A\mathbf{A} 公式 (4) 的 A[t]\mathbf{A}_{[t]}C×CC\times C 的块内注意力矩阵,公式 (5) 的 AhA_h 是每头的标量 log 尺度。上一篇 NoPE 那段的 Aj=Diag(αj)A_j = \operatorname{Diag}(\alpha_j) 又是第三个。
  • 16 和 64。 块大小 C=64C = 64 来自 Kimi Linear 的 FLOPs 分析和伪代码,16-token 二级 tile 是 FLA 内核的实现细节,K3 论文第一次把它写进正文。
  • HF 权重里 A_log 的初始化log(Uniform(1,16))\log(\mathrm{Uniform}(1,16)),那是 Kimi Linear 的初始化写法留在代码里;K3 论文说 AhA_h 初始化为 0。推理时无所谓,复现训练时以论文为准。

下一篇

序列维度还剩四分之一:每四层一层的 Gated MLA。它负责 KDA 做不了的事,无损的全局内容检索。下一篇讲 MLA 的低秩压缩、K3 为什么敢让它完全不带位置编码、满秩输出门,以及 3

这个比例是怎么选出来的。