0. 出发点:固定大小的记忆
KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。它要解决的问题是把随序列线性膨胀的 KV cache 换成一个固定大小的状态矩阵:O(n⋅d) 变 O(d2),上下文翻倍时状态大小不变。
代价是这个 dk×dv 的矩阵要承载任意长度的序列,记忆必须可写、可改、可遗忘。演进路线就是一步步补齐这三件事:
- 线性注意力:St=St−1+ϕ(kt)vt⊤,只能累加(同一个 key 写两次读出的是叠加,即记忆碰撞);
- Mamba-2:加标量衰减门 St=αtSt−1+ktvt⊤,会遗忘但衰减全局统一;
- DeltaNet:改为差值写入 St=(I−βtktkt⊤)St−1+βtktvt⊤,支持定点改写;
- GDN:两者结合 St=αt(I−βtktkt⊤)St−1+βtktvt⊤;
- KDA:把标量门拆成逐通道门 Diag(αt),各维度独立决定衰减速度,并加数值下界。
线性注意力与 SSM 两条路线的推导见前置篇《线性注意力与 SSM:两条技术路线的完整推导》,本文直接用其结论。
记号约定:两套写法互为转置
文献里状态矩阵有两种摆法,内容等价、互为转置,混用会让推导看起来「中途换了个式子」:
|
主约定(与 KDA 论文一致) |
转置约定(chunkwise 推导常用) |
| 状态形状 |
St∈Rdk×dv |
S^t∈Rdv×dk |
| 读出 |
ot=St⊤qt |
ot=S^tqt |
| 写入项 |
ktvt⊤ |
vtkt⊤ |
| 擦除算子 |
左乘 (I−βtktkt⊤)St−1 |
右乘 S^t−1(I−βtktkt⊤) |
关系就是一次转置 S^t=St⊤。由于擦除算子 Ht=I−βtktkt⊤ 对称,转置可以直接穿过它。后文若看到 kv⊤ 与 vk⊤ 互换、擦除算子从左跑到右,那是切换了约定,不是等式变了。
1. DeltaNet:从加法到差值写入
线性注意力的 S 只能叠加,修正这一点正是 DeltaNet 的动机。写入差值,不写全值——写入前先用当前 key 把状态里已存的内容读一遍:
vold=St−1⊤kt,ut=βt(vt−vold),St=St−1+ktut⊤
其中 kt 经 L2 归一化(∥kt∥2=1),βt∈(0,1) 是写入强度。
1.1 差值写入 = 先删后写
上式看不出「擦除」在哪,代入 ut 展开即可,每一步只用外积的结合律:
St=St−1+kt[βt(vt−vold)]⊤=St−1+βtktvt⊤−βtkt(St−1⊤kt)⊤=St−1+βtktvt⊤−βtktkt⊤St−1=删除:沿 kt 方向擦除(I−βtktkt⊤)St−1+新值βtktvt⊤
关键是倒数第二步:vold 自己就是由 St−1 算出来的,所以「减去旧值」必然能写成一个作用在 St−1 上的线性算子,而不是额外的加项。提公因式后 −βtktkt⊤ 并入单位阵,变成擦除算子。预测误差 vt−vold 就是 delta,Delta Rule 由此得名。
差值不是选的,是推出来的。把记忆看成一个在线回归问题:希望 S⊤kt≈vt,写成平方损失 L(S)=21∥S⊤kt−vt∥2。它的梯度是
∇SL=kt(S⊤kt−vt)⊤
即「钥匙 ⊗ 残差」——和一维情形 21(wx−y)2⇒(wx−y)x(误差乘输入)一模一样,只是乘法变外积。以 βt 为步长走一步 SGD:
St=St−1−βtkt(St−1⊤kt−vt)⊤=(I−βtktkt⊤)St−1+βtktvt⊤
与前面逐字相同。于是 βt 的角色也明确了:它就是学习率。写完立即读一次可以验证:
St⊤kt=(1−βt)St−1⊤kt+βtvt
读出是旧值与新值的凸组合,βt=1 时完全覆写。这也解释了 L2Norm 为何是前提:只有 ∥kt∥=1 时系数才是干净的插值,否则变成 1−βt∥kt∥2,可能跌出 [0,1]。
一句话:记忆 = 在线回归 → 平方损失 → 梯度 = 残差 ⊗ 钥匙 → 一步 SGD = 先删后写。
1.2 擦除算子 Ht:秩 1 与特征值
Ht=I−βtktkt⊤ 是全文频率最高的矩阵。kk⊤ 的第 j 列是 kj⋅k——换 j 只换倍率、方向永远是 k,所以它秩为 1。作为变换,
(kk⊤)x=k(k⊤x)=(k⊤x)⋅k
不管输入什么,输出永远落在 k 张成的直线上。
放回 Ht:I 让所有方向原样保留,减去 βtktkt⊤ 只在 kt 一个方向上动刀,其余 d−1 个方向碰都不碰。特征值分两种:
Htkt=(1−βt∥kt∥2)kt,Htx=x(∀x⊥kt)
配合 L2Norm 就是沿 kt 缩放 1−βt、正交补方向为 1,即特征值落在 [1−βt,1]⊂(0,1]——不放大、不翻转、不发散,这就是数值稳定性的特征值表述。反例很直白:不归一化时取 k=(1,2)⊤、β=0.6,则 1−β∥k∥2=−2,反复作用必然发散。
特征值之所以到处出现,是因为它回答了迭代系统最关心的问题:一个变换反复作用很多次之后会怎样——在特征向量方向上就是 λn,∣λ∣>1 爆炸、∣λ∣<1 衰减。这条线上的约束几乎都在围着它转(SSM 的 Aˉ 模长 ≤1、GDN/KDA 的 α∈(0,1)、Ht 的 [1−βt,1]、以及后面 1/Γ 的溢出)。
Ht 与 Householder 变换 I−2uu⊤ 同型:取 β∥k∥2=2 时两者相同,β∈(0,1) 的 delta rule 相当于「没照到底的半面镜子」,代价是不再正交,好处是擦除强度成了可学习的连续量。Householder 及其 WY 表示的线性代数细节见《线性注意力的线性代数前置知识》。
1.3 转移矩阵形式与并行的可能性
定义 Ht=I−βtktkt⊤,更新写成
St=HtSt−1+βtktvt⊤
这暴露了 DeltaNet 与 SSM 的同构:Ht 是随输入变化的状态转移矩阵(SSM 里是 Aˉ),βtktvt⊤ 是写入项(SSM 里是 Bˉxt)。记 Bt=βtktvt⊤ 逐层代入,递归可以彻底消掉:
St=(i=t∏1Hi)S0+i≤t∑(j=t∏i+1Hj)Bi
只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化之一——换成二叉树式两两合并,就能在 O(logn) 深度内并行完成。这是 parallel scan 与 chunkwise 并行化共同的理论根基。
2. GDN:把遗忘门与定点改写拼在一起
Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。核心洞察是两者互补:gating 是板擦(大面积擦除,但无法定点修改),delta rule 是铅笔(定点覆写,但无法快速清空)。
St=αt(I−βtktkt⊤)St−1+βtktvt⊤
其中 αt∈(0,1) 是数据相关的标量门(α=exp(−Softplus(Linear(xt))),在 log 空间算以保数值稳定)。三种极限:αt→1 退回 DeltaNet;βt→1 且 k 与已有记忆正交时退回 Mamba-2;αt→0 是整表清零再写入——两者单独都做不到的新能力。
几何上,(I−βkk⊤) 沿 k 方向压缩状态(定向),α 把整个状态矩阵均匀缩小(全局),作用于不同自由度,所以可以叠加。
注意作用顺序:α 乘的是整个 (I−βkk⊤)。把 α 只乘到擦除项上,chunkwise 形式会与递归形式对不上——这是实现时最容易踩的坑。
2.1 统一视角:四个模型是同一个在线优化问题
至此四个模型都出现了,它们是同一个在线优化问题的闭式解,差别只在目标函数。把每步看成一次在线学习:已有 St−1,新到样本 (kt,vt),求 St:
L(St)=正则:记忆保留∥St−At∥F2−拟合:关联学习2⟨Stkt, ut⟩
正则项惩罚状态的改动量,锚点 At 取 St−1 表示尽量不动、取 αtSt−1 表示容忍遗忘;拟合项要求用 kt 检索的结果朝写入目标 ut 对齐。目标对 St 是二次的,∇=2(St−At)−2utkt⊤=0,闭式解统一为 St=At+utkt⊤。于是差别完全归结为两个选择:锚点决定怎么遗忘,写入目标决定怎么写入(下表用转置约定)。
| 模型 |
锚点 At |
写入目标 ut |
闭式解 |
| Linear Attn |
St−1 |
vt |
St−1+vtkt⊤ |
| Mamba-2 |
αtSt−1 |
vt |
αtSt−1+vtkt⊤ |
| DeltaNet |
St−1 |
βt(vt−St−1kt) |
St−1(I−βtktkt⊤)+βtvtkt⊤ |
| GDN |
αtSt−1 |
βt(vt−αtSt−1kt) |
St−1(αt(I−βtktkt⊤))+βtvtkt⊤ |
| KDA |
St−1Diag(αt) |
βt(vt−St−1Diag(αt)kt) |
St−1Diag(αt)(I−βtktkt⊤)+βtvtkt⊤ |
说到底,这就是把「k→v 这条记忆是否还对得上」写成损失函数:对不上就修(拟合项),但别为了修这一条把整张表推翻(正则项)。GDN 在两处都取强化版本,KDA 再把锚点的标量收缩换成逐通道的 Diag(αt)。
2.2 Chunkwise 并行:WY 表示与 UT 变换
推理用递归(O(1)/token),但训练和 prefill 必须并行。设 chunk 大小 C、入口状态 S0,部分展开递归(转置约定):
Sr=γrS0Pr+i=1∑rγiγru~iki⊤,γr=j≤r∏αj,Pr=i≤r∏(I−βikiki⊤)
Pr 是 C 个 dk×dk 矩阵逐个相乘,O(Cdk3) 且严格顺序——比原递推还贵,必须处理。
关键观察:这种乘积永远不膨胀。每个因子都是「单位阵减秩 1」,连乘结果仍是「单位阵减一个低秩矩阵」:
Pr:=i=1∏r(I−βikiki⊤)=I−i≤r∑wiki⊤
证明(对 r 归纳)。P0=I 成立。设 Pr−1=I−∑i<rwiki⊤,右乘第 r 个因子:
Pr=(I−i<r∑wiki⊤)−βr(I−i<r∑wiki⊤)krkr⊤
展开后第三、四项都以 kr⊤ 结尾,合并同类项:
Pr=I−i<r∑wiki⊤−=wrβr(kr−i<r∑wi(ki⊤kr))kr⊤
归纳完成。wr 不是发明的技巧,而是「乘积保持低秩」逼出来的——要让结果保持 I−∑iwiki⊤ 的形式,括号里那一坨只能是 wr。它的直觉是修正后的擦除向量:第 r 步本想擦除 kr 方向,但若 kr 与之前的 ki 有重叠,连乘展开时前面的擦除已经顺带擦过这部分,wr 把已擦的量减掉以避免重复擦除,权重恰是重叠度 ki⊤kr。value 侧同理,只多一个衰减比:
u~r=βr(vr−i<r∑u~iγiγr(ki⊤kr))
γ 只出现在 u~ 里而不在 w 里,因为 GDN 的衰减是标量、与一切矩阵可交换,擦除连乘里的衰减可整体外提成 γr;但 value 侧每次写入的「存活时长」不同,这个相对衰减无法外提。对照 KDA:衰减变成向量后与擦除不可交换,γ 再也提不出去,只能渗进内积本身。
UT 变换:递归变成一次下三角求解。wr 只依赖 wi<r,是严格下三角依赖。把递归移项 wr+∑i<rβr(ki⊤kr)wi=βrkr,令 W 的第 r 行为 wr⊤、L=strictLower(diag(β)KK⊤),则 C 个方程一次写成 (I+L)W=diag(β)K:
W=(I+L)−1diag(β)K
求这个逆很便宜,因为严格下三角矩阵幂零(LC=0),Neumann 级数有限项精确截断:(I+L)−1=I−L+L2−⋯,一次前代法即可,O(C2)。(I+L) 是单位下三角(Unit Triangular,对角为 1 因为 wr 完整依赖自己),这就是「UT」的来源。所以 UT 不是额外发明的东西,它就是这两条递归的矩阵形态。
乘法链 → 加法链,这才是并行的真正来源。本质是把擦除矩阵的连乘 ∏r(I−βrkrkr⊤) 换成求和 I−∑rwrkr⊤,代价是求和项带修正。求和为什么就是胜利:加法可交换、可结合,因而可任意分组——树形归约、分块、稠密 matmul,GPU 的全部并行性都在奖励求和结构;而连乘必须一步一步来。同样的手法在这条路线上反复出现:∏αs=exp(∑gs)(累积衰减变 cumsum)、SSM 的卷积形式、Mamba-2 的 SSD 半可分矩阵,以及本节的 WY-UT。
辨析:因果掩码 ≠ UT 变换。两者都是下三角,但管的是两件事:因果掩码把 score 的严格上三角置零,管「不许看未来」;UT 变换处理块内历史写入之间的相互影响——delta rule 每次写入都「先读再改」,块内第 2 次写入读到了第 1 次的结果,UT 负责解耦。都是下三角不是巧合,是同一个原因:因果性使干扰系数矩阵天然下三角,对角线天然为 1。因果性决定了它是三角的,但做它的目的是解耦,不是掩码。
3. 从 GDN 到 KDA
| 维度 |
GDN |
KDA(→ K3) |
动机 |
| 更新式 |
St−1(αt(I−βkk⊤))+βvk⊤ |
(I−βkk⊤)Diag(αt)St−1+βvk⊤ |
α 标量 → 逐通道向量 |
| 门粒度 |
标量(整头同一衰减率) |
向量 αt∈(0,1)dk |
长期记忆通道与短期工作区分离 |
| 作用顺序 |
α 在外乘整个更新 |
Diag(α) 在内侧先衰减 S |
与逐通道参数化配套 |
| α 参数化 |
负 Softplus,(−∞,0) |
缩放 sigmoid,下界 gmin=−5 |
1/Γ 有界 → 全走 Tensor Core |
| chunkwise |
Γ 是标量比 |
Γ 是向量累积比 |
表达力↑ 数值难度↑ |
3.1 KDA 的递推公式
回到主约定(St∈Rdk×dv,读出 St⊤qt),与论文写法一致:
St=(I−βtktkt⊤)Diag(αt)St−1+βtktvt⊤
按从右往左三步:通道衰减 Diag(αt)St−1 对每一行分别乘对应通道的 α;方向擦除 (I−βtktkt⊤) 做定点擦除;写入 βtktvt⊤。
标量门的瓶颈在于所有 key 通道以同一比例衰减——要么一起记住,要么一起忘记。但有的通道在存「长期主题」(希望 α≈1),有的在存「临时指针」(希望快速腾出容量)。Diag(αt) 相当于给每一列配一个独立的遗忘速度,GDN 是 Diag(αt)=αtI 的特例。在线学习视角下,这就是把 §2.1 的正则项换成逐通道版:第 j 列的信任区域宽度正比于 αt(j),α 小的通道旧记忆先被压缩、新写入几乎没有阻力。KDA = 逐通道权重衰减 + delta rule。
作用顺序不可颠倒:Dt=Diag(αt) 与 (I−βkk⊤) 不可交换,写成 St−1(I−βkk⊤)Dt 或只衰减单位阵部分都是错的。擦除时检索用的也是已衰减的状态。
3.2 逐通道衰减下的 chunkwise
核心记号是累积 log 衰减 γi=∑s≤igs∈Rdk。带衰减的 key-key 内积:
Mci=(kc⊙eγc−γi)⊤ki,1≤i<c≤C
含义是 ki 写入的记忆衰减到第 c 步时与 kc 的重叠程度。衰减「长」在 M 内部,无法外提——这是向量门控带来的结构性变化,也是 KDA chunkwise 推导的核心难点。其余步骤与 GDN 同构:构造 Lci=βiMci,求 T=(I+L)−1,得 A^=Tdiag(β),然后
W=A^(eγi⊙ki)i=1..C,U=A^V,V~=U−WS[0]⊤
V~ 是伪值:它不是真实的 v,而是扣除了「历史记忆已有部分 + 块内其他位置已写部分」之后可直接累加、无需再修正的净增量。这与单步 delta rule 的 vt−St−1⊤kt 一致,V~ 是它的 chunk 并行版。输出与出口状态:
Acjqk=(qc⊙eγc−γj)⊤kj,oc=S[0](qc⊙eγc)+j≤c∑Acjqkv~j
S[C]=S[0]ΓC←0+c=1∑Cv~c(kc⊙eγC−γc)⊤
所有指数都是 ≤0 的差,整条链路没有一个 ≥1 的因子——这是刻意保持的,原因见 §3.4。
3.3 数值验算(dk=dv=2,C=2)
零初始状态,每步恒定衰减 α=(0.5,0.25),β=1:
k1=(11), v1=(02);k2=(21), v2=(20);q1=q2=(11)
递推式。第 1 步(S0=0):S1=v1k1⊤=(2020),o1=(4,0)。第 2 步先衰减 S1D=(100.50)(第 1 列 ×0.5、第 2 列 ×0.25——逐通道在动),再擦除写入:
S1D(I−k2k2⊤)=(−10−3.50),S2=(−12−3.54),o2=(−4.5, 6)
Chunkwise。eγ1=(0.5,0.25),eγ2−γ1=(0.5,0.25),衰减 KKT M21=(0.5,0.5)⋅(1,1)=1。UT:L=(0100),T=I−L,A^=T。WY:
W=A^(0.50.250.250.125)=(0.5−0.250.25−0.125),U=A^(2002)=(2−202)
伪值(S[0]=0⇒V~=U):v~1=(2,0),v~2=(−2,2)。注意 v~2=v2:因为 k2 与衰减后的 k1 写入重叠(M21=1),WY 把 v2 修正为扣除重叠后真正的新增。衰减注意力与输出:
Aqk=(20.7503),o1=2v~1=(4,0) ✓,o2=0.75v~1+3v~2=(−4.5, 6) ✓
块末状态 S[2]=v~1(k1⊙eγ2−γ1)⊤+v~2k2⊤=(−12−3.54) ✓,与递推式逐项一致。
3.4 下界衰减:1/Γ 溢出
GDN 的标量衰减在 chunkwise 里只以比值出现(∏s=j+1iαs≤1),天然安全。KDA 把 Γ 变成向量累积后,某些等价写法(如论文公式 (4) 的 K/Γ 因式分解)就得真的算出
Γt−1=exp(−s≤t∑gs)(逐通道)
衰减越快的通道,1/Γ 呈指数膨胀:恒定 α=0.5 时 t=500 已达 10150,t=2000 溢出 float64;BF16 训练下数十步即溢出。标量推广为向量后衰减率的动态范围被放大,溢出由个别情形变为普遍现象。
Kimi Linear 的规避方案是在对数空间算相对衰减、并把 chunk 再切成 16 token 二级瓦片:瓦片之间可交给 Tensor Core,但瓦片内部仍需按位置对显式计算,成为块内瓶颈。K3 的解决方式是不改公式、改参数化,用有界缩放 sigmoid 使溢出在数学上不可能发生:
gth=gmin⋅Sigmoid(eAhzth)∈(gmin,0),gmin=−5
于是 α>e−5≈6.7×10−3,16 token 块的累积对数衰减落在 (−80,0),重缩放因子小于 e80,在 BF16 动态范围内。表达力不受损:快通道在瓦片内仍可衰减到 e−80≈10−35,长期遗忘靠多步连乘即可,不需要单步清零的能力。以一个衰减下界换取全链路 Tensor Core 化,是有利的取舍。
4. 网络里的其他部件
q,k,v,α,β 本身怎么来、o~ 怎么变成层输出:GDN/KDA 采用 Llama 式宏架构,把 attention 换成 gated delta token mixer。
1 2 3 4 5 6
| x ─┬─ W_q/W_k ─ ShortConv ─ SiLU ─ L2Norm ──> q, k ├─ W_v ──── ShortConv ─ SiLU ───────────> v ├─ W_α(线性投影, log 空间)──────────────> α ├─ W_β(线性投影 + Sigmoid)─────────────> β │ 递归/分块计算 gated delta rule ──────> õ └─ y = W_o( Sigmoid(W_g x) ⊙ RMSNorm(õ) ), 再接 SwiGLU MLP + 残差
|
三个部件各负责一部分数值稳定性:
- ShortConv:窗口 4 的因果 depthwise 卷积,ShortConv(x)t=∑j<Wwj⊙xt−j。补足线性 RNN 缺失的局部建模能力——递归状态只携带压缩后的全局历史,最近几个 token 的精细模式由这层负责。depthwise 使通道间不混,与逐通道衰减语义配套;
- SiLU(即 β=1 的 Swish):x⋅σ(x),自门控、免参数,平滑非单调、梯度处处非零;
- L2Norm:只用在 Q/K。∥kt∥=1 是 delta rule 擦除稳定的前提(§1.2);Q 归一化使 QK⊤ 的 score 限制在 [−1,1]。V 不做,因为写入内容的幅值本身携带信息。
K3 另将输出门升级为输入相关的满秩 sigmoid 门 yt=Wo[Sigmoid(Wgxt)⊙RMSNorm(o~t)],允许每个 token 独立调制从循环状态读取的通道。
整体骨架自始至终没变:
状态×(衰减/删除算子)+(写入项)
参考: