Reed's News
← 返回精选

DeltaNet 系列线性注意力变体的推导之旅

AI 73 AnhTho_FR 2026/7/28 2235 字 原文 ↗

符号说明:本文默认采用bra-ket符号(狄拉克符号),因为在我受量子力学启发的视角中,它能让推导过程的结构更清晰。上方的"数学符号切换"功能可将所有公式转换为常规的粗体向量与显式转置形式。在bra-ket模式下,|·⟩为列向量,⟨·|为行向量,⟨·|·⟩为标量,|·⟩⟨·|为矩阵。向量默认朝右书写,而在写入线性注意力状态时,键(key)朝左书写。本文仅讨论单个因果注意力头与实值向量,假设DeltaNet的键已归一化,且状态为从键空间到值空间的映射。

现代线性注意力变体的结构十分复杂,初看很难理解其设计目标。作为参考,以下是Kimi Delta Attention(KDA,Kimi delta注意力)的状态更新公式:

这类变体难以理解的原因在于,它们是过去几年间迭代出的线性注意力家族的最新成果,复杂度不可避免地不断累积,导致外界难以理解最新版本的原理。

本文将梳理DeltaNet系列线性注意力变体的演进路径(其中两款被最新的Qwen与Kimi模型家族采用),展示如何通过对隐藏状态做出简单假设,推导出相同的公式。

我们的推导路径如下: softmax注意力 → 线性注意力 → DeltaNet门控DeltaNetKDA

推导完KDA后,我们再介绍实现它的循环式与分块式Triton程序。

对于第$t$个token的查询,常规因果softmax注意力的计算方式为:

每个注意力权重都是一个标量,用于衡量单个键与单个查询的相似度,随后softmax将该查询的所有相似度得分转化为概率分布,最终输出为值向量的加权和。

对于长度为$n$的序列,键-查询对的数量为$n^2$。自回归推理时,我们可以缓存键和值以避免重复计算,但缓存仍会随序列长度增长,且每个新查询仍需遍历整个历史序列。

阻碍这类计算重构的核心是softmax函数,其分母同时依赖当前查询与所有历史键。因此,我们先暂时移除softmax。

为简化表达,将常数缩放因子融入查询中,得到简化版注意力公式:

标量内积可移至右侧:

此时,所有依赖历史信息的项可整合为一个固定大小的矩阵$S_t$:

注意力计算由此转化为"循环写入+读取"的过程:

核心技巧在于利用恒等式:

其中外积为矩阵,内积为标量。我们无需再存储所有历史键与值,只需将它们的外积之和存入固定大小的状态$S_t$即可。

这一方法的时间复杂度为序列长度的线性函数,而非平方函数:只需遍历一次token,每一步更新同一个状态。但为了效率,我们舍弃了softmax的归一化与选择性。更复杂的线性注意力方法会使用特征映射与归一化器,但这种简化形式恰好暴露了DeltaNet试图解决的内存问题。

假设我们写入一对$(k_t, v_t)$,并立即用同一个键查询新状态:

写入操作不会让内存返回$v_t$,而是将$k_t v_t^\top$加到内存已有的返回值上。

如果旧状态原本就能输出正确值,累加式写入会让新状态的输出变为原来的两倍。更普遍的情况是,键之间并非相互正交,因此每次写入都会干扰之前的写入。线性注意力为我们提供了紧凑的关联内存,但其更新方式是+=,而我们实际需要的更接近=

DeltaNet用delta规则修正替代了线性注意力中无条件的写入操作,有两种实用的推导方式。

在写入第$t$个token前,先查询当前内存中与新键关联的值:

如果我们希望内存返回$v_t$,就不应写入完整的$v_t$,而只需写入差值:

引入一个可学习的写入强度$\beta_t$,定义:

随后将该误差写入当前键对应的位置:

此时立即用同一个键查询:

当$\beta_t=1$时,结果恰好为$v_t$;$\beta_t$越小,旧预测值向目标值调整的幅度就越小。

这种修正在键空间中是局部的:对于任何与当前键正交的查询$q$,有:

因此,秩1写入仅会改变所选键方向上的响应,而不影响所有正交方向。

同样的更新规则也可从在线学习目标推导得出:将当前键值对视为线性映射$S_t: k \mapsto S_t k$的一个训练样本:

其关于状态$S_{t-1}$的梯度为:

这显然是一个外积:值空间的预测误差乘以观测到该误差的键bra。从$S_{t-1}$出发,进行步长为$\beta_t$的梯度下降更新:

这与我们要求"立即重构"得到的更新规则完全一致。两种解读本质相同:

  • 从内存操作的角度,$\beta_t$控制旧关联被替换的强度;
  • 从在线学习的角度,$\beta_t$是梯度下降的步长;
  • 从线性代数的角度,状态变化是一个秩1外积。

展开误差项后,DeltaNet可被视为结构化状态转移加上新输入:

对于单位键$k_t$,$I - \beta_t k_t k_t^\top$在当前键方向上的特征值为$1-\beta_t$,在所有正交方向上的特征值为1。它会先移除当前键方向上的旧关联,再添加新关联。

DeltaNet解决了写入问题,但尚未解决状态的生命周期问题。

线性状态将整个历史压缩为一个矩阵,读取操作:

无法在某个旧token被整合进$S_t$后,单独跳过该token的影响。所有与查询重叠的存储方向都会贡献结果。delta规则可以修正当前键附近的状态,但其他方向上的陈旧信息仍会保留,可能干扰未来的读取操作。

因此,我们需要一种在使用旧状态前将其遗忘的方法。设$\gamma_t$为可学习的标量保留门:

基于这个带门控的状态运行delta规则:

这就是门控DeltaNet。操作顺序至关重要:先遗忘,再基于保留的状态预测,最后修正预测结果。如果先预测再遗忘,误差描述的将是与我们要更新的内存不同的另一个内存。

展开循环式更新可得:

delta规则实现针对性替换,标量门控实现全局擦除,二者解决不同问题,互为补充。

但$\gamma_t$仍会对整个矩阵做出统一决策,模型必须以相同速率保留或遗忘所有键通道。

Kimi Delta Attention将门控DeltaNet的标量保留门替换为向量$\boldsymbol{\gamma}_t \in \mathbb{R}^K$,并将该向量置于对角矩阵中:

由于我们的状态是从键到值的映射,键通道对应$S_t$的列。右乘$\Gamma_t$会为每一列应用不同的保留因子:

其余部分仍沿用我们已推导的delta规则:

这就是KDA。与门控DeltaNet相比,核心变化仅在于:

这一改动的影响显著:一个通道可被清除,而另一个通道可被保留。

展开KDA的修正项:

键空间的转移矩阵为:

其中

因此$\Gamma_t - \beta_t \boldsymbol{\gamma}_t k_t^\top k_t$是一个对角矩阵减去一个秩1矩阵,即**对角加低秩(DPLR)**转移矩阵。"DPLR"描述的是作用于键空间的转移,而内存状态本身仍是矩阵$S_t$。

我们可将整个演进路径简洁总结如下:

机制 状态更新 新增特性
线性注意力 固定大小的循环内存 线性时间复杂度
DeltaNet 针对性替换 解决写入干扰问题
门控DeltaNet 先应用$\gamma_t$,再执行delta更新 全局状态遗忘
KDA 先应用$\boldsymbol{\gamma}_t$,再执行delta更新 按键通道遗忘

实现时通常将$\boldsymbol{\gamma}_t$存储为$\log \boldsymbol{\gamma}_t$,再通过$\exp$得到保留因子。在参考代码使用的转置布局中,循环式更新仅需5行代码:

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参考实现。

上述循环式实现是自回归解码的自然选择。KDA主要有两种执行模式:

模式 最佳适用场景 并行单元
融合循环式 解码、短序列、有状态服务 单序列、值头、值分块
分块式 训练与长序列预填充 序列块、token子块、键/值分块

循环式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,但在训练与长序列预填充场景中吸引力较低,因为这些向量操作无法转化为张量核最高效处理的大型矩阵乘法。

这促使我们以另一种视角重新审视同一个循环式更新。

分块式KDA将多个token批量处理,必须生成与逐token循环式更新完全相同的状态与输出,但会将计算重组为矩阵乘法。

对于每个块$C$,我们需要两个结果:

  • 给定输入状态$S_{\text{in}}$,处理完整个块后的状态;
  • 块内每个因果token的输出。

唯一的难点在于,第$t$个token的delta误差依赖同一块中更早token的写入操作。一个四token的例子可以清晰展示这些依赖关系。

取token$1,2,3,4$,定义:

块边界到第$i$个token的累积衰减为:

第$j$个token的写入传递到第$i$个token时的衰减为:

当无中间衰减时$\Gamma_{j\to i}=I$。所有这些矩阵都是对角矩阵,因此彼此可交换。

首先假设每个token仅能看到经过适当衰减的输入状态,无法看到同一块内其他token的写入:

对于四个token,可得到四个临时值空间误差ket:

这些值可并行计算,但除第一个外其余都是错误的——同一块中更早的写入也会影响它们的预测结果。

第1个token没有块内前置写入,因此:

第2个token会看到经过$\Gamma_{1\to2}$衰减后的第1个token的写入:

第3个token会看到前两个token的写入:

第4个token会看到前三个token的写入:

每个括号内都是标量。定义因果键-键系数:

则四个方程可简化为统一形式:

将系数整理为严格下三角矩阵:

将误差ket作为列堆叠:

$\boldsymbol{V}$与$\boldsymbol{K}$同理。此时因果替换可表示为:

实现时无需构造通用的稠密逆矩阵。由于$M$是对角元为1的三角矩阵,该操作可转化为因果三角求解,且可在每个值通道上独立执行。

在块处理结束时,输入状态经过了四次衰减,而每个块内写入仅会经过其之后的衰减:

定义矩阵$\tilde{\boldsymbol{K}}$,其行为到达块末尾时的键:

由于$\boldsymbol{E}$将误差ket作为列堆叠,四个外积写入可合并为一次矩阵乘法:

这是块处理的第一个所需结果:一次性将循环状态推进四个token。

KDA先写入后读取。设$S_i$为第$i$个token写入后的局部状态,则:

展开四个输出:

定义因果查询-键系数:

并将系数放入下三角读取矩阵:

零元素用于保证因果性,包含对角元是因为第$i$个token会在自身写入后执行读取。

将经过块边界衰减的查询ket作为列堆叠:

输出ket同理堆叠。此时四个输出可表示为:

第一个矩阵乘积读取经过适当衰减的输入状态,第二个矩阵乘积则添加块内写入的因果贡献。这是块处理的第二个所需结果。

分块式实现将上述方程转化为一系列内核启动的流水线,而非单个大型内核。

它首先计算块内累积对数衰减,通过两个前缀和的差值来表示$\Gamma_{j\to i}$,无需显式相乘一长串保留向量。随后构造因果矩阵与交互矩阵,利用$M^{-1}$生成块修正写入的WY式表示。

状态内核执行唯一的块间扫描,生成每个块的输入状态并求解其delta误差。得到这些输入状态后,输出内核即可并行计算不同块与分块中的token。

实际源码包含这些步骤的融合与分块变体,尤其是会先计算16token的对角交互块,再执行融合的非对角与三角求解内核。前向数据流大致如下:

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_intrachunk_gated_delta_rule_fwd_hchunk_gla_fwd_o_gk。代码中的v_newhkg等名称分别对应求解后的误差、块输入状态,以及衰减到块末尾的键。

因此,循环式与分块式程序并非两种不同的注意力机制,而是同一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},
}