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

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

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

目录10 节

对应论文 §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词嵌入永远是一个来源= 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 里。这一篇讲的是同一个模块在训练和 prefill 时的算法形式。

起点:公式 (1) 的递推,为什么训练时用不了

先把上一篇的公式 (1) 摆出来,这一篇所有推导都从它出发:

St=(I−βtktkt⊤)Diag⁡(αt) St−1+βtktvt⊤,ot=St⊤qt.(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 o_t = S_t^\top q_t . \tag{1}
新状态② 擦除① 衰减旧状态③ 写入=··+从右往左读:先把旧状态每一行按各自衰减,再把落在方向上的内容擦掉的比例,最后写入倍的新关联。读出:。行 = 键通道(,各自的衰减率),列 = 值通道()。

一个 token 走一步:先把旧状态每一行按 αt\alpha_t 衰减,再沿 ktk_t 方向擦掉 βt\beta_t 的比例,最后写入 βtktvt⊤\beta_t k_t v_t^\top。每步代价固定,对序列长度已经是线性的了,为什么还要改?答案在硬件上。

Tensor Core 只接受矩阵乘矩阵。H100 上 BF16 矩阵乘的峰值约 990 TFLOPS,通用 CUDA core 上的 FP32 约 67 TFLOPS,差十几倍。公式 (1) 里的两个操作,外积 ktvt⊤k_t v_t^\top 和矩阵乘向量 S⊤qtS^\top q_t,都只有一个向量维,填不满一个 tile,只能走慢的那一边。更糟的是每个 token 都要把整个 128×128 的状态读一遍再写一遍,而且第 tt 步必须等第 t−1t-1 步做完,1M 的序列就是一百万步串行。

把 CC 个 token 摞在一起,向量就变成 C×128C \times 128 的矩阵,乘法有了第二个矩阵维度,可以上 Tensor Core。状态每 CC 个 token 读写一次,串行链从一百万步变成 106/C10^6 / C 步。代价是块内 CC 个 token 之间的依赖要另外算,这就是下面全部推导要解决的事。结果和逐 token 递推完全一致,变的只是硬件调度。

逐 token:一个向量维一个 tile(示意 16 × 16)是外积,是矩阵乘向量都只有一个向量维,填不满 tile走 CUDA core,FP32 约 67 TFLOPS每个 token 重新读一遍整个 128 × 128 状态16 个 token 摞起来:两个矩阵维同一个 tile与,矩阵乘矩阵,tile 填满走 Tensor Core,BF16 约 990 TFLOPS状态每 16 个 token 读写一次;1M 序列串行链步
Tensor Core 只接受矩阵乘矩阵。递推形式每个 token 的代价已经是常数,但两个操作都只有一个向量维;把 个 token 摞成矩阵,乘法才有第二个矩阵维度。数字是 H100 的峰值。

终点:块间一步递推,块内一次矩阵乘

把序列切成长度 CC 的块(Kimi Linear 按 C=64C = 64 分析)。记 Q[t],K[t],V[t]∈RC×d\mathbf{Q}_{[t]}, \mathbf{K}_{[t]}, \mathbf{V}_{[t]} \in \mathbb{R}^{C \times d} 为第 tt 块里的 token 堆成的矩阵,一行一个 token;S[t]\mathbf{S}_{[t]} 为进入第 tt 块时的状态。还要一个记号来表示「从块内位置 ii 走到位置 jj 一共衰减了多少」,逐通道:

γi→j:=∏r=i+1jαr∈Rdk,γr:=γ0→r=α1⊙⋯⊙αr,\gamma^{i \to j} := \prod_{r=i+1}^{j} \alpha^{r} \in \mathbb{R}^{d_k}, \qquad \gamma^{r} := \gamma^{0 \to r} = \alpha^1 \odot \cdots \odot \alpha^r,

Γ1→C∈RC×dk\Gamma^{1\to C} \in \mathbb{R}^{C \times d_k} 把 γ1,…,γC\gamma^1, \dots, \gamma^C 按行堆起来。γr\gamma^r 是「从块开头到位置 rr」的累积衰减,γi→j=γj/γi\gamma^{i\to j} = \gamma^j / \gamma^i(逐元素除)。

这一篇的核心公式是块间的状态更新(Kimi Linear 公式 (8)):

  S[t+1]=Diag⁡(γ[t]C) S[t]+(Γ[t]i→C⊙K[t]) ⁣⊤ V~[t],V~[t]:=U[t]−W[t]S[t]  (8)\boxed{\; \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]}, \qquad \widetilde{\mathbf{V}}_{[t]} := \mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]} \;} \tag{8}

配套的是块内的输出(Kimi Linear 公式 (9),K3 论文公式 (4)):

A[t]=Tril⁡ ⁣[(Γ[t]1→C⊙Q[t])(K[t]/Γ[t]1→C) ⁣⊤],O[t]=(Γ[t]1→C⊙Q[t]) S[t]⏟块间+A[t] V~[t]⏟块内.(9)\mathbf{A}_{[t]} = \operatorname{Tril}\!\Big[(\Gamma_{[t]}^{1\to C} \odot \mathbf{Q}_{[t]})\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{块内}}. \tag{9}

两条公式里的 U\mathbf{U}、W\mathbf{W} 由一次 C×CC \times C 的三角求逆给出(Kimi Linear 公式 (6)、(7)),这一步叫 UT 变换:

M=(I+StrictTril⁡(Diag⁡(β) (Γ⊙K)(KΓ) ⁣⊤))−1Diag⁡(β),W=M(Γ⊙K),U=MV.(6, 7)\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}. \tag{6, 7}
块间通道:状态逐块递推(每块一次)整块的逐通道衰减+本块的写入,一次矩阵乘块内通道:本块所有输出并行算(一次掩码矩阵乘)读进入本块时的状态+,块内 token 之间的因果交互也喂给块内通道的下三角,,16-token 二级 tile查询键非对角 tile:,两个因子都,直接 BF16 矩阵乘对角 tile:可能Kimi Linear:逐位置对、FP32,是瓶颈K3:16 步累计 log-decay 大于,在 BF16 范围内,也走矩阵乘

读法:

  • 块间那一行,先把整个状态按这一块的总衰减 γC\gamma^C 逐通道缩一下,再加上这一块的写入。写入是「每个键按它到块尾的剩余距离衰减」乘「修正后的伪值」V~\widetilde{\mathbf{V}}。伪值的意思是:U\mathbf{U} 是 V\mathbf{V} 减掉块内更早的键已经解释掉的部分,也就是 delta rule 的「先擦再写」被限制在块内的版本;WS[t]\mathbf{W}\mathbf{S}_{[t]} 是这些擦除作用到进入块时的状态上的那一部分。
  • 块内那一行,A\mathbf{A} 是一个 C×CC \times C 的下三角矩阵,第 (i,j)(i, j) 项是 qiq_i 和 kjk_j 的内积,每个通道再乘上 γi/γj\gamma^i / \gamma^j,也就是从 jj 到 ii 之间的衰减。保留对角线,因为每个输出读的是「当前 token 更新之后」的状态。
  • 哪些是矩阵乘:A\mathbf{A} 是 (C×dk)(dk×C)(C \times d_k)(d_k \times C),AV~\mathbf{A}\widetilde{\mathbf{V}} 是 (C×C)(C×dv)(C \times C)(C \times d_v),块间写入是 (dk×C)(C×dv)(d_k \times C)(C \times d_v),WS\mathbf{W}\mathbf{S} 是 (C×dk)(dk×dv)(C \times d_k)(d_k \times d_v)。除了求逆,全是 Tensor Core 能吃的形状。
  • Kimi Linear 的 FLOPs 公式(每头,C=64C = 64):6Tdh2+3TCdh+TC26Td_h^2 + 3TCd_h + TC^2,对 TT 线性;softmax 注意力是 2T2dh2T^2 d_h。

公式 (1) 是怎么变成公式 (8) 和 (9) 的,完整推导折叠在下面。只想知道结论的话可以跳过,读后面 K3 的改动不需要它;想跟着算一遍的话,只用到矩阵乘法的结合律、对角矩阵的性质和数学归纳法。

从公式 (1) 到公式 (8)、(9):完整推导

记号。 只看一个块,省掉下标 [t][t]。块内位置 r=1,…,Cr = 1, \dots, C,kr,vr,qr,αrk^r, v^r, q^r, \alpha^r 是第 rr 个 token 的向量,βr\beta^r 是标量。S0S^0 是进入块时的状态(也就是 S[t]\mathbf{S}_{[t]}),SrS^r 是处理完第 rr 个 token 之后的状态,所以 SC=S[t+1]S^C = \mathbf{S}_{[t+1]}。把公式 (1) 每一步的转移矩阵记成

Tr:=(I−βrkrkr⊤)Diag⁡(αr)∈Rdk×dk,Sr=TrSr−1+βrkrvr⊤.T^r := \big(I - \beta^r k^r k^{r\top}\big)\operatorname{Diag}(\alpha^r) \in \mathbb{R}^{d_k \times d_k}, \qquad S^r = T^r S^{r-1} + \beta^r k^r v^{r\top}.

第 1 步:把递推在块内展开,得到 PP 和 HH。

一步步代入。r=1r = 1 时就是定义:

S1=T1S0+β1k1v1⊤.S^1 = T^1 S^0 + \beta^1 k^1 v^{1\top}.

r=2r = 2 时把 S1S^1 代进去,用矩阵乘法的分配律拆开:

S2=T2S1+β2k2v2⊤=T2T1S0+T2 β1k1v1⊤+β2k2v2⊤.S^2 = T^2 S^1 + \beta^2 k^2 v^{2\top} = T^2 T^1 S^0 + T^2\,\beta^1 k^1 v^{1\top} + \beta^2 k^2 v^{2\top}.

规律出来了:每个 S0S^0 前面是所有转移矩阵的乘积,每个第 ii 步的写入 βikivi⊤\beta^i k^i v^{i\top} 前面是它之后的转移矩阵 Ti+1,…,TrT^{i+1}, \dots, T^r 的乘积(后发生的在左边,因为每次都是左乘)。一般地

Sr=(∏i=1rTi)⏟Pr S0 + ∑i=1r(∏j=i+1rTj)βikivi⊤⏟Hr,S^{r} = \underbrace{\Big(\prod_{i=1}^{r} T^i\Big)}_{P^r}\ S^0 \ +\ \underbrace{\sum_{i=1}^{r} \Big(\prod_{j=i+1}^{r} T^j\Big)\beta^i k^i v^{i\top}}_{H^r},

这里的 ∏\prod 约定按 TrTr−1⋯T1T^r T^{r-1} \cdots T^1 的顺序排,空乘积是 II。这就是 Kimi Linear 的公式 (2)。PrP^r 是「进入块时的状态被怎样变换」,HrH^r 是「块内写入的贡献」,两者满足和 SS 一样的递推:

Pr=TrPr−1,P0=I;Hr=TrHr−1+βrkrvr⊤,H0=0.P^r = T^r P^{r-1},\quad P^0 = I; \qquad H^r = T^r H^{r-1} + \beta^r k^r v^{r\top},\quad H^0 = 0 .

到这里只是换了个写法,PrP^r 和 HrH^r 里的连乘还是顺序的,没有省任何计算。

第 2 步:先看只有衰减的情形,看 γ\gamma 是从哪来的。

假设所有 βr=0\beta^r = 0,即没有擦除,那么 Tr=Diag⁡(αr)T^r = \operatorname{Diag}(\alpha^r)。对角矩阵相乘等于对角线逐元素相乘:

Diag⁡(αr)⋯Diag⁡(αi+1)=Diag⁡(αi+1⊙⋯⊙αr)=Diag⁡(γi→r).\operatorname{Diag}(\alpha^r)\cdots\operatorname{Diag}(\alpha^{i+1}) = \operatorname{Diag}(\alpha^{i+1}\odot\cdots\odot\alpha^r) = \operatorname{Diag}(\gamma^{i\to r}).

于是 Pr=Diag⁡(γr)P^r = \operatorname{Diag}(\gamma^r),Hr=∑iDiag⁡(γi→r) βikivi⊤H^r = \sum_i \operatorname{Diag}(\gamma^{i\to r})\,\beta^i k^i v^{i\top}。γi→r\gamma^{i\to r} 就是「从位置 ii 写进去的东西,到位置 rr 时还剩多少」,逐通道。i=ri = r 时是空乘积,γr→r=1\gamma^{r\to r} = \mathbf{1},刚写进去的还没来得及衰减。

后面还会反复用到两条对角矩阵的性质:

  • Diag⁡(αr)Diag⁡(γi→r−1)=Diag⁡(γi→r)\operatorname{Diag}(\alpha^r)\operatorname{Diag}(\gamma^{i\to r-1}) = \operatorname{Diag}(\gamma^{i\to r}),衰减累乘一步。
  • x⊤Diag⁡(γ)y=∑cγcxcycx^\top \operatorname{Diag}(\gamma) y = \sum_c \gamma_c x_c y_c,是一个标量,且等于 y⊤Diag⁡(γ)xy^\top\operatorname{Diag}(\gamma)x。

第 3 步:加回擦除,PrP^r 的 WY 表示。

现在 βr≠0\beta^r \ne 0。要证的命题(Kimi Linear 附录 B,命题 1)是:一串带衰减的 Householder 矩阵的乘积,等于「一个对角矩阵减去 rr 个秩 1 项」:

Pr=Diag⁡(γr)−∑i=1rDiag⁡(γi→r) kiwi⊤,wr:=βr(Diag⁡(γr)kr−∑i=1r−1(ki⊤Diag⁡(γi→r)kr) wi).P^r = \operatorname{Diag}(\gamma^r) - \sum_{i=1}^{r} \operatorname{Diag}(\gamma^{i\to r})\, k^i w^{i\top}, \qquad w^r := \beta^r \Big(\operatorname{Diag}(\gamma^r) k^r - \sum_{i=1}^{r-1} \big(k^{i\top}\operatorname{Diag}(\gamma^{i\to r}) k^r\big)\, w^i\Big).

这种「对角减低秩和」的写法叫 WY 表示,来自数值线性代数里把一串 Householder 反射合并成 I−YWY⊤I - YWY^\top 的技巧。wi∈Rdkw^i \in \mathbb{R}^{d_k} 是要构造的辅助向量,它的定义看起来循环(wrw^r 用到 w1,…,wr−1w^1,\dots,w^{r-1}),但这是一个三角递推,从 w1w^1 开始一个个算就行。

用归纳法。

起点 r=1r = 1。 直接乘开 T1T^1:

P1=T1=(I−β1k1k1⊤)Diag⁡(α1)=Diag⁡(α1)−β1k1(k1⊤Diag⁡(α1)).P^1 = T^1 = \big(I - \beta^1 k^1 k^{1\top}\big)\operatorname{Diag}(\alpha^1) = \operatorname{Diag}(\alpha^1) - \beta^1 k^1 \big(k^{1\top}\operatorname{Diag}(\alpha^1)\big).

γ1=α1\gamma^1 = \alpha^1,而 k1⊤Diag⁡(α1)k^{1\top}\operatorname{Diag}(\alpha^1) 是行向量,它的转置是 Diag⁡(α1)k1\operatorname{Diag}(\alpha^1)k^1(对角矩阵是对称的)。所以第二项是 k1w1⊤k^1 w^{1\top},其中 w1=β1Diag⁡(γ1)k1w^1 = \beta^1\operatorname{Diag}(\gamma^1)k^1,和定义里 r=1r=1 的情形(求和为空)一致。前面的系数 Diag⁡(γ1→1)=I\operatorname{Diag}(\gamma^{1\to 1}) = I。命题在 r=1r = 1 成立。

归纳步。 假设 Pr−1P^{r-1} 满足命题,计算 Pr=TrPr−1P^r = T^r P^{r-1}。TrT^r 有两个因子,先乘右边的 Diag⁡(αr)\operatorname{Diag}(\alpha^r):

Y:=Diag⁡(αr)Pr−1=Diag⁡(αr)Diag⁡(γr−1)−∑i=1r−1Diag⁡(αr)Diag⁡(γi→r−1) kiwi⊤=Diag⁡(γr)−∑i=1r−1Diag⁡(γi→r) kiwi⊤.Y := \operatorname{Diag}(\alpha^r) P^{r-1} = \operatorname{Diag}(\alpha^r)\operatorname{Diag}(\gamma^{r-1}) - \sum_{i=1}^{r-1} \operatorname{Diag}(\alpha^r)\operatorname{Diag}(\gamma^{i\to r-1})\, k^i w^{i\top} = \operatorname{Diag}(\gamma^r) - \sum_{i=1}^{r-1} \operatorname{Diag}(\gamma^{i\to r})\, k^i w^{i\top}.

用的是第 2 步的第一条性质:衰减多累乘了一步。注意 YY 已经具有命题要的形状,只是求和还差 i=ri = r 这一项。再乘左边的 (I−βrkrkr⊤)(I - \beta^r k^r k^{r\top}):

Pr=Y−βrkr(kr⊤Y).P^r = Y - \beta^r k^r \big(k^{r\top} Y\big).

关键在于 kr⊤Yk^{r\top}Y 是一个行向量(1×dk1 \times d_k),所以 βrkr(kr⊤Y)\beta^r k^r (k^{r\top}Y) 是一个秩 1 矩阵,正好是命题里缺的那一项。把它算出来:

kr⊤Y=kr⊤Diag⁡(γr)−∑i=1r−1(kr⊤Diag⁡(γi→r)ki)⏟标量 wi⊤.k^{r\top} Y = k^{r\top}\operatorname{Diag}(\gamma^r) - \sum_{i=1}^{r-1} \underbrace{\big(k^{r\top}\operatorname{Diag}(\gamma^{i\to r}) k^i\big)}_{\text{标量}}\, w^{i\top}.

第二项里 kr⊤Diag⁡(γi→r)kik^{r\top}\operatorname{Diag}(\gamma^{i\to r})k^i 是标量,可以挪到前面;按第 2 步第二条性质它等于 ki⊤Diag⁡(γi→r)krk^{i\top}\operatorname{Diag}(\gamma^{i\to r})k^r。把整个行向量转置成列向量,乘上 βr\beta^r:

βr(kr⊤Y) ⁣⊤=βr(Diag⁡(γr)kr−∑i=1r−1(ki⊤Diag⁡(γi→r)kr) wi)=wr.\beta^r \big(k^{r\top} Y\big)^{\!\top} = \beta^r\Big(\operatorname{Diag}(\gamma^r)k^r - \sum_{i=1}^{r-1} \big(k^{i\top}\operatorname{Diag}(\gamma^{i\to r})k^r\big)\, w^i\Big) = w^r .

这正是 wrw^r 的定义。于是

Pr=Y−krwr⊤=Diag⁡(γr)−∑i=1r−1Diag⁡(γi→r) kiwi⊤−Diag⁡(γr→r) krwr⊤,P^r = Y - k^r w^{r\top} = \operatorname{Diag}(\gamma^r) - \sum_{i=1}^{r-1} \operatorname{Diag}(\gamma^{i\to r})\, k^i w^{i\top} - \operatorname{Diag}(\gamma^{r\to r})\, k^r w^{r\top},

最后一项补上了 Diag⁡(γr→r)=I\operatorname{Diag}(\gamma^{r\to r}) = I,求和就凑齐到 i=ri = r。命题得证。

回头看 wrw^r 的含义:Diag⁡(γr)kr\operatorname{Diag}(\gamma^r)k^r 是「从块开头看过去的键 krk^r」,减去的每一项是块内更早的键 kik^i 与 krk^r 的(带衰减的)内积乘上 wiw^i,也就是更早的擦除已经覆盖掉的方向。wrw^r 告诉你第 rr 步的擦除要怎么作用到进入块时的状态 S0S^0 上。

第 4 步:HrH^r 的 WY 表示,同样的归纳。

命题 2:

Hr=∑i=1rDiag⁡(γi→r) kiui⊤,ur:=βr(vr−∑i=1r−1(ki⊤Diag⁡(γi→r)kr) ui).H^r = \sum_{i=1}^{r} \operatorname{Diag}(\gamma^{i\to r})\, k^i u^{i\top}, \qquad u^r := \beta^r \Big(v^r - \sum_{i=1}^{r-1} \big(k^{i\top}\operatorname{Diag}(\gamma^{i\to r}) k^r\big)\, u^i\Big).

起点 r=1r = 1。 H1=β1k1v1⊤=k1u1⊤H^1 = \beta^1 k^1 v^{1\top} = k^1 u^{1\top},u1=β1v1u^1 = \beta^1 v^1,Diag⁡(γ1→1)=I\operatorname{Diag}(\gamma^{1\to 1}) = I。成立。

归纳步。 Hr=TrHr−1+βrkrvr⊤H^r = T^r H^{r-1} + \beta^r k^r v^{r\top}。和第 3 步一样先乘 Diag⁡(αr)\operatorname{Diag}(\alpha^r),衰减累乘一步:

Diag⁡(αr)Hr−1=∑i=1r−1Diag⁡(γi→r) kiui⊤=:Z.\operatorname{Diag}(\alpha^r) H^{r-1} = \sum_{i=1}^{r-1} \operatorname{Diag}(\gamma^{i\to r})\, k^i u^{i\top} =: Z .

再乘 (I−βrkrkr⊤)(I - \beta^r k^r k^{r\top}),并把本步写入加上:

Hr=Z−βrkr(kr⊤Z)+βrkrvr⊤=Z+kr βr(vr⊤−kr⊤Z)⏟行向量.H^r = Z - \beta^r k^r \big(k^{r\top} Z\big) + \beta^r k^r v^{r\top} = Z + k^r\,\underbrace{\beta^r\big(v^{r\top} - k^{r\top} Z\big)}_{\text{行向量}} .

算 kr⊤Z=∑i<r(kr⊤Diag⁡(γi→r)ki) ui⊤k^{r\top}Z = \sum_{i<r} \big(k^{r\top}\operatorname{Diag}(\gamma^{i\to r})k^i\big)\,u^{i\top},标量系数和第 3 步一样对称。把行向量转置:

βr(vr⊤−kr⊤Z) ⁣⊤=βr(vr−∑i=1r−1(ki⊤Diag⁡(γi→r)kr) ui)=ur.\beta^r\big(v^{r\top} - k^{r\top}Z\big)^{\!\top} = \beta^r\Big(v^r - \sum_{i=1}^{r-1}\big(k^{i\top}\operatorname{Diag}(\gamma^{i\to r})k^r\big)\,u^i\Big) = u^r .

所以 Hr=Z+krur⊤H^r = Z + k^r u^{r\top},补上 Diag⁡(γr→r)=I\operatorname{Diag}(\gamma^{r\to r}) = I 后求和凑齐到 i=ri = r。得证。

uru^r 是伪值:vrv^r 减掉块内更早的键已经解释掉的部分。对比公式 (1) 里逐 token 的擦除 (I−βkk⊤)(I - \beta k k^\top) 作用在整个状态上,这里的减法只针对块内前 r−1r-1 个写入,剩下针对 S0S^0 的那部分由 wrw^r 负责。

到这一步,PrP^r 和 HrH^r 都不再是连乘了,但 wrw^r 和 uru^r 的定义还是一个个算的三角递推。第 5 步把它们变成矩阵运算。

第 5 步:UT 变换,把两个三角递推写成一次 C×CC\times C 求逆。

两个递推里出现的标量系数是同一个:

ari:=ki⊤Diag⁡(γi→r)kr=∑ckci γcrγci kcr=∑c(γcrkcr)(kciγci),i<r.a_{ri} := k^{i\top}\operatorname{Diag}(\gamma^{i\to r})k^r = \sum_{c} k^i_c\, \frac{\gamma^r_c}{\gamma^i_c}\, k^r_c = \sum_{c} \big(\gamma^r_c k^r_c\big)\Big(\frac{k^i_c}{\gamma^i_c}\Big), \qquad i < r .

第二个等号用了 γi→r=γr/γi\gamma^{i\to r} = \gamma^r/\gamma^i。最后的形式是两个向量的内积:γr⊙kr\gamma^r \odot k^r 和 ki/γik^i / \gamma^i。把所有 rr 的 γr⊙kr\gamma^r \odot k^r 按行堆起来就是 Γ⊙K\Gamma \odot \mathbf{K}(C×dkC \times d_k),所有 ki/γik^i/\gamma^i 堆起来是 K/Γ\mathbf{K}/\Gamma,于是

ari=[(Γ⊙K)(K/Γ) ⁣⊤]ri,a_{ri} = \big[(\Gamma \odot \mathbf{K})(\mathbf{K}/\Gamma)^{\!\top}\big]_{ri},

这是一个 C×CC \times C 的矩阵乘。我们只用到 i<ri < r 的元素,即严格下三角部分,记 L:=StrictTril⁡[(Γ⊙K)(K/Γ) ⁣⊤]\mathbf{L} := \operatorname{StrictTril}\big[(\Gamma \odot \mathbf{K})(\mathbf{K}/\Gamma)^{\!\top}\big]。

现在把 wrw^r 的递推按行堆起来。令 W∈RC×dk\mathbf{W} \in \mathbb{R}^{C \times d_k} 的第 rr 行是 wr⊤w^{r\top},Diag⁡(β)\operatorname{Diag}(\beta) 是 C×CC \times C 的对角矩阵,对角线是 β1,…,βC\beta^1, \dots, \beta^C。递推 wr=βr(γr⊙kr−∑i<rariwi)w^r = \beta^r(\gamma^r\odot k^r - \sum_{i<r} a_{ri} w^i) 的第 rr 行写成矩阵是

W=Diag⁡(β)(Γ⊙K−LW),\mathbf{W} = \operatorname{Diag}(\beta)\big(\Gamma \odot \mathbf{K} - \mathbf{L}\mathbf{W}\big),

检查一下:(LW)(\mathbf{L}\mathbf{W}) 的第 rr 行是 ∑iLri wi⊤=∑i<rariwi⊤\sum_i \mathbf{L}_{ri}\, w^{i\top} = \sum_{i<r} a_{ri} w^{i\top},正是递推里的求和。把 W\mathbf{W} 移到一边:

(I+Diag⁡(β)L)W=Diag⁡(β)(Γ⊙K)⟹W=(I+Diag⁡(β)L)−1Diag⁡(β)⏟M (Γ⊙K).\big(I + \operatorname{Diag}(\beta)\mathbf{L}\big)\mathbf{W} = \operatorname{Diag}(\beta)(\Gamma \odot \mathbf{K}) \quad\Longrightarrow\quad \mathbf{W} = \underbrace{\big(I + \operatorname{Diag}(\beta)\mathbf{L}\big)^{-1}\operatorname{Diag}(\beta)}_{\mathbf{M}}\,(\Gamma \odot \mathbf{K}).

Diag⁡(β)L\operatorname{Diag}(\beta)\mathbf{L} 仍是严格下三角(对角矩阵左乘只缩放每一行),所以 I+Diag⁡(β)LI + \operatorname{Diag}(\beta)\mathbf{L} 是对角线全为 1 的下三角矩阵,行列式为 1,一定可逆。这就是公式 (6) 的 M\mathbf{M}(论文把 Diag⁡(β)\operatorname{Diag}(\beta) 写进了 StrictTril 里面,是同一个矩阵)。

uru^r 的递推形状完全一样,只是 γr⊙kr\gamma^r \odot k^r 换成 vrv^r:U=Diag⁡(β)(V−LU)\mathbf{U} = \operatorname{Diag}(\beta)(\mathbf{V} - \mathbf{L}\mathbf{U}),解出 U=MV\mathbf{U} = \mathbf{M}\mathbf{V}。公式 (7) 得证。一次求逆,两个递推一起解掉。

第 6 步:代回去,得到公式 (8) 和 (9)。

状态更新。 块末的状态是 SC=PCS0+HCS^C = P^C S^0 + H^C。把两个 WY 表示代入:

SC=Diag⁡(γC)S0−∑i=1CDiag⁡(γi→C)ki(wi⊤S0)+∑i=1CDiag⁡(γi→C)kiui⊤=Diag⁡(γC)S0+∑i=1C(γi→C⊙ki)(ui−S0⊤wi) ⁣⊤.S^C = \operatorname{Diag}(\gamma^C) S^0 - \sum_{i=1}^{C}\operatorname{Diag}(\gamma^{i\to C})k^i \big(w^{i\top}S^0\big) + \sum_{i=1}^{C}\operatorname{Diag}(\gamma^{i\to C})k^i u^{i\top} = \operatorname{Diag}(\gamma^C) S^0 + \sum_{i=1}^{C}\big(\gamma^{i\to C}\odot k^i\big)\big(u^i - S^{0\top}w^i\big)^{\!\top}.

第二个等号把两个求和合并,并用 Diag⁡(γ)k=γ⊙k\operatorname{Diag}(\gamma)k = \gamma \odot k。求和是 CC 个秩 1 项相加,写成矩阵乘:左边把 γi→C⊙ki\gamma^{i\to C}\odot k^i 按行堆成 Γi→C⊙K\Gamma^{i\to C}\odot\mathbf{K}(C×dkC \times d_k)再转置,右边把 (ui−S0⊤wi)⊤(u^i - S^{0\top}w^i)^\top 按行堆起来就是 U−WS0\mathbf{U} - \mathbf{W}S^0(C×dvC \times d_v,因为 (S0⊤wi)⊤=wi⊤S0(S^{0\top}w^i)^\top = w^{i\top}S^0 是 WS0\mathbf{W}S^0 的第 ii 行)。于是

S[t+1]=SC=Diag⁡(γC) S[t]+(Γi→C⊙K) ⁣⊤(U−WS[t]),\mathbf{S}_{[t+1]} = S^C = \operatorname{Diag}(\gamma^C)\,\mathbf{S}_{[t]} + \big(\Gamma^{i\to C}\odot\mathbf{K}\big)^{\!\top}\big(\mathbf{U} - \mathbf{W}\mathbf{S}_{[t]}\big),

这就是公式 (8)。

输出。 位置 rr 的输出是 or=Sr⊤qro^r = S^{r\top}q^r,转置成行向量 or⊤=qr⊤Sr=qr⊤PrS0+qr⊤Hro^{r\top} = q^{r\top}S^r = q^{r\top}P^r S^0 + q^{r\top}H^r。分别代入 WY 表示:

qr⊤PrS0=(γr⊙qr) ⁣⊤S0−∑i=1r(qr⊤Diag⁡(γi→r)ki)⏟bri wi⊤S0,qr⊤Hr=∑i=1rbri ui⊤.q^{r\top}P^r S^0 = \big(\gamma^r\odot q^r\big)^{\!\top}S^0 - \sum_{i=1}^{r}\underbrace{\big(q^{r\top}\operatorname{Diag}(\gamma^{i\to r})k^i\big)}_{b_{ri}}\,w^{i\top}S^0, \qquad q^{r\top}H^r = \sum_{i=1}^{r} b_{ri}\,u^{i\top}.

合并:

or⊤=(γr⊙qr) ⁣⊤S0+∑i=1rbri(ui−S0⊤wi) ⁣⊤.o^{r\top} = \big(\gamma^r\odot q^r\big)^{\!\top}S^0 + \sum_{i=1}^{r} b_{ri}\big(u^i - S^{0\top}w^i\big)^{\!\top}.

标量 brib_{ri} 和第 5 步的 aria_{ri} 是同一种东西,只是把 krk^r 换成 qrq^r:bri=∑c(γcrqcr)(kci/γci)=[(Γ⊙Q)(K/Γ)⊤]rib_{ri} = \sum_c (\gamma^r_c q^r_c)(k^i_c/\gamma^i_c) = [(\Gamma\odot\mathbf{Q})(\mathbf{K}/\Gamma)^\top]_{ri}。这次求和到 i=ri = r,包含对角线(brr=qr⊤krb_{rr} = q^{r\top}k^r,因为 γr→r=1\gamma^{r\to r} = \mathbf{1}),所以取的是 Tril⁡\operatorname{Tril} 而不是 StrictTril⁡\operatorname{StrictTril}。把 CC 行堆起来:

O=(Γ⊙Q) S[t]+Tril⁡ ⁣[(Γ⊙Q)(K/Γ) ⁣⊤](U−WS[t]),\mathbf{O} = (\Gamma\odot\mathbf{Q})\,\mathbf{S}_{[t]} + \operatorname{Tril}\!\big[(\Gamma\odot\mathbf{Q})(\mathbf{K}/\Gamma)^{\!\top}\big]\big(\mathbf{U} - \mathbf{W}\mathbf{S}_{[t]}\big),

这就是公式 (9)。

回顾整条链。 公式 (1) 展开成 PP、HH(第 1 步,纯改写);只看衰减时连乘变成 γ\gamma(第 2 步);加回擦除后,归纳法证明连乘等于「对角减秩 1 之和」(第 3、4 步,WY 表示),代价是引入两个三角递推定义的 ww、uu;把三角递推按行堆起来就是一个单位下三角线性方程组,一次求逆解掉(第 5 步,UT 变换);最后代回 SCS^C 和 oro^r,秩 1 之和恰好能写成矩阵乘(第 6 步)。整个过程没有任何近似。

内核里怎么求逆

公式 (6) 的求逆,教科书做法是前向替换,代价 O(C2)O(C^2),但一行要等上一行,在 GPU 上不好并行。FlashKDA 的做法更适合硬件。要逆的是 I+LI + L,LL 严格下三角。FlashKDA 的块是 16 个 token,LL 是 16×16,于是 L16=0L^{16} = 0,Neumann 级数是有限的:(I+L)−1=∑i=015(−L)i(I + L)^{-1} = \sum_{i=0}^{15} (-L)^i,精确,没有截断。再把这个和写成乘积 (I−L)(I+L2)(I+L4)(I+L8)(I - L)(I + L^2)(I + L^4)(I + L^8),三轮平方就得到全部 16 项。每一轮都是 16×16 的矩阵乘,前向替换那种一行等一行的顺序依赖没有了。

数值问题: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⁡(gi−gj)\gamma^i/\gamma^j = \exp(g_i - g_j),对 j≤ij \le i 永远不超过 1。问题是矩阵乘要求把 exp⁡(gi−gj)\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⁡(gi−gj)=exp⁡(gi−gs)⋅exp⁡(gs−gj)\exp(g_i - g_j) = \exp(g_i - g_s)\cdot\exp(g_s - g_j),i≥s>ji \ge s > j,两个因子都不超过 1,可以放心拆成两边、用 BF16 矩阵乘。
  • 对角 tile(ii、jj 同在一个 tile):同样拆法里 exp⁡(gs−gj)\exp(g_s - g_j) 是 ≥1\ge 1 的,16 步累计能有多大取决于每步的 log-decay 有没有下界。Kimi Linear 的映射没有下界,所以对角 tile 只能逐位置对、用 FP32 算 exp⁡(gi−gj)\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=gmin⁡ Sigmoid⁡ ⁣(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:Kimi K3:( 最小 )

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

后果一:数值范围有界。 每步 α>e−5≈6.7×10−3\alpha > e^{-5} \approx 6.7\times10^{-3},16 步累计的 log-decay 在 (−80,0)(-80, 0) 里,exp⁡(gs−gj)<e80≈5.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 里它只能改横向的斜率,衰减的上限被 gmin⁡g_{\min} 钉死。

为什么是 16,为什么拆两个内核

flash-linear-attention 的默认块大小是 64,Kimi Linear 的 FLOPs 分析和伪代码也按 C=64C = 64 写,16 只是它内核里的二级 tile。FlashKDA 直接把块定成 16,仓库的设计文档给了三个理由,每条接一个后果。

  • BF16 范围。 下界衰减保证 16 步累乘的 exp⁡(gs−gj)<e80\exp(g_s - g_j) < e^{80} 在 BF16 之内,这是上面的后果一。64 步就是 e320e^{320},溢出,块内要再做一次缩放。gmin⁡=−5g_{\min} = -5 和 C=16C = 16 是配套选的。
  • 求逆。 16×16 的逆用 Neumann 三轮就精确。64×64 要六轮,中间矩阵大 16 倍。
  • 指令形状。 块内所有运算都落在 SM80 就有的 MMA 指令形状上,内核的数学部分不依赖 Hopper 专用指令。完整实现仍要 SM90 或更新,因为装载和存储路径用了 Hopper 的 TMA 和 STSM。

块内计算按 token 并行,块间递推按序列串行、按头并行。融合在一个内核里,并行的部分要等串行的部分。FlashKDA 拆成两次 launch。Kernel 1 每个 16-token 块每个头一个 CTA,算块内的矩阵写到工作区,满占用。Kernel 2 每条序列每个头一个 CTA,四个 MMA warp 加一个装载 warp 和一个存储 warp,走块间递推,128×128 的状态以 BF16 常驻共享内存,更新用 FP32 FMA 再转回 BF16。仓库文档说拆分比融合原型端到端快至少 15%。两个阶段有各自的 launch 形状,可以分开调度和调优。

下界作为一个 host 端常数进内核:

cpp
// FlashKDA/csrc/flash_kda.cpp
float gate_scale = float(lower_bound * 1.4426950408889634);

乘的是 log⁡2e\log_2 e,因为内核用以 2 为底的指数。

开源边界。 公开的是 forward-only 的 prefill,头维固定 128,在 torch.inference_mode() 下且输入为 BF16、safe_gate 打开时自动分发为 FLA 的一个后端,所以有下界的门是快路径的前提。backward、decode 内核和 SM 级的上下文并行规划器没有公开。第 7 篇讲三种执行形态时会再回到这里。

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

上一篇提到 KDA 的转移是 D−atbt⊤D - a_t b_t^\top,且 at=βtkta_t = \beta_t k_t、bt=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_kernel、use_beta_sigmoid_in_kernel、use_gate_in_kernel),Python 侧只算到 logit。

术语坑

  • γi→j\gamma^{i\to j} 的端点。 Kimi Linear 公式 (3) 把 γi→j\gamma^{i\to j} 写成从 k=ik = i 起连乘到 jj,但附录 B 的两个命题和公式 (8) 只有在「从 k=i+1k = i+1 起乘」时才成立(拿 r=1r = 1 验算:P1=T1P^1 = T^1 的秩 1 项前面没有 Diag⁡(α1)\operatorname{Diag}(\alpha^1))。本文统一按后者写,γi→j\gamma^{i\to j} 是「从位置 ii 走到位置 jj 经历的衰减」,γr→r=1\gamma^{r\to r} = \mathbf{1}。两种写法在 γr\gamma^r 上一致。
  • 两个 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 分析和伪代码,那是 FLA 内核里 64 的块套 16 的二级 tile;FlashKDA 的块就是 16,K3 论文第一次把它写进正文。这一篇的公式对两种块大小都成立,只是 CC 不同。
  • HF 权重里 A_log 的初始化是 log⁡(Uniform(1,16))\log(\mathrm{Uniform}(1,16)),那是 Kimi Linear 的初始化写法留在代码里;K3 论文说 AhA_h 初始化为 0。推理时无所谓,复现训练时以论文为准。

下一篇

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

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

评论