卡尔曼增量网络:不确定性感知的联想记忆 Kalman Delta Networks: Uncertainty-aware Associative Memory
将Delta规则线性注意力重构为线性高斯状态空间模型,用卡尔曼增益按不确定性自适应调节记忆写入
前置知识
线性注意力
Softmax 注意力保留全部历史键值对,计算量随序列长度二次增长;线性注意力则把前缀压缩为固定尺寸的循环状态 $S_t \in \mathbb{R}^{d_k \times d_v}$,每个 token 做加法写入 $S_t = S_{t-1} + k_t v_t^\top$,查询读取 $o_t = S_t^\top q_t$。状态可用并行扫描训练,解码只需常数内存,是长上下文高效建模的主流路线。
KDN 是线性注意力家族的新成员,全文的效率约束(固定状态、scan 并行训练、常数解码内存)都由这一范式决定,也是所有基线共处的技术坐标系。
Delta 规则(DeltaNet)
DeltaNet 发现加性写入无法删除旧关联,于是先读出记忆对当前键的预测 $S_{t-1}^\top k_t$,只写入残差:$S_t = (I - \beta_t k_t k_t^\top) S_{t-1} + \beta_t k_t v_t^\top$,可解释为对瞬时预测损失做一步在线梯度下降;$\beta_t$ 是由当前 token 表示预测的写入强度门控。
本文证明 Delta 规则正是卡尔曼滤波更新中用各向同性协方差替代真实不确定性后的固定增益特例,这一统一视角是全文定位和创新的出发点。
卡尔曼滤波
线性高斯状态空间模型的最优递归估计器。每步先“预测”($\hat{S}_t = D_t S_{t-1}$,$\hat{P}_t = D_t P_{t-1} D_t^\top + \Omega_t$),再用新观测“更新”,增益 $\kappa_t = \hat{P}_t k_t / (r_t + k_t^\top \hat{P}_t k_t)$ 自动平衡先验不确定性与观测噪声:记忆越不确定越信任新证据,越确定越保护旧关联。
KDN 的全部新颖性都来自把线性注意力的写入增益换成这个由协方差递推导出的卡尔曼增益,理解它是读懂本文的前提。
选择性状态空间模型(Mamba 系列)
Mamba 用循环 $h_t = A_t h_{t-1} + B_t x_t$ 处理序列,$A_t, B_t, C_t$ 随输入变化(选择性),配合硬件感知的并行扫描实现线性复杂度。Mamba-2 建立了与结构化掩码注意力的对偶(SSD),Mamba-3 加入指数梯形离散化、复值动力学与 SISO/MIMO 变体。
KDN 与 Mamba 同属 SSM 谱系但写入方式不同:Mamba 直接加控制输入,KDN 做键条件化的残差修正;实验中 Mamba-3 也是最主要的对比基线之一。
关联扫描(Associative Scan)
若每步状态更新可写成满足结合律的映射(如仿射 $S_t = A_t S_{t-1} + b_t$),就能像前缀和一样用并行扫描在 $O(\log n)$ 深度内算出所有时间步的状态,这是线性注意力能高效利用 GPU 的关键。KDN 把不确定性递推表示为 Möbius 映射,用 $2 \times 2$ 矩阵乘法复合,从而同样可扫描并行。
精确卡尔曼滤波的 Riccati 递推是状态依赖的、不满足该结构,如何设计 scan 兼容的近似正是本文要解决的核心工程难题。
平均场变分推断
当精确后验难以表示时,在易处理的分布族(如对角高斯)中寻找最小化反向 KL 散度 $KL(q \| p)$ 的近似后验。平均场假设各坐标独立,代价是丢弃相关性、逐坐标过度自信。
对角 KDN 每步把稠密的精确后验协方差投影回对角族用的正是这一工具,其过度自信缺陷直接催生了本文的信息缩放因子 $\mu$。
研究动机
Softmax 注意力的 KV 缓存随上下文线性增长、query-key 交互二次方增长,前沿模型因此越来越依赖线性注意力:把历史压缩进固定尺寸循环记忆,换取常数内存解码与线性吞吐。但这把注意力变成了在线内存管理问题——每个 token 必须在还不知道未来查询需要什么的条件下,决定往记忆里写什么、用多大力气覆盖旧关联。现有 Delta 规则模型(DeltaNet、Gated DeltaNet、KDA)的写入强度 $\beta_t$ 只由当前 token 表示预测,是与历史证据无关的“固定增益”:它无法区分一条被反复确认的可靠关联和一条只见过一次的临时关联,既可能在证据充分时过度覆盖,也可能在记忆本不可靠时写入不足。此外 Gated DeltaNet 的标量衰减 $\alpha_t$ 对同一头内所有键通道一视同仁,无法同时保存长寿命与短寿命信息;KDA 虽给出每通道衰减 $D_t = \mathrm{diag}(\alpha_t)$,写入门控依旧只是 token 预测的标量。
本文的目标是本文要为循环联想记忆补上“不确定性”这个缺失的状态变量:把记忆建模为对一个潜在、非平稳键值映射的在线估计,每个 token 只在单一键方向上提供一次带噪观测,从而用线性高斯状态空间模型的最优递归估计器——卡尔曼滤波——统一推导写入规则,让每次写入的增益由累积证据与观测噪声共同决定:记忆不确定时多采信新证据,关联被反复验证时则加以保护。同时作者要求算法不牺牲线性注意力的硬件效率:不确定性状态要小(每头 $O(d_k)$ 或 $O(1)$),递推必须能写成满足结合律的映射以支持 GPU 并行扫描训练;最后在 750M/50B 与 1.3B/100B 的受控预训练中,与 Mamba-3、KDA、GDN-2 等最强循环混频器做数据、骨干、优化配方完全对齐的公平比较。
与已有工作不同的是,此前工作用两个互不相通的视角理解 Delta 混频器:状态空间动力学(衰减、门控)与在线优化(快速权重的一步梯度)。本文的独特切入点是第三个视角——贝叶斯滤波:对潜在记忆假设 $\tilde{S}_t = D_t \tilde{S}_{t-1} + W_t$、$W_t \sim \mathcal{N}_{col}(0, \Omega_t)$,观测为 $v_t = \tilde{S}_t^\top k_t + e_t$,则 Delta 规则的残差写入恰好是卡尔曼滤波的新息更新,而 DeltaNet/GDN/KDA 都成为丢弃协方差、用各向同性代理 $\hat{P}_t \approx \hat{b}_t I$ 的固定增益特例;遗忘也不再是附加门控,而是转移模型对非平稳记忆的预测。与同样跟踪不确定性的 Kalman Linear Attention(时间不变 OU 动力学下沿特征坐标分解信念)不同,KDN 以键条件化的向量观测跟踪键空间中跨值通道共享的稠密后验协方差,并给出 scan 兼容的对角/各向同性近似。
核心方法
直觉上,记忆应当“对没把握的关联多听新证据,对反复验证的关联保持固执”。作者把循环状态看作潜在键值映射的贝叶斯信念:每步先预测 $\hat{S}_t = D_t S_{t-1}$、$\hat{P}_t = D_t P_{t-1} D_t^\top + \Omega_t$,再用 $(k_t, v_t)$ 更新 $S_t = \hat{S}_t + \kappa_t (v_t - \hat{S}_t^\top k_t)^\top$,增益 $\kappa_t = \hat{P}_t k_t / (r_t + k_t^\top \hat{P}_t k_t)$ 由预测不确定性与观测噪声 $r_t$ 共同决定。精确执行需每头维护稠密 $d_k \times d_k$ 协方差并递推状态依赖的 Riccati 方程,无法并行扫描,因此 KDN 家族给出两个 scan 兼容近似:各向同性 KDN 每头只跟踪一个标量不确定性 $\hat{b}_t$($O(1)$ 辅助状态),对角 KDN 每个键通道跟踪一个不确定性值($O(d_k)$ 辅助状态)。两者的不确定性递推均为 Möbius 映射,可写成 $2 \times 2$ 矩阵乘法的关联扫描、以对数并行深度求解,再接常规的仿射记忆扫描完成训练与推理。
核心创新是把写入增益从“token 预测的独立门控”变成“由协方差递推导出的卡尔曼增益”。理论上,若用各向同性代理 $\hat{P}_t \approx \hat{b}_t I$ 并丢弃协方差跟踪,卡尔曼增益退化为 $\beta_t k_t$,正好还原 DeltaNet/GDN/KDA 的共享残差形式,三者只在过程模型上不同:$D_t = I$、$\alpha_t I$、$\mathrm{diag}(\alpha_t)$。KDN 恢复了被丢弃的不确定性动力学:衰减 $\alpha_t$、过程噪声 $\omega_t$、观测噪声 $r_t$ 都由 token 表示经 $\sigma$/softplus 参数化,增益因此反映“这个键方向上已积累多少证据”。对角 KDN 另有两个精细设计:其一,用在线平均场变分推断把稠密的一步后验投影回对角族,且保持精确后验均值不变;其二,针对反向 KL 的逐坐标过度自信——丢弃跨通道相关性使模型沿已观测的联合键方向反而过于不确定、导致重复键触发过度写入——引入信息缩放因子 $\mu$,只缩放写后精度增量 $u_t = (\mu / r_t) k_t \odot k_t$:增大 $\mu$ 不影响当前写入、只削弱未来覆盖,$\mu = d_k$ 补偿归一化稠密键的 $1/d_k$ 信息稀释。
方法步骤详情
以对角 KDN 为主(每头 $d_k = 128$)。第一步,由 token 隐状态 $x_t$ 预测每通道衰减 $\alpha_t = \sigma(W_\alpha x_t + b_\alpha)$、过程噪声 $\omega_t = \mathrm{softplus}(W_\omega x_t + b_\omega)$、观测噪声 $r_t = r_{\min} + \mathrm{softplus}(w_r^\top x_t + b_r)$($r_{\min} = 0.01$,初始化 $W_\omega = 0$,每头学习标量初始协方差 $c_{0,h}$)。第二步预测:$\hat{p}_t = \alpha_t^2 \odot p_{t-1} + \omega_t$。第三步算增益:$\kappa_{t,i} = \hat{p}_{t,i} k_{t,i} / (r_t + \sum_i \hat{p}_{t,i} k_{t,i}^2)$。第四步写入并读出:$S_t = (I - \kappa_t k_t^\top) D_t S_{t-1} + \kappa_t v_t^\top$,$o_t = S_t^\top q_t$。第五步更新不确定性:$u_t = (\mu / r_t) k_t \odot k_t$,$p_t = \hat{p}_t / (1 + u_t \odot \hat{p}_t)$,每通道满足 Möbius 递推 $p_{t,i} = (\alpha_{t,i}^2 p_{t-1,i} + \omega_{t,i}) / (u_{t,i} \alpha_{t,i}^2 p_{t-1,i} + 1 + u_{t,i} \omega_{t,i})$,写成 $2 \times 2$ 矩阵做关联扫描并行求解。实现上先跑协方差扫描得到全部增益,再用 compact-WY 核做记忆扫描;各向同性 KDN 改用迹投影 $\hat{b}_t = a_t b_{t-1} + \omega_t / d_k$($a_t = \frac{1}{d_k} \sum_i \alpha_{t,i}^2$)与标量增益 $\beta_t = \hat{b}_t / (r_t + \hat{b}_t \|k_t\|_2^2)$,全程固定 $\mu = d_k$。
技术新颖性
技术新颖性有四层。第一,理论统一:首次把 DeltaNet、Gated DeltaNet、KDA 统一为线性高斯滤波器的固定增益特例,把残差写入解释为卡尔曼新息、遗忘解释为转移模型,从而把 Delta 混频器与 Mamba 谱系接进同一数学框架——但与 Mamba 的控制驱动加性写入 $B_t x_t$ 不同,KDN 是键条件化的残差修正。第二,算法结构:作者发现不确定性递推是 Möbius 映射,因此可表示为 $2 \times 2$ 矩阵复合并纳入关联扫描,绕开了精确 Riccati 递推的状态依赖性;对角情形的在线平均场投影保持精确后验均值(命题 4.1)。第三,校准机制:识别出对角近似沿已观测键方向过度不确定、重复键过度写入的失效模式,提出只作用于写后精度增量的信息缩放 $\mu$,兼顾当前写入与未来覆盖。第四,与近邻工作的区分:Preconditioned DeltaNet 的精确在线最小二乘更新对应本文框架的静态固定噪声极限,其增益仍是 token 预测的;Kalman Linear Attention 在时间不变 OU 动力学下沿特征坐标分解信念,而 KDN 用键条件化观测跟踪键空间共享协方差,并以各向同性/对角近似换取 scan 兼容性。
实验结果
在 GDN-2 训练协议(AdamW、峰值学习率 $4 \times 10^{-4}$、全局 batch 0.5M token、序列 4K、余弦调度)下的受控实验结论一致。语言建模与常识推理(表 1):750M/50B 循环组中,对角 KDN 达 WikiText 18.64、LAMBADA 14.15、六任务平均 54.97%,优于 KDA(18.85/15.06/53.87)、Mamba-3 MIMO(18.99/15.67/54.39)与 GDN-2(21.20/17.88/51.45);各向同性 KDN 拿下最佳 WikiText 18.42。1.3B/100B 组对角 KDN 15.04/9.75/60.45 全面领先 KDA(15.40/10.09/60.28)。混合组(9 循环 + 9 层 2K 滑窗注意力)对角 KDN+SWA 平均 60.11% 最高,超过 KDA+SWA 59.94% 与纯 Transformer 56.23%。上下文检索(表 2):750M 对角 KDN 是唯一在 S-NIAH-1 全部长度(含 8K)保持 100% 的模型(KDA 82.6%、GDN-2 47.2%),S-NIAH-2@4K 71.2% 对 KDA 63.2%;1.3B 下 S-NIAH-3 达 96.4/88.6/48.6(GDN-2 为 88.0/64.4/29.4),两个规模均取得 14 格 RULER 总分第一。真实检索(表 3):循环组对角 KDN 平均 34.86% 最佳(KDA 33.76%),领先 FDA 30.61%、DROP 24.10%;混合组各向同性 KDN+SWA 45.94% 最佳,远超 Transformer 40.02%。消融(表 4/5):$\mu = d_k$ 综合最佳(54.97%),学习型 $\mu$ 的 WikiText 最佳(18.51)但整体略逊;固定 $r_t = 1$ 或 $r_t, \omega_t$ 均为 1 结果好坏参半。吞吐量(图 4):各向同性 KDN 与 KDA 几乎重合,对角 KDN 接近 GDN-2 且保持线性扩展(需 FP32/TF32x3)。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| 语言建模 WikiText(750M/50B 循环模型) | 困惑度 ↓ | 对角 KDN 18.64;各向同性 KDN 18.42 | KDA 18.85;Gated DeltaNet 19.50;GDN-2 21.20 | 比最强基线 KDA 低 0.21,比 GDN-2 低 2.56 |
| 语言建模 LAMBADA(1.3B/100B 循环模型) | 困惑度 ↓ | 对角 KDN 9.75 | KDA 10.09;Mamba-3 MIMO 10.49;GDN-2 11.29 | 比 KDA 低 0.34(约 3.4%) |
| 六任务零样本平均(LAMBADA/PIQA/HellaSwag/WinoGrande/ARC-e/ARC-c,1.3B 循环) | 平均准确率 ↑ | 对角 KDN 60.45% | KDA 60.28%;Mamba-3 MIMO 59.85%;GDN-2 58.51% | +0.17pt(vs KDA),+2.60pt(vs GDN-2) |
| RULER S-NIAH-1 大海捞针(750M 循环,8K 上下文) | 检索准确率 ↑ | 对角 KDN 100.0%(唯一全长度满分) | KDA 82.6%;GDN-2 47.2%;Mamba-3 SISO 33.2% | +17.4pt(vs KDA) |
| RULER S-NIAH-2 干扰检索(750M 循环,4K 上下文) | 检索准确率 ↑ | 对角 KDN 71.2% | KDA 63.2%;GDN-2 36.0% | +8.0pt(vs KDA),+35.2pt(vs GDN-2) |
| JRT 真实世界检索六任务平均(1.3B 混合模型,2K 输入) | 零样本准确率 ↑ | 各向同性 KDN+SWA 45.94%;对角 KDN+SWA 45.55% | KDA+SWA 45.09%;纯 Transformer(2K SWA)40.02% | +0.85pt(vs KDA+SWA),+5.92pt(vs Transformer) |
| 混频器层吞吐量(单张 H200,序列 2K–32K) | tokens/s/GPU | 各向同性 KDN 与 KDA 几乎重合;对角 KDN 接近 GDN-2、线性扩展(FP32/TF32x3) | KDA、GDN-2;全注意力随长度急剧衰减,Mamba-3 SISO 长上下文领先 | 不确定性跟踪的额外开销很小,保持线性复杂度 |
局限与改进
作者明确承认:KDN 只是“迈向”而非完整实现卡尔曼联想记忆——精确滤波需要稠密键空间协方差与状态依赖的 Riccati 递推,各向同性和对角近似都是为保住并行训练所做的压缩;对角投影丢弃跨通道相关性,必须额外引入信息缩放 $\mu$ 抑制过度覆盖,而表 4 显示 $\mu$ 的最优值依赖评测指标($d_k$ 综合最佳、$4d_k$ 的 LAMBADA 最佳、学习型 $\mu$ 的 WikiText 最佳),削弱了“原则性推导”叙事的完整性;表 5 显示把 $r_t, \omega_t$ 固定为 1 结果好坏参半,说明可学习噪声并非必要组件。我自己的观察:LM 与零样本指标上的领先幅度其实不大(困惑度约 0.2–0.4、平均准确率约 0.1–0.6pt),最亮眼的收益集中在 RULER 合成检索;真实世界检索上循环版仍然偏弱(JRT 平均 34.86%),要靠混合滑窗注意力才大幅提升到 45.94%,说明固定尺寸记忆的读出端仍是瓶颈;评测上下文最长仅 8K(RULER)/2K(JRT),未验证 32K 以上场景;对角 KDN 需 FP32/TF32x3 计算协方差扫描,数值稳定性与效率是部署隐患。
独立分析的弱点
第一,对角近似的信息丢失靠人工常数 $\mu = d_k$ 补救,且不同指标的最优值不同,说明校准并不彻底;改进方向是维护“对角 + 低秩”协方差,只对少量高频键方向保留稠密相关性,并让 $\mu$ 可学习或随记忆置信度自适应。第二,对角 KDN 的不确定性扫描需要 FP32/TF32x3,精度与速度双受损;可研究在 log 空间参数化精度、或设计专为 Möbius 扫描优化的低精度核。第三,写入端引入了不确定性,读出端却仍是简单的 $o_t = S_t^\top q_t$:JRT 结果表明固定记忆读出是真实检索的短板,可让查询端也利用不确定性,例如对高不确定度方向回退到全注意力或外置检索。第四,实验规模止步 1.3B/100B token,且基线限于循环与简单混合模型,缺少 7B 级、32K+ 长上下文与 MoE 配置的验证;混合比例也沿用 GDN-2 未做消融。第五,观测噪声假设各向同性 $r_t I_{d_v}$,对“值的哪些维度是上下文噪声”建模粗糙,可按值坐标或语义子空间分解 $r_t$,让“什么值得写入”的判断更细粒度。
未来方向
作者提出的方向:把对角衰减扩展为 Mamba-3 式的阻尼旋转,让已存关联既能衰减也能旋转(对应更丰富的转移与复值动力学),同时保持协方差更新的扫描效率,这是一个公开难题。基于本文成果还可延伸:其一,把 $\mu$ 从固定常数变为输入相关或按头学习的量,或把信息缩放推广为结构化的精度增量;其二,探索“对角 + 秩-$k$”协方差以部分恢复跨通道相关性,逼近精确卡尔曼滤波;其三,把不确定性显式用于下游行为——解码时按记忆置信度决定是否触发全注意力回看、作为检索式 Agent 的记忆可信度信号、或用于数据质量过滤;其四,KDN 的不确定性状态或许能支撑更激进的“长循环 + 短注意力”混合配置,混合比例搜索值得系统研究;其五,理论层面分析变分投影偏差随序列长度的累积行为,以及 $\mu = d_k$ 补偿的适用边界。
复现评估
复现条件较好:论文给出开源代码仓库(github.com/ngocbh/kalman-delta-networks),训练数据为公开的 FineWeb-Edu,评测全部是公开基准(WikiText、LAMBADA、PIQA、HellaSwag、WinoGrande、ARC、RULER、SWDE/SQuAD/FDA/TriviaQA/NQ/DROP)。训练配方交代完整:AdamW、峰值学习率 $4 \times 10^{-4}$、$\beta = (0.9, 0.95)$、权重衰减 0.1、梯度裁剪 1.0、余弦调度带 1% warmup、全局 batch 0.5M token、序列 4K;关键超参有 $r_{\min} = 0.01$、初始 $W_\omega = 0$、每头可学习初始协方差 $c_{0,h}$、固定 $\mu = d_k$,附录 C/D 给出投影与 chunkwise 实现推导。主要门槛在算力与工程:需要 750M×50B 与 1.3B×100B token 量级的预训练(估计数百 GPU 日,吞吐测试用 H200),且对角 KDN 的 Möbius 协方差扫描 + compact-WY 记忆扫描是自定义 chunkwise 核、需 FP32/TF32x3 保证数值稳定,复现者需要 Triton/CUDA 级实现能力;论文未公布逐 run 随机种子与完整训练曲线。
论文图表