对应论文 §2.1.1 的后半:公式 (3) 到 (5) 和 Figure 3。上一篇的递推形式一个 token 走一步,全是向量和矩阵的逐元素运算,训练时用不上 Tensor Core。这一篇讲怎么把一个块(C=64 个 token)的计算改写成几次矩阵乘,然后你会看到 K3 为什么要动衰减函数。
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
还在 KDA 里。这一篇讲的是同一个模块在训练和 prefill 时的算法形式。
目标:块间递推,块内并行
把序列切成长度 C 的块。记 X[t] 为第 t 块里的 token 堆成的矩阵(Q,K,V∈RC×d),S[t] 为进入第 t 块时的状态。想要的形式是:
- 块间:状态只在块边界更新一次,S[t]→S[t+1]。
- 块内:块里 C 个输出一起算,用矩阵乘。
先定义逐通道的累积衰减(论文公式 (3)):
γ[t]i→j:=r=i∏jα[t]r,γ[t]r:=γ[t]1→r,
Γ[t]1→C∈RC×dk 把 γ1,…,γC 按行堆起来。γr 是「从块开头到位置 r 一共衰减了多少」,逐通道。
块内展开:P 和 H
把递推在块内展开 r 步(Kimi Linear 公式 (2)):
S[t]r=Pri=1∏r(I−βikiki⊤)Diag(αi) S[t]0 + Hri=1∑r(j=i+1∏r(I−βjkjkj⊤)Diag(αj))βikivi⊤.
Pr 是进入块时的状态被怎样变换,Hr 是块内写入的贡献。两项都是一串带衰减的 Householder 矩阵的乘积,直接算还是顺序的。
WY 表示:一串 Householder 乘积等于「对角减低秩和」
Kimi Linear 附录 B 的两个命题(归纳法证):
Pr=Diag(γr)−i=1∑rDiag(γi→r)kiwi⊤,Hr=i=1∑rDiag(γi→r)kiui⊤,
其中辅助向量由两个三角递推给出:
wr=βr(Diag(γr)kr−i=1∑r−1wi(ki⊤Diag(γi→r)kr)),ur=βr(vr−i=1∑r−1ui(ki⊤Diag(γi→r)kr)).
怎么理解 u 和 w:ur 是伪值,等于 vr 减掉块内更早的键已经解释掉的部分,也就是 delta rule 的「先擦再写」被限制在块内的版本。wr 是配套的衰减键,告诉你这个修正要怎么作用到进入块时的状态上。
UT 变换:一次 C×C 三角求逆
两个递推可以对整个块一次解出(Kimi Linear 公式 (6)、(7)):
M=(I+StrictTril(Diag(β)(Γ⊙K)(ΓK)⊤))−1Diag(β),W=M(Γ⊙K),U=MV.
单位下三角矩阵的逆用前向替换做,代价 O(C2) 一次,之后全是矩阵乘。这一步把非矩阵乘的顺序工作变成了矩阵乘,Kimi Linear 说它「对硬件利用率至关重要」。
块间递推与块内输出
有了 U、W,定义伪值项 V[t]:=U[t]−W[t]S[t],K3 论文公式 (4):
A[t]=Tril[(Q[t]⊙Γ[t]1→C)(K[t]/Γ[t]1→C)⊤],O[t]=块间(Γ[t]1→C⊙Q[t])S[t]+块内A[t]V[t].
状态更新(Kimi Linear 公式 (8)):
S[t+1]=Diag(γ[t]C)S[t]+(Γ[t]i→C⊙K[t])⊤V[t].
读法:
- 块间那一行,先把整个状态按这一块的总衰减 γC 逐通道缩一下,再加上这一块的写入。写入是「每个键按它到块尾的剩余距离衰减」乘「修正后的伪值」。
- 块内那一行,A 是一个 C×C 的下三角矩阵,第 (i,j) 项是 qi 和 kj 的内积,每个通道再乘上 γi/γj,也就是从 j 到 i 之间的衰减。保留对角线,因为每个输出读的是「当前 token 更新之后」的状态。
- Kimi Linear 的 FLOPs 公式(每头,C=64):6Tdh2+3TCdh+TC2,对 T 线性;softmax 注意力是 2T2dh。
数值问题:1/Γ 会溢出
公式 (4) 里 K/Γ 把键除以累积衰减。Γ 是一串 (0,1) 里的数的乘积,倒数可以任意大,半精度下溢出。
标准解法(GLA 的做法,Kimi Linear 沿用)是进 log 空间:令 gr=logγr(对每步的 log-decay 做前缀和),那么 γi/γj=exp(gi−gj),对 j≤i 永远不超过 1。问题是矩阵乘要求把 exp(gi−gj) 拆成「只依赖行」乘「只依赖列」的两个因子,拆开就又回到了 exp(gi)⋅exp(−gj),后者会炸。
于是再把 C=64 的块切成 16-token 的二级 tile,以 tile 起点 s 为参考:
- 非对角 tile(查询 i 在后面的 tile,键 j 在前面的 tile):exp(gi−gj)=exp(gi−gs)⋅exp(gs−gj),i≥s>j,两个因子都不超过 1,可以放心拆成两边、用 BF16 矩阵乘。
- 对角 tile(i、j 同在一个 tile):同样拆法里 exp(gs−gj) 是 ≥1 的,16 步累计能有多大取决于每步的 log-decay 有没有下界。Kimi Linear 的映射没有下界,所以对角 tile 只能逐位置对、用 FP32 算 exp(gi−gj),不走 Tensor Core。论文说这是块内计算的主要瓶颈。
K3 的改动:下界衰减
从 logit z 到每步 log-decay g 的映射,Kimi Linear 沿用 GDN 和 Mamba-2:
gth=−eAhSoftplus(zth)∈(−∞,0)dk.
K3 换成一个带尺度的 sigmoid(论文公式 (5)):
gth=gminSigmoid(eAhzth)∈(gmin,0)dk,αth=exp(gth)∈(egmin,1)dk,
gmin=−5 固定,Ah 是每头一个的可学习 log 尺度,初始化为 0。config 里就是 gate_lower_bound = -5.0。
两条曲线的差别:负 softplus 在 z→+∞ 时线性往下走,没有底;scaled sigmoid 在 z→+∞ 时贴到 gmin。Ah 只改横轴的伸缩。
后果一:数值范围有界。 每步 α>e−5≈6.7×10−3,16 步累计的 log-decay 在 (−80,0) 里,exp(gs−gj)<e80≈5.5×1034,在 BF16 的动态范围(最大约 3.4×1038)之内。于是对角 tile 也能拆成两个因子做 BF16 矩阵乘,Figure 3b 说的「消除逐位置对的对角路径」就是这个意思。gmin=−5 和 tile 宽度 16 是配套选的:5×16=80 正好卡在 BF16 指数范围以内。
后果二:记忆不会靠衰减被瞬间清空。 一个通道靠衰减能做到的最快遗忘是每步剩 0.67%。真正需要精确擦除的时候,delta rule 的 (I−βkk⊤) 还在。论文没有讨论这一点对表达力的影响,只提了这种下界门在 RWKV-7、Griffin、HGRN2 里有先例。
后果三:Ah 的作用变了。 在负 softplus 里 eAh 是纵向的尺度,直接决定衰减能多狠;在 scaled sigmoid 里它只能改横向的斜率,衰减的上限被 gmin 钉死。
为什么 KDA 内核比通用 DPLR 快一倍
上一篇提到 KDA 的转移是 D−atbt⊤,且 at=βtkt、bt=kt⊙αt 都绑在 kt 上。在 chunkwise 形式里这意味着:通用 DPLR 内核需要四张二级 tile 矩阵(Aab,Aak,Aqb,Aqk),KDA 只需要两张(Aqk,Akk);输出阶段 DPLR 的三项输出和两次状态更新,在 KDA 里合成一行输出和一次状态更新。Kimi Linear Figure 2 显示在 2K 到 64K 长度上 KDA 内核约为 DPLR 内核的 2 倍速度。K3 在此基础上做的 FlashKDA 内核放到第 7 篇。
训练、prefill、decode 各用哪种形式
| 阶段 | 形式 | 代码 |
|---|
| 训练 | chunkwise | chunk_kda(..., safe_gate=True, lower_bound=-5.0) |
| prefill | chunkwise | 同上 |
| decode(每次 1 个 token) | 递推 | fused_recurrent_kda(..., lower_bound=-5.0) |
HF 代码里 mode = 'fused_recurrent' if use_cache and q_len == 1 else 'chunk'。L2 归一化、β 的 sigmoid、α 的映射都在内核里做(use_qk_l2norm_in_kernel、use_beta_sigmoid_in_kernel、use_gate_in_kernel),Python 侧只算到 logit。
术语坑
- 两个 A。 公式 (4) 的 A[t] 是 C×C 的块内注意力矩阵,公式 (5) 的 Ah 是每头的标量 log 尺度。上一篇 NoPE 那段的 Aj=Diag(αj) 又是第三个。
- 16 和 64。 块大小 C=64 来自 Kimi Linear 的 FLOPs 分析和伪代码,16-token 二级 tile 是 FLA 内核的实现细节,K3 论文第一次把它写进正文。
- HF 权重里
A_log 的初始化是 log(Uniform(1,16)),那是 Kimi Linear 的初始化写法留在代码里;K3 论文说 Ah 初始化为 0。推理时无所谓,复现训练时以论文为准。
下一篇
序列维度还剩四分之一:每四层一层的 Gated MLA。它负责 KDA 做不了的事,无损的全局内容检索。下一篇讲 MLA 的低秩压缩、K3 为什么敢让它完全不带位置编码、满秩输出门,以及 3
这个比例是怎么选出来的。