← 文章 / AI技术
Hacker News 2小时前 · 2026-07-29 02:02:19 · 1 阅读

Kimi Delta Attention:从线性注意力到状态更新的推导

数学符号 ⟨k|q⟩ kᵀq

符号说明:本文默认使用 bra-ket 符号,因为(在我受量子启发的观点看来)它能让推导中的形状关系非常清晰。上方的数学符号开关会将所有公式改写为传统的粗体向量和显式转置。在 bra-ket 模式下,∣q⟩\lvert q\rangle∣q⟩ 是列向量,⟨k∣\langle k\rvert⟨k∣ 是行向量,⟨k∣q⟩\langle k\rvert q\rangle⟨k∣q⟩ 是一个数值,而 ∣v⟩⟨k∣\lvert v\rangle\langle k\rvert∣v⟩⟨k∣ 是一个矩阵。向量默认朝右,而键在写入线性注意力状态时朝左。我们处理一个因果注意力头,使用实值向量,假设 DeltaNet 的键已归一化,并让状态从键空间映射到值空间。

现代线性注意力变体很复杂,乍一看很难理解它们的设计目标。以下是 Kimi Delta Attention(KDA)的状态更新方程供参考:

S~t=St−1Diag⁡(αt)\widetilde S_t = S_{t-1}\operatorname{Diag}(\alpha_t)St​=St−1​Diag(αt​) ∣v^t⟩=S~t∣kt⟩\lvert\widehat v_t\rangle = \widetilde S_t\lvert k_t\rangle∣vt​⟩=St​∣kt​⟩ ∣et⟩=βt(∣vt⟩−∣v^t⟩)\lvert e_t\rangle = \beta_t \left( \lvert v_t\rangle-\lvert\widehat v_t\rangle \right)∣et​⟩=βt​(∣vt​⟩−∣vt​⟩) St=S~t+∣et⟩⟨kt∣S_t = \widetilde S_t+\lvert e_t\rangle\langle k_t\rvertSt​=St​+∣et​⟩⟨kt​∣ ∣ot⟩=St(dk−1/2∣qt⟩)\lvert o_t\rangle = S_t\left(d_k^{-1/2}\lvert q_t\rangle\right)∣ot​⟩=St​(dk−1/2​∣qt​⟩)

它们之所以难以理解,是因为这是过去几年发展起来的线性注意力变体家族中的最新成员,其复杂度不可避免地膨胀,导致从外部看最新的变体显得难以接近。

在这篇文章中,我们将逐步讲解 DeltaNet 系列的线性注意力变体(最新版 Qwen 和 Kimi 模型家族使用了其中两种),并展示如何通过对隐藏状态做出简单假设,就能推导出同样的方程。

这就是我们将采取的路径:

softmax 注意力  →  线性注意力  →  DeltaNet  →  Gated DeltaNet  →  KDA

只有在推导出 KDA 之后,我们才会转向执行它的循环和分块 Triton 程序。

1. 从二次型注意力开始

对于 token ttt 上的查询,普通的因果 softmax 注意力为

ati=exp⁡ ⁣(s⟨ki∣qt⟩)∑j≤texp⁡ ⁣(s⟨kj∣qt⟩),s=dk−1/2,∣ot⟩=∑i≤tati∣vi⟩.\begin{aligned} 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. \end{aligned}ati​∣ot​⟩​=∑j≤t​exp(s⟨kj​∣qt​⟩)exp(s⟨ki​∣qt​⟩)​,s=dk−1/2​,=i≤t∑​ati​∣vi​⟩.​

每个注意力权重都是一个标量,衡量一个键与一个查询之间的相似度,然后 softmax 将针对该查询的所有分数转化为一个分布。输出是值向量的加权和。

在长度为 TTT 的序列中,共有 T2T^2T2 个键-查询对。在自回归推理过程中,我们可以缓存键和值而无需重新计算,但缓存会随序列增长,且每个新查询仍需检查整个历史记录。

重组这种计算的障碍在于 softmax。它的分母同时依赖于当前查询和所有之前的键。因此,我们暂时先去掉它。

1.1 去掉 softmax

为清晰起见,将常数缩放因子 sss 吸收到查询中。那么注意力最简版本为

∣ot⟩=∑i≤t⟨ki∣qt⟩∣vi⟩.\lvert o_t\rangle = \sum_{i\leq t} \langle k_i\rvert q_t\rangle \lvert v_i\rangle.∣ot​⟩=i≤t∑​⟨ki​∣qt​⟩∣vi​⟩.

标量内积可以移到右边:

∣ot⟩=∑i≤t∣vi⟩⟨ki∣qt⟩=(∑i≤t∣vi⟩⟨ki∣)∣qt⟩.\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}∣ot​⟩​=i≤t∑​∣vi​⟩⟨ki​∣qt​⟩=(i≤t∑​∣vi​⟩⟨ki​∣)∣qt​⟩.​

所有依赖过去的内容现在可以收集到一个固定大小的 V×KV \times KV×K 矩阵中:

St=∑i≤t∣vi⟩⟨ki∣\boxed{ S_t = \sum_{i\leq t} \lvert v_i\rangle\langle k_i\rvert }St​=i≤t∑​∣vi​⟩⟨ki​∣​

注意力由此变为一个循环写入后接读取的过程:

St=St−1+∣vt⟩⟨kt∣,∣ot⟩=St∣qt⟩.\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} }St​∣ot​⟩​=St−1​+∣vt​⟩⟨kt​∣,=St​∣qt​⟩.​​

恒等式

核心技巧全在这里了。外积得到矩阵,内积得到数值。我们不再存储所有历史键值对,而是将它们的累加外积保存在固定大小的状态矩阵 \(S_t\) 中。 这种方法的计算复杂度是序列长度的线性而非平方级:只需扫描一次 token,每一步更新同一个 \(d_v \times d_k\) 的状态矩阵。但效率提升的代价是丢弃了 softmax 的归一化和选择性。更复杂的线性注意力方法会使用特征映射和归一化器,但这种朴素形式暴露了 DeltaNet 要解决的内存问题。 ### 1.2 加法不等于赋值 假设我们写入一对 \(\lvert v_t\rangle\langle k_t\rvert\),然后立即用同一个 key 查询新状态: \[ \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\)。 如果旧状态已经给出了正确值,这次加法写入会让新状态输出两倍的值。更普遍的问题是,key 之间并不正交,因此每次写入都可能干扰之前的写入。线性注意力虽然给了我们一个紧凑的联想记忆,但它的更新行为像 `+=`,而我们真正需要的是接近 `=` 的效果。 ## 2. DeltaNet:写入误差,而非数值 [DeltaNet](https://arxiv.org/abs/2406.06484) 将无条件的线性注意力写入替换为 delta 规则修正。有两种实用的推导方式。 ### 2.1 推导一:要求写入能被读回 在写入第 \(t\) 个 token 之前,先询问内存当前对新的 key 有什么关联: ∣v̂_t⟩ = S_{t-1} |k_t⟩ 如果希望记忆返回的是 |v_t⟩,我们就不应该把整个值都写进去,而只应该写入差值: |v_t⟩ - |v̂_t⟩ 引入一个可学习的写入强度 β_t ∈ [0,1],定义: |e_t⟩ = β_t (|v_t⟩ - S_{t-1} |k_t⟩) 然后将这个误差写入当前键的位置: S_t = S_{t-1} + |e_t⟩⟨k_t| 现在立即读取同一个键: S_t |k_t⟩ = S_{t-1} |k_t⟩ + |e_t⟩⟨k_t|k_t⟩ = (1-β_t) S_{t-1} |k_t⟩ + β_t |v_t⟩ 当 β_t=1 时,结果恰好就是 |v_t⟩。β_t 越小,旧预测就越向目标值靠近。 这个修正也是键空间局部的。对于任何与当前键正交的查询 |x⟩: ⟨k_t|x⟩=0 ⇒ (S_t - S_{t-1})|x⟩ = |e_t⟩ ⟨k_t|x⟩ = 0 因此,这个秩一写入只改变选定键方向上的响应,而所有正交方向保持不变。 ### 2.2 推导二:对重建损失走一步 同样的更新规则也可以从在线学习的目标函数中推导出来。把当前键值对当作线性映射 S 的一个训练样本: L_t(S) = ½ ‖S|k_t⟩ - |v_t⟩‖₂² 它对状态的梯度是: ∇_S L_t(S) = (S|k_t⟩ - |v_t⟩) ⟨k_t| 这显然是一个外积:值空间的预测误差乘以观测到该误差的键的 bra。从 S_{t-1} 出发,以 β_t 为步长走一步梯度下降: St = S_{t-1} - β_t ∇_S L_t(S_{t-1}) = S_{t-1} - β_t (S_{t-1}|k_t⟩ - |v_t⟩)⟨k_t| = S_{t-1} + β_t (|v_t⟩ - S_{t-1}|k_t⟩)⟨k_t|. 这正是我们通过要求即时重建得到的更新。两种解读是等价的: - 从记忆操作的角度看,β_t 控制旧关联被替换的强度; - 从在线学习的角度看,β_t 是步长; - 从线性代数的角度看,变化是一个秩一外积。 ### 2.3 DeltaNet 的状态转移 展开误差项后,DeltaNet 表现为一个结构化状态转移加上新输入: S_t = S_{t-1} + β_t (|v_t⟩ - S_{t-1}|k_t⟩)⟨k_t| = S_{t-1} (I - β_t |k_t⟩⟨k_t|) + β_t |v_t⟩⟨k_t|. 对于单位键,I - β_t |k_t⟩⟨k_t| 在当前键方向上的特征值为 1 - β_t,在所有正交方向上的特征值为 1。它在添加新关联之前,先沿当前键方向移除旧关联。 DeltaNet 解决了写入问题,但尚未解决状态的寿命问题。 ## 3. 门控 DeltaNet:旧信息有时应该消失 线性状态将整个历史压缩到一个矩阵中。一次读取 S_t|q⟩ = Σ_{i ≤ t} ⟨k_i|q⟩|v_i⟩ 无法在某个旧 token 被折叠进 S_t 后跳过它。所有与查询重叠的存储方向都会参与贡献。Delta 规则可以修正当前键附近的状态,但其他方向上的陈旧信息仍然可用,并可能扭曲未来的读取结果。 因此,我们需要一种方法,在使用旧状态之前先将其遗忘。设 αt∈[0,1]\alpha_t\in[0,1]αt​∈[0,1] 为一个可学习的标量保留门: S~t=αtSt−1.\widetilde S_t = \alpha_t S_{t-1}.St​=αt​St−1​. 对这个门控状态应用同样的 delta 规则: S~t=αtSt−1,forget,∣v^t⟩=S~t∣kt⟩,predict,∣et⟩=βt(∣vt⟩−∣v^t⟩),correct,St=S~t+∣et⟩⟨kt∣,write.\boxed{ \begin{aligned} \widetilde S_t &= \alpha_tS_{t-1}, &&\text{forget},\\ \lvert\widehat v_t\rangle &= \widetilde S_t\lvert k_t\rangle, &&\text{predict},\\ \lvert e_t\rangle &= \beta_t \left( \lvert v_t\rangle-\lvert\widehat v_t\rangle \right), &&\text{correct},\\ S_t &= \widetilde S_t+\lvert e_t\rangle\langle k_t\rvert, &&\text{write}. \end{aligned} }St​∣vt​⟩∣et​⟩St​​=αt​St−1​,=St​∣kt​⟩,=βt​(∣vt​⟩−∣vt​⟩),=St​+∣et​⟩⟨kt​∣,​​forget,predict,correct,write.​ 这就是 Gated DeltaNet。顺序很重要:先遗忘,然后从保留状态进行预测,最后修正预测。如果我们在遗忘之前预测,误差描述的将是与更新时不同的记忆。 展开递推式得到: St=αtSt−1(I−βt∣kt⟩⟨kt∣)+βt∣vt⟩⟨kt∣.S_t = \alpha_tS_{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.St​=αt​St−1​(I−βt​∣kt​⟩⟨kt​∣)+βt​∣vt​⟩⟨kt​∣. delta 规则实现精准替换,标量门实现全局擦除。它们解决不同的问题,互为补充。 但 αt\alpha_tαt​ 仍然对整个矩阵做单一决策。模型必须以相同速率保留或遗忘所有键通道。 ## 4. Kimi Delta Attention:独立遗忘每个通道 Kimi Delta Attention 将 Gated DeltaNet 的标量保留替换为向量 αt∈[0,1]dk\alpha_t\in[0,1]^{d_k}αt​∈[0,1]dk​。将该向量放在对角线上: Dt=Diag⁡(αt)∈Rdk×dk.D_t = \operatorname{Diag}(\alpha_t) \in\mathbb R^{d_k\times d_k}.Dt​=Diag(αt​)∈Rdk​×dk​. 我们的状态将键映射到值,因此键通道就是 SSS 的列。右乘对每个通道应用不同的保留因子: S~t=St−1Dt.\widetilde S_t = S_{t-1}D_t.St​=St−1​Dt​. 其余部分就是我们已推导出的 delta 规则: S̃ₜ = Sₜ₋₁Dₜ, 遗忘每个键通道, |v̂ₜ⟩ = S̃ₜ|kₜ⟩, 预测, |eₜ⟩ = βₜ(|vₜ⟩ - |v̂ₜ⟩), 修正, Sₜ = S̃ₜ + |eₜ⟩⟨kₜ|, 写入, |oₜ⟩ = Sₜ(s|qₜ⟩), s = dₖ⁻¹/², 读取. 这就是 KDA。与 Gated DeltaNet 相比,概念上的变化仅在于将 αₜ ⟶ Dₜ = Diag(αₜ) 这一提升。效果是显著的:一个通道可以被清空,而另一个通道则得以保留。

4.1 为什么转移矩阵是对角加低秩

展开 KDA 的修正步骤: Sₜ = Sₜ₋₁Dₜ + βₜ(|vₜ⟩ - Sₜ₋₁Dₜ|kₜ⟩)⟨kₜ| = Sₜ₋₁Dₜ(I - βₜ|kₜ⟩⟨kₜ|) + βₜ|vₜ⟩⟨kₜ|.  ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ ̄ Aₜ
原始来源: Hacker News

评论 (0)