卡尔曼Delta网络:不确定性感知的联想记忆
Kalman Delta Networks: Uncertainty-aware Associative Memory
- Yale University(耶鲁大学)
机构由 AI 辅助整理,请以论文原文为准。
中文总结 AI 辅助
针对线性注意力固定记忆无法感知不确定性的问题,提出卡尔曼Delta网络,将联想记忆建模为线性高斯状态空间模型,用卡尔曼滤波跟踪不确定性,并给出两种可并行扫描的近似,提升预训练困惑度和下游准确率。
中文摘要 AI 辅助
线性注意力越来越多地用于前沿语言模型,以实现高效的长上下文推理和恒定内存解码。然而,其固定大小的循环记忆在每个词元处需要一个在线决策:在知道未来查询需要哪些信息之前,写入什么以及以多强的程度覆盖现有关联。Delta规则模型从当前词元嵌入中学习这种强度,但不跟踪记忆估计中的置信度,从而阻止每次写入适应累积的证据。为了显式表示这种不确定性,我们将循环联想记忆重新表述为线性-高斯状态空间模型,对于该模型,卡尔曼滤波器是最优递归估计器,并引入了一个新的模型家族:卡尔曼Delta网络(KDNs)。在KDNs中,转移同时传播记忆状态及其不确定性,使得卡尔曼增益能够根据累积证据和观测可靠性对每次残差写入进行加权。在此公式下,Delta式更新作为一个特例出现,该特例用逐词元各向同性替代品替代预测协方差,并省略协方差跟踪。然而,精确跟踪需要密集的、状态相关的Riccati递归,这不太适合GPU并行的线性注意力扫描。为解决此问题,我们引入了两种兼容扫描的KDN近似。对角KDN通过在线平均场变分推断将每个单步后验投影到对角高斯族,而各向同性KDN使用各向同性近似,每个头具有单个不确定性标量。它们的不确定性递归是Mobius映射,使得联想扫描具有对数并行深度。在750M和1.3B参数的受控预训练中,KDN变体在困惑度和平均下游准确率上持续优于最先进的线性注意力模型。
英文摘要
Linear attention enables efficient long-context inference by compressing token history into a fixed-size recurrent memory. This compression makes each update a trade-off between incorporating new information and preserving useful associations. Models such as DeltaNet, Gated DeltaNet, and KDA predict write strength from the current token representation, without explicitly tracking uncertainty in the stored memory. Yet this uncertainty matters: a new observation should have greater influence when the existing association is uncertain and less when it is already well supported. We introduce Kalman Delta Networks (KDNs), a family of linear-attention models that explicitly track memory uncertainty to guide each update. By formulating associative memory as a linear-Gaussian state-space model, KDNs propagate both the memory estimate and its uncertainty, using the Kalman gain to balance accumulated evidence against the reliability of new observations. This formulation also recovers standard delta-rule updates by replacing tracked covariance with a token-predicted isotropic surrogate. To support hardware-efficient training and inference, we derive Diagonal KDN and Isotropic KDN, which retain one uncertainty value per key channel and per head, respectively. Their uncertainty updates admit associative scans with logarithmic parallel depth, requiring only $O(d_k)$ and $O(1)$ auxiliary state per head. Across controlled pretraining at 750M and 1.3B parameters, both variants consistently improve perplexity and mean downstream accuracy over the evaluated state-of-the-art linear-attention baselines.