对应论文 §2.1.1 的后半:公式 (3) 到 (5) 和 Figure 3。上一篇的递推形式一个 token 走一步,全是向量和矩阵的逐元素运算,训练时用不上 Tensor Core。这一篇先给出结论:一个块(C=64 个 token)的计算可以改写成几次矩阵乘,状态只在块边界更新一次。推导折叠起来放在结论后面,可以跳过也可以逐步跟着算。最后你会看到 K3 为什么要动衰减函数。
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
还在 KDA 里。这一篇讲的是同一个模块在训练和 prefill 时的算法形式。
起点:公式 (1) 的递推,为什么训练时用不了
先把上一篇的公式 (1) 摆出来,这一篇所有推导都从它出发:
St=(I−βtktkt⊤)Diag(αt)St−1+βtktvt⊤,ot=St⊤qt.(1)
新状态② 擦除① 衰减旧状态③ 写入=··+从右往左读:先把旧状态每一行按各自衰减,再把落在方向上的内容擦掉的比例,最后写入倍的新关联。读出:。行 = 键通道(,各自的衰减率),列 = 值通道()。
一个 token 走一步:先把旧状态每一行按 αt 衰减,再沿 kt 方向擦掉 βt 的比例,最后写入 βtktvt⊤。每步代价固定,对序列长度已经是线性的了,为什么还要改?答案在硬件上。
Tensor Core 只接受矩阵乘矩阵。H100 上 BF16 矩阵乘的峰值约 990 TFLOPS,通用 CUDA core 上的 FP32 约 67 TFLOPS,差十几倍。公式 (1) 里的两个操作,外积 ktvt⊤ 和矩阵乘向量 S⊤qt,都只有一个向量维,填不满一个 tile,只能走慢的那一边。更糟的是每个 token 都要把整个 128×128 的状态读一遍再写一遍,而且第 t 步必须等第 t−1 步做完,1M 的序列就是一百万步串行。
把 C 个 token 摞在一起,向量就变成 C×128 的矩阵,乘法有了第二个矩阵维度,可以上 Tensor Core。状态每 C 个 token 读写一次,串行链从一百万步变成 106/C 步。代价是块内 C 个 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 的代价已经是常数,但两个操作都只有一个向量维;把 C 个 token 摞成矩阵,乘法才有第二个矩阵维度。数字是 H100 的峰值。
终点:块间一步递推,块内一次矩阵乘
把序列切成长度 C 的块(Kimi Linear 按 C=64 分析)。记 Q[t],K[t],V[t]∈RC×d 为第 t 块里的 token 堆成的矩阵,一行一个 token;S[t] 为进入第 t 块时的状态。还要一个记号来表示「从块内位置 i 走到位置 j 一共衰减了多少」,逐通道:
γi→j:=r=i+1∏jαr∈Rdk,γr:=γ0→r=α1⊙⋯⊙αr,
Γ1→C∈RC×dk 把 γ1,…,γC 按行堆起来。γr 是「从块开头到位置 r」的累积衰减,γi→j=γj/γ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)
配套的是块内的输出(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)
两条公式里的 U、W 由一次 C×C 的三角求逆给出(Kimi Linear 公式 (6)、(7)),这一步叫 UT 变换:
M=(I+StrictTril(Diag(β)(Γ⊙K)(ΓK)⊤))−1Diag(β),W=M(Γ⊙K),U=MV.(6, 7)
块间通道:状态逐块递推(每块一次)整块的逐通道衰减+本块的写入,一次矩阵乘块内通道:本块所有输出并行算(一次掩码矩阵乘)读进入本块时的状态+,块内 token 之间的因果交互也喂给块内通道的下三角,,16-token 二级 tile查询键非对角 tile:,两个因子都,直接 BF16 矩阵乘对角 tile:可能Kimi Linear:逐位置对、FP32,是瓶颈K3:16 步累计 log-decay 大于,在 BF16 范围内,也走矩阵乘
读法:
- 块间那一行,先把整个状态按这一块的总衰减 γC 逐通道缩一下,再加上这一块的写入。写入是「每个键按它到块尾的剩余距离衰减」乘「修正后的伪值」V。伪值的意思是:U 是 V 减掉块内更早的键已经解释掉的部分,也就是 delta rule 的「先擦再写」被限制在块内的版本;WS[t] 是这些擦除作用到进入块时的状态上的那一部分。
- 块内那一行,A 是一个 C×C 的下三角矩阵,第 (i,j) 项是 qi 和 kj 的内积,每个通道再乘上 γi/γj,也就是从 j 到 i 之间的衰减。保留对角线,因为每个输出读的是「当前 token 更新之后」的状态。
- 哪些是矩阵乘:A 是 (C×dk)(dk×C),AV 是 (C×C)(C×dv),块间写入是 (dk×C)(C×dv),WS 是 (C×dk)(dk×dv)。除了求逆,全是 Tensor Core 能吃的形状。
- Kimi Linear 的 FLOPs 公式(每头,C=64):6Tdh2+3TCdh+TC2,对 T 线性;softmax 注意力是 2T2dh。
公式 (1) 是怎么变成公式 (8) 和 (9) 的,完整推导折叠在下面。只想知道结论的话可以跳过,读后面 K3 的改动不需要它;想跟着算一遍的话,只用到矩阵乘法的结合律、对角矩阵的性质和数学归纳法。
从公式 (1) 到公式 (8)、(9):完整推导
记号。 只看一个块,省掉下标 [t]。块内位置 r=1,…,C,kr,vr,qr,αr 是第 r 个 token 的向量,βr 是标量。S0 是进入块时的状态(也就是 S[t]),Sr 是处理完第 r 个 token 之后的状态,所以 SC=S[t+1]。把公式 (1) 每一步的转移矩阵记成
Tr:=(I−βrkrkr⊤)Diag(αr)∈Rdk×dk,Sr=TrSr−1+βrkrvr⊤.第 1 步:把递推在块内展开,得到 P 和 H。
一步步代入。r=1 时就是定义:
S1=T1S0+β1k1v1⊤.r=2 时把 S1 代进去,用矩阵乘法的分配律拆开:
S2=T2S1+β2k2v2⊤=T2T1S0+T2β1k1v1⊤+β2k2v2⊤.规律出来了:每个 S0 前面是所有转移矩阵的乘积,每个第 i 步的写入 βikivi⊤ 前面是它之后的转移矩阵 Ti+1,…,Tr 的乘积(后发生的在左边,因为每次都是左乘)。一般地
Sr=Pr(i=1∏rTi) S0 + Hri=1∑r(j=i+1∏rTj)βikivi⊤,这里的 ∏ 约定按 TrTr−1⋯T1 的顺序排,空乘积是 I。这就是 Kimi Linear 的公式 (2)。Pr 是「进入块时的状态被怎样变换」,Hr 是「块内写入的贡献」,两者满足和 S 一样的递推:
Pr=TrPr−1,P0=I;Hr=TrHr−1+βrkrvr⊤,H0=0.到这里只是换了个写法,Pr 和 Hr 里的连乘还是顺序的,没有省任何计算。
第 2 步:先看只有衰减的情形,看 γ 是从哪来的。
假设所有 βr=0,即没有擦除,那么 Tr=Diag(αr)。对角矩阵相乘等于对角线逐元素相乘:
Diag(αr)⋯Diag(αi+1)=Diag(αi+1⊙⋯⊙αr)=Diag(γi→r).于是 Pr=Diag(γr),Hr=∑iDiag(γi→r)βikivi⊤。γi→r 就是「从位置 i 写进去的东西,到位置 r 时还剩多少」,逐通道。i=r 时是空乘积,γr→r=1,刚写进去的还没来得及衰减。
后面还会反复用到两条对角矩阵的性质:
- Diag(αr)Diag(γi→r−1)=Diag(γi→r),衰减累乘一步。
- x⊤Diag(γ)y=∑cγcxcyc,是一个标量,且等于 y⊤Diag(γ)x。
第 3 步:加回擦除,Pr 的 WY 表示。
现在 βr=0。要证的命题(Kimi Linear 附录 B,命题 1)是:一串带衰减的 Householder 矩阵的乘积,等于「一个对角矩阵减去 r 个秩 1 项」:
Pr=Diag(γr)−i=1∑rDiag(γi→r)kiwi⊤,wr:=βr(Diag(γr)kr−i=1∑r−1(ki⊤Diag(γi→r)kr)wi).这种「对角减低秩和」的写法叫 WY 表示,来自数值线性代数里把一串 Householder 反射合并成 I−YWY⊤ 的技巧。wi∈Rdk 是要构造的辅助向量,它的定义看起来循环(wr 用到 w1,…,wr−1),但这是一个三角递推,从 w1 开始一个个算就行。
用归纳法。
起点 r=1。 直接乘开 T1:
P1=T1=(I−β1k1k1⊤)Diag(α1)=Diag(α1)−β1k1(k1⊤Diag(α1)).γ1=α1,而 k1⊤Diag(α1) 是行向量,它的转置是 Diag(α1)k1(对角矩阵是对称的)。所以第二项是 k1w1⊤,其中 w1=β1Diag(γ1)k1,和定义里 r=1 的情形(求和为空)一致。前面的系数 Diag(γ1→1)=I。命题在 r=1 成立。
归纳步。 假设 Pr−1 满足命题,计算 Pr=TrPr−1。Tr 有两个因子,先乘右边的 Diag(αr):
Y:=Diag(αr)Pr−1=Diag(αr)Diag(γr−1)−i=1∑r−1Diag(αr)Diag(γi→r−1)kiwi⊤=Diag(γr)−i=1∑r−1Diag(γi→r)kiwi⊤.用的是第 2 步的第一条性质:衰减多累乘了一步。注意 Y 已经具有命题要的形状,只是求和还差 i=r 这一项。再乘左边的 (I−βrkrkr⊤):
Pr=Y−βrkr(kr⊤Y).关键在于 kr⊤Y 是一个行向量(1×dk),所以 βrkr(kr⊤Y) 是一个秩 1 矩阵,正好是命题里缺的那一项。把它算出来:
kr⊤Y=kr⊤Diag(γr)−i=1∑r−1标量(kr⊤Diag(γi→r)ki)wi⊤.第二项里 kr⊤Diag(γi→r)ki 是标量,可以挪到前面;按第 2 步第二条性质它等于 ki⊤Diag(γi→r)kr。把整个行向量转置成列向量,乘上 βr:
βr(kr⊤Y)⊤=βr(Diag(γr)kr−i=1∑r−1(ki⊤Diag(γi→r)kr)wi)=wr.这正是 wr 的定义。于是
Pr=Y−krwr⊤=Diag(γr)−i=1∑r−1Diag(γi→r)kiwi⊤−Diag(γr→r)krwr⊤,最后一项补上了 Diag(γr→r)=I,求和就凑齐到 i=r。命题得证。
回头看 wr 的含义:Diag(γr)kr 是「从块开头看过去的键 kr」,减去的每一项是块内更早的键 ki 与 kr 的(带衰减的)内积乘上 wi,也就是更早的擦除已经覆盖掉的方向。wr 告诉你第 r 步的擦除要怎么作用到进入块时的状态 S0 上。
第 4 步:Hr 的 WY 表示,同样的归纳。
命题 2:
Hr=i=1∑rDiag(γi→r)kiui⊤,ur:=βr(vr−i=1∑r−1(ki⊤Diag(γi→r)kr)ui).起点 r=1。 H1=β1k1v1⊤=k1u1⊤,u1=β1v1,Diag(γ1→1)=I。成立。
归纳步。 Hr=TrHr−1+βrkrvr⊤。和第 3 步一样先乘 Diag(αr),衰减累乘一步:
Diag(αr)Hr−1=i=1∑r−1Diag(γi→r)kiui⊤=:Z.再乘 (I−βrkrkr⊤),并把本步写入加上:
Hr=Z−βrkr(kr⊤Z)+βrkrvr⊤=Z+kr行向量βr(vr⊤−kr⊤Z).算 kr⊤Z=∑i<r(kr⊤Diag(γi→r)ki)ui⊤,标量系数和第 3 步一样对称。把行向量转置:
βr(vr⊤−kr⊤Z)⊤=βr(vr−i=1∑r−1(ki⊤Diag(γi→r)kr)ui)=ur.所以 Hr=Z+krur⊤,补上 Diag(γr→r)=I 后求和凑齐到 i=r。得证。
ur 是伪值:vr 减掉块内更早的键已经解释掉的部分。对比公式 (1) 里逐 token 的擦除 (I−βkk⊤) 作用在整个状态上,这里的减法只针对块内前 r−1 个写入,剩下针对 S0 的那部分由 wr 负责。
到这一步,Pr 和 Hr 都不再是连乘了,但 wr 和 ur 的定义还是一个个算的三角递推。第 5 步把它们变成矩阵运算。
第 5 步:UT 变换,把两个三角递推写成一次 C×C 求逆。
两个递推里出现的标量系数是同一个:
ari:=ki⊤Diag(γi→r)kr=c∑kciγciγcrkcr=c∑(γcrkcr)(γcikci),i<r.第二个等号用了 γi→r=γr/γi。最后的形式是两个向量的内积:γr⊙kr 和 ki/γi。把所有 r 的 γr⊙kr 按行堆起来就是 Γ⊙K(C×dk),所有 ki/γi 堆起来是 K/Γ,于是
ari=[(Γ⊙K)(K/Γ)⊤]ri,这是一个 C×C 的矩阵乘。我们只用到 i<r 的元素,即严格下三角部分,记 L:=StrictTril[(Γ⊙K)(K/Γ)⊤]。
现在把 wr 的递推按行堆起来。令 W∈RC×dk 的第 r 行是 wr⊤,Diag(β) 是 C×C 的对角矩阵,对角线是 β1,…,βC。递推 wr=βr(γr⊙kr−∑i<rariwi) 的第 r 行写成矩阵是
W=Diag(β)(Γ⊙K−LW),检查一下:(LW) 的第 r 行是 ∑iLriwi⊤=∑i<rariwi⊤,正是递推里的求和。把 W 移到一边:
(I+Diag(β)L)W=Diag(β)(Γ⊙K)⟹W=M(I+Diag(β)L)−1Diag(β)(Γ⊙K).Diag(β)L 仍是严格下三角(对角矩阵左乘只缩放每一行),所以 I+Diag(β)L 是对角线全为 1 的下三角矩阵,行列式为 1,一定可逆。这就是公式 (6) 的 M(论文把 Diag(β) 写进了 StrictTril 里面,是同一个矩阵)。
ur 的递推形状完全一样,只是 γr⊙kr 换成 vr:U=Diag(β)(V−LU),解出 U=MV。公式 (7) 得证。一次求逆,两个递推一起解掉。
第 6 步:代回去,得到公式 (8) 和 (9)。
状态更新。 块末的状态是 SC=PCS0+HC。把两个 WY 表示代入:
SC=Diag(γC)S0−i=1∑CDiag(γi→C)ki(wi⊤S0)+i=1∑CDiag(γi→C)kiui⊤=Diag(γC)S0+i=1∑C(γi→C⊙ki)(ui−S0⊤wi)⊤.第二个等号把两个求和合并,并用 Diag(γ)k=γ⊙k。求和是 C 个秩 1 项相加,写成矩阵乘:左边把 γi→C⊙ki 按行堆成 Γi→C⊙K(C×dk)再转置,右边把 (ui−S0⊤wi)⊤ 按行堆起来就是 U−WS0(C×dv,因为 (S0⊤wi)⊤=wi⊤S0 是 WS0 的第 i 行)。于是
S[t+1]=SC=Diag(γC)S[t]+(Γi→C⊙K)⊤(U−WS[t]),这就是公式 (8)。
输出。 位置 r 的输出是 or=Sr⊤qr,转置成行向量 or⊤=qr⊤Sr=qr⊤PrS0+qr⊤Hr。分别代入 WY 表示:
qr⊤PrS0=(γr⊙qr)⊤S0−i=1∑rbri(qr⊤Diag(γi→r)ki)wi⊤S0,qr⊤Hr=i=1∑rbriui⊤.合并:
or⊤=(γr⊙qr)⊤S0+i=1∑rbri(ui−S0⊤wi)⊤.标量 bri 和第 5 步的 ari 是同一种东西,只是把 kr 换成 qr:bri=∑c(γcrqcr)(kci/γci)=[(Γ⊙Q)(K/Γ)⊤]ri。这次求和到 i=r,包含对角线(brr=qr⊤kr,因为 γr→r=1),所以取的是 Tril 而不是 StrictTril。把 C 行堆起来:
O=(Γ⊙Q)S[t]+Tril[(Γ⊙Q)(K/Γ)⊤](U−WS[t]),这就是公式 (9)。
回顾整条链。 公式 (1) 展开成 P、H(第 1 步,纯改写);只看衰减时连乘变成 γ(第 2 步);加回擦除后,归纳法证明连乘等于「对角减秩 1 之和」(第 3、4 步,WY 表示),代价是引入两个三角递推定义的 w、u;把三角递推按行堆起来就是一个单位下三角线性方程组,一次求逆解掉(第 5 步,UT 变换);最后代回 SC 和 or,秩 1 之和恰好能写成矩阵乘(第 6 步)。整个过程没有任何近似。
内核里怎么求逆
公式 (6) 的求逆,教科书做法是前向替换,代价 O(C2),但一行要等上一行,在 GPU 上不好并行。FlashKDA 的做法更适合硬件。要逆的是 I+L,L 严格下三角。FlashKDA 的块是 16 个 token,L 是 16×16,于是 L16=0,Neumann 级数是有限的:(I+L)−1=∑i=015(−L)i,精确,没有截断。再把这个和写成乘积 (I−L)(I+L2)(I+L4)(I+L8),三轮平方就得到全部 16 项。每一轮都是 16×16 的矩阵乘,前向替换那种一行等一行的顺序依赖没有了。
数值问题: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 钉死。
为什么是 16,为什么拆两个内核
flash-linear-attention 的默认块大小是 64,Kimi Linear 的 FLOPs 分析和伪代码也按 C=64 写,16 只是它内核里的二级 tile。FlashKDA 直接把块定成 16,仓库的设计文档给了三个理由,每条接一个后果。
- BF16 范围。 下界衰减保证 16 步累乘的 exp(gs−gj)<e80 在 BF16 之内,这是上面的后果一。64 步就是 e320,溢出,块内要再做一次缩放。gmin=−5 和 C=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 端常数进内核:
// FlashKDA/csrc/flash_kda.cpp
float gate_scale = float(lower_bound * 1.4426950408889634);
乘的是 log2e,因为内核用以 2 为底的指数。
开源边界。 公开的是 forward-only 的 prefill,头维固定 128,在 torch.inference_mode() 下且输入为 BF16、safe_gate 打开时自动分发为 FLA 的一个后端,所以有下界的门是快路径的前提。backward、decode 内核和 SM 级的上下文并行规划器没有公开。第 7 篇讲三种执行形态时会再回到这里。
为什么 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。
术语坑
- γi→j 的端点。 Kimi Linear 公式 (3) 把 γi→j 写成从 k=i 起连乘到 j,但附录 B 的两个命题和公式 (8) 只有在「从 k=i+1 起乘」时才成立(拿 r=1 验算:P1=T1 的秩 1 项前面没有 Diag(α1))。本文统一按后者写,γi→j 是「从位置 i 走到位置 j 经历的衰减」,γr→r=1。两种写法在 γr 上一致。
- 两个 A。 公式 (4) 的 A[t] 是 C×C 的块内注意力矩阵,公式 (5) 的 Ah 是每头的标量 log 尺度。上一篇 NoPE 那段的 Aj=Diag(αj) 又是第三个。
- 16 和 64。 公式里的 C=64 来自 Kimi Linear 的 FLOPs 分析和伪代码,那是 FLA 内核里 64 的块套 16 的二级 tile;FlashKDA 的块就是 16,K3 论文第一次把它写进正文。这一篇的公式对两种块大小都成立,只是 C 不同。
- HF 权重里
A_log 的初始化是 log(Uniform(1,16)),那是 Kimi Linear 的初始化写法留在代码里;K3 论文说 Ah 初始化为 0。推理时无所谓,复现训练时以论文为准。
下一篇
序列维度还剩四分之一:每四层一层的 Gated MLA。它负责 KDA 做不了的事,无损的全局内容检索。下一篇讲 MLA 的低秩压缩、K3 为什么敢让它完全不带位置编码、满秩输出门,以及 3
这个比例是怎么选出来的。
评论