pydata: Huiming's learning notes

Keep Looking, Don't Settle

KDA - 从线性注意力到 Kimi Delta Attention

本文从标准自注意力出发,依次说明核化线性注意力、循环状态、门控与 Delta 规则,并最终落到 Kimi Linear 的核心组件 Kimi Delta Attention(KDA)。重点不是把所有模型强行归为同一种结构,而是追踪一个共同问题:如何用固定大小的状态高效地读写长序列记忆,同时尽量降低信息干扰。

0. 统一符号与维度

全文采用列向量约定。第 \(t\) 个 token 的输入表示统一记为

$$ \mathbf{x}_t\in\mathbb{R}^{d_{\mathrm{model}}}. $$

对第 \(h\) 个注意力头,投影统一写成

$$ \mathbf{q}_t^{h}=\mathbf{W}_{h}^{Q}\mathbf{x}_t,\qquad \mathbf{k}_t^{h}=\mathbf{W}_{h}^{K}\mathbf{x}_t,\qquad \mathbf{v}_t^{h}=\mathbf{W}_{h}^{V}\mathbf{x}_t, $$

其中

$$ \mathbf{W}_{h}^{Q},\mathbf{W}_{h}^{K}\in \mathbb{R}^{d_k\times d_{\mathrm{model}}},\qquad \mathbf{W}_{h}^{V}\in \mathbb{R}^{d_v\times d_{\mathrm{model}}}. $$

为简化公式,讨论单头时省略上标 \(h\)。其他符号如下:

符号 含义 维度
\(N\) 序列长度 标量
\(H\) 注意力头数 标量
\(\mathbf{Q},\mathbf{K}\) 按行堆叠的 query/key \(N\times d_k\)
\(\mathbf{V},\mathbf{O}\) 按行堆叠的 value/output \(N\times d_v\)
\(\mathbf{S}_t\) 关联记忆状态 \(d_k\times d_v\)
\(\mathbf{z}_t\) 归一化状态 \(d_k\)
\(\alpha_t,\mathbf{\alpha}_t\) 标量/通道级遗忘门 标量或 \(d_k\)
\(\beta_t\) Delta 更新强度 标量

输出投影统一记为 \(\mathbf{W}^{O}\)。矩阵 \(\operatorname{Diag}(\mathbf{\alpha}_t)\) 以向量 \(\mathbf{\alpha}_t\) 为对角线。

1. 从标准注意力到线性注意力

1.1 标准自注意力

单头标准注意力为

$$ \mathbf{O} = \operatorname{softmax}\!\left( \frac{\mathbf{Q}\mathbf{K}^{\top}}{\sqrt{d_k}} +\mathbf{M} \right)\mathbf{V}, $$

其中 \(\mathbf{M}\) 是因果掩码;不需要因果约束时可省略。多头结果拼接后再经过输出投影。按行堆叠整个序列时:

$$ \mathbf{Y} = \operatorname{Concat}\!\left( \mathbf{O}^{1},\ldots,\mathbf{O}^{H} \right) \left(\mathbf{W}^{O}\right)^{\top}. $$

\(\mathbf{Q}\mathbf{K}^{\top}\in\mathbb{R}^{N\times N}\),因此核心计算量为 \(O(N^2d_k)\)。显式实现的注意力矩阵空间开销为 \(O(N^2)\);FlashAttention 等 IO-aware 实现可以避免把整张矩阵持久写入显存,但不会改变全局注意力的二次算术复杂度。自回归推理虽然不必重新计算整张矩阵,但每个新 query 仍要读取全部历史 key/value,单步代价随上下文长度 \(N\) 增长。

1.2 为什么不能直接改变 Softmax 的乘法顺序

如果去掉 Softmax,则矩阵结合律给出

$$ (\mathbf{Q}\mathbf{K}^{\top})\mathbf{V} = \mathbf{Q}(\mathbf{K}^{\top}\mathbf{V}). $$

为什么右侧的复杂度是线性的?设

$$ \mathbf{Q},\mathbf{K}\in\mathbb{R}^{N\times d_k}, \qquad \mathbf{V}\in\mathbb{R}^{N\times d_v}. $$

右侧按两步计算:

  1. 先计算 \(\mathbf{K}^{\top}\mathbf{V}\)。其维度变化为 \((d_k\times N)(N\times d_v)\rightarrow d_k\times d_v\),因此需要 \(O(Nd_kd_v)\) 次运算。
  2. 再计算 \(\mathbf{Q}(\mathbf{K}^{\top}\mathbf{V})\)。其维度变化为 \((N\times d_k)(d_k\times d_v)\rightarrow N\times d_v\),同样需要 \(O(Nd_kd_v)\) 次运算。

两步合计约为 \(2Nd_kd_v\),忽略常数后即

$$ O(Nd_kd_v). $$

\(d_k=d_v=d\) 时,它就是 \(O(Nd^2)\)。关键在于中间结果 \(\mathbf{K}^{\top}\mathbf{V}\) 的大小是 \(d_k\times d_v\),整个过程中没有生成 \(N\times N\) 矩阵。相反,若按左侧顺序先计算 \(\mathbf{Q}\mathbf{K}^{\top}\),就会产生 \(N\times N\) 中间矩阵,复杂度回到关于 \(N\) 的二次量级。

以上结合律推导适用于没有 Softmax 且没有因果掩码阻断的可结合形式。因果线性注意力需要使用前缀状态或并行扫描,不能直接用整段序列的 \(\mathbf{K}^{\top}\mathbf{V}\)

但标准注意力中的 Softmax 是逐行归一化的非线性算子:

$$ \operatorname{softmax}(\mathbf{Q}\mathbf{K}^{\top})\mathbf{V} \neq \mathbf{Q}(\mathbf{K}^{\top}\mathbf{V}). $$

因此,线性注意力不是简单地给原公式“换括号”,而是要把相似度改写为可分解的核。为简化符号,本文将核特征映射吸收到 query/key 的生成过程,映射后的向量仍记为 \(\mathbf{q}_t,\mathbf{k}_t\),不再额外添加函数或波浪号:

$$ \operatorname{sim}(\mathbf{q}_t,\mathbf{k}_j) = \mathbf{q}_t^{\top}\mathbf{k}_j. $$

归一化的核注意力输出为

$$ \mathbf{o}_t = \frac{ \sum_{j=1}^{N} \left(\mathbf{q}_t^{\top}\mathbf{k}_j\right)\mathbf{v}_j }{ \sum_{j=1}^{N} \mathbf{q}_t^{\top}\mathbf{k}_j+\varepsilon }. $$

将只依赖 \(t\) 的 query 移到求和外:

$$ \mathbf{o}_t = \frac{ \left(\sum_{j=1}^{N}\mathbf{k}_j\mathbf{v}_j^{\top}\right)^{\top} \mathbf{q}_t }{ \left(\sum_{j=1}^{N}\mathbf{k}_j\right)^{\top} \mathbf{q}_t+\varepsilon }. $$

这里不再生成 \(N\times N\) 的中间矩阵,而是先形成固定大小的 \(d_k\times d_v\) 统计量。

1.3 复杂度为什么对 \(N\) 是线性的

$$ \mathbf{A}=\mathbf{K}^{\top}\mathbf{V} \in\mathbb{R}^{d_k\times d_v}. $$

矩阵 \(\mathbf{A}\)\(d_kd_v\) 个元素,每个元素需要沿序列维度做长度为 \(N\) 的内积,因此计算量是

$$ O(Nd_kd_v). $$

\(d_k\sim d_v\sim d\) 时,通常简写为 \(O(Nd^2)\)。相比之下,标准注意力先生成 \(N\times N\) 分数矩阵,计算量是 \(O(N^2d)\)

“线性”只表示复杂度关于 \(N\) 线性,并不表示关于隐藏维度也是线性。实际性能还取决于头维度、特征映射、分块算法、内存访问和硬件利用率。

1.4 因果场景需要逐时间步的前缀聚合

上一节的全局聚合

$$ \mathbf{K}^{\top}\mathbf{V} = \sum_{j=1}^{N}\mathbf{k}_j\mathbf{v}_j^{\top} $$

包含整段序列的全部 token。如果时间步 \(t\) 直接使用这个全局状态,求和中就会包含 \(j>t\) 的 key/value,即模型会提前看到未来信息。因果注意力要求时间步 \(t\) 只能使用 \(j\le t\) 的前缀,因此每个时间步需要不同的聚合状态。

先从多头分解形式看这一点。对第 \(h\) 个头,仍采用全文统一的列向量约定:

$$ \mathbf{q}_t^h=\mathbf{W}_h^Q\mathbf{x}_t,\qquad \mathbf{k}_j^h=\mathbf{W}_h^K\mathbf{x}_j,\qquad \mathbf{v}_j^h=\mathbf{W}_h^V\mathbf{x}_j. $$

原稿的“分块/分解形式”可更严谨地改写为下面的未归一化多头线性注意力:

$$ \mathbf{y}_t = \sum_{h=1}^{H} \mathbf{W}_h^O \left[ \sum_{j=1}^{t} \frac{(\mathbf{q}_t^h)^{\top}\mathbf{k}_j^h}{\sqrt{d_k}} \mathbf{v}_j^h \right] \tag{2} $$

其中 \(\mathbf{W}_h^O\in\mathbb{R}^{d_{\mathrm{model}}\times d_v}\) 是总输出投影 \(\mathbf{W}^O\) 中对应第 \(h\) 个头的分块。这里的上限是 \(t\) 而不是 \(N\),正是因果约束。原稿使用的索引 \(c\) 容易与后文的序列分块混淆,因此这里改用头索引 \(h\)。若采用归一化核注意力,还需要将每个头的结果除以对应的权重和;上式只保留与状态递推最相关的未归一化部分。

进一步把同一头中的 Q/K 投影与 V/O 投影分别合并。定义

$$ \mathbf{W}_h^{QK} = (\mathbf{W}_h^Q)^{\top}\mathbf{W}_h^K \in\mathbb{R}^{d_{\mathrm{model}}\times d_{\mathrm{model}}}, \qquad \mathbf{W}_h^{OV} = \mathbf{W}_h^O\mathbf{W}_h^V \in\mathbb{R}^{d_{\mathrm{model}}\times d_{\mathrm{model}}}. $$

由于

$$ (\mathbf{q}_t^h)^{\top}\mathbf{k}_j^h = \mathbf{x}_t^{\top}\mathbf{W}_h^{QK}\mathbf{x}_j, \qquad \mathbf{W}_h^O\mathbf{v}_j^h = \mathbf{W}_h^{OV}\mathbf{x}_j, $$

公式 (2) 可写成下面的合并权重形式:

$$ \mathbf{y}_t = \sum_{h=1}^{H} \sum_{j=1}^{t} \frac{ \mathbf{x}_t^{\top}\mathbf{W}_h^{QK}\mathbf{x}_j }{\sqrt{d_k}} \mathbf{W}_h^{OV}\mathbf{x}_j \tag{3} $$

公式 (3) 主要用于说明代数结构:相关性是关于 \(\mathbf{x}_t\)\(\mathbf{x}_j\) 的双线性形式,value 与输出投影也可以合并。这一步依赖 \(\mathbf{q}_t^h,\mathbf{k}_j^h\) 是上面定义的线性投影;如果 query/key 的生成中还包含不可合并的非线性特征映射,就不能把它们压缩成单个 \(\mathbf{W}_h^{QK}\),但后面的前缀状态递推仍然成立。

实际实现通常仍保留低秩分解后的 \(\mathbf{W}_h^Q,\mathbf{W}_h^K,\mathbf{W}_h^V,\mathbf{W}_h^O\),避免显式构造较大的 \(\mathbf{W}_h^{QK}\)\(\mathbf{W}_h^{OV}\)

对因果计算而言,更关键的改写是为每个头定义前缀状态:

$$ \mathbf{S}_t^h = \sum_{j=1}^{t}\mathbf{k}_j^h(\mathbf{v}_j^h)^{\top}, \qquad \mathbf{o}_t^h = (\mathbf{S}_t^h)^{\top}\mathbf{q}_t^h. $$

这个状态随时间步递推:

$$ \mathbf{S}_t^h = \mathbf{S}_{t-1}^h + \mathbf{k}_t^h(\mathbf{v}_t^h)^{\top}. $$

因此,因果线性注意力仍然具有关于序列长度的线性复杂度,但不能让所有时间步共享同一个全局 \(\mathbf{K}^{\top}\mathbf{V}\)。推理时应逐 token 更新前缀状态;训练时则可以用并行前缀扫描或分块算法同时计算各时间步对应的前缀状态。下一节将在加入归一化项后完整推导这一 RNN 形式。

2. 因果线性注意力的 RNN 形式

2.1 从前缀和到循环状态

因果核注意力为

$$ \mathbf{o}_t = \frac{ \sum_{j=1}^{t} \left(\mathbf{q}_t^{\top}\mathbf{k}_j\right)\mathbf{v}_j }{ \sum_{j=1}^{t} \mathbf{q}_t^{\top}\mathbf{k}_j+\varepsilon }. $$

定义矩阵状态和归一化状态

$$ \mathbf{S}_t = \sum_{j=1}^{t}\mathbf{k}_j\mathbf{v}_j^{\top}, \qquad \mathbf{z}_t = \sum_{j=1}^{t}\mathbf{k}_j. $$

于是

$$ \mathbf{o}_t = \frac{ \mathbf{S}_t^{\top}\mathbf{q}_t }{ \mathbf{z}_t^{\top}\mathbf{q}_t+\varepsilon }. $$

两个状态都可递推:

$$ \begin{aligned} \mathbf{S}_0&=\mathbf{0},& \mathbf{S}_t&=\mathbf{S}_{t-1} +\mathbf{k}_t\mathbf{v}_t^{\top},\\ \mathbf{z}_0&=\mathbf{0},& \mathbf{z}_t&=\mathbf{z}_{t-1} +\mathbf{k}_t. \end{aligned} $$

若加入残差与层变换,可写为

$$ \mathbf{y}_t = f_{\mathrm{layer}}\!\left( \mathbf{x}_t+\mathbf{W}^{O}\mathbf{o}_t \right). $$

2.2 这种等价带来什么

训练时,整段长度为 \(N\) 的序列可以通过分块矩阵乘法或并行扫描计算,线性注意力核心的总时间复杂度为 \(O(Nd_kd_v)\);当 \(d_k=d_v=d\) 时,可简写为 \(O(Nd^2)\)。推理时,每生成一个 token,只需更新并读取 \(\mathbf{S}_{t-1}\)\(\mathbf{z}_{t-1}\),单步时间复杂度为 \(O(d_kd_v)\);生成完整的 \(N\) 个 token 总计仍为 \(O(Nd_kd_v)\)

推理时不再保存随 \(N\) 增长的 KV cache,只需维护 \(\mathbf{S}_{t-1}\)\(\mathbf{z}_{t-1}\),状态空间复杂度为 \(O(d_kd_v+d_k)\)。在模型维度固定时,单步时间和状态空间都不随上下文长度增长,所以常被口语化地称为“单步 \(O(1)\)”;更严谨的写法是“关于 \(N\)\(O(1)\)”。

固定状态也有代价:所有历史被压入有限容量的关联记忆,不能保证无损保存。序列变长后,不同 key/value 可能发生干扰或碰撞,这正是后续门控与 Delta 规则要解决的问题。

3. 从纯累加到可选择的状态更新

3.1 纯累加的局限

省略核映射与归一化后,最简单的线性注意力状态为

$$ \mathbf{S}_t = \mathbf{S}_{t-1} +\mathbf{k}_t\mathbf{v}_t^{\top}, \qquad \mathbf{o}_t=\mathbf{S}_t^{\top}\mathbf{q}_t. $$

它把每个关联永久写入状态,没有显式遗忘机制。有限状态在长序列中会逐渐混入过时信息,降低精确检索能力。

3.2 Mamba-2:标量衰减

用与本文一致的状态方向,Mamba-2 的核心递推可概括为

$$ \mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} +\mathbf{k}_t\mathbf{v}_t^{\top}, \qquad \alpha_t\in(0,1]. $$

\(\alpha_t\) 可以快速清除旧记忆,但同一状态头的所有通道共享一个衰减率,遗忘较粗粒度。

若定义累计衰减

$$ \gamma_t=\prod_{i=1}^{t}\alpha_i, $$

则时刻 \(j\) 写入的记忆在时刻 \(t\) 的权重为 \(\gamma_t/\gamma_j\)。这揭示了递推形式与带衰减因果掩码之间的对应关系。

3.3 DeltaNet:定向覆盖

纯累加更新

$$ \mathbf{S}_t = \mathbf{S}_{t-1} +\mathbf{k}_t\mathbf{v}_t^\top $$

只负责把新关联叠加到状态中,却没有回答一个重要问题:如果同一个 key 再次出现,应该把新 value 继续累加,还是用新 value 修正旧关联?把状态看成一个从 key 空间到 value 空间的线性映射后,当前 key 在写入前读出的旧值为

$$ \widehat{\mathbf{v}}_t = \mathbf{S}_{t-1}^{\top}\mathbf{k}_t. $$

若仍采用加法写入,则写入后在同一 key 上读到

$$ \left( \mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top \right)^\top\mathbf{k}_t = \widehat{\mathbf{v}}_t +\beta_t\|\mathbf{k}_t\|_2^2\mathbf{v}_t. $$

旧值 \(\widehat{\mathbf{v}}_t\) 并没有被消除;同一关联被反复写入时,状态会不断累加, 而不是趋近于最新目标。DeltaNet 的设计动机正是把“写入完整的新 value”改成“只写入 预测误差”。先定义读写残差

$$ \mathbf{e}_t = \mathbf{v}_t-\widehat{\mathbf{v}}_t = \mathbf{v}_t-\mathbf{S}_{t-1}^{\top}\mathbf{k}_t. $$

我们希望寻找一个状态修正 \(\Delta\mathbf{S}_t\),使它在当前 key 上产生 \(\mathbf{e}_t\)

$$ \Delta\mathbf{S}_t^\top\mathbf{k}_t = \mathbf{e}_t. $$

满足这一约束的最小 Frobenius 范数秩一修正为

$$ \Delta\mathbf{S}_t = \frac{\mathbf{k}_t\mathbf{e}_t^\top} {\|\mathbf{k}_t\|_2^2}. $$

当 key 已做 L2 归一化时,分母为 \(1\)。再引入更新强度 \(\beta_t\),便得到 DeltaNet 更新:

$$ \mathbf{S}_t = \mathbf{S}_{t-1} +\beta_t\mathbf{k}_t \left( \mathbf{v}_t-\mathbf{S}_{t-1}^{\top}\mathbf{k}_t \right)^{\top}. $$

这一步可以直接验证。更新后在当前 key 上的读出为

$$ \begin{aligned} \mathbf{S}_t^\top\mathbf{k}_t &= \widehat{\mathbf{v}}_t +\beta_t\|\mathbf{k}_t\|_2^2 \left( \mathbf{v}_t-\widehat{\mathbf{v}}_t \right). \end{aligned} $$

\(\|\mathbf{k}_t\|_2=1\),则

$$ \mathbf{S}_t^\top\mathbf{k}_t = (1-\beta_t)\widehat{\mathbf{v}}_t +\beta_t\mathbf{v}_t. $$

因此 \(\beta_t=0\) 表示保留旧关联,\(0<\beta_t<1\) 表示向新 value 做部分修正,\(\beta_t=1\) 则在当前 key 方向上完成精确覆盖。对任意其他 query \(\mathbf{q}\),更新引起的读出变化为

$$ \mathbf{S}_t^\top\mathbf{q} - \mathbf{S}_{t-1}^\top\mathbf{q} = \beta_t \left( \mathbf{k}_t^\top\mathbf{q} \right)\mathbf{e}_t. $$

也就是说,修正强度由 \(\mathbf{q}\) 与当前 key 的相似度控制:与 \(\mathbf{k}_t\) 正交的方向不受影响,相似方向共享这次误差修正。这就是“定向覆盖” 而不是“全局遗忘”。

将残差形式展开,可以进一步看出一次更新由“擦除旧关联”和“写入新关联”组成:

$$ \begin{aligned} \mathbf{S}_t &= \mathbf{S}_{t-1} -\beta_t\mathbf{k}_t\mathbf{k}_t^\top\mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top\\ &= \left( \mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^{\top} \right)\mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^{\top}. \end{aligned} $$

其中 \(-\beta_t\mathbf{k}_t\mathbf{k}_t^\top\mathbf{S}_{t-1}\) 只擦除旧状态在 \(\mathbf{k}_t\) 方向上的分量,而 \(\beta_t\mathbf{k}_t\mathbf{v}_t^\top\) 随后写入新关联。

Delta 规则还可以从在线学习角度得到。定义当前样本的重建损失:

$$ \mathcal{L}_t(\mathbf{S}) = \frac{1}{2} \left\| \mathbf{S}^{\top}\mathbf{k}_t-\mathbf{v}_t \right\|_2^2 $$

其关于状态的梯度为

$$ \nabla_{\mathbf{S}}\mathcal{L}_t = \mathbf{k}_t \left( \mathbf{S}^{\top}\mathbf{k}_t-\mathbf{v}_t \right)^\top. $$

\(\mathbf{S}_{t-1}\) 出发执行一步学习率为 \(\beta_t\) 的在线梯度下降:

$$ \begin{aligned} \mathbf{S}_t &= \mathbf{S}_{t-1} -\beta_t \nabla_{\mathbf{S}}\mathcal{L}_t \big|_{\mathbf{S}=\mathbf{S}_{t-1}}\\ &= \mathbf{S}_{t-1} +\beta_t\mathbf{k}_t \left( \mathbf{v}_t-\mathbf{S}_{t-1}^\top\mathbf{k}_t \right)^\top. \end{aligned} $$

因此,DeltaNet 不是额外附加的一条启发式擦除规则,而是在固定大小的关联记忆上, 用当前 key/value 对执行一步数据依赖的在线误差校正。它能精确修改被当前 key 访问到的方向,但不会主动清理长期未再次出现的旧方向;这正是下一节加入遗忘门的原因。

3.4 Gated DeltaNet:全局遗忘与定向编辑

DeltaNet 解决了“如何覆盖当前 key 对应的旧值”,但它的擦除矩阵

$$ \mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top $$

只作用于当前 key 方向。若某些旧记忆长期不再被相似 key 访问,它们仍会保留在有限状态 中并持续占用容量。反过来,Mamba-2 式标量门 \(\alpha_t\mathbf{S}_{t-1}\) 可以快速衰减全部历史,却无法只修改某一个具体关联。 Gated DeltaNet(GDN)把这两种互补机制组合起来:\(\alpha_t\) 决定旧状态整体保留多少, \(\beta_t\) 决定当前 key 方向修改多少。

这个公式可以从两个连续步骤推导。第一步只对旧记忆施加标量遗忘门:

$$ \overline{\mathbf{S}}_{t-1} = \alpha_t\mathbf{S}_{t-1}, \qquad \alpha_t\in(0,1]. $$

此时在当前 key 上读出的、经过遗忘后的旧值为

$$ \widehat{\mathbf{v}}_t^{\,\mathrm{g}} = \overline{\mathbf{S}}_{t-1}^{\top}\mathbf{k}_t = \alpha_t\mathbf{S}_{t-1}^{\top}\mathbf{k}_t. $$

第二步以衰减后的状态为起点,执行一次 Delta 误差校正:

$$ \begin{aligned} \mathbf{S}_t &= \overline{\mathbf{S}}_{t-1} +\beta_t\mathbf{k}_t \left( \mathbf{v}_t-\widehat{\mathbf{v}}_t^{\,\mathrm{g}} \right)^\top\\ &= \overline{\mathbf{S}}_{t-1} +\beta_t\mathbf{k}_t \left( \mathbf{v}_t-\overline{\mathbf{S}}_{t-1}^{\top}\mathbf{k}_t \right)^\top. \end{aligned} $$

\(\overline{\mathbf{S}}_{t-1}=\alpha_t\mathbf{S}_{t-1}\) 代入并展开:

$$ \begin{aligned} \mathbf{S}_t &= \alpha_t\mathbf{S}_{t-1} -\alpha_t\beta_t \mathbf{k}_t\mathbf{k}_t^\top\mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top\\ &= \alpha_t \left( \mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^{\top} \right)\mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^{\top}. \end{aligned} $$

这里 \(\alpha_t\) 只乘在旧状态的转移项上,不乘新写入的 \(\beta_t\mathbf{k}_t\mathbf{v}_t^\top\)。这样,新 value 不会在刚写入时就被遗忘门 再次缩小。由于 \(\alpha_t\) 是标量,它与擦除矩阵可交换;但它不能跨过加号作用到新写入项。

\(\|\mathbf{k}_t\|_2=1\) 时,GDN 对当前 key 的读出为

$$ \mathbf{S}_t^\top\mathbf{k}_t = (1-\beta_t)\alpha_t \mathbf{S}_{t-1}^\top\mathbf{k}_t +\beta_t\mathbf{v}_t. $$

这个式子把两个门的职责分开得很清楚:

  1. \(\alpha_t\) 先决定旧读出还能保留多少,较小的 \(\alpha_t\) 可以在边界变化或上下文切换时快速重置整头记忆。
  2. \(\beta_t\) 再在当前 key 方向上混合旧读出与新目标;当 \(\beta_t=1\) 时,无论 \(\alpha_t\) 取何值,当前关联都被新 value 覆盖。
  3. 对与当前 key 不相似的方向,Delta 项影响很小,但它们仍会受到 \(\alpha_t\) 的全局衰减。

对任意 query \(\mathbf{q}\),同样可以写出

$$ \mathbf{S}_t^\top\mathbf{q} = \alpha_t\mathbf{S}_{t-1}^\top\mathbf{q} +\beta_t \left( \mathbf{k}_t^\top\mathbf{q} \right) \left( \mathbf{v}_t -\alpha_t\mathbf{S}_{t-1}^\top\mathbf{k}_t \right). $$

第一项是全局保留的历史,第二项是按 key 相似度路由的局部误差修正。因此 GDN 同时具备可学习的时间尺度与关联级编辑能力:它比纯累加更能抑制状态污染,比单独使用 标量衰减更擅长精确更新,也补上了 DeltaNet 无法主动清理未访问旧方向的缺口。不过, 同一个状态头仍共享单个 \(\alpha_t\),所有 key 通道只能同步遗忘;第 4 节的 KDA 将进一步把这个标量门提升为通道级向量门。

3.5 与状态空间模型的关系

经典离散线性时不变状态空间模型(LTI SSM)为

$$ \mathbf{h}_t=\mathbf{A}\mathbf{h}_{t-1}+\mathbf{B}\mathbf{x}_t, \qquad \mathbf{y}_t=\mathbf{C}\mathbf{h}_t. $$

\(\mathbf{A},\mathbf{B},\mathbf{C}\) 固定时,它可展开为卷积核

$$ \overline{\mathbf{K}} = \left( \mathbf{C}\mathbf{B}, \mathbf{C}\mathbf{A}\mathbf{B}, \mathbf{C}\mathbf{A}^2\mathbf{B}, \ldots \right). $$

线性注意力、Mamba-2、GDN 与 KDA 都可视为“输入驱动的线性状态更新”,并可借用递推、分块和扫描算法。但 KDA 的转移矩阵依赖当前输入,因此不是固定的 LTI 系统,也不能简单化为一个数据无关的全局卷积核。二者是相关的状态序列模型,而不是完全相同的数学对象。

4. Kimi Delta Attention

4.1 GDN 的剩余瓶颈

GDN 的 \(\alpha_t\) 是头级标量。若不同记忆通道需要不同时间尺度,一个标量只能让整头状态同步遗忘,限制了记忆控制的粒度。

KDA 将它改为通道级向量

$$ \mathbf{\alpha}_t\in(0,1]^{d_k}, \qquad \mathbf{D}_t=\operatorname{Diag}(\mathbf{\alpha}_t). $$

每个 key 通道因此拥有独立衰减率。

4.2 KDA 的核心递推

KDA 的状态更新与读取为

$$ \mathbf{S}_t = \left( \mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^{\top} \right) \mathbf{D}_t\mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^{\top} $$
$$ \mathbf{o}_t=\mathbf{S}_t^{\top}\mathbf{q}_t $$

其中

$$ \mathbf{S}_t\in\mathbb{R}^{d_k\times d_v},\quad \mathbf{k}_t,\mathbf{q}_t\in\mathbb{R}^{d_k},\quad \mathbf{v}_t,\mathbf{o}_t\in\mathbb{R}^{d_v}. $$

可以把 \(\mathbf{S}_{t-1}\) 理解为一个固定大小的“压缩记忆”:它不是逐条保存历史 key/value,而是把历史关联压入一个从 key 空间到 value 空间的线性映射。于是 \(\mathbf{k}_t\) 决定当前信息写到哪个方向,\(\mathbf{v}_t\) 决定写入什么内容, \(\mathbf{q}_t\) 则在更新完成后决定从哪个方向读取。为了把更新公式与“选择性遗忘、 修正与调整、写入新信息”三步对应起来,可以引入两个中间状态。

第一步:选择性遗忘。 先对旧状态应用通道级门:

$$ \widetilde{\mathbf{S}}_{t-1} = \mathbf{D}_t\mathbf{S}_{t-1}. $$

因为 \(\mathbf{D}_t\) 是对角矩阵,第 \(i\) 个 key 通道满足

$$ \left[ \widetilde{\mathbf{S}}_{t-1} \right]_{i,:} = \alpha_{t,i} \left[ \mathbf{S}_{t-1} \right]_{i,:}. $$

\(\alpha_{t,i}\) 接近 \(1\) 时,该通道中的旧信息基本保留;接近 \(0\) 时,该通道被快速衰减。这里的“选择性”不是从历史 token 列表中删除某一条记录, 而是当前输入为每个 key 特征通道生成不同的保留率。不同通道因此可以学习不同的 记忆时间尺度:有些通道长期积累信息,有些通道快速响应局部变化。

第二步:修正与调整。 在写入新 value 之前,先擦除衰减后状态在当前 key 方向上的旧关联。定义经过遗忘后的旧读出

$$ \widehat{\mathbf{v}}_t^{\,\mathrm{old}} = \widetilde{\mathbf{S}}_{t-1}^{\top}\mathbf{k}_t, $$

再令

$$ \begin{aligned} \overline{\mathbf{S}}_{t-1} &= \left( \mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top \right) \widetilde{\mathbf{S}}_{t-1}\\ &= \widetilde{\mathbf{S}}_{t-1} -\beta_t\mathbf{k}_t \left( \widehat{\mathbf{v}}_t^{\,\mathrm{old}} \right)^\top. \end{aligned} $$

第二行说明,这一步实际减去的是“当前 key 方向 \(\mathbf{k}_t\)”与“该方向原来 读出的 value”所形成的秩一关联。若 \(\|\mathbf{k}_t\|_2=1\),则修正后在同一 key 上的读出变为

$$ \overline{\mathbf{S}}_{t-1}^{\top}\mathbf{k}_t = (1-\beta_t) \widehat{\mathbf{v}}_t^{\,\mathrm{old}}. $$

因此 \(\beta_t=1\) 会完全清空当前 key 方向上的旧值, \(0<\beta_t<1\) 则只做部分调整。对任意其他 query \(\mathbf{q}\),这一步引起的变化为

$$ \overline{\mathbf{S}}_{t-1}^{\top}\mathbf{q} - \widetilde{\mathbf{S}}_{t-1}^{\top}\mathbf{q} = -\beta_t \left( \mathbf{k}_t^\top\mathbf{q} \right) \widehat{\mathbf{v}}_t^{\,\mathrm{old}}. $$

只有与 \(\mathbf{k}_t\) 内积绝对值较大的读取方向会受到明显影响,所以这种“修正”不同于第一步 对所有旧记忆通道施加的遗忘:第一步按通道控制时间尺度,第二步按当前 key 定位需要覆盖的关联。

第三步:写入新信息。 最后把当前 key/value 关联加入修正后的状态:

$$ \mathbf{S}_t = \overline{\mathbf{S}}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top. $$

把第二步与第三步合并,可以重新得到更直观的残差校正形式:

$$ \begin{aligned} \mathbf{S}_t &= \widetilde{\mathbf{S}}_{t-1} +\beta_t\mathbf{k}_t \left( \mathbf{v}_t -\widetilde{\mathbf{S}}_{t-1}^{\top}\mathbf{k}_t \right)^\top\\ &= \widetilde{\mathbf{S}}_{t-1} +\beta_t\mathbf{k}_t \left( \mathbf{v}_t-\widehat{\mathbf{v}}_t^{\,\mathrm{old}} \right)^\top. \end{aligned} $$

这表明“修正”与“写入”不是两次无关操作:二者共同把当前 key 上的旧读出推向新目标 \(\mathbf{v}_t\)。当 \(\|\mathbf{k}_t\|_2=1\) 时,

$$ \mathbf{S}_t^\top\mathbf{k}_t = (1-\beta_t)\widehat{\mathbf{v}}_t^{\,\mathrm{old}} +\beta_t\mathbf{v}_t. $$

\(\beta_t\) 同时控制擦除强度和写入强度,使旧值减少多少与新值增加多少保持匹配; 当 \(\beta_t=1\) 时,新 value 精确覆盖经过选择性遗忘后的旧关联。

完成这三步后, \(\mathbf{o}_t=\mathbf{S}_t^\top\mathbf{q}_t\) 使用 query 从新状态中提取输出。 在线性映射的意义下,可以把 \(\mathbf{k}_t\) 看作写入地址、\(\mathbf{v}_t\) 看作写入内容、\(\mathbf{q}_t\) 看作读取地址;但这些“地址”是连续向量方向, 相似方向之间会共享和干扰记忆,并不是离散存储槽。

对角门的位置不能随意移动。一般情况下,

$$ \left(\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top\right)\mathbf{D}_t \neq \mathbf{D}_t\left(\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top\right), $$

因为矩阵乘法不可交换。KDA 采用前一种次序,并围绕该结构设计了稳定的分块 WY/UT 算法。

4.3 KDA 的神经参数化

对第 \(h\) 个头,Kimi Linear 使用

$$ \begin{aligned} \mathbf{q}_t^h &= \operatorname{L2Norm}\!\left( \operatorname{Swish}\!\left( \operatorname{ShortConv}(\mathbf{W}_h^Q\mathbf{x}_t) \right)\right),\\ \mathbf{k}_t^h &= \operatorname{L2Norm}\!\left( \operatorname{Swish}\!\left( \operatorname{ShortConv}(\mathbf{W}_h^K\mathbf{x}_t) \right)\right),\\ \mathbf{v}_t^h &= \operatorname{Swish}\!\left( \operatorname{ShortConv}(\mathbf{W}_h^V\mathbf{x}_t) \right),\\ \mathbf{\alpha}_t^h &= f\!\left( \mathbf{W}_{\uparrow}^{\alpha} \mathbf{W}_{\downarrow}^{\alpha} \mathbf{x}_t \right),\\ \beta_t^h &= \operatorname{sigmoid}\!\left( \mathbf{W}_h^{\beta}\mathbf{x}_t \right). \end{aligned} $$

\(\mathbf{W}_{\downarrow}^{\alpha}\)\(\mathbf{W}_{\uparrow}^{\alpha}\) 构成低秩瓶颈,在控制参数量的同时生成 \(d_k\) 维通道级遗忘门。query 与 key 的 L2 归一化用于改善状态转移的数值稳定性。

多头 KDA 的结果经过逐头 RMSNorm、数据依赖输出门和输出投影:

$$ \mathbf{y}_t = \mathbf{W}^{O} \left[ \operatorname{sigmoid}\!\left( \mathbf{W}_{\uparrow}^{G} \mathbf{W}_{\downarrow}^{G}\mathbf{x}_t \right) \odot \operatorname{RMSNorm}\!\left( \operatorname{Concat}(\mathbf{o}_t^1,\ldots,\mathbf{o}_t^H) \right) \right]. $$

4.4 为什么仍能高效训练

逐 token 递推适合推理,但会限制训练并行度。KDA 将序列切成长度为 \(C\) 的块:

  • 块之间递推传递固定大小的状态;
  • 块内部把连续的秩一更新压缩为 WY 表示;
  • 使用 UT 变换减少非矩阵乘法操作;
  • 通过 Tensor Core 友好的矩阵乘法并行计算输出。

因此 KDA 同时保留了循环推理的固定状态和训练阶段的块内并行能力。它的优势并非只来自渐近复杂度,也来自对 GPU 计算路径的专门设计。

5. Kimi Linear 的混合架构

纯线性状态模型在精确长程检索上仍可能受有限状态容量约束。Kimi Linear 因此没有把全部层都替换成 KDA,而是采用层级混合:

$$ \underbrace{\mathrm{KDA}\rightarrow\mathrm{KDA}\rightarrow\mathrm{KDA}}_{3\ \text{层}} \rightarrow \underbrace{\mathrm{Full\ MLA}}_{1\ \text{层}}, $$

即重复使用 \(3:1\) 的 KDA/全局 MLA 比例。

KDA 层承担高效的长序列状态建模与位置信息编码;少量 Full MLA 层保留全局 token-to-token 检索能力。论文在 Full MLA 层中使用 NoPE,把位置与近因偏置主要交给 KDA 的数据依赖衰减机制。

相对于每层都使用 Full MLA,\(3:1\) 配比使需要维护全局 KV cache 的层数直观上减少约 \(75\%\);实际显存收益还取决于 MLA 潜变量维度、层配置、batch size 与实现。

需要特别区分:

  • 线性注意力/KDA:通过固定大小的循环状态避免随 \(N\) 增长的 KV cache。
  • MLA:通过低秩潜变量压缩 KV cache,但全局注意力本身仍包含 token 两两交互。
  • Kimi Linear:用大多数 KDA 层获得效率,再用少数 MLA 层补足精确全局检索。

因此,MLA 不是线性注意力;二者在 Kimi Linear 中是互补组件。

6. 复杂度与内存总结

设单头状态为 \(d_k\times d_v\),忽略 batch、层数和头数等公共因子:

结构与阶段 时间复杂度 持久状态/缓存 关于 \(N\)
标准注意力训练 \(O(N^2d_k)\) 显式实现为 \(O(N^2)\);FlashAttention 可降低 二次
标准注意力单步生成 \(O(Nd_k)\) KV cache \(O(N(d_k+d_v))\) 线性
核线性注意力训练 \(O(Nd_kd_v)\) 分块实现相关 线性
核线性注意力单步生成 \(O(d_kd_v)\) \(O(d_kd_v+d_k)\) 常数
KDA 训练 典型为 \(O(Nd_kd_v)\) 分块临时量 + 状态 线性
KDA 单步生成 \(O(d_kd_v)\) \(O(d_kd_v)\) 常数
Kimi Linear KDA 与少量 Full MLA 的加权组合 KDA 状态 + MLA cache 由混合比例决定

这里的“常数”均指不随上下文长度 \(N\) 增长,而不是与模型维度无关。对于完整模型,还需乘上 batch size、层数和头数。

7. 演进脉络与关键结论

可以用一条清晰的记忆更新主线总结全文:

$$ \begin{aligned} \text{线性注意力:}\quad &\mathbf{S}_t = \mathbf{S}_{t-1} +\mathbf{k}_t\mathbf{v}_t^\top, \\[2mm] \text{Mamba-2:}\quad &\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} +\mathbf{k}_t\mathbf{v}_t^\top, \\[2mm] \text{DeltaNet:}\quad &\mathbf{S}_t = \left(\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top\right) \mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top, \\[2mm] \text{Gated DeltaNet:}\quad &\mathbf{S}_t = \alpha_t \left(\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top\right) \mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top, \\[2mm] \text{KDA:}\quad &\mathbf{S}_t = \left(\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top\right) \operatorname{Diag}(\mathbf{\alpha}_t) \mathbf{S}_{t-1} +\beta_t\mathbf{k}_t\mathbf{v}_t^\top. \end{aligned} $$

核心变化依次是:

  1. 用固定大小的矩阵状态替代显式 token-to-token 注意力矩阵。
  2. 用遗忘门控制旧信息寿命。
  3. 用 Delta 规则定向擦除并覆盖 key 关联。
  4. 用通道级遗忘门让不同记忆通道拥有不同时间尺度。
  5. 用 KDA/MLA 混合架构在效率与精确全局检索之间折中。

参考资料