缓存时间:
2026/07/28 18:27
# 你本来也能想到 Kimi Delta Attention | Doubleword
来源:https://blog.doubleword.ai/you-could-have-come-up-with-kimi-delta-attention
*记法说明:本文默认使用 bra‑ket 记法,因为(在我深受量子启发的观点看来)它使推导中的形状非常清晰。文中的**数学记法**开关会用传统的粗体向量和显式转置重写每个方程。在 bra‑ket 模式下,$\lvert q\rangle$ 是列向量,$\langle k\rvert$ 是行向量,$\langle k\rvert q\rangle$ 是一个数,而 $\lvert v\rangle\langle k\rvert$ 是一个矩阵。向量默认朝右,而键在写入线性注意力状态时朝左。我们处理一个因果注意力头和实值向量,假设 DeltaNet 的键已归一化,并让状态从键空间映射到值空间。*
现代线性注意力变体很复杂,乍一看很难看出它们的设计目标。作为参考,以下是 Kimi Delta Attention (KDA) 的状态更新方程:
$$\\widetilde S_t = S_{t-1}\\operatorname{Diag}(\\alpha_t)$$
$$\\lvert\\widehat v_t\\rangle = \\widetilde S_t\\lvert k_t\\rangle$$
$$\\lvert e_t\\rangle = \\beta_t \\left( \\lvert v_t\\rangle-\\lvert\\widehat v_t\\rangle \\right)$$
$$S_t = \\widetilde S_t+\\lvert e_t\\rangle\\langle k_t\\rvert$$
$$\\lvert o_t\\rangle = S_t\\left(d_k^{-1/2}\\lvert q_t\\rangle\\right)$$
它们之所以如此难以理解,是因为这是过去几年发展起来的线性注意力变体家族中的最新一员,其复杂性不可避免地膨胀,使得从外部看最新变体显得难以接近。
在这篇文章中,我们将逐步解析 DeltaNet 系列的线性注意力变体(最新 Qwen 和 Kimi 模型家族使用了其中两种),并展示如果你对自己的隐藏状态做出简单断言,是如何得出相同方程的。
我们将遵循以下路线:
softmax 注意力 → 线性注意力 → [DeltaNet](https://arxiv.org/abs/2406.06484) → [Gated DeltaNet](https://arxiv.org/abs/2412.06464) → [KDA](https://arxiv.org/abs/2510.26692)
只有在推导出 KDA 之后,我们才会转向执行它的循环和分块 Triton 程序。
## 1. 从二次注意力开始
对于第 $t$ 个 token 的查询,普通的因果 softmax 注意力为:
$$
a_{ti} = \frac{ \exp\!\left(s\langle k_i\rvert q_t\rangle\right) }{ \sum_{j\leq t} \exp\!\left(s\langle k_j\rvert q_t\rangle\right) }, \qquad s=d_k^{-1/2},
$$
$$
\lvert o_t\rangle = \sum_{i\leq t}a_{ti}\lvert v_i\rangle.
$$
每个注意力权重都是一个标量。它测量一个键和一个查询之间的相似性,然后 softmax 将该查询的所有得分转换为一个分布。输出是值向量的加权和。
在长度为 $T$ 的序列上,有 $T^2$ 个键‑查询对。在自回归推理期间,我们可以缓存键和值而不是重新计算,但缓存仍然随序列增长,并且每个新查询仍然需要检查整个历史。
阻碍重新排列此计算的是 softmax。它的分母同时依赖于当前查询和每个较早的键。因此,目前我们先去掉它。
### 1.1 去掉 softmax
为清晰起见,将常数标量 $s$ 吸收到查询中。然后,故意简化的注意力版本变为:
$$
\lvert o_t\rangle = \sum_{i\leq t} \langle k_i\rvert q_t\rangle \lvert v_i\rangle.
$$
标量内积可以移到右边:
$$
\begin{aligned}
\lvert o_t\rangle &= \sum_{i\leq t} \lvert v_i\rangle \langle k_i\rvert q_t\rangle \\
&= \left( \sum_{i\leq t} \lvert v_i\rangle\langle k_i\rvert \right) \lvert q_t\rangle.
\end{aligned}
$$
所有依赖过去的部分现在可以收集到一个固定大小 $V \times K$ 的矩阵中:
$$
\boxed{ S_t = \sum_{i\leq t} \lvert v_i\rangle\langle k_i\rvert }
$$
注意力变成一个循环写入接一个读取:
$$
\boxed{
\begin{aligned}
S_t &= S_{t-1} + \lvert v_t\rangle\langle k_t\rvert,\\
\lvert o_t\rangle &= S_t\lvert q_t\rangle.
\end{aligned}
}
$$
恒等式
$$
\left(\lvert v\rangle\langle k\rvert\right)\lvert q\rangle = \langle k\rvert q\rangle\lvert v\rangle
$$
就是全部窍门。外积是一个矩阵;内积是一个数。我们不再存储每个过去的键和值。我们存储它们的外积之和,放在固定大小的状态 $S_t$ 中。
这在序列长度上是线性的而非二次的:扫描所有 token 一次,每一步更新相同的 $d_v \times d_k$ 状态。我们通过抛弃 softmax 的归一化和选择性来为这种效率付出代价。更复杂的线性注意力方法使用特征映射和归一化器,但这种未经修饰的形式暴露了促使 DeltaNet 产生的记忆问题。
### 1.2 加法不是赋值
假设我们写入一对 $\lvert v_t\rangle\langle k_t\rvert$,然后用同一个键立即查询新状态:
$$
\begin{aligned}
S_t\lvert k_t\rangle &= \left( S_{t-1} + \lvert v_t\rangle\langle k_t\rvert \right) \lvert k_t\rangle \\
&= S_{t-1}\lvert k_t\rangle + \lvert v_t\rangle \underbrace{\langle k_t\rvert k_t\rangle}_{1} \\
&= S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle.
\end{aligned}
$$
写入**并不**使记忆返回 $\lvert v_t\rangle$。它是在记忆已经返回的内容上加上 $\lvert v_t\rangle$。
如果旧状态已经产生了正确的值,那么加法写入会使新状态产生两倍的值。更一般地说,键并不相互正交,所以每次写入都可能干扰之前的写入。线性注意力给了我们一个紧凑的关联记忆,但它的更新行为类似于 `+=`,而我们想要的是更接近 `=` 的东西。
## 2. DeltaNet:写入误差,而不是值
[DeltaNet](https://arxiv.org/abs/2406.06484) 用 delta 规则修正取代了无条件的线性注意力写入。有两种有用的推导方式。
### 2.1 推导一:要求写入可以被读回
在写入第 $t$ 个 token 之前,询问记忆当前与新键关联着什么:
$$
\lvert\widehat v_t\rangle = S_{t-1}\lvert k_t\rangle.
$$
如果我们希望记忆返回 $\lvert v_t\rangle$,我们不应该加上整个值。我们应该只加上差值:
$$
\lvert v_t\rangle - \lvert\widehat v_t\rangle.
$$
引入一个学习到的写入强度 $\beta_t \in [0,1]$,并定义
$$
\lvert e_t\rangle = \beta_t \left( \lvert v_t\rangle - S_{t-1}\lvert k_t\rangle \right).
$$
然后在当前键处写入这个误差:
$$
\boxed{ S_t = S_{t-1} + \lvert e_t\rangle\langle k_t\rvert. }
$$
现在立即读取同一个键:
$$
\begin{aligned}
S_t\lvert k_t\rangle &= S_{t-1}\lvert k_t\rangle + \lvert e_t\rangle \langle k_t\rvert k_t\rangle \\
&= (1-\beta_t)S_{t-1}\lvert k_t\rangle + \beta_t\lvert v_t\rangle.
\end{aligned}
$$
当 $\beta_t=1$ 时,结果恰好是 $\lvert v_t\rangle$。较小的 $\beta_t$ 将旧预测部分地移向目标。
修正也是在键空间局部进行的。对于任何与当前键正交的查询 $\lvert x\rangle$,
$$
\langle k_t\rvert x\rangle=0 \quad\Longrightarrow\quad (S_t-S_{t-1})\lvert x\rangle = \lvert e_t\rangle \underbrace{\langle k_t\rvert x\rangle}_{0}=0.
$$
因此,秩一写入在选定的键方向上改变响应,而保持每个正交方向不变。
### 2.2 推导二:在重构损失上走一步
同样的更新从一个在线学习目标中得出。将当前的键值对视为线性映射 $S$ 的一个训练样本:
$$
\mathcal L_t(S) = \frac12 \left\| S\lvert k_t\rangle-\lvert v_t\rangle \right\|_2^2.
$$
它对状态的梯度是
$$
\nabla_S\mathcal L_t(S) = \left( S\lvert k_t\rangle-\lvert v_t\rangle \right) \langle k_t\rvert.
$$
这显然是一个外积:值空间预测误差乘以观察该误差的键的 bra。从 $S_{t-1}$ 处迈出大小为 $\beta_t$ 的一步梯度下降:
$$
\begin{aligned}
S_t &= S_{t-1} - \beta_t\nabla_S\mathcal L_t(S_{t-1}) \\
&= S_{t-1} - \beta_t \left( S_{t-1}\lvert k_t\rangle-\lvert v_t\rangle \right) \langle k_t\rvert \\
&= S_{t-1} + \beta_t \left( \lvert v_t\rangle - S_{t-1}\lvert k_t\rangle \right) \langle k_t\rvert.
\end{aligned}
$$
这与我们通过要求立即重构得到的更新完全一致。这两种解释是相同的:
- 作为记忆操作,$\beta_t$ 控制替换旧关联的强度;
- 作为在线学习,$\beta_t$ 是步长;
- 作为线性代数,变化是一个秩一外积。
### 2.3 DeltaNet 的状态转换
展开误差项,揭示了 DeltaNet 是一种结构化状态转换加上一个新输入:
$$
\begin{aligned}
S_t &= S_{t-1} + \beta_t \left( \lvert v_t\rangle - S_{t-1}\lvert k_t\rangle \right) \langle k_t\rvert \\
&= S_{t-1} \left( I - \beta_t\lvert k_t\rangle\langle k_t\rvert \right) + \beta_t\lvert v_t\rangle\langle k_t\rvert.
\end{aligned}
$$
对于一个单位键,$I - \beta_t\lvert k_t\rangle\langle k_t\rvert$ 在当前键方向上的特征值为 $1 - \beta_t$,在每个正交方向上的特征值为 $1$。它在添加新关联之前,沿着当前键移除旧关联。
DeltaNet 解决了写入的问题。但它还没有解决状态的生命周期问题。
## 3. Gated DeltaNet:有时旧信息应该消失
线性状态将整个历史压缩到一个矩阵中。一次读取
$$
S_t\lvert q\rangle = \sum_{i\leq t} \langle k_i\rvert q\rangle\lvert v_i\rangle
$$
无法在某个旧 token 被折叠进 $S_t$ 之后选择跳过它。每个与查询重叠的存储方向都会贡献。delta 规则可以修正当前键周围的状态,但其他方向上的过时信息仍然可用,并且可能扭曲未来的读取。
因此,我们需要一种在用旧状态之前遗忘它的方法。令 $\alpha_t \in [0,1]$ 为一个学习到的标量保留门:
$$
\widetilde S_t = \alpha_t S_{t-1}.
$$
对这个门控状态运行同样的 delta 规则:
$$
\boxed{
\begin{aligned}
\widetilde S_t &= \alpha_t S_{t-1}, &&\text{遗忘},\\
\lvert\widehat v_t\rangle &= \widetilde S_t\lvert k_t\rangle, &&\text{预测},\\
\lvert e_t\rangle &= \beta_t \left( \lvert v_t\rangle - \lvert\widehat v_t\rangle \right), &&\text{修正},\\
S_t &= \widetilde S_t + \lvert e_t\rangle\langle k_t\rvert, &&\text{写入}.
\end{aligned}
}
$$
这就是 [Gated DeltaNet](https://arxiv.org/abs/2412.06464)。顺序很重要:先遗忘,然后从保留的状态预测,再修正该预测。如果我们在遗忘之前预测,误差描述的记忆将与我们要更新的记忆不同。
展开循环得到:
$$
S_t = \alpha_t S_{t-1} \left( I - \beta_t\lvert k_t\rangle\langle k_t\rvert \right) + \beta_t\lvert v_t\rangle\langle k_t\rvert.
$$
delta 规则提供针对性替换;标量门提供全局擦除。它们解决不同的问题,并且是互补的。
但 $\alpha_t$ 仍然对整个矩阵做出一个决定。模型必须以相同的速率保留或遗忘每个键通道。
## 4. Kimi Delta Attention:独立遗忘每个通道
[Kimi Delta Attention](https://arxiv.org/abs/2510.26692) 将 Gated DeltaNet 的标量保留替换为一个向量 $\alpha_t \in [0,1]^{d_k}$。将向量放在对角线上:
$$
D_t = \operatorname{Diag}(\alpha_t) \in \mathbb R^{d_k \times d_k}.
$$
我们的状态将键映射到值,因此键通道是 $S$ 的列。右乘对每一个应用不同的保留因子:
$$
\widetilde S_t = S_{t-1} D_t.
$$
其他一切都是我们已经推导出的 delta 规则:
$$
\boxed{
\begin{aligned}
\widetilde S_t &= S_{t-1} D_t, &&\text{遗忘每个键通道},\\
\lvert\widehat v_t\rangle &= \widetilde S_t\lvert k_t\rangle, &&\text{预测},\\
\lvert e_t\rangle &= \beta_t \left( \lvert v_t\rangle - \lvert\widehat v_t\rangle \right), &&\text{修正},\\
S_t &= \widetilde S_t + \lvert e_t\rangle\langle k_t\rvert, &&\text{写入},\\
\lvert o_t\rangle &= S_t(s\lvert q_t\rangle), \qquad s = d_k^{-1/2}, &&\text{读取}.
\end{aligned}
}
$$
这就是 KDA。与 Gated DeltaNet 相比,概念上的变化仅在于提升
$$
\alpha_t \quad\longrightarrow\quad D_t = \operatorname{Diag}(\alpha_t).
$$
但效果是显著的:一个通道可以被清除而另一个通道被保留。
### 4.1 为什么转换是对角加低秩
展开 KDA 的修正:
$$
\begin{aligned}
S_t &= S_{t-1} D_t + \beta_t \left( \lvert v_t\rangle - S_{t-1} D_t \lvert k_t\rangle \right) \langle k_t\rvert \\
&= S_{t-1} \underbrace{ D_t \left( I - \beta_t\lvert k_t\rangle\langle k_t\rvert \right) }_{A_t} + \beta_t \lvert v_t\rangle\langle k_t\rvert.
\end{aligned}
$$