KDA 的来龙去脉:从线性注意力到 Kimi Delta Attention
KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。本文不铺开 K3 的全部架构,只讲一条演化线:KDA 之前的几个模型(线性注意力、Mamba、DeltaNet/GDN)分别解决了什么问题,KDA 又在它们的基础上改了什么 。
故事线很直白:线性注意力先把「固定大小的记忆」造出来,但这个记忆只会往里堆、不会清理;Mamba 教会它遗忘,但遗忘只能整片擦;DeltaNet 让它能精准改写某条记忆;GDN 把「擦除」和「改写」两样拼在一起;KDA 最后一步,让每个维度各自决定忘得多快。每一步都有公式推导、手动验算和工程上的理由。
0. 先算一笔账:为什么需要「固定大小的 KV 记忆」
标准 softmax 注意力在生成文本时,必须把每个历史 token 的 Key 和 Value 都存下来(KV cache)–因为每生成一个新 token,都要和前面所有 token 重新算一遍注意力。序列越长,这笔账越贵:
开销
softmax 注意力(KV cache)
固定状态(线性注意力一族)
解码显存
O ( n ⋅ d ) O(n \cdot d) O ( n ⋅ d ) (全部历史的 KV 都得缓存)
O ( d 2 ) O(d^2) O ( d 2 ) (一个固定矩阵 S,与 n 无关)
单步解码
O ( n ⋅ d ) O(n \cdot d) O ( n ⋅ d ) (扫全部历史)
O ( d 2 ) O(d^2) O ( d 2 ) (更新 S + 读出,两次矩阵乘)
预填充
O ( n 2 d ) O(n^2 d) O ( n 2 d )
O ( n d 2 ) O(n d^2) O ( n d 2 )
上下文长度翻倍
显存、延迟一起翻倍
完全不变
(单头视角,d 为特征维度。)
这笔账具体有多痛:几十万 token 的上下文就能把 KV cache 撑到几十 GB。长上下文场景(长文档、agent 记忆、仓库级代码理解)里,第一瓶颈往往不是模型权重,而是 KV cache 的显存和带宽。K3 把「序列维度缩放」列为头号工程目标,93 层里 69 层换成 KDA,买的就是右列这张表。
但固定大小不是白来的。一个 d v × d k d_v \times d_k d v × d k 的矩阵要装下任意长的序列,记忆就必须会写、会改、会忘 –只进不出的仓库,序列一长就糊成一团。怎么设计这块「会自我管理的固定记忆」,就是本文主线,按出场顺序:
线性注意力 :先把固定状态本身造出来(用外积累积出一个矩阵 S),但只会做加法,只进不出;
Mamba/SSM :加上一个全局衰减系数 α–会忘了,但只能整片擦;
DeltaNet :加上定点改写(delta rule)–会改了,写入的是新旧值之差;
GDN :把「板擦」(α)和「铅笔」(delta)拼在一起,先删后写;
KDA :把一个总的衰减系数拆成每个维度一个,各维度自己决定记多久,再加固数值下界,进 K3。
线性注意力和 SSM 这两条技术路线的完整推导–softmax 为什么挡路、核化怎么搬、积分因子法、ZOH 离散化、卷积形式、S4 到 Mamba-2 的演化–单独成了一篇:《线性注意力与 SSM:两条技术路线的完整推导》 。本文只要求读者记得结论,主线是:每个模型解决前一个的什么问题、留下什么新问题,讲完之后再回头看总表。
1. KDA 是什么:先看递推式
KDA 是 K3 序列维度缩放的核心,基于 delta rule 循环,并加上了通道维度的遗忘门 。
路线一:线性注意力–固定状态是怎么来的(概述)
线性注意力的完整推导(softmax 为什么挡路、核函数怎么把非线性搬走、结合律怎么把 O ( n 2 ) O(n^2) O ( n 2 ) 降到 O ( n ) O(n) O ( n ) )在前置篇《线性注意力与 SSM:两条技术路线的完整推导》 里一步步展开,这里只留骨架:
去掉 softmax ,( Q K ⊤ ) V → Q ( K ⊤ V ) (QK^\top)V \to Q(K^\top V) ( Q K ⊤ ) V → Q ( K ⊤ V ) :换一种括号方式,先算 K ⊤ V K^\top V K ⊤ V ,复杂度 O ( d n 2 ) → O ( n d 2 ) O(dn^2) \to O(nd^2) O ( d n 2 ) → O ( n d 2 ) ;
核化 s i m ( q , k ) = ϕ ( q ) ⊤ ϕ ( k ) \mathrm{sim}(q,k) = \phi(q)^\top\phi(k) sim ( q , k ) = ϕ ( q ) ⊤ ϕ ( k ) :非线性提前到 Q/K 各自身上,n × n n\times n n × n 的注意力矩阵就消失了,换成一个固定大小的状态 S ∈ R d × d S \in \mathbb{R}^{d\times d} S ∈ R d × d ;
状态可以逐步累积:S t = S t − 1 + ϕ ( k t ) v t ⊤ S_t = S_{t-1} + \phi(k_t)v_t^\top S t = S t − 1 + ϕ ( k t ) v t ⊤ ,每来一个 token 更新一次、读一次(fast weights 视角:S 是被逐 token 编程的快速权重,查询是读出)。
代价也在这三行里:S 只会加法 。写入只有「往上堆」,没有「删掉」–序列长度远超状态容量时,不同 k → v k \to v k → v 的关联互相干扰,旧信息永远赖在状态里,新信息覆盖不掉。前置篇里有个数值演示:同一个 key 写两次,线性注意力读出来的是两个 value 的叠加 ,而 delta rule 读出来的是新的那个 。这个「记忆碰撞」就是线性注意力最大的问题,修复它是 DeltaNet 的全部动机。
DeltaNet:从加法到差值写入
线性注意力留下的 S 是个只会「贴便利贴」的状态,修正它就是 DeltaNet 的全部动机。解法:写入差值,不写全值 。写入前,先用当前 key 把状态里已存的内容读一遍:
v o l d = S t − 1 ⊤ k t ( 当前 key 指向的旧内容 ) , u t = β t ( v t − v o l d ) v_{\mathrm{old}} = S_{t-1}^\top k_t \quad (\text{当前 key 指向的旧内容}), \qquad u_t = \beta_t (v_t - v_{\mathrm{old}})
v old = S t − 1 ⊤ k t ( 当前 key 指向的旧内容 ) , u t = β t ( v t − v old )
S t = S t − 1 + k t u t ⊤ = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤ S_t = S_{t-1} + k_t u_t^\top = (I - \beta_t k_t k_t^\top)\, S_{t-1} + \beta_t k_t v_t^\top
S t = S t − 1 + k t u t ⊤ = ( I − β t k t k t ⊤ ) S t − 1 + β t k t v t ⊤
符号
形状
含义
v o l d v_{\mathrm{old}} v old
[ d v ] [d_v] [ d v ]
当前 key 从状态里读出的已有内容
β t \beta_t β t
标量 ∈ ( 0 , 1 ) \in (0,1) ∈ ( 0 , 1 )
内容替换强度
u t u_t u t
[ d v ] [d_v] [ d v ]
实际写入的差值
展开后两项分解:删除项 ( I − β t k t k t ⊤ ) (I - \beta_t k_t k_t^\top) ( I − β t k t k t ⊤ ) 沿 k t k_t k t 方向抹掉旧内容,再加回新值项 β t k t v t ⊤ \beta_t k_t v_t^\top β t k t v t ⊤ –先删旧、再写新,一次外积同时完成两件事。预测误差 v t − v o l d v_t - v_{\mathrm{old}} v t − v old 就是 delta,Delta Rule 由此得名 :每次写入的从来不是新值本身,而是新值与旧值之差。
为什么恰好是差值?因为它就是对回归损失做一步 SGD 的结果 。注意 DeltaNet 没有动线性注意力的两个前置条件–ϕ \phi ϕ 、外积状态、递推形式全部保留(实现里常取 $\phi = $ L2Norm,即下文单位球约束),它改的只是写入规则 。而且这个规则不是拍脑袋的设计,它可以从回归损失严格推出来:
L ( S ) = 1 2 ∥ S k t − v t ∥ 2 , ∇ S L = ( S k t − v t ) k t ⊤ \mathcal{L}(S) = \tfrac{1}{2}\|S k_t - v_t\|^2, \qquad \nabla_S\mathcal{L} = (Sk_t - v_t)k_t^\top
L ( S ) = 2 1 ∥ S k t − v t ∥ 2 , ∇ S L = ( S k t − v t ) k t ⊤
S t = S t − 1 − β t ∇ S L = S t − 1 ( I − β t k t k t ⊤ ) ⏟ 先删:擦掉 S t − 1 k t 方向的旧值 + β t v t k t ⊤ ⏟ 后写 S_t = S_{t-1} - \beta_t\nabla_S\mathcal{L} = \underbrace{S_{t-1}(I - \beta_tk_tk_t^\top)}_{\text{先删:擦掉 } S_{t-1}k_t \text{ 方向的旧值}} + \underbrace{\beta_tv_tk_t^\top}_{\text{后写}}
S t = S t − 1 − β t ∇ S L = 先删 : 擦掉 S t − 1 k t 方向的旧值 S t − 1 ( I − β t k t k t ⊤ ) + 后写 β t v t k t ⊤
β t \beta_t β t 就是学习率。论文原版写法还提供了一个插值视角:S t = S t − 1 − v t o l d k t ⊤ + v t n e w k t ⊤ S_t = S_{t-1} - v_t^{\mathrm{old}}k_t^\top + v_t^{\mathrm{new}}k_t^\top S t = S t − 1 − v t old k t ⊤ + v t new k t ⊤ ,其中 v t n e w = β t v t + ( 1 − β t ) S t − 1 k t v_t^{\mathrm{new}} = \beta_tv_t + (1-\beta_t)S_{t-1}k_t v t new = β t v t + ( 1 − β t ) S t − 1 k t –δ-rule 是软覆写 ,β 控制新旧比例(β=1 且 k 为单位向量时完全覆写)。q/k 做 L2 归一化(钉在单位球上)保证 k k ⊤ kk^\top k k ⊤ 特征值 ≤ 1 \le 1 ≤ 1 ,更新稳定。
直观类比:state 是白板,key 是指针 。纯加性写入是往白板上不停贴便利贴,贴满了就糊;DeltaNet 是先擦掉指针 k t k_t k t 指向的那块,再贴新的。β t = 0 \beta_t = 0 β t = 0 :完全不写入,状态不动;β t = 1 \beta_t = 1 β t = 1 :完全替换,指针指向的内容整个换掉。这种「先擦后写」的机制让状态能覆盖错误记忆,这是 GDN/DeltaNet 一线的核心改进。
关键整理:写成转移矩阵形式 。定义 H t = I − β t k t k t ⊤ H_t = I - \beta_t k_t k_t^\top H t = I − β t k t k t ⊤ ,状态更新变成
S t = H t S t − 1 + β t k t v t ⊤ S_t = H_t\, S_{t-1} + \beta_t k_t v_t^\top
S t = H t S t − 1 + β t k t v t ⊤
这个形式之所以非常关键,是因为它暴露了 DeltaNet 和 SSM 的同构:H t H_t H t 正是随输入变化的状态转移矩阵 (SSM 里是 A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A ),β t k t v t ⊤ \beta_t k_t v_t^\top β t k t v t ⊤ 正是写入项(SSM 里是 B ˉ x t \bar B x_t B ˉ x t )。
把递归展开,消掉时间依赖 。记 B t = β t k t v t ⊤ B_t = \beta_t k_t v_t^\top B t = β t k t v t ⊤ ,递推 S t = S t − 1 H t + B t S_t = S_{t-1}H_t + B_t S t = S t − 1 H t + B t 逐层代入:
S 1 = S 0 H 1 + B 1 S_1 = S_0 H_1 + B_1
S 1 = S 0 H 1 + B 1
S 2 = S 0 H 1 H 2 + B 1 H 2 + B 2 S_2 = S_0 H_1 H_2 + B_1 H_2 + B_2
S 2 = S 0 H 1 H 2 + B 1 H 2 + B 2
S 3 = S 0 H 1 H 2 H 3 + B 1 H 2 H 3 + B 2 H 3 + B 3 S_3 = S_0 H_1 H_2 H_3 + B_1 H_2 H_3 + B_2 H_3 + B_3
S 3 = S 0 H 1 H 2 H 3 + B 1 H 2 H 3 + B 2 H 3 + B 3
S 4 = S 0 H 1 H 2 H 3 H 4 + B 1 H 2 H 3 H 4 + B 2 H 3 H 4 + B 3 H 4 + B 4 S_4 = S_0 H_1 H_2 H_3 H_4 + B_1 H_2 H_3 H_4 + B_2 H_3 H_4 + B_3 H_4 + B_4
S 4 = S 0 H 1 H 2 H 3 H 4 + B 1 H 2 H 3 H 4 + B 2 H 3 H 4 + B 3 H 4 + B 4
规律一目了然:S t = S 0 ∏ i ≤ t H i + ∑ i ≤ t B i ∏ i < j ≤ t H j S_t = S_0 \prod_{i \le t} H_i + \sum_{i \le t} B_i \prod_{i < j \le t} H_j S t = S 0 ∏ i ≤ t H i + ∑ i ≤ t B i ∏ i < j ≤ t H j 。展开后递归被彻底消掉了 :每个 S t S_t S t 都是初始状态、各步写入项与转移矩阵连乘的线性组合,只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化方案之一–换个括号方式(比如二叉树式两两合并),H H H 的连乘与写入项的累积可以在 O ( log n ) O(\log n) O ( log n ) 深度内并行完成。这就是 parallel scan / associative scan(并行扫描/结合扫描) 类方法的核心思想,也是 chunkwise 并行化和 SSM 训练并行(如 Mamba 的 selective scan)共同的理论根基。
但问题还没解决:H t H_t H t 是什么? 回看定义 H t = I − β t k t k t ⊤ H_t = I - \beta_t k_t k_t^\top H t = I − β t k t k t ⊤ 。设 k t ∈ R d k_t \in \mathbb{R}^d k t ∈ R d ,则 k t k t ⊤ k_t k_t^\top k t k t ⊤ 是 d × d d \times d d × d 的秩 1(rank-1)矩阵 –所有列都是 k t k_t k t 的标量倍,只张出一维。所以 H t H_t H t 本质是单位矩阵减去一个秩 1 矩阵 :k t k_t k t 方向被缩放 1 − β t ∥ k t ∥ 2 1 - \beta_t\|k_t\|^2 1 − β t ∥ k t ∥ 2 ,正交补方向原封不动。这类「I 减秩 1」的结构与**豪斯霍尔德变换(Householder transformation)**同族–QR 分解里用它做反射/消元。区别在于:Householder 取 β = 2 / ∥ k ∥ 2 \beta = 2/\|k\|^2 β = 2/∥ k ∥ 2 且要求 ∥ k ∥ = 1 \|k\|=1 ∥ k ∥ = 1 时是精确反射(范数保长),DeltaNet 的 β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 ) 是「部分反射」,软化为可学习的写入强度。秩 1 结构也解释了删除项为什么便宜:抹掉一个方向不需要满秩运算,一次外积就够。
到这里两条路线可以拼在一起了:DeltaNet 的 S t = H t S t − 1 + β t k t v t ⊤ S_t = H_t S_{t-1} + \beta_t k_t v_t^\top S t = H t S t − 1 + β t k t v t ⊤ 与 SSM 的 h t = A ˉ h t − 1 + B ˉ x t h_t = \bar A h_{t-1} + \bar B x_t h t = A ˉ h t − 1 + B ˉ x t 结构完全同构,差的只有一件事–H t H_t H t 的遗忘是「沿 k t k_t k t 方向删一块」,没有 SSM 那种全通道的指数衰减。
拼在一起:Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。论文的核心洞察:gating 和 delta rule 是互补 的两种记忆管理机制:
机制
比喻
能力
缺陷
Gating(Mamba-2)
板擦
一键大面积擦除(全局衰减)
无法定点修改
Delta rule(DeltaNet)
铅笔
定点覆写某个 key 的关联
无法快速清空
序列长度超过状态容量时,记忆碰撞必然发生;GDN 用板擦 + 铅笔的组合管理这块固定大小的白板:
S t = S t − 1 ( α t ( I − β t k t k t ⊤ ) ) + β t v t k t ⊤ S_t = S_{t-1}\big(\alpha_t(I - \beta_tk_tk_t^\top)\big) + \beta_tv_tk_t^\top
S t = S t − 1 ( α t ( I − β t k t k t ⊤ ) ) + β t v t k t ⊤
其中 α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 ) 是数据相关的标量门(Mamba-2 的参数化:α = exp ( \alpha = \exp( α = exp ( 负 Softplus(Linear( x t ) ) ) (x_t))) ( x t ))) ,在 log 空间计算以保证数值稳定)。三个极限情形读懂这个式子:
极限
行为
对应模型
α t → 1 \alpha_t \to 1 α t → 1
纯 delta rule,只定点改写
DeltaNet
β t → 1 \beta_t \to 1 β t → 1 ,k ⊥ 已有记忆
退化为 S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_tS_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤
Mamba-2
α t → 0 \alpha_t \to 0 α t → 0
整表清零再写入(硬重置)
新能力:两者都做不到
几何直觉 :( I − β k k ⊤ ) (I - \beta kk^\top) ( I − β k k ⊤ ) 是广义 Householder 反射–把状态往 k 方向「压扁擦除」;再乘标量 α \alpha α 把整张表均匀缩小。一个是定向 操作,一个是全局 操作,作用在不同维度上,所以可以叠加。
注意作用顺序 :擦除量是 β t ( S t − 1 k t ) \beta_t(S_{t-1}k_t) β t ( S t − 1 k t ) ,用的是未衰减 的旧读出;而 delta 对照的是衰减后 的旧值(v t − α t S t − 1 k t v_t - \alpha_tS_{t-1}k_t v t − α t S t − 1 k t )。α \alpha α 乘的是整个 ( I − β k k ⊤ ) (I-\beta kk^\top) ( I − β k k ⊤ ) –这一点写代码时极易弄错(把 α 只乘到擦除项上,chunkwise 形式会与递归形式对不上)。读取侧同样随时间累积衰减:token 在时间步 x x x 写入,在 x + t x+t x + t 读取时已经被 α x α x + 1 … α x + t \alpha_x\alpha_{x+1}\dots\alpha_{x+t} α x α x + 1 … α x + t 衰减过。实现中通过 γ r / γ i \gamma^r/\gamma^i γ r / γ i 项修正–分子分母都是 α \alpha α 连乘,相除就是区间衰减,本质是乘法形式的前缀和(prefix-sum) ,与 SSM 的 A ˉ \bar A A ˉ 连乘完全同源。
回头看:一张图 + 一张表把全家统一 (论文 Table 1)。到这里四个模型都出场了,回头看,它们其实是同一个在线优化问题的闭式解 ,差别只在目标函数:
Linear Attn ( 纯加性 ) → Mamba-2 ( 全局遗忘门 ) ↘ DeltaNet ( 定点覆写 ) ↗ GDN ( 板擦 + 铅笔 ) → KDA ( 逐通道门 + 下界 ) \text{Linear Attn}(纯加性) \to \text{Mamba-2}(全局遗忘门) \searrow \\
\text{DeltaNet}(定点覆写) \nearrow \;\; \text{GDN}(板擦+铅笔) \to \text{KDA}(逐通道门+下界)
Linear Attn ( 纯加性 ) → Mamba-2 ( 全局遗忘门 ) ↘ DeltaNet ( 定点覆写 ) ↗ GDN ( 板擦 + 铅笔 ) → KDA ( 逐通道门 + 下界 )
模型
在线学习目标
闭式解(状态更新)
Linear Attn
∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩ |S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle ∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩
S t = S t − 1 + v t k t ⊤ S_t = S_{t-1} + v_tk_t^\top S t = S t − 1 + v t k t ⊤
Mamba-2
∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩ |S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle ∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , v t ⟩
S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_tS_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤
DeltaNet
∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − S t − 1 k t ) ⟩ |S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - S_{t-1}k_t)\rangle ∣ S t − S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − S t − 1 k t )⟩
S t = S t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤ S_t = S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top S t = S t − 1 ( I − β t k t k t ⊤ ) + β t v t k t ⊤
GDN
∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − α t S t − 1 k t ) ⟩ |S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - \alpha_tS_{t-1}k_t)\rangle ∣ S t − α t S t − 1 ∣ F 2 − 2 ⟨ S t k t , β t ( v t − α t S t − 1 k t )⟩
S t = S t − 1 ( α t ( I − β t k t k t ⊤ ) ) + β t v t k t ⊤ S_t = S_{t-1}(\alpha_t(I-\beta_tk_tk_t^\top)) + \beta_tv_tk_t^\top S t = S t − 1 ( α t ( I − β t k t k t ⊤ )) + β t v t k t ⊤
KDA
(GDN 的逐通道化)
( I − β t k t k t ⊤ ) D i a g ( α t ) S t − 1 + β t k t v t ⊤ (I-\beta_tk_tk_t^\top)\,\mathrm{Diag}(\alpha_t)S_{t-1} + \beta_tk_tv_t^\top ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤
读法:第一项是正则 (别离上一刻状态太远 = 记忆保留),第二项是拟合 (S k ≈ v Sk \approx v S k ≈ v = 关联学习)。Mamba-2 与 DeltaNet 是两条独立的进化路线–前者把正则锚点换成可收缩的 α t S t − 1 \alpha_tS_{t-1} α t S t − 1 (选择性遗忘),后者把拟合项换成回归目标(修正误差);GDN 把两条路线拼在了一起,KDA 再把拼完之后的标量门拆成逐通道门。沿着两个轴看这张表的演化:
拟合轴 :线性注意力/Mamba-2 的拟合项 − 2 ⟨ S k , v ⟩ -2\langle Sk, v\rangle − 2 ⟨ S k , v ⟩ 只要求内积大,太弱,所以只能「越叠越厚」;DeltaNet/GDN 换成回归目标 ∥ S k − v ∥ 2 \|Sk - v\|^2 ∥ S k − v ∥ 2 类的构造,才获得「修正误差」的闭式解–delta 写入是回归目标的直接推论;
正则轴 :Mamba-2/GDN 的正则锚点是 α t S t − 1 \alpha_tS_{t-1} α t S t − 1 (可收缩的锚):正则放松 ⇒ \Rightarrow ⇒ 允许状态偏离旧值 ⇒ \Rightarrow ⇒ 选择性遗忘 。
所以这棵「进化树」可以反过来读:不是「谁加了什么模块」,而是目标函数从弱拟合 + 硬正则,逐级加强到强拟合 + 软正则 ,每一级都是上一级的严格泛化(GDN 令 β t \beta_t β t 退化为 0/边界值就退回全部下游模型)。再往前推一步就是**测试时训练(Test-Time Training, TTT)**一族:delta rule 是对 L ( S ) = 1 2 ∥ S k t − v t ∥ 2 \mathcal{L}(S)=\frac12\|Sk_t-v_t\|^2 L ( S ) = 2 1 ∥ S k t − v t ∥ 2 逐步做 SGD(β=学习率),GDN 等价于 SGD + 自适应权重衰减(α)–把「序列处理」看成「在线优化」,线性 RNN 全家都是这个视角的特例。
KDA 要做的最后一步就清楚了。Kimi Linear(后演化为 KDA)在 Gated DeltaNet 基础上的核心改进是细粒度门控 :不再是每个注意力头一个标量 α \alpha α ,而是每个通道一个独立衰减值:
α t ∈ ( 0 , 1 ) d k ( channel-wise 遗忘 ) \alpha_t \in (0,1)^{d_k} \quad (\text{channel-wise 遗忘})
α t ∈ ( 0 , 1 ) d k ( channel-wise 遗忘 )
作用是模型可以对不同维度做不同程度的记忆衰减 :部分通道 α \alpha α 接近 1,保留长期信息;部分通道 α \alpha α 接近 0,快速遗忘。打个比方,Gated DeltaNet 的标量门是「全屋一个总闸」,KDA 的逐通道门是「每盏灯一个旋钮」–同一个状态里,慢通道记长程依赖,快通道只管局部上下文,记忆容量按维度重新分配。再加上衰减下界(log \log log -decay 限制在 ( g min , 0 ) (g_{\min}, 0) ( g m i n , 0 ) )保证数值稳定,就得到 KDA 的完整递推式–正是下一小节的内容。
下面先把 GDN(Gated DeltaNet)的机制吃透:手算一遍、推导 chunkwise 并行形式,最后对照 KDA 看它继承了什么、改了什么。
GDN 数值验算:手算一遍
设定 d k = d v = 2 d_k = d_v = 2 d k = d v = 2 ,C = 3 C = 3 C = 3 个 token,S 0 = 0 S_0 = 0 S 0 = 0 。精心设计的数据–k 1 = k 2 = e 1 k_1 = k_2 = e_1 k 1 = k 2 = e 1 制造 key 碰撞 :
K = [ 1 0 1 0 0 1 ] , V = [ 1 0 2 0 3 1 ] , α = [ 0.8 , 0.5 , 0.9 ] , β = [ 1 , 1 , 0.6 ] K = \begin{bmatrix}1&0\\1&0\\0&1\end{bmatrix},\ V = \begin{bmatrix}1&0\\2&0\\3&1\end{bmatrix},\ \alpha = [0.8,\,0.5,\,0.9],\ \beta = [1,\,1,\,0.6]
K = 1 1 0 0 0 1 , V = 1 2 3 0 0 1 , α = [ 0.8 , 0.5 , 0.9 ] , β = [ 1 , 1 , 0.6 ]
顺序递归(ground truth) 。累积衰减 γ j = ∏ i ≤ j α i = [ 0.8 , 0.4 , 0.36 ] \gamma_j = \prod_{i\le j}\alpha_i = [0.8, 0.4, 0.36] γ j = ∏ i ≤ j α i = [ 0.8 , 0.4 , 0.36 ] 。
t=1 (k = e 1 k=e_1 k = e 1 , v = [ 1 , 0 ] v=[1,0] v = [ 1 , 0 ] , α = 0.8 \alpha=0.8 α = 0.8 , β = 1 \beta=1 β = 1 ):S 0 k = 0 S_0k = 0 S 0 k = 0 ,无旧记忆可删:
S 1 = 0.8 ⋅ 0 + 1 ⋅ ( [ 1 , 0 ] − 0 ) e 1 ⊤ = [ 1 0 0 0 ] , o 1 = S 1 q 1 = [ 1 , 0 ] S_1 = 0.8\cdot 0 + 1\cdot([1,0]-0)\,e_1^\top = \begin{bmatrix}1&0\\0&0\end{bmatrix}, \qquad o_1 = S_1q_1 = [1,0]
S 1 = 0.8 ⋅ 0 + 1 ⋅ ([ 1 , 0 ] − 0 ) e 1 ⊤ = [ 1 0 0 0 ] , o 1 = S 1 q 1 = [ 1 , 0 ]
t=2 (k = e 1 k=e_1 k = e 1 , v = [ 2 , 0 ] v=[2,0] v = [ 2 , 0 ] , α = 0.5 \alpha=0.5 α = 0.5 , β = 1 \beta=1 β = 1 ):旧读出 S 1 k 2 = [ 1 , 0 ] = v 1 S_1k_2 = [1,0] = v_1 S 1 k 2 = [ 1 , 0 ] = v 1 –碰撞发生 。擦除 − 1 ⋅ [ 1 , 0 ] e 1 ⊤ -1\cdot[1,0]e_1^\top − 1 ⋅ [ 1 , 0 ] e 1 ⊤ 把 v 1 v_1 v 1 完全擦掉;写入 delta = [ 2 , 0 ] − 0.5 ⋅ [ 1 , 0 ] = [ 1.5 , 0 ] = [2,0] - 0.5\cdot[1,0] = [1.5, 0] = [ 2 , 0 ] − 0.5 ⋅ [ 1 , 0 ] = [ 1.5 , 0 ] :
S 2 = 0.5 S 1 − [ 1 , 0 ] e 1 ⊤ + [ 1.5 , 0 ] e 1 ⊤ = [ 2 0 0 0 ] S_2 = 0.5S_1 - [1,0]e_1^\top + [1.5,0]e_1^\top = \begin{bmatrix}2&0\\0&0\end{bmatrix}
S 2 = 0.5 S 1 − [ 1 , 0 ] e 1 ⊤ + [ 1.5 , 0 ] e 1 ⊤ = [ 2 0 0 0 ]
现在查 e 1 e_1 e 1 得 [2,0] = v 2 v_2 v 2 ,v 1 v_1 v 1 已被覆写(对照线性注意力:纯加性会得到 [3,0] 的污染叠加)。
t=3 (k = e 2 k=e_2 k = e 2 , v = [ 3 , 1 ] v=[3,1] v = [ 3 , 1 ] , α = 0.9 \alpha=0.9 α = 0.9 , β = 0.6 \beta=0.6 β = 0.6 ):旧读出 S 2 k 3 = 0 S_2k_3 = 0 S 2 k 3 = 0 ,正交无碰撞。写入 0.6 ⋅ [ 3 , 1 ] e 2 ⊤ 0.6\cdot[3,1]e_2^\top 0.6 ⋅ [ 3 , 1 ] e 2 ⊤ ,同时全表再乘 α \alpha α (此处也乘了已写入的 e 1 e_1 e 1 行):
S 3 = [ 1.8 1.8 0 0.6 ] , o 3 = S 3 q 3 = [ 1.8 , 0.6 ] S_3 = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}, \qquad o_3 = S_3q_3 = [1.8,\,0.6]
S 3 = [ 1.8 0 1.8 0.6 ] , o 3 = S 3 q 3 = [ 1.8 , 0.6 ]
检查点:为什么 S 3 [ 0 , 0 ] = 1.8 S_3[0,0] = 1.8 S 3 [ 0 , 0 ] = 1.8 不是 2.0? t=3 的 α 3 = 0.9 \alpha_3=0.9 α 3 = 0.9 作用在整张表 上:第 1 行 2.0 × 0.9 = 1.8 2.0\times0.9 = 1.8 2.0 × 0.9 = 1.8 。这就是 gating 与 delta 的交互:即使 token 3 的 key 与 e 1 e_1 e 1 正交,它的遗忘门仍然衰减了 e 1 e_1 e 1 通道上的记忆。逐通道门(KDA)与下界衰减都是围绕这一约束做文章。
GDN 的 Chunkwise 并行形式:递归怎么塞进 GPU
推理用递归(O ( 1 ) O(1) O ( 1 ) /token),但训练/prefill 必须并行。GDN 的贡献是把 gating 并入 DeltaNet 的 WY 表示 chunkwise 框架。设 chunk 大小 C,chunk 入口状态 S 0 S_0 S 0 ,目标:一次矩阵乘算出整个 chunk 的 O 和 chunk 出口状态 S C S_C S C 。
第一步:部分展开递归 :
S r = γ r S 0 P r ⏟ F r + ∑ i = 1 r γ r γ i u ~ i k i ⊤ ⏟ G r S_r = \underbrace{\gamma_r S_0 P_r}_{F_r} + \underbrace{\sum_{i=1}^r\frac{\gamma_r}{\gamma_i}\,\tilde u_i k_i^\top}_{G_r}
S r = F r γ r S 0 P r + G r i = 1 ∑ r γ i γ r u ~ i k i ⊤
γ r = ∏ j ≤ r α j \gamma_r = \prod_{j\le r}\alpha_j γ r = ∏ j ≤ r α j (α \alpha α 是标量,可提到矩阵连乘外面);
P r = ∏ i ≤ r ( I − β i k i k i ⊤ ) P_r = \prod_{i\le r}(I - \beta_ik_ik_i^\top) P r = ∏ i ≤ r ( I − β i k i k i ⊤ ) :纯 Householder 连乘,与 gating 无关 ;
u ~ i \tilde u_i u ~ i :吸收了 β \beta β 和衰减修正的「伪 value」。
第二步:WY 表示–秩 1 连乘压缩成两个小矩阵 。关键观察:P r = I − ∑ i ≤ r w i k i ⊤ P_r = I - \sum_{i\le r}w_ik_i^\top P r = I − ∑ i ≤ r w i k i ⊤ (乘积不膨胀,始终是 I 减一个秩 ≤ r \le r ≤ r 的项)。w r w_r w r 、u ~ r \tilde u_r u ~ r 的递归(W 不带 γ \gamma γ 、U ~ \tilde U U ~ 带 γ \gamma γ ,这是 gating 并入的位置 ):
w r = β r ( k r − ∑ i < r w i ( k i ⊤ k r ) ) , u ~ r = β r ( v r − ∑ i < r u ~ i γ r γ i ( k i ⊤ k r ) ) w_r = \beta_r\Big(k_r - \sum_{i<r}w_i\,(k_i^\top k_r)\Big), \qquad \tilde u_r = \beta_r\Big(v_r - \sum_{i<r}\tilde u_i\,\tfrac{\gamma_r}{\gamma_i}(k_i^\top k_r)\Big)
w r = β r ( k r − i < r ∑ w i ( k i ⊤ k r ) ) , u ~ r = β r ( v r − i < r ∑ u ~ i γ i γ r ( k i ⊤ k r ) )
直觉:w r w_r w r 是「擦除向量」–前面的写入若与 k r k_r k r 重叠,先从自己的擦除量里扣掉(避免重复擦除);u ~ r \tilde u_r u ~ r 是「修正量 value 版」。两者形式相同,差别只在 u ~ \tilde u u ~ 的内积多了一个衰减比 γ r / γ i \gamma_r/\gamma_i γ r / γ i :i 越久远、衰减越多,它对当前擦除的干扰越小 。
第三步:UT 变换–递归变成一次下三角方程求解 。w 和 u ~ \tilde u u ~ 的递归是「自回归」的(依赖前面的 w/u ~ \tilde u u ~ ),但系数矩阵是单位下三角,可整体解方程。
辨析:因果掩码 ≠ UT 变换 。两者容易混,因为碰巧都是下三角,但要管的完全是两件事:
因果掩码 (tril):把 score 矩阵的严格上三角直接置零,一行掩码操作,没有任何「变换」可言。它管的是「第 i 个输出不许看未来的 token」;
UT 变换 (求 ( I + L ) − 1 (I+L)^{-1} ( I + L ) − 1 ):管的是已经看过的历史之间怎么算账 。delta rule 每次写入是「先读、再改」,块内第 2 次写入读到了第 1 次的结果,第 3 次读到前两次–历史写入之间互相纠缠,UT 把这团纠缠解开。
两者都是下三角不是巧合,是同一个原因:因果性 。位置 i 的写入只能影响 i 之后的位置,所以「干扰系数矩阵」天然下三角;对角线天然是 1(自己的写入自己完整可见),于是要逆的矩阵恰好是单位 下三角–这就是「UT」(Unit Triangular)名字的由来。一句话:因果性决定了它是三角的,但做它的目的是解耦,不是掩码 。
为什么必须做 (最短版本):并行化要求把 C 次顺序写入折成一锤子外积累加 ∑ i v ~ i k i ⊤ \sum_i \tilde v_i k_i^\top ∑ i v ~ i k i ⊤ 。如果直接用原始 v i v_i v i 累加,重叠部分会被重复计算–第 1 次写入的内容会透过后续写入的「先读」环节被间接再写一遍。UT 变换算出每个 v i v_i v i 该扣除多少,使等式精确成立。不做的代价:要么结果错,要么退回逐 token 循环。
旁注:这类下三角变换是个大家族 。「顺序递推 ↔ 三角矩阵求逆」是个通用模式,KDA 的 UT 只是其中一员:
家族成员
三角结构
与 UT 的关系
三角方程组求解(前代/回代)
LU、Cholesky 分解之后的三角系统
UT 变换的计算过程就是 一次前代法–同一个算法
Householder QR 的紧凑 WY 表示
反射连乘 ∏ ( I − β k k ⊤ ) = I + Y T Y ⊤ \prod(I - \beta kk^\top) = I + YTY^\top ∏ ( I − β k k ⊤ ) = I + Y T Y ⊤ ,T T T 三角
「WY」「UT」两个名字的学术出处(Schreiber-Van Loan 1989);DeltaNet 把同样的打包思想借到 delta rule
因果卷积的逆(去卷积)
因果线性系统 = 下三角 Toeplitz,其逆也是下三角 Toeplitz
信号处理经典:「用三角逆矩阵解顺序依赖」
Mamba-2/SSD 的半可分矩阵
块内注意力 ( C B ⊤ ) ⊙ L (CB^\top)\odot L ( C B ⊤ ) ⊙ L ,L L L 下三角衰减
同一枚硬币另一面:顺序 SSM 递推等价于带结构下三角矩阵,KDA 的 A q k A^{qk} A q k 衰减注意力项完全是这个结构
幂零矩阵 Neumann 级数
( I + L ) − 1 = I − L + L 2 − ⋯ (I+L)^{-1} = I - L + L^2 - \cdots ( I + L ) − 1 = I − L + L 2 − ⋯
一般矩阵是无穷级数,但严格下三角矩阵幂零(L C = 0 L^C = 0 L C = 0 ),有限项精确截断 –UT 能精确、便宜算出的数学原因
放进大图景里记:凡是「顺序执行的因果更新」,分块并行化时都会变成一个下三角矩阵的求逆/求解问题 –SSM、delta rule、因果卷积、QR 分解,全是这个模式的化身。KDA chunkwise 里的 WY 表示(第二步)与 UT 变换(第三步),分别是这个模式在「写入打包」与「依赖解耦」上的具体化。
回到公式。定义衰减感知掩码 Γ i j = γ i / γ j ( i > j ) \Gamma_{ij} = \gamma_i/\gamma_j\ (i>j) Γ ij = γ i / γ j ( i > j ) ,则:
T plain = [ I + s t r i c t L o w e r ( d i a g ( β ) K K ⊤ ) ] − 1 d i a g ( β ) , W = T plain K T_{\text{plain}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)\,KK^\top)\big]^{-1}\mathrm{diag}(\beta), \qquad W = T_{\text{plain}}K
T plain = [ I + strictLower ( diag ( β ) K K ⊤ ) ] − 1 diag ( β ) , W = T plain K
T gated = [ I + s t r i c t L o w e r ( d i a g ( β ) ( Γ ⊙ K K ⊤ ) ) ] − 1 d i a g ( β ) , U ~ = T gated V T_{\text{gated}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)\,(\Gamma\odot KK^\top))\big]^{-1}\mathrm{diag}(\beta), \qquad \tilde U = T_{\text{gated}}V
T gated = [ I + strictLower ( diag ( β ) ( Γ ⊙ K K ⊤ )) ] − 1 diag ( β ) , U ~ = T gated V
W 和 U ~ \tilde U U ~ 只差内积处的一个 γ \gamma γ 比。下三角求逆(C × C C\times C C × C ,实践中 C=64)用前代法(forward substitution)完成。
第四步:chunk 输出与出口状态 。记号(论文的箭头约定):( ⋅ ) ← r = γ r ( ⋅ ) r \overleftarrow{(\cdot)}_r = \gamma_r(\cdot)_r ( ⋅ ) r = γ r ( ⋅ ) r (衰减到 chunk 首端),( ⋅ ) → r = γ C γ r ( ⋅ ) r \overrightarrow{(\cdot)}_r = \frac{\gamma_C}{\gamma_r}(\cdot)_r ( ⋅ ) r = γ r γ C ( ⋅ ) r (衰减到 chunk 末端)。
输出 (两项:读 chunk 外旧状态 + chunk 内交互):
O = Q ← S 0 ⊤ + ( Q K ⊤ ⊙ Γ causal ) ( U ~ − W ← S 0 ⊤ ) O = \overleftarrow{Q}\,S_0^\top + \big(QK^\top\odot\Gamma_{\text{causal}}\big)\big(\tilde U - \overleftarrow W S_0^\top\big)
O = Q S 0 ⊤ + ( Q K ⊤ ⊙ Γ causal ) ( U ~ − W S 0 ⊤ )
出口状态 :
S C = γ C S 0 + ( U ~ → − W → S 0 ⊤ ) ⊤ K ( U ~ → r = γ C γ r u ~ r ) S_C = \gamma_C S_0 + \big(\overrightarrow{\tilde U} - \overrightarrow W S_0^\top\big)^\top K \qquad(\overrightarrow{\tilde U}_r = \tfrac{\gamma_C}{\gamma_r}\tilde u_r)
S C = γ C S 0 + ( U ~ − W S 0 ⊤ ) ⊤ K ( U ~ r = γ r γ C u ~ r )
结构读法:
输出的第二项是「chunk 内小注意力」:Q K ⊤ ⊙ Γ causal QK^\top\odot\Gamma_{\text{causal}} Q K ⊤ ⊙ Γ causal 就是带衰减的因果注意力矩阵 ,attend 的对象不是 V 而是修正后的伪 value U ~ − W ← S 0 ⊤ \tilde U - \overleftarrow WS_0^\top U ~ − W S 0 ⊤ (后者是「旧状态在这个 chunk 里该被擦掉的部分」);
出口状态 = 旧状态整体衰减 γ C \gamma_C γ C + 修正量加权写入(权重 γ C / γ i \gamma_C/\gamma_i γ C / γ i :越早写入衰减越多);
全部计算都是 C × C C\times C C × C 、C × d k C\times d_k C × d k 、C × d v C\times d_v C × d v 的稠密矩阵乘–Tensor Core 友好 ;chunk 间只传一个 d v × d k d_v\times d_k d v × d k 矩阵。
数值验算:拿顺序递归的例子走一遍 chunkwise 。
Step 1 :K K ⊤ = [ 1 1 0 1 1 0 0 0 1 ] KK^\top = \begin{bmatrix}1&1&0\\1&1&0\\0&0&1\end{bmatrix} K K ⊤ = 1 1 0 1 1 0 0 0 1 ,Γ strict ⊙ K K ⊤ \Gamma_{\text{strict}}\odot KK^\top Γ strict ⊙ K K ⊤ 只有 (2,1) 处非零 = γ 2 / γ 1 × 1 = 0.5 = \gamma_2/\gamma_1\times1 = 0.5 = γ 2 / γ 1 × 1 = 0.5 。
Step 2 :解两个下三角方程(β = [ 1 , 1 , 0.6 ] \beta = [1,1,0.6] β = [ 1 , 1 , 0.6 ] ):
T plain = [ 1 0 0 − 1 1 0 0 0 0.6 ] , T gated = [ 1 0 0 − 0.5 1 0 0 0 0.6 ] T_{\text{plain}} = \begin{bmatrix}1&0&0\\-1&1&0\\0&0&0.6\end{bmatrix}, \qquad T_{\text{gated}} = \begin{bmatrix}1&0&0\\-0.5&1&0\\0&0&0.6\end{bmatrix}
T plain = 1 − 1 0 0 1 0 0 0 0.6 , T gated = 1 − 0.5 0 0 1 0 0 0 0.6
对比唯一差别 (2,1):-1 -> -0.5,正是 γ 2 / γ 1 = 0.5 \gamma_2/\gamma_1 = 0.5 γ 2 / γ 1 = 0.5 的衰减 –t=2 时 v 1 v_1 v 1 已被衰减一半,擦除它的需求也减半。
Step 3 :W = T plain K = [ 1 0 0 0 0 0.6 ] W = T_{\text{plain}}K = \begin{bmatrix}1&0\\0&0\\0&0.6\end{bmatrix} W = T plain K = 1 0 0 0 0 0.6 ,U ~ = T gated V = [ 1 0 1.5 0 1.8 0.6 ] \tilde U = T_{\text{gated}}V = \begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix} U ~ = T gated V = 1 1.5 1.8 0 0 0.6 。
看 u ~ \tilde u u ~ 的第二行:$v_2 - 0.5\cdot\tilde u_1(k_1^\top k_2) = [2,0] - 0.5[1,0] = $ [1.5, 0] –与顺序递归里手算的 delta 完全一致 ✓(W 2 = 0 W_2 = 0 W 2 = 0 则因为 w 1 w_1 w 1 已把 e 1 e_1 e 1 方向擦干净,β = 1 \beta=1 β = 1 时无需再擦)。
Step 4 (S 0 = 0 S_0=0 S 0 = 0 ,修正项消失):U ~ → = γ C γ i u ~ i \overrightarrow{\tilde U} = \frac{\gamma_C}{\gamma_i}\tilde u_i U ~ = γ i γ C u ~ i 逐行乘 [ 0.36 / 0.8 , 0.36 / 0.4 , 1 ] = [ 0.45 , 0.9 , 1 ] [0.36/0.8, 0.36/0.4, 1] = [0.45, 0.9, 1] [ 0.36/0.8 , 0.36/0.4 , 1 ] = [ 0.45 , 0.9 , 1 ] :
U ~ → = [ 0.45 0 1.35 0 1.8 0.6 ] , S C = U ~ → ⊤ K = [ 1.8 1.8 0 0.6 ] ✓ \overrightarrow{\tilde U} = \begin{bmatrix}0.45&0\\1.35&0\\1.8&0.6\end{bmatrix}, \qquad S_C = \overrightarrow{\tilde U}^\top K = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}\ \checkmark
U ~ = 0.45 1.35 1.8 0 0 0.6 , S C = U ~ ⊤ K = [ 1.8 0 1.8 0.6 ] ✓
与顺序递归的 S 3 S_3 S 3 完全一致 。
Step 5 :Γ causal = [ 1 0 0 0.5 1 0 0.45 0.9 1 ] \Gamma_{\text{causal}} = \begin{bmatrix}1&0&0\\0.5&1&0\\0.45&0.9&1\end{bmatrix} Γ causal = 1 0.5 0.45 0 1 0.9 0 0 1 ,O = ( Q K ⊤ ⊙ Γ causal ) U ~ O = (QK^\top\odot\Gamma_{\text{causal}})\tilde U O = ( Q K ⊤ ⊙ Γ causal ) U ~ :
O = [ 1 0 0 0.5 1 0 0 0 1 ] [ 1 0 1.5 0 1.8 0.6 ] = [ 1 0 2 0 1.8 0.6 ] ✓ O = \begin{bmatrix}1&0&0\\0.5&1&0\\0&0&1\end{bmatrix}\begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix} = \begin{bmatrix}1&0\\2&0\\1.8&0.6\end{bmatrix}\ \checkmark
O = 1 0.5 0 0 1 0 0 0 1 1 1.5 1.8 0 0 0.6 = 1 2 1.8 0 0 0.6 ✓
与顺序递归的 o 1 = [ 1 , 0 ] o_1=[1,0] o 1 = [ 1 , 0 ] 、o 2 = [ 2 , 0 ] o_2=[2,0] o 2 = [ 2 , 0 ] 、o 3 = [ 1.8 , 0.6 ] o_3=[1.8,0.6] o 3 = [ 1.8 , 0.6 ] 逐步精确一致 (含非零 S 0 S_0 S 0 的一般情形同样对拍通过)。
完整 block:GDN 在网络里长什么样
Llama 式宏架构,把 attention 换成 gated delta token mixer:
1 2 3 4 5 6 7 x ─┬─ W_q/W_k ─ ShortConv ─ SiLU ─ L2Norm ─-> q, k ├─ W_v ──── ShortConv ─ SiLU ───────────-> v ├─ W_α(线性投影,log 空间负 Softplus)──-> α ∈ (0,1) 标量 ├─ W_β(线性投影 + Sigmoid)────────────-> β ∈ (0,1) │ 递归/分块计算 gated delta rule ─-> õ ├─ W_g ─ SiLU ─────────────────────────-> 输出门 └─ y = W_o( Sigmoid(W_g x) ⊙ RMSNorm(õ) ),再走 SwiGLU MLP + 残差
要点:q/k 走 ShortConv+SiLU+L2Norm(局部上下文 + 单位球稳定擦除几何),α/β 只走线性投影 (它们是标量,无需卷积);输出门与 Mamba 的 SiLU 门一脉相承(K3 把它升级为全秩)。
旁注:ShortConv 的数学定义
上面流程图里的 ShortConv(Short Convolution,短因果深度卷积)出自 Kimi Linear(K3 技术报告引用 [64]),位于 Q/K/V 线性投影之后、进入 KDA 递归之前 –顺序是 x → W q / k / v → S h o r t C o n v → … x \to W_{q/k/v} \to \mathrm{ShortConv} \to \dots x → W q / k / v → ShortConv → … ,即卷积是「KDA 前置」而非「投影前置」。它只捕捉当前 token 前面少量局部上下文,无未来 token 泄露;且在 Kimi Linear / FLA 实现里卷积不是裸用的,后面紧跟 SiLU:s i l u ( c o n v ( x ) ) \mathrm{silu}(\mathrm{conv}(x)) silu ( conv ( x )) (见下一个旁注)。
输入单头向量 x t ∈ R d x_t \in \mathbb{R}^d x t ∈ R d ,卷积核窗口大小 W(标准取 4),深度可分离(depthwise)且因果:
S h o r t C o n v ( x ) t = ∑ j = 0 W − 1 w j ⊙ x t − j \mathrm{ShortConv}(x)_t = \sum_{j=0}^{W-1} w_j \odot x_{t-j}
ShortConv ( x ) t = j = 0 ∑ W − 1 w j ⊙ x t − j
记号
含义
w j ∈ R d w_j \in \mathbb{R}^d w j ∈ R d
每通道独立卷积权重–depthwise,每个特征通道一套独立卷积核
⊙ \odot ⊙
逐元素相乘
因果约束
t − j ≥ 0 t - j \ge 0 t − j ≥ 0 ;t < j t < j t < j 时零填充,看不到未来 token
三个性质决定了它为什么放在这个位置:
时序维度滑动,只混合最近 W 个历史 token ,开销远小于全局注意力;
depthwise :w k w_k w k 与 x t − k x_{t-k} x t − k 逐通道相乘,通道间不混–局部上下文的注入不破坏各通道独立的衰减语义(与 D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 逐通道门配套);
因果 :与上一节因果时序卷积同理,LLM 自回归必备。
直觉上,它给 KDA 补上了线性 RNN 天生缺失的「局部听觉」–递归状态只携带压缩后的全局历史,最近几个 token 的精细局部模式(词内字符、短程搭配)由这层小卷积负责。
旁注:Swish(SiLU,自门控激活函数)
流程图里 ShortConv 之后、L2Norm 之前的 SiLU,即 Swish–严格说 Swish 带可学参数时是 x ⋅ σ ( β x ) x \cdot \sigma(\beta x) x ⋅ σ ( β x ) ,论文和实现里固定 β = 1 \beta = 1 β = 1 ,两个名字就此等价:
S w i s h ( x ) = x ⋅ σ ( x ) , σ ( x ) = 1 1 + e − x \mathrm{Swish}(x) = x \cdot \sigma(x), \qquad \sigma(x) = \frac{1}{1 + e^{-x}}
Swish ( x ) = x ⋅ σ ( x ) , σ ( x ) = 1 + e − x 1
分层计算只有两步:1对输入的每个元素算 σ ( x ) = 1 / ( 1 + exp ( − x ) ) \sigma(x) = 1/(1+\exp(-x)) σ ( x ) = 1/ ( 1 + exp ( − x )) ;2原输入 x 与 sigmoid 结果逐元素相乘 得到输出。Swish = 输入 × 输入自己的 sigmoid 门–自门控 (self-gated):门控信号不是外部来的,是输入自身,免参数。
为什么这层用它:1平滑非单调(负区间有个小下凹,x → − ∞ x \to -\infty x → − ∞ 时输出趋于 0 而非 ReLU 的硬截断),梯度处处非零,深网训练稳;2门控形式与整个 block 的「信息流过 vs 被抑制」语义一致–Q/K/V 进递归状态前先过一道自门控,相当于对局部卷积混合后的特征做一次软筛选,再交给 L2Norm 钉上单位球。K3 把输出侧的这门升级为全秩输入相关门(满秩门控一节),源头就是这里的 SiLU。
旁注:L2Norm(逐头 L2 归一化,Q/K 专用)
流程图最后一框。对单头向量 z ∈ R d k \bm{z} \in \mathbb{R}^{d_k} z ∈ R d k ,逐通道 L2 标准化到单位范数:
L 2 N o r m ( z ) = z ∥ z ∥ 2 + ϵ , ∥ z ∥ 2 = ∑ i = 1 d k z i 2 \mathrm{L2Norm}(\bm{z}) = \frac{\bm{z}}{\|\bm{z}\|_2 + \epsilon}, \qquad \|\bm{z}\|_2 = \sqrt{\sum_{i=1}^{d_k} z_i^2}
L2Norm ( z ) = ∥ z ∥ 2 + ϵ z , ∥ z ∥ 2 = i = 1 ∑ d k z i 2
记号
含义
ϵ \epsilon ϵ
极小防除零常数(10 − 6 10^{-6} 1 0 − 6 左右)
操作维度
每个注意力头独立归一化,跨头不共享统计
为什么只有 Q/K 用、V 不用?回到 DeltaNet 一节埋下的伏笔:归一化后 ∥ k t ∥ 2 = 1 \|k_t\|_2 = 1 ∥ k t ∥ 2 = 1 ,于是 delta rule 的擦除矩阵 I − β k k ⊤ I - \beta k k^\top I − β k k ⊤ 的特征值被夹在 [ 1 − β , 1 ] [1-\beta,\ 1] [ 1 − β , 1 ] –写入步长天然稳定,β=1 时是精确保长反射。Q 同样归一化 则是让读出 o = S ⊤ q o = S^\top q o = S ⊤ q 的尺度可控,查询和 key 在同一单位球上内积才有可比性,chunkwise 公式里 Q K ⊤ QK^\top Q K ⊤ 的 score 变成有余弦界的小量。V 不归一化:写入内容的幅值本身携带信息(value 的模长是特征,不是 bug),而且 β t k t v t ⊤ \beta_t k_t v_t^\top β t k t v t ⊤ 的稳定性由 k 侧保证,与 v 的尺度无关。这与 softmax 注意力里的 QK-norm(防 logit 爆炸)动机同源,但线性 RNN 里它 additionally 承担擦除几何 的稳定性–归一化不是锦上添花,是 delta rule 能正常工作的前提。
至此三个旁注凑齐流程图的 Q/K 支路:Linear 投影 -> ShortConv(局部混合) -> SiLU(软筛选) -> L2Norm(单位球),三步全部因果、逐通道、免跨头统计,正好匹配 D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 的逐通道语义。更重要的是三者各管一段数值安全:ShortConv 管局部混合、L2Norm 管幅度稳定、逐通道衰减差 γ i − γ j ≤ 0 \gamma_i - \gamma_j \le 0 γ i − γ j ≤ 0 管指数安全 (下界衰减一节)–全链路数值安全正是 KDA 能在 BF16 上跑起来的主线。
混合架构 (论文的另一贡献):GDN + 滑窗注意力(H1)或 Mamba-2 + GDN + SWA(H2)交错堆叠,取长补短–与 K3 的「3 KDA + 1 MLA」同款思路,GDN 论文可视为这个混合范式的先声。
从 GDN 到 KDA:K3 继承了什么、改了什么
维度
GDN
KDA(Kimi Linear -> K3)
动机
更新式
S t − 1 ( α t ( I − β k k ⊤ ) ) + β v k ⊤ S_{t-1}(\alpha_t(I-\beta kk^\top)) + \beta vk^\top S t − 1 ( α t ( I − β k k ⊤ )) + β v k ⊤
( I − β k k ⊤ ) D i a g ( α t ) S t − 1 + β v k ⊤ (I-\beta kk^\top)\,\mathrm{Diag}(\alpha_t)S_{t-1} + \beta vk^\top ( I − β k k ⊤ ) Diag ( α t ) S t − 1 + β v k ⊤
α 从标量 -> 逐通道向量
遗忘门粒度
标量(整头同一衰减率)
向量 α t ∈ ( 0 , 1 ) d k \alpha_t\in(0,1)^{d_k} α t ∈ ( 0 , 1 ) d k
长期记忆通道与短期工作区分离
门作用顺序
α 在外乘整个更新
Diag(α) 在内侧先衰减 S,再删写
与逐通道参数化配套(KCP 推导需要)
α 参数化
负 Softplus,值域 ( − ∞ , 0 ) (-\infty,0) ( − ∞ , 0 )
K3:缩放 sigmoid,下界 g min = − 5 g_{\min}=-5 g m i n = − 5
1 / Γ 1/\Gamma 1/Γ 有界 -> 对角 tile 全走 Tensor Core
chunkwise
WY + UT + 衰减箭头
同框架,Γ 从标量比变成向量累积比
表达力↑ 数值难度↑(K3 用下界解决)
输出门
低秩/简单门
K3:输入相关全秩 sigmoid 门
逐通道调节读出
两处值得注意的「暗改」:
作用顺序变了 :GDN 是 S ( α ( I − β k k ⊤ ) ) S(\alpha(I-\beta kk^\top)) S ( α ( I − β k k ⊤ )) (α 吸进 Householder 连乘),KDA 是 ( I − β k k ⊤ ) D i a g ( α ) S (I-\beta kk^\top)\mathrm{Diag}(\alpha)S ( I − β k k ⊤ ) Diag ( α ) S (α 先作用、再删写)。KDA 的逐通道门是矩阵 D i a g ( α ) \mathrm{Diag}(\alpha) Diag ( α ) ,与 Householder 不交换,顺序成为实质设计选择–它使 KCP 的「段转移分解」成为可能。
衰减在 Γ 处的复杂化 :GDN 的 γ 是标量连乘,Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j 只是数;KDA 的 Γ 是向量 逐元素累积,chunk 公式里 K/Γ、Q⊙Γ 等运算随之复杂化,数值范围问题(1 / Γ 1/\Gamma 1/Γ 爆炸)由此而生–K3 的下界衰减正是对 GDN->KDA 这一步引入的新问题的修复。
SSM 部分速览,详细推导见前置篇
KDA 递推式的 SSM 部分来自一条完整的推导链:连续状态方程 h ′ ( t ) = A h ( t ) + B x ( t ) h'(t) = Ah(t) + Bx(t) h ′ ( t ) = A h ( t ) + B x ( t ) ,经积分因子法得通解,ZOH 离散化得递推 h t = A ˉ h t − 1 + B ˉ x t h_t = \bar A h_{t-1} + \bar B x_t h t = A ˉ h t − 1 + B ˉ x t ,其中 A ˉ = e Δ A \bar A = e^{\Delta A} A ˉ = e Δ A 、B ˉ = A − 1 ( e Δ A − I ) B \bar B = A^{-1}(e^{\Delta A} - I)B B ˉ = A − 1 ( e Δ A − I ) B ;递归展开等价于卷积(S4 双形式);Mamba 让 Δ t \Delta_t Δ t 依赖输入,A ˉ t = e − Δ t a ∈ ( 0 , 1 ) \bar A_t = e^{-\Delta_t a} \in (0,1) A ˉ t = e − Δ t a ∈ ( 0 , 1 ) 变成看内容的遗忘门 ;Mamba-2 把状态写成外积形式 S t = α t S t − 1 + v t k t ⊤ S_t = \alpha_t S_{t-1} + v_tk_t^\top S t = α t S t − 1 + v t k t ⊤ ,和线性注意力在数学上就是一回事了(SSD 对偶)。每一步的严格推导、数值验算(ZOH 精确性 1.2642 = 1.2642)和 S4/Mamba/Mamba-2 三站的演化,见前置篇《线性注意力与 SSM:两条技术路线的完整推导》 。
这里只需要记住结论:SSM 这条路线最终给出的,是乘在旧状态上的标量衰减系数 α t \alpha_t α t –它回答了线性注意力一族「状态该怎么遗忘」的问题,但答案是整片擦除、不分通道。它和 DeltaNet(定点覆写)在 GDN 处碰头;GDN 又被 KDA 逐通道化。骨架始终是同一条:
状态 × ( 衰减/删除算子 ) + ( 写入项 ) \text{状态} \times (\text{衰减/删除算子}) + (\text{写入项})
状态 × ( 衰减 / 删除算子 ) + ( 写入项 )
递推公式(论文 2.1.1 节,核心式 1)
S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤ \mathbf{S}_t = \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) \operatorname{Diag}(\bm{\alpha}_t)\, \mathbf{S}_{t-1} + \beta_t \bm{k}_t \bm{v}_t^\top
S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤
o ~ t = S t ⊤ q t \tilde{\bm{o}}_t = \mathbf{S}_t^\top \bm{q}_t
o ~ t = S t ⊤ q t
符号逐个过一遍 (论文原定义):
符号
含义
α t ∈ R d k \bm{\alpha}_t \in \mathbb{R}^{d_k} α t ∈ R d k
逐通道一维保留因子向量(channel-wise one-step retention factor)
Diag ( α t ) \operatorname{Diag}(\bm{\alpha}_t) Diag ( α t )
向量转对角矩阵算子,把通道级衰减系数变成矩阵乘
I \mathbf{I} I
同维度单位矩阵
β t ∈ ( 0 , 1 ) \beta_t \in (0,1) β t ∈ ( 0 , 1 )
delta rule 写入强度
读这个式子的三层动作 (从右往左):
通道衰减 :Diag ( α t ) S t − 1 \operatorname{Diag}(\bm{\alpha}_t)\mathbf{S}_{t-1} Diag ( α t ) S t − 1 等价于对 S t − 1 \mathbf{S}_{t-1} S t − 1 的每一行分别乘以对应通道的 α \alpha α 系数–逐通道缩放历史记忆。普通线性注意力/SSM 大多全局统一衰减(α t \alpha_t α t 是标量),KDA 给每一个 key 通道分配独立衰减权重:不同语义通道可以选择更快/更慢遗忘 S t − 1 \mathbf{S}_{t-1} S t − 1 ;
方向擦除 :( I − β t k t k t ⊤ ) \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) ( I − β t k t k t ⊤ ) 再对衰减后的状态做当前 token 的定点擦除(秩 1 Householder 型,详见 DeltaNet 一节);
新信息写入 :叠加 β t k t v t ⊤ \beta_t \bm{k}_t \bm{v}_t^\top β t k t v t ⊤ 。
三步合起来:先通道衰减,再方向擦除,最后写入 –带通道精细遗忘的 delta 递推。作用顺序是设计选择:Diag 与 Householder 不交换,这个顺序使 KCP 的段转移分解成为可能(见下界衰减与基础设施篇)。
记号约定 :论文里大写 Diag ( ⋅ ) \operatorname{Diag}(\cdot) Diag ( ⋅ ) 是「向量 -> 对角矩阵」算子;小写 diag ( ⋅ ) \operatorname{diag}(\cdot) diag ( ⋅ ) 有时指反向操作(输入矩阵、提取对角线为向量)。本文全程大写表示向量转对角矩阵。
为什么是逐通道门:动机与在线学习视角
标量门的表达力瓶颈 。GDN 的 α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 ) 是标量:每一步遗忘时,所有 key 通道以同一个比例 衰减–要么一起记住,要么一起忘记。但不同通道承担的角色不同:
有的通道在存「长期主题」(希望 α ≈ 1 \alpha \approx 1 α ≈ 1 ,几乎不遗忘);
有的通道在存「临时指针」(希望快速衰减,腾出容量)。
KDA 的核心改动只有一处:把标量换成向量 α t ∈ ( 0 , 1 ) d k \bm{\alpha}_t \in (0,1)^{d_k} α t ∈ ( 0 , 1 ) d k (β t \beta_t β t 仍是标量,这是 KDA 的选择而非必须)。直觉:S 的第 j 列对应 key 空间的第 j 个通道,D i a g ( α t ) \mathrm{Diag}(\alpha_t) Diag ( α t ) 作用上去就是给每一列配一个独立的遗忘速度 。GDN 是 KDA 在 D i a g ( α t ) = α t I \mathrm{Diag}(\alpha_t) = \alpha_t I Diag ( α t ) = α t I 时的特例。
在线学习视角:逐通道权重衰减的 delta rule 。与 GDN 一节的在线学习表同构,只需把正则项换成逐通道版。每步给定新样本 ( k t , v t ) (k_t, v_t) ( k t , v t ) ,希望新状态 S:1拟合新样本 S k t ≈ v t Sk_t \approx v_t S k t ≈ v t ;2不偏离「衰减后的旧记忆」太远 S ≈ S t − 1 D i a g ( α t ) S \approx S_{t-1}\mathrm{Diag}(\alpha_t) S ≈ S t − 1 Diag ( α t ) (以下记 D t = D i a g ( α t ) D_t = \mathrm{Diag}(\alpha_t) D t = Diag ( α t ) ,α t = exp ( g t ) \alpha_t = \exp(g_t) α t = exp ( g t ) ,g t ∈ R < 0 d k g_t \in \mathbb{R}_{<0}^{d_k} g t ∈ R < 0 d k ):
L t ( S ) = 1 2 ∥ S k t − v t ∥ 2 ⏟ 拟合 + 1 2 ∥ S − S t − 1 D t ∥ F 2 ⏟ 逐通道正则 \mathcal{L}_t(S) = \underbrace{\tfrac12 \|S k_t - v_t\|^2}_{\text{拟合}} + \underbrace{\tfrac12 \|S - S_{t-1} D_t\|_F^2}_{\text{逐通道正则}}
L t ( S ) = 拟合 2 1 ∥ S k t − v t ∥ 2 + 逐通道正则 2 1 ∥ S − S t − 1 D t ∥ F 2
从 S t − 1 D t S_{t-1}D_t S t − 1 D t 出发对第一项做一步梯度下降(步长 β t \beta_t β t ):
S t = S t − 1 D t − β t ( S t − 1 D t k t − v t ) k t ⊤ = S t − 1 D t ( I − β t k t k t ⊤ ) + β t v t k t ⊤ S_t = S_{t-1}D_t - \beta_t\,(S_{t-1}D_t\, k_t - v_t)\,k_t^\top = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top
S t = S t − 1 D t − β t ( S t − 1 D t k t − v t ) k t ⊤ = S t − 1 D t ( I − β t k t k t ⊤ ) + β t v t k t ⊤
正是递推式(转置约定下)。逐通道正则的含义 :第 j 列的「信任区域」宽度正比于 α t ( j ) \alpha_t^{(j)} α t ( j ) –α \alpha α 小的通道,旧记忆先被压缩,新写入覆盖几乎没有阻力(快遗忘);α ≈ 1 \alpha \approx 1 α ≈ 1 的通道,旧记忆原样进入下一步,delta rule 只做精细增量(慢遗忘)。这就是「细粒度记忆控制」的优化论表述:KDA = 逐通道权重衰减 + delta rule。
实现注意:作用顺序是坑 。D t D_t D t 作用于整个 S t − 1 S_{t-1} S t − 1 、先于擦除项 –擦除时检索用的也是已衰减 的状态:
S t = S t − 1 D t ⏟ 先衰减 ( I − β t k t k t ⊤ ) ⏟ 再擦除/写入 S_t = \underbrace{S_{t-1} D_t}_{\text{先衰减}}\underbrace{(I - \beta_t k_t k_t^\top)}_{\text{再擦除/写入}}
S t = 先衰减 S t − 1 D t 再擦除 / 写入 ( I − β t k t k t ⊤ )
写成 S t − 1 ( I − β k k ⊤ ) D t S_{t-1}(I-\beta kk^\top)D_t S t − 1 ( I − β k k ⊤ ) D t (顺序颠倒)或只衰减单位阵部分都是错的–D t D_t D t 与 ( I − β k k ⊤ ) (I - \beta kk^\top) ( I − β k k ⊤ ) 不可交换 ,顺序错了结果就错(与 GDN 一节「α 乘整个括号」的警告同源,但逐通道化之后错误更隐蔽)。FLA 参考实现里对应 S = S * g.exp() 之后立刻用衰减后的 S 做检索 v - k^T S。
下界衰减(Lower-bounded decay)
KDA 的衰减参数化是一个关键改进。Kimi Linear 使用无界的负 Softplus 映射 g = − e A Softplus ( z ) ∈ ( − ∞ , 0 ) g = -e^A \text{Softplus}(z) \in (-\infty, 0) g = − e A Softplus ( z ) ∈ ( − ∞ , 0 ) ,而 K3 改用有界缩放 sigmoid :
g t h = g min ⋅ Sigmoid ( e A h z t h ) ∈ ( g min , 0 ) g_t^h = g_{\min} \cdot \text{Sigmoid}(e^{A_h} z_t^h) \in (g_{\min}, 0)
g t h = g m i n ⋅ Sigmoid ( e A h z t h ) ∈ ( g m i n , 0 )
其中 g min = − 5 g_{\min} = -5 g m i n = − 5 固定,A h A_h A h 是可学习的每头对数尺度。这意味着每个保留因子满足 α > e − 5 ≈ 6.7 × 10 − 3 \alpha > e^{-5} \approx 6.7 \times 10^{-3} α > e − 5 ≈ 6.7 × 1 0 − 3 ,16 token 块的累积对数衰减落在 ( − 80 , 0 ) (-80, 0) ( − 80 , 0 ) 内,对应的重缩放因子小于 e 80 e^{80} e 80 ,在 BF16 动态范围内。
计算收益 :有限范围使得因果对角块和离对角块都能用密集 Tensor Core 矩阵乘法,消除了 Kimi Linear 中需要的 position-pair 对角计算路径。
问题从哪来:1 / Γ 1/\Gamma 1/Γ 爆炸是向量门控自带的代价
为什么这个修复在 GDN 上不必要、在 KDA 上变成刚需?GDN 的标量衰减在 chunkwise 里只以比值 出现(∏ s = j + 1 i α s ≤ 1 \prod_{s=j+1}^{i}\alpha_s \le 1 ∏ s = j + 1 i α s ≤ 1 ,i ≥ j i \ge j i ≥ j ),天然安全。KDA 把 Γ \Gamma Γ 变成向量累积 Γ t = exp ( ∑ s ≤ t g s ) \Gamma_t = \exp(\sum_{s\le t} g_s) Γ t = exp ( ∑ s ≤ t g s ) ,逐通道独立:第 4 步式的推导可以刻意只用 i ≥ j i \ge j i ≥ j 的差(指数 ≤ 0 \le 0 ≤ 0 ,安全);但只要换一种等价写法–把状态「反归一化」回 chunk 起点、或把衰减从 key 上整体外提(论文公式 (4) 的 K / Γ K/\Gamma K /Γ 因式分解正是这种写法)–就得真的算出
Γ t − 1 = exp ( − ∑ s ≤ t g s ) (逐通道) \Gamma_t^{-1} = \exp\Big(-\sum_{s \le t} g_s\Big) \quad \text{(逐通道)}
Γ t − 1 = exp ( − s ≤ t ∑ g s ) ( 逐通道 )
衰减越快的通道,− γ t -\gamma_t − γ t 越大,1 / Γ 1/\Gamma 1/Γ 指数膨胀。量级感受(实测):
恒定 α \alpha α
t = 100 t=100 t = 100
t = 500 t=500 t = 500
t = 1000 t=1000 t = 1000
t = 2000 t=2000 t = 2000
0.9
3.8 × 10 4 3.8\times10^{4} 3.8 × 1 0 4
7.6 × 10 22 7.6\times10^{22} 7.6 × 1 0 22
5.7 × 10 45 5.7\times10^{45} 5.7 × 1 0 45
3.3 × 10 91 3.3\times10^{91} 3.3 × 1 0 91
0.5
1.3 × 10 30 1.3\times10^{30} 1.3 × 1 0 30
3.3 × 10 150 3.3\times10^{150} 3.3 × 1 0 150
1.1 × 10 301 1.1\times10^{301} 1.1 × 1 0 301
溢出 float64
0.1
10 100 10^{100} 1 0 100
溢出
溢出
溢出
float64 上界约 e 709 ≈ 1.8 × 10 308 e^{709} \approx 1.8\times10^{308} e 709 ≈ 1.8 × 1 0 308 ;BF16 训练下几十步就出事。标量 -> 向量的推广放大了衰减率的动态范围 ,爆炸从例外变成常态–所以必须给衰减率本身夹界。
Kimi Linear 的绕行方案(补丁,不是修复) :在对数空间算相对衰减(减法代替除法,不溢出),并把每个 chunk 再切成 16 token 的二级瓦片。效果:瓦片之间 (非对角块)可以安全交给 Tensor Core 稠密矩阵乘;但瓦片内部 (对角块)衰减可能极端,仍需逐位置对显式计算–position-pair 路径无法组织成大矩阵乘,吃不满 Tensor Core,成为块内主要瓶颈。
K3 的根治 :不改公式,改参数化,让爆炸在数学上不可能发生(负 Softplus 允许 α \alpha α 任意接近 0,即「一步清零」;缩放 sigmoid 不允许)。
表达力为什么不受损 :快通道在瓦片内照样能把记忆衰减到 e − 80 ≈ 10 − 35 e^{-80} \approx 10^{-35} e − 80 ≈ 1 0 − 35 (约等于零,只是不真的为零);长期遗忘靠多步连乘(每步 × 0.0067 \times 0.0067 × 0.0067 ,几十步后照样忘干净),不需要单步清零的能力。用一个下界换来全 Tensor Core 化,划算的买卖。
满秩门控(Full-rank gate)
K3 将 KDA 的输出门从低秩参数化改为输入相关的满秩投影 。在递推输出经过 head-wise RMSNorm 后,应用数据相关的输出门控:
y t = W o [ Sigmoid ( W g x t ) ⊙ RMSNorm ( o ~ t ) ] y_t = W_o [\text{Sigmoid}(W_g x_t) \odot \text{RMSNorm}(\tilde{o}_t)]
y t = W o [ Sigmoid ( W g x t ) ⊙ RMSNorm ( o ~ t )]
满秩门控允许每个 token 独立调制从循环状态读取的通道。
Chunkwise 并行形式
递推形式在推理(decode)时是优势–O ( 1 ) O(1) O ( 1 ) 状态更新;但在训练/prefill 时是灾难:每个 token 的状态依赖前一个 token,顺序循环让 GPU 的数千个核心集体围观一个 for 循环。Chunkwise 并行化 把序列切成长度 C 的块:块内矩阵运算并行,块间只传状态。先看通用的复杂度框架,再看 KDA 逐通道门带来的推导细节。
通用框架 。Chunk 内(intra-chunk) :token 交互用带衰减的因果注意力直接算,O ( C 2 d ) O(C^2 d) O ( C 2 d ) –C 是常数(64/128),对序列长度 N 线性。Chunk 间(inter-chunk) :每个 chunk 对状态做一次递推更新,chunk 内外积先归约成固定大小的 state 增量,块间只传 d × d d \times d d × d 的 state。总计算量:
O ( N ⋅ C ⋅ d ) ⏟ chunk 内注意力 + O ( N / C ⋅ d 2 ) ⏟ chunk 间状态更新 = 2 N d 2 ( 固定项 ) + 2 N C d \underbrace{O\big(N \cdot C \cdot d\big)}_{\text{chunk 内注意力}} + \underbrace{O\big(N/C \cdot d^2\big)}_{\text{chunk 间状态更新}} \;=\; 2Nd^2(\text{固定项}) + 2NCd
chunk 内注意力 O ( N ⋅ C ⋅ d ) + chunk 间状态更新 O ( N / C ⋅ d 2 ) = 2 N d 2 ( 固定项 ) + 2 N C d
C
形态
复杂度
C = 1 C = 1 C = 1
每个 token 一个 chunk,chunk 内注意力消失
纯线性注意力递推,FLOPs 最少但不一定最快(GPU 对小矩阵乘利用率低)
C = N C = N C = N
整个序列一个 chunk,递推消失
标准 O ( N 2 ) O(N^2) O ( N 2 ) 注意力
实践中 C 取 64/128:小到让 chunk 间项不贵,大到让 C × C C\times C C × C 注意力矩阵铺满 Tensor Core 的 tile。这与 S4 时代「训练用卷积、推理用递归」双形式一脉相承:选择性打破 LTI 后,chunkwise 就是卷积的继任者。
KDA 版推导:逐通道衰减下的 WY / UT
记 chunk 起点传入状态 S [ 0 ] S_{[0]} S [ 0 ] (以下用局部下标 i = 1.. C i = 1..C i = 1.. C ;乘积按时间倒序)。核心记号是累积 log 衰减 :
γ i = ∑ s = 1 i g s ∈ R d k , Γ i ← j = diag ( exp ( γ i − γ j ) ) ( i ≥ j ) \gamma_i = \sum_{s=1}^{i} g_s \in \mathbb{R}^{d_k}, \qquad \Gamma_{i \leftarrow j} = \operatorname{diag}\big(\exp(\gamma_i - \gamma_j)\big) \quad (i \ge j)
γ i = s = 1 ∑ i g s ∈ R d k , Γ i ← j = diag ( exp ( γ i − γ j ) ) ( i ≥ j )
Γ i ← j \Gamma_{i\leftarrow j} Γ i ← j 是「从第 j 步衰减到第 i 步」的逐通道算子。两条性质:i ≥ j i \ge j i ≥ j 时 γ i − γ j ≤ 0 \gamma_i - \gamma_j \le 0 γ i − γ j ≤ 0 逐分量成立(g s < 0 g_s < 0 g s < 0 ),元素都在 ( 0 , 1 ] (0,1] ( 0 , 1 ] ,数值安全;反向 Γ j ← i − 1 \Gamma_{j\leftarrow i}^{-1} Γ j ← i − 1 的元素 ≥ 1 \ge 1 ≥ 1 且随 chunk 长度指数增长 –这是 1 / Γ 1/\Gamma 1/Γ 爆炸的根源,先记住这个观察。
第一步:衰减 KKT 矩阵 M 。类比 GDN chunkwise 里的普通 k i ⊤ k j k_i^\top k_j k i ⊤ k j ,逐通道版需要带衰减的 key-key 内积:
M c i = ( k c ⊙ e γ c − γ i ) ⊤ k i , 1 ≤ i < c ≤ C M_{ci} = \big(k_c \odot e^{\gamma_c - \gamma_i}\big)^\top k_i, \qquad 1 \le i < c \le C
M c i = ( k c ⊙ e γ c − γ i ) ⊤ k i , 1 ≤ i < c ≤ C
含义:k i k_i k i 写入的记忆衰减到第 c 步时,与 k c k_c k c 的重叠程度。Hadamard 积 ⊙ \odot ⊙ 作用在 key 通道维–正是逐通道衰减出现的位置。
第二步:UT 变换 。构造严格下三角 L L L :L c i = β i M c i ( c > i ) L_{ci} = \beta_i M_{ci}\ (c > i) L c i = β i M c i ( c > i ) ,求 T = ( I + L ) − 1 T = (I+L)^{-1} T = ( I + L ) − 1 (幂零,有限项截断;实践中不显式求逆,前代法逐行解),再每列乘 β j \beta_j β j 得 A ^ = T diag ( β ) \hat{A} = T\,\operatorname{diag}(\beta) A ^ = T diag ( β ) 。FLA 实现里那个三行循环,数学上就是解 ( I + L ) T = I (I+L)T = I ( I + L ) T = I 。
与 GDN 对照 :GDN 这一步的 M 是普通 k c ⊤ k i k_c^\top k_i k c ⊤ k i (标量衰减被拆成比值吸收进别的项);KDA 里衰减「长」在 M 内部 ,无法外提–这是向量门控带来的结构性变化,KDA chunkwise 推导的核心难点。
第三步:WY 表示 :
W = A ^ ( e γ i ⊙ k i ) i = 1.. C ∈ R C × d k , U = A ^ V ∈ R C × d v , V ~ = U − W S [ 0 ] ⊤ W = \hat{A}\,\big(e^{\gamma_i} \odot k_i\big)_{i=1..C} \in \mathbb{R}^{C \times d_k}, \qquad
U = \hat{A}\,V \in \mathbb{R}^{C \times d_v}, \qquad \tilde{V} = U - W\,S_{[0]}^\top
W = A ^ ( e γ i ⊙ k i ) i = 1.. C ∈ R C × d k , U = A ^ V ∈ R C × d v , V ~ = U − W S [ 0 ] ⊤
伪值 V ~ \tilde V V ~ 到底是什么:一本账 。UT 变换产出 U 和 W,用它们定义 V ~ : = U − W S \tilde V := U - WS V ~ := U − W S 。三个符号各司其职:
符号
形状
角色
S [ 0 ] S_{[0]} S [ 0 ]
d k × d v d_k \times d_v d k × d v
旧账本 :chunk 之前所有 token 写入的记忆,块内计算时固定不变
W W W
C × d k C \times d_k C × d k
翻账本的钥匙 :每行是某位置的衰减 key 修正组合;W S WS W S 即「这个位置能从旧记忆里读到什么」
U U U
C × d v C \times d_v C × d v
想记的新账(已和同伴对过账) :A ^ \hat A A ^ 作用在 V 上,块内写入重叠已扣除
V ~ \tilde V V ~
C × d v C \times d_v C × d v
真正上账的净额 = 新账 - (旧记忆已有的 + 同伴已写的)
为什么叫「伪」值:它们不是真实的 v(手算例子里 v ~ 2 = ( − 2 , 2 ) ≠ v 2 \tilde v_2 = (-2,2) \ne v_2 v ~ 2 = ( − 2 , 2 ) = v 2 ),而是预扣完所有重叠后可以不做任何修正直接累加 的修正值。这正是 delta rule 的灵魂的单步版:单步写入 v t − S t − 1 ⊤ k t v_t - S_{t-1}^\top k_t v t − S t − 1 ⊤ k t = 目标值 - 旧记忆读数;V ~ \tilde V V ~ 是它的 chunk 并行版,每行回答「这个位置真正新增 了多少」。账在 V ~ \tilde V V ~ 里算清了,后面的块内注意力 A V ~ A\tilde V A V ~ 和状态更新才敢放心一锤子矩阵乘。
第四步:输出与跨 chunk 状态 :
A c j q k = ( q c ⊙ e γ c − γ j ) ⊤ k j ( j ≤ c ) , o c = S [ 0 ] ( q c ⊙ e γ c ) ⏟ 块间:query 打折查旧状态 + ∑ j ≤ c A c j q k v ~ j ⏟ 块内:衰减注意力 A^{qk}_{cj} = \big(q_c \odot e^{\gamma_c - \gamma_j}\big)^\top k_j \ (j \le c), \qquad o_c = \underbrace{S_{[0]}\,(q_c \odot e^{\gamma_c})}_{\text{块间:query 打折查旧状态}} + \underbrace{\sum_{j \le c} A^{qk}_{cj}\,\tilde v_j}_{\text{块内:衰减注意力}}
A c j q k = ( q c ⊙ e γ c − γ j ) ⊤ k j ( j ≤ c ) , o c = 块间 :query 打折查旧状态 S [ 0 ] ( q c ⊙ e γ c ) + 块内 : 衰减注意力 j ≤ c ∑ A c j q k v ~ j
S [ C ] = S [ 0 ] Γ C ← 0 + ∑ c = 1 C v ~ c ( k c ⊙ e γ C − γ c ) ⊤ S_{[C]} = S_{[0]}\,\Gamma_{C\leftarrow 0} + \sum_{c=1}^{C} \tilde v_c \big(k_c \odot e^{\gamma_C - \gamma_c}\big)^\top
S [ C ] = S [ 0 ] Γ C ← 0 + c = 1 ∑ C v ~ c ( k c ⊙ e γ C − γ c ) ⊤
出口状态读法:旧状态整体按整 chunk 累积衰减缩小;每个伪值以「写入时刻衰减到块末」的 key 为地址写入。所有指数都是 ≤ 0 \le 0 ≤ 0 的差,整条链路没有一个 ≥ 1 \ge 1 ≥ 1 的因子 –刻意保持,原因见下界衰减一节。
手算一遍(d k = d v = 2 d_k = d_v = 2 d k = d v = 2 ,C = 2)
零初始状态,每步恒定衰减 α = ( 0.5 , 0.25 ) \alpha = (0.5, 0.25) α = ( 0.5 , 0.25 ) (即 g = ( ln 0.5 , ln 0.25 ) g = (\ln 0.5, \ln 0.25) g = ( ln 0.5 , ln 0.25 ) ),β 取 1:
k 1 = ( 1 1 ) , v 1 = ( 2 0 ) ; k 2 = ( 1 2 ) , v 2 = ( 0 2 ) ; q 1 = q 2 = ( 1 1 ) k_1 = \binom{1}{1},\ v_1 = \binom{2}{0};\quad k_2 = \binom{1}{2},\ v_2 = \binom{0}{2};\quad q_1 = q_2 = \binom{1}{1}
k 1 = ( 1 1 ) , v 1 = ( 0 2 ) ; k 2 = ( 2 1 ) , v 2 = ( 2 0 ) ; q 1 = q 2 = ( 1 1 )
递推式 。第 1 步(S 0 = 0 S_0 = 0 S 0 = 0 ):S 1 = v 1 k 1 ⊤ = ( 2 2 0 0 ) S_1 = v_1k_1^\top = \begin{pmatrix}2&2\\0&0\end{pmatrix} S 1 = v 1 k 1 ⊤ = ( 2 0 2 0 ) ,o 1 = ( 4 , 0 ) o_1 = (4,0) o 1 = ( 4 , 0 ) 。第 2 步,D = d i a g ( 0.5 , 0.25 ) D = \mathrm{diag}(0.5, 0.25) D = diag ( 0.5 , 0.25 ) ,先衰减 S 1 D = ( 1 0.5 0 0 ) S_1D = \begin{pmatrix}1&0.5\\0&0\end{pmatrix} S 1 D = ( 1 0 0.5 0 ) (第 1 列 ×0.5、第 2 列 ×0.25–逐通道在动 ),再擦除写入:
S 1 D ( I − k 2 k 2 ⊤ ) = ( − 1 − 3.5 0 0 ) , S 2 = ( − 1 − 3.5 2 4 ) , o 2 = ( − 4.5 , 6 ) S_1D(I - k_2k_2^\top) = \begin{pmatrix}-1&-3.5\\0&0\end{pmatrix}, \qquad S_2 = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix}, \qquad o_2 = (-4.5,\ 6)
S 1 D ( I − k 2 k 2 ⊤ ) = ( − 1 0 − 3.5 0 ) , S 2 = ( − 1 2 − 3.5 4 ) , o 2 = ( − 4.5 , 6 )
Chunkwise 。e γ 1 = ( 0.5 , 0.25 ) e^{\gamma_1} = (0.5, 0.25) e γ 1 = ( 0.5 , 0.25 ) ,e γ 2 = ( 0.25 , 0.0625 ) e^{\gamma_2} = (0.25, 0.0625) e γ 2 = ( 0.25 , 0.0625 ) ,e γ 2 − γ 1 = ( 0.5 , 0.25 ) e^{\gamma_2-\gamma_1} = (0.5, 0.25) e γ 2 − γ 1 = ( 0.5 , 0.25 ) 。衰减 KKT:M 21 = ( k 2 ⊙ e γ 2 − γ 1 ) ⊤ k 1 = ( 0.5 , 0.5 ) ⋅ ( 1 , 1 ) = 1 M_{21} = (k_2 \odot e^{\gamma_2-\gamma_1})^\top k_1 = (0.5, 0.5)\cdot(1,1) = 1 M 21 = ( k 2 ⊙ e γ 2 − γ 1 ) ⊤ k 1 = ( 0.5 , 0.5 ) ⋅ ( 1 , 1 ) = 1 。UT:L = ( 0 0 1 0 ) L = \begin{pmatrix}0&0\\1&0\end{pmatrix} L = ( 0 1 0 0 ) ,T = I − L = ( 1 0 − 1 1 ) T = I - L = \begin{pmatrix}1&0\\-1&1\end{pmatrix} T = I − L = ( 1 − 1 0 1 ) ,A ^ = T \hat A = T A ^ = T 。WY:
W = A ^ ( 0.5 0.25 0.25 0.125 ) = ( 0.5 0.25 − 0.25 − 0.125 ) , U = A ^ ( 2 0 0 2 ) = ( 2 0 − 2 2 ) W = \hat A\begin{pmatrix}0.5&0.25\\0.25&0.125\end{pmatrix} = \begin{pmatrix}0.5&0.25\\-0.25&-0.125\end{pmatrix}, \quad U = \hat A\begin{pmatrix}2&0\\0&2\end{pmatrix} = \begin{pmatrix}2&0\\-2&2\end{pmatrix}
W = A ^ ( 0.5 0.25 0.25 0.125 ) = ( 0.5 − 0.25 0.25 − 0.125 ) , U = A ^ ( 2 0 0 2 ) = ( 2 − 2 0 2 )
伪值(S [ 0 ] = 0 ⇒ V ~ = U S_{[0]}=0 \Rightarrow \tilde V = U S [ 0 ] = 0 ⇒ V ~ = U ):v ~ 1 = ( 2 , 0 ) \tilde v_1 = (2,0) v ~ 1 = ( 2 , 0 ) ,v ~ 2 = ( − 2 , 2 ) \tilde v_2 = (-2,2) v ~ 2 = ( − 2 , 2 ) 。注意 v ~ 2 ≠ v 2 \tilde v_2 \ne v_2 v ~ 2 = v 2 :因为 k 2 k_2 k 2 与衰减后的 k 1 k_1 k 1 写入重叠(M 21 = 1 ≠ 0 M_{21} = 1 \ne 0 M 21 = 1 = 0 ),WY 把 v 2 v_2 v 2 修正为扣除重叠后真正的新增–单步 delta rule 的 chunk 版样子。衰减注意力与输出:
A q k = ( 2 0 0.75 3 ) , o 1 = 2 v ~ 1 = ( 4 , 0 ) ✓ , o 2 = 0.75 v ~ 1 + 3 v ~ 2 = ( − 4.5 , 6 ) ✓ A^{qk} = \begin{pmatrix}2&0\\0.75&3\end{pmatrix}, \qquad o_1 = 2\tilde v_1 = (4,0)\ \checkmark, \qquad o_2 = 0.75\,\tilde v_1 + 3\,\tilde v_2 = (-4.5,\ 6)\ \checkmark
A q k = ( 2 0.75 0 3 ) , o 1 = 2 v ~ 1 = ( 4 , 0 ) ✓ , o 2 = 0.75 v ~ 1 + 3 v ~ 2 = ( − 4.5 , 6 ) ✓
块末状态:S [ 2 ] = v ~ 1 ( k 1 ⊙ e γ 2 − γ 1 ) ⊤ + v ~ 2 k 2 ⊤ = ( − 1 − 3.5 2 4 ) S_{[2]} = \tilde v_1(k_1 \odot e^{\gamma_2-\gamma_1})^\top + \tilde v_2 k_2^\top = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix} S [ 2 ] = v ~ 1 ( k 1 ⊙ e γ 2 − γ 1 ) ⊤ + v ~ 2 k 2 ⊤ = ( − 1 2 − 3.5 4 ) ✓ \checkmark ✓ 与递推式逐项一致。(随机对拍:非零初始状态、随机门控下递推 vs chunkwise 最大误差 10 − 9 10^{-9} 1 0 − 9 量级。)
与论文公式 (4) 的对照:K/Γ 因式分解
K3 报告(沿用 Kimi Linear)用乘积记号:γ i → j = ∏ r = i j α r = exp ( ∑ r = i j g r ) \gamma_{i\to j} = \prod_{r=i}^{j}\alpha_r = \exp(\sum_{r=i}^j g_r) γ i → j = ∏ r = i j α r = exp ( ∑ r = i j g r ) ;Γ ∈ R C × d k \Gamma \in \mathbb{R}^{C\times d_k} Γ ∈ R C × d k 把各步 γ \gamma γ 按行堆叠(注意 Γ \Gamma Γ 本身不是对角阵 ,每个位置 r r r 的 d i a g ( γ r ) \mathrm{diag}(\gamma_r) diag ( γ r ) 才是遗忘对角阵,Γ \Gamma Γ 是 C 个对角阵的打包)。论文状态约定 S ∈ R d k × d v S \in \mathbb{R}^{d_k \times d_v} S ∈ R d k × d v (本文的转置),其公式 (4):
A = T r i l [ ( Q ⊙ Γ ) ( K / Γ ) ⊤ ] , O = ( Γ ⊙ Q ) S ⏟ 块间 + A V ~ ⏟ 块内 A = \mathrm{Tril}\big[(Q\odot\Gamma)(K/\Gamma)^\top\big], \qquad O = \underbrace{(\Gamma\odot Q)S}_{\text{块间}} + \underbrace{A\,\tilde V}_{\text{块内}}
A = Tril [ ( Q ⊙ Γ ) ( K /Γ ) ⊤ ] , O = 块间 ( Γ ⊙ Q ) S + 块内 A V ~
为什么 A 能这样拆 :看 ( i , j ) (i,j) ( i , j ) 元素 A i j = ∑ d q i , d γ i , d ⋅ k j , d / γ j , d = ( q i ⊙ γ j → i ) ⊤ k j A_{ij} = \sum_d q_{i,d}\gamma_{i,d}\cdot k_{j,d}/\gamma_{j,d} = (q_i \odot \gamma_{j\to i})^\top k_j A ij = ∑ d q i , d γ i , d ⋅ k j , d / γ j , d = ( q i ⊙ γ j → i ) ⊤ k j ,与 A q k A^{qk} A q k 逐元素相同–逐对位置的衰减比值被因式分解成 query 侧乘 Γ、key 侧除以 Γ 两个逐位置操作,整个 C × C C\times C C × C 矩阵 = 一次稠密 matmul + 两次逐元素乘,完全并行。Tril 保留对角线 :delta rule 里 o i o_i o i 读的是写入当前 token 之后 的状态。
用手算例子核对 :Q ⊙ Γ = ( 0.5 0.25 0.25 0.0625 ) Q\odot\Gamma = \begin{pmatrix}0.5&0.25\\0.25&0.0625\end{pmatrix} Q ⊙ Γ = ( 0.5 0.25 0.25 0.0625 ) ,K / Γ = ( 2 4 4 32 ) K/\Gamma = \begin{pmatrix}2&4\\4&32\end{pmatrix} K /Γ = ( 2 4 4 32 ) ,乘积 Tril 后 A = ( 2 0 0.75 3 ) A = \begin{pmatrix}2&0\\0.75&3\end{pmatrix} A = ( 2 0.75 0 3 ) ,与上面 A q k A^{qk} A q k 一致。注意 K / Γ K/\Gamma K /Γ 里已出现 32 这种被放大的数–1 / Γ 1/\Gamma 1/Γ 的膨胀就藏在这一步 ,数值后果见下界衰减一节。
常见疑问:C C C 与 d k d_k d k 有倍数关系吗? 没有。C 是序列轴的切分(一个 chunk 装多少 token),d k d_k d k 是特征轴的宽度,两根轴独立。有整除要求的是:K3 把 chunk 再切成 16 token 二级瓦片,故 C 需是 16 的倍数;d k d_k d k 、d v d_v d v 按 Tensor Core 友好尺寸取是工程约束,不是数学要求。
KDA 的核心优势是:相比传统 softmax attention,它是线性复杂度 的;相比纯粹的线性 RNN,它通过 delta rule 实现了更强大的信息写入和遗忘控制。
一页纸速查(KDA/GDN 全家)
维度
内容
状态
S ∈ R^{d_v×d_k}(论文记法),固定大小,每序列每头一份
更新
S = S(α(I-βkkT)) + βvkT:板擦 α + 铅笔 δ
读出
o = Sq,后接 RMSNorm + sigmoid 输出门
α
GDN 标量 / KDA 逐通道,数据相关;β:sigmoid
q/k
Linear -> ShortConv -> SiLU -> L2Norm
等价视角
在线最小二乘 SGD(β=学习率)+ 自适应权重衰减(α)
并行训练
chunkwise:WY 表示(W 无 γ / Ũ 有 γ)+ UT 变换(下三角求逆)+ 衰减箭头
chunk 输出
O = Q⃖S0T + (QKT⊙Γ因果)(Ũ - W⃖S0T)
chunk 状态
S_C = γ_C S0 + (Ũ-> - W->S0T)TK
记忆瓶颈修复
覆写(δ)治叠加污染;清场(α)治长期堆积
后继
KDA:α->逐通道、换作用顺序、加下界 -> 进 K3
参考 :
Yang, Kautz & Hatamizadeh, Gated Delta Networks: Improving Mamba2 with Delta Rule , arXiv:2412.06464 (GDN 出处,ICLR 2025)
Kimi K3 Technical Report, arXiv:2607.24653 (KDA 出处)
Gu & Dao, Mamba: Linear-Time Sequence Modeling with Selective State Spaces , arXiv:2312.00752
Gu et al., Efficiently Modeling Long Sequences with Structured State Spaces (S4), arXiv:2111.00396
Katharopoulos et al., Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention , arXiv:2006.16236