你本可以提出 Kimi Delta 注意力
关于符号的说明:本文默认为 Bra-ket 符号,因为(在我以量子为灵感的观点中)它使得此推导中的形状非常清晰。上面的数学符号切换将每个方程重写为常规的粗体向量和显式的转置。在 Bra-ket 模式中,∣q⟩(lvert q angle)是列向量,⟨k∣(langle k vert)是行向量,⟨k∣q⟩(langle k vert q angle)是一个数字,∣v⟩⟨k∣(lvert v angle riangleleft k angle)是一个矩阵。向量默认面向右,而当写入线性注意力状态时,键面向左。我们将使用一个因果注意力头和实值向量,假设 DeltaNet 的键已标准化,并让状态从键空间映射到值空间。现代线性注意力变体是复杂的,乍一看并不容易理解它们旨在实现什么。作为参考,这里是 Kimi Delta 注意力(KDA)的状态更新方程:S~t=St−1Diag(αt) ilde S_t = S_{t-1} ext{Diag}(eta_t) ∣v^t⟩=S~t∣kt⟩ angle= ilde S_t| k_t angle ∣et⟩=βt(∣vt⟩−∣v^t⟩) vert e_t angle = eta_t igg( vert v_t angle - vert ilde v_t angle igg) St=S~t+∣et⟩⟨kt∣S_t = ilde S_t+ vert e_t angleigg angle⟨ k_t vert ∣ot⟩=St(dk−1/2∣qt⟩) vert o_t angle = S_tigg(d_k^{-1/2} vert q_t angleigg)它们之所以如此难以理解,是因为这是过去几年开发的最新一系列线性注意力变体,复杂性不可避免地膨胀,以至于从外部看,最新的变体似乎无法访问。在这篇文章中,我们将逐步探讨 DeltaNet 系列的线性注意力变体,其中两者被最新的 Qwen 和 Kimi 模型系列使用,并展示如何通过对隐藏状态进行简单的假设,你可能会到达相同的方程。这是我们将采取的路径:softmax 注意力 → 线性注意力 → DeltaNet → 门控 DeltaNet → KDA 只有在推导出 KDA 之后我们才会转向执行它的递归和分块 Triton 程序。 1. 从二次注意力开始 针对标记 tt 的查询,普通的因果 softmax 注意力为 ati=exp (s⟨ki∣qt⟩)∑j≤texp (s⟨kj∣qt⟩),s=dk−1/2,∣ot⟩=∑i≤tati∣vi⟩.egin{aligned} a_{ti} &= rac{ ext{exp}igg(sigg ag{k_i}{q_t}iggigg)}{ ext{Σ}_{j ext{≤}t} ext{exp}igg(s igg ag{k_j}{q_t}igg)igg}, ext{。},s=d_k^{-1/2},\ vert o_t angle &= ext{Σ}_{i ext{≤}t}a_{ti} vert v_i angle。 egin{aligned} 每个注意权重都是标量。它度量一个键和一个查询之间的相似性,然后 softmax 会将该查询的所有分数转化为一个分布。输出是值向量的加权和。在长度为 TT 的序列中,有 T2T^2 对键-查询。在自回归推理的过程中,我们可以缓存键和值,而不是重新计算它们,但缓存仍会随着序列的增长而增长,每个新的查询仍然必须检查整个历史。重组此计算的障碍是 softmax。它的分母同时依赖于当前查询和所有早期的键。所以,暂时去掉它。 1.1 去掉 softmax 为了清楚起见,将常量规模 ss 吸收到查询中。于是基于注意力的简化版本为 ∣ot⟩=∑i≤t⟨ki∣qt⟩∣vi⟩。 vert o_t angle = ext{Σ}_{i ext{≤}t} igg ag{k_i}{q_t} angle igg ag{v_i}{rangle} 社内积标量向右移动: ∣ot⟩=∑i≤t∣vi⟩⟨ki∣qt⟩=(∑i≤t∣vi⟩⟨ki∣)∣qt⟩.egin{aligned} vert o_t angle &= ext{Σ}_{i ext{≤}t} vert v_i angle ag{k_i} angle ag{q_t}\ &= igg( ext{Σ}_{i ext{≤}t} vert v_i angle ag{k_i}iggiggl) ag{q_t}。 egin{aligned} 所有依赖于过去的东西现在可以汇聚成一个固定大小 V×KV imes K 的矩阵: St=∑i≤t∣vi⟩⟨ki∣oxed{ S_t = ext{Σ}_{i ext{≤}t} vert v_i angle ag{k_i} } 注意力变成了一个递归写入,随后是读取: St=St−1+∣vt⟩⟨kt∣,∣ot⟩=St∣qt⟩。oxed{ egin{aligned} S_t &= S_{t-1} + vert v_t ag{k_t} angle,\ vert o_t ag angle &= S_t vert q_t ag{q} angle。 egin{aligned} 身份 (∣v⟩⟨k∣)∣q⟩=⟨k∣q⟩∣v⟩igg ag{v angle ag{k angle} vert q angle = ⟨k vert q angleigg ag{k angleiggtag{v angle是整个操作。外积是一个矩阵;内积是一个数字。我们不再存储每个过去的键和值。我们存储它们累加的外积在固定大小的状态 StS_t 中。这是序列长度的线性,而不是二次:扫描每个 token,一步更新相同的 d_v imes d_k 状态。我们为这样的效率付出了代价,舍弃了 softmax 的规范和选择性。更加复杂的线性注意力方法使用特征图和标准化器,但这种简单形式揭示了 DeltaNet 的动机记忆问题。 1.2 加法不是赋值 假设我们写了一对 ∣vt⟩⟨kt∣ vert v_t angleigg ag{k_t} 并立即使用同一个键查询新的状态: St∣kt⟩=(St−1+∣vt⟩⟨kt∣)∣kt⟩=St−1∣kt⟩+∣vt⟩ ag{k_t}kt⟩⏟1=St−
本站免费、广告极少。如果觉得有帮助,可以请我们喝杯咖啡 —— 任何金额都对持续运营有实际帮助。
☕请我喝杯咖啡