对应论文 §2.1.1 的前半,公式 (1)、(2) 和 (6)。K3 的 93 层里有 69 层是 KDA,它来自 Kimi 团队 2025 年 10 月的 Kimi Linear。这一篇只讲递推形式,也就是「一个 token 进来,状态怎么变」;chunkwise 并行和 K3 的下界衰减放到下一篇。
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
你现在在这里:每个重复块里的三层 KDA,以及第 1 层。
为什么要线性注意力
softmax 注意力的代价是 KV cache 随上下文线性增长:1M 上下文下,即便用 MLA 压到每 token 每层 576 个数,24 层也要 27 GB 的 BF16 cache,每生成一个 token 都要全读一遍。
线性注意力能把这个 cache 压成每头一个 128×128 的矩阵,原因只有一步改写:
- 两种注意力都是 ot=∑i≤twt,ivi,区别只在权重。softmax 的权重是 exp(qt⋅ki) 再归一化,线性注意力直接用裸点积 wt,i=qt⋅ki。
- 线性的一项 (qt⋅ki)vi 可以改写成 qt⊤(kivi⊤)。查询跑到了括号外面,括号里只依赖 token i,于是求和可以搬进去:ot=qt⊤∑i≤tkivi⊤=qt⊤St。每个 token 到来时把自己的外积加进 S,不需要知道以后谁来查。
- softmax 做不到这一步:exp(qt⋅ki) 拆不成 f(qt)⊤g(ki),分母又是对所有 j 求和,查询被钉在权重内部,只能等查询来了再对每个 key 算一次。这就是 KV cache 存在的原因。
softmax:查询在权重里面分母:也含,且要遍历全部的指数和分母都钉在上,求和搬不进去只能等来了再对每个算一次,所以要一直存着这就是 KV cache线性:查询在求和外面128 × 128128 × 128128 × 128128 × 128每个 token 到来时把自己的外积加进去,不必知道以后谁来查16,384 个滑动求和本身不存在任何地方只和连一根线,读出精确,每步代价固定同样是对过去的 value 加权求和。softmax 的权重里有查询,求和只能等查询来了再做,所以每个 key、value 都要留着;线性注意力的一项 (qt⋅ki)vi 可以写成 qt⊤(kivi⊤),括号里只依赖 token i,可以提前累加。
代价是记忆有损:128 维空间里最多 128 个正交方向,序列一过 128 个 token 键必然重叠,取出来的总是一堆相似键下面的值的混合。
St∈Rdk×dv 就是这一族模型的记忆,读出是 ot=St⊤qt。对一张固定大小的表,能做的操作只有三种:往里加一条新记录,把某个键下的旧记录擦掉再写,让所有记录随时间忘掉。这三个动作各自做不做、做到什么粒度,就是从线性注意力到 KDA 的全部设计空间。下文一直用「加、擦、忘」指代它们。
一步在线学习:公式 (1) 从哪来
S 是一张查找表:行是键通道,列是值通道,用 q 去查得到 S⊤q。第 t 个 token 带来一对 (kt,vt),意思是「以后拿 kt 来查,应该查到 vt」。Kimi Linear(Table 7)用一个统一的视角回答「S 该怎么变」:写一个只看当前 token 的损失 Lt(S),做一步梯度下降,步长取 1:
St=St−1−∇SLt(St−1).
等号右边只有 St−1 和当前的 (kt,vt),所以这就是一个 RNN,标题里的「递推形式」指的就是它。三个动作各对应损失里的一项,梯度对 S 都是线性的,所以每一步都能整理成 St=AtSt−1+Bt:
| 动作 | 目标项 | 梯度 ∇SLt | 落在递推里的位置 |
|---|
| 加 | −⟨S⊤kt,vt⟩ | −ktvt⊤ | Bt=ktvt⊤ |
| 擦 | 2βt∥S⊤kt−vt∥2 | βtktkt⊤S−βtktvt⊤ | At 里减 βtktkt⊤,Bt=βtktvt⊤ |
| 忘 | 21∥Diag(1−αt)S∥F2 | Diag(1−αt)S | At 里乘 Diag(αt) |
「加」的目标对 S 是线性的,没有最小值,只会往上堆。「擦」把它换成回归目标,最小值处 S⊤kt=vt 精确成立,梯度里自带写入项,所以「加」和「擦」不会同时出现。「忘」是把 S 往零拉的正则,αt 取标量就是所有通道同速衰减,取向量就是每个键通道各有自己的保留率。
KDA 取「擦 + 逐通道忘」,并且先衰减、再对衰减后的状态做回归那一步:
St=(I−βtktkt⊤)Diag(αt)St−1+βtktvt⊤,o~t=St⊤qt,(K3 Eq. 1)
其中 αt∈(0,1)dk,βt∈(0,1)。先衰减再擦,意思是擦的量按衰减后剩下的状态算,和一步到位地把两项梯度相加差一项 βt(1−αt)ktkt⊤St−1。GDN 和 KDA 用的都是这个顺序,下一篇的 chunkwise 形式依赖它。
三项梯度的完整推导
对矩阵求梯度只用一招:把损失展开成对每个元素 Sij 的求和,对 Sij 求偏导,再按 (i,j) 摆回矩阵。反复用到 (S⊤k)j=∑iSijki。
加。
Lt(S)=−⟨S⊤kt,vt⟩=−j∑(i∑Sijkt,i)vt,j,∂Sij∂Lt=−kt,ivt,j,∇SLt=−ktvt⊤.代回梯度步:St=St−1+ktvt⊤,即 At=I,Bt=ktvt⊤。这是 Hebb 规则,也是最早的线性注意力。
擦。 记残差 e=S⊤kt−vt,ej=∑iSijkt,i−vt,j:
Lt(S)=2βtj∑ej2,∂Sij∂Lt=βtejkt,i⇒∇SLt=βtkte⊤=βtktkt⊤S−βtktvt⊤.第二个等号是把 e⊤=kt⊤S−vt⊤ 代进去。代回梯度步并把含 St−1 的项归到一起:
St=At(I−βtktkt⊤)St−1+Btβtktvt⊤.同一个式子的三种读法:只写入残差 βtkt(vt−St−1⊤kt)⊤,记忆已经答对就什么都不写,这是 delta rule 名字的来源;先把 kt 方向上原来存的东西擦掉 βt 的比例再写新的,这是「先擦再写」;∥kt∥=1 时 I−βtktkt⊤ 是沿 kt 的一个收缩,βt=1 时完全替换。
忘。 正则按元素独立,逐通道版本第 i 行有自己的系数:
Lt(S)+=21ij∑(1−αt,i)Sij2,∂Sij∂Lt=(1−αt,i)Sij,∇SLt=Diag(1−αt)S.代回梯度步:S−Diag(1−αt)S=Diag(αt)S。标量版本把 αt,i 换成同一个 αt 即可。让 αt 依赖输入,模型就能逐 token 决定「前面的东西还要不要」,这是 Mamba 一系强调的选择性。
拼起来。 「擦 + 忘」放进同一个目标,梯度相加,一步走完:
St=St−1−(βtktkt⊤St−1−βtktvt⊤)−Diag(1−αt)St−1=(Diag(αt)−βtktkt⊤)St−1+βtktvt⊤.GDN 和 KDA 实际用的是串行版本:令 S~t−1=Diag(αt)St−1,目标是 2βt∥S~t−1⊤kt−vt∥2,更新 St=(I−βtktkt⊤)S~t−1+βtktvt⊤,展开即公式 (1)。两者差一项 βt(1−αt)ktkt⊤St−1。
谱系:每一步加了哪块积木
「忘」路径:加衰减,不擦「擦」路径:delta rule,不忘旁支+ 标量衰减(忘)内积目标 → 回归目标(擦)标量 → 逐通道+ 标量衰减(忘)+ delta rule(擦)标量 → 逐通道隐式解擦除 ×放开绑定线性注意力 2020线性注意力 2020只加不擦,目标无下界RetNet 2023 · Mamba2 2024RetNet 2023 · Mamba2 2024每头一个标量Mamba2:三个合成任务全挂DeltaNet 2021DeltaNet 2021在方向先擦再写Longhorn 2024Longhorn 2024一步 SGD 换成闭式解GLA 2023 · HGRN2 2024GLA 2023 · HGRN2 2024每个键通道一个Gated DeltaNet 2024Gated DeltaNet 2024先忘,再擦,再写收敛慢于 KDAComba 2025Comba 2025擦得比写得少(),读出用KDA 2025KDA 2025先逐通道忘,再擦,再写RWKV-7 2025RWKV-7 2025一般 DPLR:擦除键、写入键各不相同「忘」路径:加衰减,不擦「擦」路径:delta rule,不忘+ 标量衰减(忘)→ 回归目标(擦)标量 → 逐通道+ 标量衰减(忘)+ delta rule(擦)标量 → 逐通道隐式解擦除 ×放开绑定线性注意力 2020线性注意力 2020只加不擦,目标无下界RetNet 2023 · Mamba2 2024RetNet 2023 · Mamba2 2024每头一个标量Mamba2:三个合成任务全挂DeltaNet 2021DeltaNet 2021在方向先擦再写Longhorn 2024Longhorn 2024一步 SGD 换成闭式解GLA 2023 · HGRN2 2024GLA 2023 · HGRN2 2024每个键通道一个Gated DeltaNet 2024Gated DeltaNet 2024先忘,再擦,再写收敛慢于 KDAComba 2025Comba 2025擦得比写得少(),读出用KDA 2025KDA 2025先逐通道忘,再擦,再写RWKV-7 2025RWKV-7 2025一般 DPLR:擦除键、写入键各不相同
从线性注意力出发有两条路:左边只加「忘」,先标量后逐通道,始终不擦;右边先加「擦」,再补「忘」。两条路在 KDA 汇合。图里已经写了每个节点的递推式,下面每个节点只说它改了什么、还缺什么。
线性注意力
Katharopoulos 等,2020。只加不擦不忘。它证明了 softmax 注意力可以换成固定大小的状态加一次递推,但旧关联堆积、互相干扰,长上下文效果远逊于 softmax。后面每一个模型都是在补这个目标函数的缺陷。
RetNet 和 Mamba2
加「忘」,标量粒度:St=αtSt−1+βtktvt⊤。RetNet(2023)的 α 是每头一个固定常数,一个头一个时间尺度。Mamba2(2024)的 αt=exp(−ΔtAh) 由输入决定,写入项也乘同一个步长,不重要的 token 可以既不忘也不写。两者都没有擦除:k 方向上原来存的东西只能等着淡出,不能被覆盖。下一节的合成任务里 Mamba2 全部失败,问题就出在这里。
GLA 和 HGRN2
「忘」从标量变成逐通道:St=Diag(αt)St−1+ktvt⊤,αt∈(0,1)dk。GLA(2023)的 αt 来自低秩投影加 sigmoid,这个参数化被 KDA 直接沿用。HGRN2(2024)把写入门和遗忘门绑在一起,键换成 1−αt,少一组参数。仍然没有擦除。
DeltaNet
Schlag 等 2021 提出,Yang 等 2024 给出并行形式。换目标而不是加项:回归目标,即上一节的「擦」。(I−βtktkt⊤) 是一个广义 Householder 变换,后面 chunkwise 并行全靠它。能做精确的键值覆盖,但没有任何遗忘:除非同一个键再来一次,写进去的东西永远留在 S 里。
Longhorn
DeltaNet 的旁支(2024)。不做一步梯度下降,而是把带正则的单步问题精确解出来,结果是 βt 换成 βt/(1+βtkt⊤kt),好处是 βt 不用再限制在 (0,1)。没有衰减。
Gated DeltaNet
Yang 等 2024。DeltaNet 加标量「忘」,顺序是先衰减再擦写。αt 每头一个,沿 Mamba2 的参数化;βt 是 sigmoid。这是第一个同时有「擦」和「忘」的模型,Kimi Linear 的所有消融都拿它当基线。剩下的问题是一个头的 128 个键通道以同一个速率衰减。
Comba
GDN 的旁支(2025)。擦除强度乘一个学习到的标量 b∈(0,1),让擦得比写得少;读出前把 query 减去 d 倍的当前键。衰减仍是每头一个标量。Kimi Linear 的 chunkwise 推导借用了它对转移矩阵累乘的写法,下一篇会碰到。
KDA
Kimi Linear,2025。把 GDN 的标量换成对角矩阵,或者说给 GLA 加上 delta rule,即公式 (1)。逐通道衰减的理由来自位置编码:RoPE 的力量在于每个维度有自己的旋转频率,标量衰减没有这种逐维的多样性。「转移矩阵的累乘就是位置编码」一节展开。
RWKV-7
和 KDA 并排的兄弟(2025):同样是逐通道衰减加擦除,但擦除键 κ^t、写入键 k~t 和逐通道学习率 at 各自独立,转移矩阵是一般的「对角加低秩」D−ab⊤。KDA 是 a、b、写入键全绑在 kt 上的特例,表达能力让了一步,换来内核少算一半东西,下一篇讲。
一句话总结
按「有没有擦」和「忘的粒度」两个坐标摆开:
| 无衰减 | 标量衰减 | 逐通道(对角)衰减 |
|---|
| 无 delta rule | 线性注意力 | RetNet、Mamba2 | GLA、HGRN2 |
| 有 delta rule | DeltaNet、Longhorn | Gated DeltaNet、Comba | KDA、RWKV-7 |
三个合成任务:擦和逐通道忘各买来了什么
语言模型的 PPL 差异只有零点几,看不出机制。Kimi Linear §5.1 用三个合成任务分别考察状态的三种能力(2 层 2 头、头维 128,序列长度 256 到 2048,比较 KDA、GDN 和 Mamba2):
| 任务 | 例子 | 考的能力 | 依赖的动作 |
|---|
| Palindrome | a b c d e | → e d c b a | 把每个 token 分别存下、按位置精确取回 | 键之间互不干扰,只叠加和衰减的记忆读出来是加权和,长了就分不开 |
| MQAR | A 4 · B 7 · C 2 · … · A ? → 4 | 同时保持很多组 k→v 绑定,同一个键再来要覆盖旧值 | 擦:查出来不等于 v 就把差值写进去 |
| Stack | push 3 a · push 3 b · pop 3 → b · pop 3 → a | 同一个键下反复改写,pop 后露出上一层 | 擦:没有擦除,pop 3 的两次回答无法不同 |
结果:KDA 在三个任务、所有长度上准确率最高,在 Palindrome 和 MQAR 上收敛明显快于 GDN;Mamba2 三个任务全部失败。对着更新规则看:Mamba2 和 KDA 的差别是有没有「擦」,它全挂,说明擦除是精确回忆和状态跟踪的前提;GDN 和 KDA 的差别只在衰减粒度,它慢,说明逐通道衰减买来的是收敛速度。这就是谱系图上两条汇入 KDA 的边的实验注脚。
公式 (1) 的图解:忘对付年龄,擦对付碰撞
新状态② 擦除① 衰减旧状态③ 写入=··+从右往左读:先把旧状态每一行按各自衰减,再把落在方向上的内容擦掉的比例,最后写入倍的新关联。读出:。行 = 键通道(,各自的衰减率),列 = 值通道()。
顺序是先衰减、再擦除、再写入。展开后是
St=(Diag(αt)−βtktkt⊤Diag(αt))St−1+βtktvt⊤,
即「对角加低秩」(DPLR)转移 D−atbt⊤,D=Diag(αt),at=βtkt,bt=kt⊙αt。两个低秩向量都绑在 kt 上,这是 KDA 内核比通用 DPLR 内核快约两倍的原因,下一篇细说。
既然有擦,为什么还要忘?两个减法项删的是不同的东西。
忘对付年龄。 S 固定大小,信息必须随时间淡出才能腾位置;有些信息也会自然失效(变量重新赋值、子任务结束、话题切换)。门控让「淡多少」由当前 token 决定。副产品是衰减编码了新近度,越早写入的东西被乘的衰减越多,这是下一节的来源。
擦对付碰撞。 所有 token 写进同一张表,一个和早先某个键相似的键,查出来是两个值的混合。kt⊤S 是「状态现在对这个键的回答」,减去 βtkt(kt⊤S) 就是写入之前把旧回答拿掉,kt 不指向的行保持原样。忘做不到这件事:忘按年龄均匀地淡,不知道哪条记录和新键撞了。
衰减在键轴。 Diag(αt) 左乘,衰减按行做,第 j 个键通道存的所有值一起以 αt,j 淡出。读出 S⊤q 沿键轴乘进去,查询触及的和键写入的是同一组分量,淡化某一行就是削弱「方向接近这个键分量的所有键」下面存的东西,这才是遗忘。若改成淡化某一列,缩的是所有存储值的同一个坐标,那是改输出尺度。内核的张量形状印证了这一点:S 是 [B,H,K,V],门 g 是 [B,T,H,K],没有 V 轴。
按行衰减:列:值通道,的第个分量行:键通道的第个分量沿键轴进入按列衰减:列:值通道,的第个分量行:键通道的第个分量沿键轴进入第行是「方向接近第个键分量的键」写进来的全部值淡这一行 = 忘掉这一类键下面存的东西这是遗忘。门的形状:每个键分量一个门第列是所有存储值的第个坐标,不管是哪个键写的淡这一列 = 每个读出向量的第维一起缩小这不是遗忘,是改输出尺度。没有轴的门S 的两个轴各由键和值的 128 个分量索引,读出 S⊤q 沿键轴乘进去,所以查询触及的和键写入的是同一组分量。只有沿键轴衰减才对应「忘掉某一类键下面的记录」。
转移矩阵的累乘就是位置编码
K3 的 24 层 MLA 全是 NoPE,整个模型里没有 RoPE。不带位置信息的注意力层分不清「A 在 B 前面」和「B 在 A 前面」,那顺序从哪来?从 KDA 的递推里来。
记每一步的转移矩阵 Tj=(I−βjkjkj⊤)Diag(αj),从 S0=0 展开公式 (1) 再用 qt 读出:
St=i=1∑t(TtTt−1⋯Ti+1)βikivi⊤,o~t=St⊤qt=i=1∑tst,i(qt⊤(Tt⋯Ti+1)ki)βivi.
这和注意力的形状完全一样:对每个历史位置 i 算一个分数 st,i,给 vi 加权求和。区别在于 qt 和 ki 中间夹着一串矩阵,隔得越远夹得越多。RoPE 做的是同一件事:q~t=Rtqt、k~i=Riki,打分时
st,i=qt⊤(Rt)⊤Riki=qt⊤(j=i+1∏tR−1)ki,
也是中间夹一串矩阵,只不过每一个都是同一个 R−1。
KDA数据相关、非正交、逐通道衰减⋯RoPE固定、正交、逐维频率⋯:每一格由第个 token 决定,颜色深浅示意该步保留了多少。:同一个分块旋转矩阵累乘次,只和距离有关,和内容无关。
两条链并排看,差别有三处:
- RoPE 的格子是固定的;KDA 的格子由第 j 个 token 的 αj、βj、kj 决定,可学习、数据相关。
- RoPE 的 R 是正交矩阵,分数随距离震荡但不衰减;KDA 的 Diag(αj) 是收缩的,分数随距离衰减,天然偏向近期。
- RoPE 每个维度有自己的频率,低频看长程、高频看短程;GDN 的标量 αj 相当于所有维度共用一个频率,KDA 的逐通道 αj 才是逐维频率的对应物。Kimi Linear 就是从这里出发把 GDN 的门换成了逐通道。
由此 K3 的两个决定就顺了:每层 MLA 前面都有三层 KDA 往残差流里写位置敏感的内容,所以 24 层 MLA 全部 NoPE;也不存在需要为长上下文重调的 RoPE 频率基或 YaRN 插值,K3 §3.4 从 8K 扩到 1M 不改任何位置编码参数。Kimi Linear Table 5 给了反面证据:MLA 层加回 RoPE 的变体短上下文打平,128K 的 RULER 从 84.3 掉到 78.8。第 3 篇讲 MLA 时再展开 NoPE 的工程收益。
参数化:从 xt 到 q,k,v,α,β
论文公式 (2),每个头 h:
qth,kthvthβthzth=L2Norm(Swish(ShortConv(Wq/khxt)))∈Rdk=Swish(ShortConv(Wvhxt))∈Rdv=Sigmoid(Wβhxt)∈(0,1)=Wα↑Wα↓xt+bαh∈Rdk
7168Linear7168 → 12288ShortConv(4) + SwishL2Norm逐头96 × 128Linear7168 → 12288ShortConv(4) + SwishL2Norm逐头96 × 128Linear7168 → 12288ShortConv(4) + Swish96 × 128Linear7168 → 96Sigmoid96 个标量Linear7168 → 128Linear128 → 12288有下界,96 × 128Linear(满秩)7168 → 12288SigmoidKDA 递推96 头 × (128 × 128)训练 / prefill:chunkwisedecode:逐 token 递推逐头 RMSNorm128 维,每头独立⊙输出门 96 × 128,Linear12288 → 71687168一层 KDA。粉色的那一行是 K3 的两个改动所在:衰减映射换成有下界的 scaled sigmoid,输出门从低秩改成满秩。
维度按 K3 的 d=7168、96 头、dk=dv=128:
- ShortConv。 深度可分的因果卷积,核大小 4(
short_conv_kernel_size),给每个 token 的 q、k、v 掺进前三个 token 的信息。Kimi Linear 的消融:去掉它 PPL 从 5.65 变 5.70。decode 时每层要为 q、k、v 各缓存最近 3 个 token 的投影,这是 KDA 除了 S 之外唯一的状态。
- q、k 的 L2 归一化。 让 ∥kt∥=1,I−βtktkt⊤ 的特征值落在 [1−βt,1],递推不会爆。
- β。 7168 → 96 的线性层加 sigmoid,每头一个标量,代码里叫
b_proj。
- α 的 logit z。 低秩投影 7168 → 128 → 12288(
f_a_proj、f_b_proj),加一个 12288 维偏置(dt_bias),每个头每个键通道一个 logit。从 logit 到 (0,1) 的映射,Kimi Linear 用 GDN 的负 softplus,K3 换成有下界的 scaled sigmoid,另有每头一个的可学习 log 尺度(A_log)。这是 K3 的改动之一,动机在 chunkwise 形式里,放到下一篇。
输出:逐头 RMSNorm、门、投影
递推读出 o~t=St⊤qt 之后,论文公式 (6):
yt=Wo[Sigmoid(Wgxt)⊙RMSNorm(o~t)].
RMSNorm 逐头,在 128 维上归一化;sigmoid 门作用在归一化后的每个通道上;然后 Wo(12288 → 7168)。代码里归一化和门融在一个 FusedRMSNormGated(activation='sigmoid') 里。
K3 的改动:门从低秩改满秩。 Kimi Linear 的门是 7168 → 128 → 12288 的低秩投影,论文自己说了原因:为了和基线做公平的参数量对比,并且「和满秩门性能相当」。K3 不需要这个约束,直接用 7168 → 12288 的满秩 Wg(use_full_rank_gate = true)。MLA 的输出门也一起改成满秩,第 3 篇会看到同样的公式。
# KimiDeltaAttention.forward 的收尾(HF modeling_kimi_linear.py)
if self.use_full_rank_gate:
g = self.g_proj(hidden_states) # 7168 -> 12288
else:
g = self.g_b_proj(self.g_a_proj(hidden_states)) # 7168 -> 128 -> 12288(Kimi Linear)
g = rearrange(g, '... (h d) -> ... h d', d=self.head_dim)
o = self.o_norm(o, g) # 逐头 RMSNorm,再乘 sigmoid(g)
o = self.o_proj(rearrange(o, 'b t h d -> b t (h d)'))
为什么是 sigmoid 门:Kimi Linear 的消融里 sigmoid 门 PPL 5.65,无门 5.67,GDN 默认的 Swish 门 5.81。sigmoid 输出在 (0,1),能把某些通道彻底关掉;引用的是 Qwen 团队关于 gated attention 的工作,门带来非线性和稀疏性,并缓解 attention sink。
每层的状态有多大
一层 KDA 在 decode 时携带的状态:
| 状态 | 形状 | 数量 |
|---|
| S | 96 头 × 128 × 128 | 1,572,864 |
| 三路 ShortConv 的窗口 | 3 × 12288 × 3 个 token | 110,592 |
约 1.68M 个数,BF16 下约 3.4 MB,69 层约 232 MB,和上下文长度无关。对照 MLA:每 token 每层 576 个数,24 层 1M 上下文约 27.6 GB。KDA 全部 69 层的状态相当于大约 8K 个 token 的 MLA cache。完整的账在第 3 篇。
术语坑
- α 在两篇论文里是同一个东西,Kimi Linear 写 αt=f(⋅) 没给 f,K3 把映射写全了(公式 (5))并且改了它。
- β 是学习率还是样本权重。 DeltaNet 原文把 βt 写成 SGD 的学习率,Kimi Linear Table 7 把它放进目标函数、步长取 1。两种写法得到同一个递推,本文跟 Table 7。
- 头数和头维。 KDA 和 MLA 共用
num_attention_heads = 96、head_dim = 128,但 MLA 的 query 和 key 实际是 192 维(128 加 64),第 3 篇解释。
- 行向量还是列向量。 论文按 S∈Rdk×dv 写,读出是 S⊤q;RWKV-7 和 Comba 原文用行向量约定,转移矩阵写在 St−1 右边,本文全部转成了列向量约定。FLA 内核里
transpose_state_layout=True 存的是转置,只是内存布局的事。
- 打分式里因子的顺序。 Kimi Linear 公式 (12) 把每一步的转移写成 Aj(I−βjkjkj⊤),先擦再衰减;本文按公式 (1) 写成 (I−βjkjkj⊤)Diag(αj),先衰减再擦。两个矩阵不相等,真正实现的是公式 (1) 的顺序。
下一篇
递推形式一个 token 一步,训练时没法用 Tensor Core。下一篇讲 Kimi Linear 的 chunkwise 并行形式(WY 表示、UT 变换、块间递推加块内并行),然后就能看懂 K3 为什么要给衰减加一个 −5 的下界,以及这个改动怎么把对角 tile 也搬上 Tensor Core。
评论