Kimi DeltaNet 线性注意力系列指南
本文详细推导了DeltaNet系列线性注意力机制,包括softmax注意力、线性注意力、DeltaNet、门控DeltaNet和Kimi Delta注意力(KDA)。解释了每个变体如何通过解决记忆、选择性和遗忘问题来改进前一个。文章使用bra-ket表示法,并强调模型背后的数学直觉。
在这篇文章中,我们将逐步推导Kimi Delta Attention(KDA)及其前身,从标准的softmax注意力开始,最终到达KDA这一高效的线性注意力变体。文章采用bra-ket表示法,因为这种表示法能清晰地展示推导中的形状。在bra-ket模式中,∣q⟩表示列向量,⟨k∣表示行向量,⟨k∣q⟩是一个数,而∣v⟩⟨k∣是一个矩阵。我们假设只有一个因果注意力头,向量为实数,DeltaNet的键已归一化,状态从键空间映射到值空间。
首先,我们从标准的softmax注意力开始。对于第t个token的查询,注意力权重计算为ati = exp(s⟨ki∣qt⟩) / ∑_{j≤t} exp(s⟨kj∣qt⟩),其中s = d_k^{-1/2}。输出是所有值的加权和。该计算需要计算所有T^2个键-查询对,在自回归推理中,缓存会随着序列增长而增加,每个新查询仍需检查整个历史。
移除softmax后,我们得到线性注意力。将尺度s吸收到查询中,输出变为∣ot⟩ = ∑_{i≤t} ⟨ki∣qt⟩∣vi⟩ = (∑_{i≤t} ∣vi⟩⟨ki∣)∣qt⟩。定义状态St = ∑_{i≤t} ∣vi⟩⟨ki∣,则注意力变为递归写入后读取:St = S_{t-1} + ∣vt⟩⟨kt∣,∣ot⟩ = St∣qt⟩。这实现了O(T)的复杂度,但加法写入存在信息重叠问题:写入并不使记忆返回∣vt⟩,而是将∣vt⟩加到已有值上,当旧状态已产生正确值时,加法写入会使新状态产生两倍值。
为了解决这个问题,DeltaNet引入delta规则修正。在写入当前键值对之前,先从旧状态中预测v̂_t = S_{t-1}∣kt⟩,然后计算误差∣et⟩ = β_t (∣vt⟩ - v̂_t),最后写入误差:St = S_{t-1} + ∣et⟩⟨kt∣。当β_t=1时,立即读取同一键能得到精确值∣vt⟩。该修正具有局部性:对于与当前键正交的查询,写入不影响其输出。DeltaNet也可以从在线学习角度理解:它是对重构损失L_t(S) = 1/2 ||S∣kt⟩ - ∣vt⟩||^2的一步梯度下降,步长为β_t。DeltaNet的更新等价于St = S_{t-1}(I - β_t∣kt⟩⟨kt∣) + β_t∣vt⟩⟨kt∣,其中I - β_t∣kt⟩⟨kt∣在当前键方向上的特征值为1-β_t,正交方向上为1,从而在添加新关联前移除旧关联。
然而,DeltaNet无法清除历史中的陈旧信息。为了管理记忆生命周期,Gated DeltaNet添加了一个遗忘门α_t,在每次更新前将状态缩放:S̃_t = α_t S_{t-1},然后应用delta规则。这允许模型遗忘过时的关联。Gated DeltaNet的完整更新包括遗忘、预测、修正和写入四步。
最后,Kimi Delta Attention(KDA)结合了门控和delta规则。KDA的更新方程如下:首先遗忘:S̃_t = S_{t-1} Diag(α_t),其中α_t是向量,控制每个维度的遗忘;然后预测:v̂_t = S̃_t∣kt⟩;接着计算误差:∣et⟩ = β_t (∣vt⟩ - v̂_t);最后写入:St = S̃_t + ∣et⟩⟨kt∣;输出:∣ot⟩ = St (d_k^{-1/2}∣qt⟩)。KDA的推导源自对记忆操作的直观要求:写操作应能覆盖旧值,遗忘机制应能控制历史的持久性。
在推导完KDA后,文章还讨论了如何通过Triton实现循环和分块计算,以提高效率。总之,DeltaNet家族通过逐步改进,从简单的线性注意力发展到具有选择性写入和遗忘能力的复杂模型,为高效序列建模提供了理论基础。