你也能提出Kimi Delta Attention

Hacker News Top 论文

摘要

这篇博客文章逐步推导了从标准softmax注意力经过线性注意力和DeltaNet变体到Kimi Delta Attention的过程,并解释了近期Qwen和Kimi模型使用的状态更新方程。

暂无内容
查看原文
查看缓存全文

缓存时间: 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} $$

相似文章

线性注意力架构:机制、权衡与跨层路由

arXiv cs.LG

本文对比了softmax注意力与四种线性注意力架构(DeltaNet、Gated DeltaNet、Kimi Delta Attention、Gated DeltaNet-2),并介绍了跨层路由机制。在350M参数规模的实验表明,使用Muon优化器的Kimi Delta Attention取得了最低的验证损失,而使用AdamW的纯Gated DeltaNet吞吐量最高。

Kimi-K3 技术报告 [pdf]

Hacker News Top

MoonshotAI发布了Kimi-K3,一个拥有2.8万亿参数的开源权重多模态智能体模型,具备100万token的上下文窗口,基于全新的Kimi Delta Attention和Attention Residuals架构,实现了显著的扩展性改进。