Kimi Delta Attention:从线性注意力到状态更新的推导
数学符号
⟨k|q⟩ kᵀq符号说明:本文默认使用 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 → Gated DeltaNet → KDA
只有在推导出 KDA 之后,我们才会转向执行它的循环和分块 Triton 程序。
1. 从二次注意力开始
对于 token $t$ 处的查询,普通的因果 softmax 注意力为
$$\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}$$ 每个注意力权重都是一个标量,它衡量一个键与一个查询之间的相似度,然后 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$,然后立即用同一个 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 用 delta 规则修正替代了无条件的线性注意力写入。有两种有用的推导方式。
2.1 推导一:要求写入能被读回
在写入 token $t$ 之前,先询问内存当前与新 key 关联的是什么:
$$\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).$$然后将这个误差写入当前 key 的位置:
$$\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.$$这显然是一个外积:值空间的预测误差乘以观察到该误差的键的左矢。从 $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. 门控 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_tS_{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} }$$这就是门控 DeltaNet。顺序很重要:先遗忘,然后从保留的状态进行预测,最后修正该预测。如果我们在遗忘之前进行预测,误差描述的记忆将与我们要更新的记忆不一致。
展开递推关系得到
$$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.$$Delta 规则提供定向替换,标量门提供全局擦除。它们解决不同的问题,并且互为补充。
但 $\alpha_t$ 仍然为整个矩阵做出单一决策。模型必须以相同的速率保留或遗忘每一个键通道。
4. Kimi Delta Attention:独立遗忘每个通道
Kimi Delta Attention 将 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}$$键空间的转移矩阵为
$$\begin{aligned} A_t &= D_t-\beta_tD_t\lvert k_t\rangle\langle k_t\rvert\\ &= D_t-\lvert b_t\rangle\langle a_t\rvert, \end{aligned}$$其中
$$\lvert b_t\rangle=D_t\lvert k_t\rangle, \qquad \langle a_t\rvert=\beta_t\langle k_t\rvert.$$因此 $A_t$ 是一个对角矩阵减去一个秩一矩阵:即对角加低秩(DPLR)的转移矩阵。“DPLR”描述的是作用于键空间的 $d_k\times d_k$ 转移矩阵。记忆状态本身仍然是 $d_v\times d_k$ 的矩阵 $S_t$。
整个流程现在可以简洁地总结如下:
| 机制 | 状态更新 | 新增功能 |
|---|---|---|
| 线性注意力 | $S+\lvert v\rangle\langle k\rvert$ | 固定大小的循环记忆 |
| DeltaNet | $S+\beta(\lvert v\rangle-S\lvert k\rangle)\langle k\rvert$ | 定向替换 |
| 门控 DeltaNet | 先应用 $\alpha S$,再执行 delta 更新 | 整体状态遗忘 |
| KDA | 先应用 $SD$,再执行 delta 更新 | 逐键通道遗忘 |
实现时通常存储 $g_t=\log\alpha_t$,其中 $g_t\leq0$,然后通过 $\exp(g_t)$ 得到保留因子。在参考代码使用的转置 $d_k\times d_v$ 布局中,循环只需五行代码:
state = state * g_t.exp().unsqueeze(-1)
prediction = einsum("bhkv,bhk->bhv", state, k_t)
residual = beta_t.unsqueeze(-1) * (v_t - prediction)
state = state + einsum("bhk,bhv->bhkv", k_t, residual)
output = einsum("bhk,bhkv->bhv", q_t * scale, state)
详情请参考官方
naive_recurrent_kda
实现。
5. 融合循环 Triton 内核
上述循环是自回归解码的自然实现方式。KDA 有两种主要的执行模式:
| 模式 | 最佳用途 | 并行单元 |
|---|---|---|
| 融合循环 | 解码、短序列、有状态服务 | 一个序列、值头、值块 |
| 分块模式 | 训练和长序列预填充 | 块、令牌子块、键/值块 |
循环 Triton 启动时,每个序列、值头和 32 宽的值块各分配一个程序:
BK = triton.next_power_of_2(K)
BV = 32
grid = (triton.cdiv(V, BV) * N * HV,)
请参阅
fused_recurrent_kda_fwd
的启动代码。
BK 覆盖了常规支持配置下的键维度。每个程序拥有实现中转置状态的一个 [BK, BV] 分块,并按顺序遍历 token。不同的值分块、注意力头和序列独立运行。
该内核几乎是递推公式的直译:
state *= tl.exp(g_t[:, None])
prediction = tl.sum(state * k_t[:, None], axis=0)
residual = beta_t * (v_t - prediction)
state += k_t[:, None] * residual[None, :]
out_t = tl.sum(state * (q_t * SCALE)[:, None], axis=0)
预测和读取是归约操作,写入则是外积。这对解码阶段非常有利,因为每次只处理一个新 token。但对于训练和长预填充来说,吸引力较小,因为这些向量运算无法变成张量核心最高效的大型矩阵乘法。
这促使我们从另一个视角审视完全相同的递推公式。
6. 分块 KDA
分块 KDA 一次性处理 $C$ 个 token。它必须产生与逐 token 递推完全相同的状态和输出,但将计算重组为矩阵乘积。
对于每个分块 $c$,我们需要两个结果:
- 给定输入状态 $S_c$ 后,整个分块处理完毕后的状态 $S_{c+1}$;
- 分块内每个因果 token 的输出。
唯一的难点在于,token $i$ 的 delta 误差依赖于同一分块中更早 token 的写入。一个四 token 的例子能清晰展示这些依赖关系。
6.1 四 token 分块的衰减记号
取 token $0,1,2,3$,定义
$$D_i=\operatorname{Diag}(\alpha_i).$$从分块边界到 token $i$ 的累积衰减为
$$D_{0:i}=D_0D_1\cdots D_i.$$将 token $j$ 处的写入传递到 token $i$ 的衰减为
$$D_{j+1:i}=D_{j+1}D_{j+2}\cdots D_i, \qquad j<i,$$当没有中间衰减时,$D_{i+1:i}=I$。所有这些矩阵都是对角矩阵,因此它们彼此可交换。
6.2 从临时误差开始
首先假设每个 token 只能看到经过适当衰减的传入状态,而看不到其所在 chunk 内的其他写入:
$$\boxed{ \lvert\bar e_i\rangle = \beta_i \left( \lvert v_i\rangle - S_cD_{0:i}\lvert k_i\rangle \right). }$$对于四个 token,我们得到四个临时的值空间误差 ket:
$$\lvert\bar e_0\rangle,\quad \lvert\bar e_1\rangle,\quad \lvert\bar e_2\rangle,\quad \lvert\bar e_3\rangle.$$这些可以并行计算,但除了第一个之外,其他都是错误的:同一 chunk 内更早的写入也会影响它们的预测。
6.3 恢复因果依赖关系
Token $0$ 在 chunk 内没有更早的写入,因此
$$\lvert e_0\rangle=\lvert\bar e_0\rangle.$$Token $1$ 看到 token $0$ 的写入经过 $D_1$ 衰减后的结果:
$$\lvert e_1\rangle = \lvert\bar e_1\rangle - \beta_1 \langle k_0\rvert D_1\lvert k_1\rangle \lvert e_0\rangle.$$Token $2$ 看到它前面的两次写入:
$$\begin{aligned} \lvert e_2\rangle &= \lvert\bar e_2\rangle\\ &\quad- \beta_2 \langle k_0\rvert D_1D_2\lvert k_2\rangle \lvert e_0\rangle\\ &\quad- \beta_2 \langle k_1\rvert D_2\lvert k_2\rangle \lvert e_1\rangle. \end{aligned}$$Token $3$ 看到它前面的所有三次写入:
$$\begin{aligned} \lvert e_3\rangle &= \lvert\bar e_3\rangle\\ &\quad- \beta_3 \langle k_0\rvert D_1D_2D_3\lvert k_3\rangle \lvert e_0\rangle\\ &\quad- \beta_3 \langle k_1\rvert D_2D_3\lvert k_3\rangle \lvert e_1\rangle\\ &\quad- \beta_3 \langle k_2\rvert D_3\lvert k_3\rangle \lvert e_2\rangle. \end{aligned}$$每个括号 $\langle k_j\rvert D_{j+1:i}\lvert k_i\rangle$ 都是一个标量。定义因果键-键系数
$$\boxed{ \rho_{ij} = \beta_i \langle k_j\rvert D_{j+1:i}\lvert k_i\rangle, \qquad j<i. }$$那么所有四个方程都可以写成紧凑形式
$$\begin{aligned} \lvert e_0\rangle &=\lvert\bar e_0\rangle,\\ \lvert e_1\rangle &=\lvert\bar e_1\rangle-\rho_{10}\lvert e_0\rangle,\\ \lvert e_2\rangle &=\lvert\bar e_2\rangle-\rho_{20}\lvert e_0\rangle -\rho_{21}\lvert e_1\rangle,\\ \lvert e_3\rangle &=\lvert\bar e_3\rangle-\rho_{30}\lvert e_0\rangle -\rho_{31}\lvert e_1\rangle-\rho_{32}\lvert e_2\rangle. \end{aligned}$$将这些系数收集成一个严格下三角矩阵:
$$R_c = \begin{bmatrix} 0&0&0&0\\ \rho_{10}&0&0&0\\ \rho_{20}&\rho_{21}&0&0\\ \rho_{30}&\rho_{31}&\rho_{32}&0 \end{bmatrix}, \qquad A^{kk}_c=(I+R_c)^{-1}.$$将误差 ket 按列堆叠:
$$\bar E_c = \begin{bmatrix} \lvert\bar e_0\rangle& \lvert\bar e_1\rangle& \lvert\bar e_2\rangle& \lvert\bar e_3\rangle \end{bmatrix},$$$E_c$ 同理。于是因果替换为
$$\boxed{ E_c = \bar E_c\left(A^{kk}_c\right)^\mathsf T. }$$实现时无需形成一般的稠密逆矩阵。因为 $I+R_c$ 是三角矩阵且对角线全为 1,该操作相当于对每个值通道独立进行因果三角求解。
6.4 快速推进状态
在 chunk 末尾,传入状态已历经全部四次衰减。每个 chunk 内的写入只经过其后的衰减:
$$\begin{aligned} S_{c+1} &= S_cD_0D_1D_2D_3\\ &\quad+ \lvert e_0\rangle\langle k_0\rvert D_1D_2D_3\\ &\quad+ \lvert e_1\rangle\langle k_1\rvert D_2D_3\\ &\quad+ \lvert e_2\rangle\langle k_2\rvert D_3\\ &\quad+ \lvert e_3\rangle\langle k_3\rvert. \end{aligned}$$定义到达末端边界时键所构成的矩阵(行向量为键):
$$K_c^{\mathrm{end}} = \begin{bmatrix} \langle k_0\rvert D_1D_2D_3\\ \langle k_1\rvert D_2D_3\\ \langle k_2\rvert D_3\\ \langle k_3\rvert \end{bmatrix}.$$由于 $E_c$ 将误差 ket 按列堆叠,四个外积写入合并为一次矩阵乘法:
$$\boxed{ S_{c+1} = S_cD_{0:3} + E_cK_c^{\mathrm{end}}. }$$这是第一个所需的 chunk 结果:将循环状态一次性推进四个 token。
6.5 计算所有因果输出
KDA 采用先写后读模式。若 $S^{[i+1]}$ 是 token $i$ 之后的局部状态,则
$$\lvert o_i\rangle = sS^{[i+1]}\lvert q_i\rangle.$$展开四个输出:
$$\begin{aligned} \lvert o_0\rangle &= sS_cD_0\lvert q_0\rangle + s\langle k_0\rvert q_0\rangle\lvert e_0\rangle,\\ \lvert o_1\rangle &= sS_cD_0D_1\lvert q_1\rangle\\ &\quad+ s\langle k_0\rvert D_1\lvert q_1\rangle\lvert e_0\rangle + s\langle k_1\rvert q_1\rangle\lvert e_1\rangle,\\ \lvert o_2\rangle &= sS_cD_0D_1D_2\lvert q_2\rangle\\ &\quad+ s\langle k_0\rvert D_1D_2\lvert q_2\rangle\lvert e_0\rangle\\ &\quad+ s\langle k_1\rvert D_2\lvert q_2\rangle\lvert e_1\rangle + s\langle k_2\rvert q_2\rangle\lvert e_2\rangle,\\ \lvert o_3\rangle &= sS_cD_0D_1D_2D_3\lvert q_3\rangle\\ &\quad+ s\langle k_0\rvert D_1D_2D_3\lvert q_3\rangle\lvert e_0\rangle\\ &\quad+ s\langle k_1\rvert D_2D_3\lvert q_3\rangle\lvert e_1\rangle\\ &\quad+ s\langle k_2\rvert D_3\lvert q_3\rangle\lvert e_2\rangle + s\langle k_3\rvert q_3\rangle\lvert e_3\rangle. \end{aligned}$$定义因果查询-键系数
$$\boxed{ \chi_{ij} = s\langle k_j\rvert D_{j+1:i}\lvert q_i\rangle, \qquad j\leq i, }$$并将这些系数放入一个下三角读取矩阵中:
$$A^{qk}_c = \begin{bmatrix} \chi_{00}&0&0&0\\ \chi_{10}&\chi_{11}&0&0\\ \chi_{20}&\chi_{21}&\chi_{22}&0\\ \chi_{30}&\chi_{31}&\chi_{32}&\chi_{33} \end{bmatrix}.$$零元素保证了因果性。对角线包含在内,因为 token $i$ 在完成自身写入后才进行读取。
现在将边界衰减后的查询 ket 向量按列堆叠:
$$Q_c^{\mathrm{boundary}} = \begin{bmatrix} D_0\lvert q_0\rangle& D_0D_1\lvert q_1\rangle& D_0D_1D_2\lvert q_2\rangle& D_0D_1D_2D_3\lvert q_3\rangle \end{bmatrix},$$输出 ket 向量也按同样方式堆叠。所有四个输出为:
$$\boxed{ O_c = sS_cQ_c^{\mathrm{boundary}} + E_c\left(A^{qk}_c\right)^\mathsf T. }$$第一个矩阵乘积读取了经过适当衰减的输入状态。第二个矩阵乘积则添加了块内写入操作的因果贡献。这是所需的第二个块结果。
6.6 Triton 管线的组织方式
分块实现将上述公式转化为一系列内核启动的管线,而非单个巨型内核。
它首先计算分块内的累积对数衰减。两个前缀和之间的差值编码了 $D_{j+1:i}$,而无需显式地乘以一长串保留向量。然后,它构建因果 $A^{qk}$ 和 $A^{kk}$ 交互矩阵。$A^{kk}$ 用于形成该分块修正写入的 WY 风格表示。
一个状态核执行唯一的分块间扫描,生成进入每个分块的状态并解决其 delta 误差。一旦这些传入状态已知,输出核就可以并行计算不同分块和 tile 中的 token。
实际源码包含这些阶段的融合和分块变体。具体来说,它在融合的非对角线和三角求解核之前,计算 16-token 的对角线交互块。前向数据流大致如下:
def chunkwise_kda(q, k, v, log_decay, beta, initial_state, scale):
# 分块内前缀和。G[i] - G[j] 编码了将状态或写入从 token j 携带到 token i 的衰减。
G = chunk_local_cumsum(log_decay)
# 构建因果查询-键交互以及解决 delta 误差之间依赖关系的三角系统。
A_qk_diag, A_kk_diag = intra_token_parallel(
q, k, G, beta, scale
)
A_qk, A_kk = inter_and_triangular_solve(
q, k, G, beta, A_qk_diag, A_kk_diag, scale
)
# 将分块转换为其 WY 伪键/伪值形式。
W, U, K_to_end = build_wy_factors(k, v, G, beta, A_kk)
# 剩下的唯一循环是在分块边界上。
H, E, final_state = scan_chunk_states(
K_to_end, W, U, G, initial_state
)
# 将每个传入分块状态的读取与该分块内写入的因果贡献相结合。
output = calculate_outputs(q, E, G, A_qk, H, scale)
return output, final_state
在参考源码中,这些阶段由
chunk_kda_fwd
编排。其主要实现入口点是 chunk_kda_fwd_intra、chunk_gated_delta_rule_fwd_h 和 chunk_gla_fwd_o_gk。代码中的 v_new、h 和 kg 等名称分别对应已解决的误差、传入的分块状态以及衰减到分块末尾的键。
因此,循环模式和分块模式并非两种不同的注意力机制。它们只是同一 KDA 递推关系的两种调度方式:串行向量运算用于低延迟解码,分块矩阵运算用于高张量核心的训练和预填充。
引用本文@misc{doubleword-you-could-have-come-up-with-kimi-delta-attention,
title = {You Could Have Come Up With Kimi Delta Attention},
author = {Jamie Dborin},
year = {2026},
howpublished = {Doubleword Blog},
url = {https://blog.doubleword.ai/you-could-have-come-up-with-kimi-delta-attention},
}