面向持续学习的快速权重注意力 Fast Weight Attention for Continual Learning
把循环状态视为读后写语义下的在线学习器,导出归一化 Falcon 更新族,提升长度外推
前置知识
快速权重与线性注意力
快速权重指在前向传播内部被输入流快速修改的临时参数(如循环状态矩阵 $S_t$),与训练中缓慢更新的慢权重相对。线性注意力用核特征映射 $\phi(\cdot)$ 把 softmax 注意力改写为 $o_t = \frac{\sum_{j\le t} \phi(q_t)^\top\phi(k_j)v_j}{\sum_{j\le t}\phi(q_t)^\top\phi(k_j)}$,利用矩阵乘法结合律把上下文累积进固定尺寸状态 $S_t = S_{t-1} + \phi(k_t)v_t^\top$,实现 $O(N)$ 训练与 $O(1)$ 每步推理。
本文的研究对象正是这类固定尺寸快速记忆状态,全文把其更新规则重新解释为在线学习规则的一步梯度下降;不熟悉线性注意力的状态累积形式,就无法理解 Falcon 各变体究竟在修改什么。
Delta 规则与 DeltaNet
DeltaNet 将状态更新视为对值重构误差 $\ell_t(S) = \frac12\|S^\top k_t - v_t\|_2^2$ 的一步在线梯度:$S_t = (I - \eta_t k_tk_t^\top)S_{t-1} + \eta_t k_tv_t^\top$。秩一项 $(I-\eta_tk_tk_t^\top)$ 先沿当前键方向削减旧记忆、再写入新值,缓解纯加性写入的干扰;Gated DeltaNet 进一步叠加了全局衰减门。
Falcon 回归族(Falcon-1/2/3)就是 Delta 规则经'读后写对齐 + NLMS 归一化'改造后的产物,与 DeltaNet 的差别仅在训练对的时间对齐和步长归一化方式,理解 DeltaNet 才能看出改动点。
NLMS 归一化最小均方
NLMS(Normalized Least Mean Squares)是自适应滤波的经典算法:把 LMS 步长除以输入能量,取 $\eta_t = \beta_t/\|x_t\|_2^2$($\beta\in(0,2)$),使更新幅度对输入尺度不敏感。理论上,对 $L$-光滑目标取 $\eta = \beta/L$ 可保证每步目标值下降:$f(S^+) \le f(S) - \frac{\eta(2-\eta L)}{2}\|\nabla f(S)\|_F^2$。
Falcon 全部步长设计($\eta_t = \beta_t/(\|x_t\|_2^2+\lambda_t+\varepsilon)$,滑窗版用谱范数 $\mu^{(B)}_t$ 作分母)都是 NLMS 原则的推广,是理解'目标匹配归一化'这一核心卖点的前提。
SSD 与块并行训练
结构化状态空间对偶(SSD,由 Mamba-2 提出)证明选择性 SSM 与一类因果线性注意力在算法上等价,使循环模型可以块并行训练:把长度 $N$ 切成 $M$ 块(块大小 $C$),块内用掩码注意力并行计算,块间只传递固定尺寸状态。WY 表示与三角求解(TriSolve)是把 Delta 类秩一修正递推精确转换为这种并行形式的代数工具。
论文的核心工程贡献是让 Falcon 六个变体都拥有精确的 SSD 式块并行实现(含 log 空间正衰减重整化);没有这类并行化,在线学习型更新规则无法高效训练。
读后写(RAW)自回归语义
指先观测并写入第 $t$ 个 token、再用更新后的状态 $S_t$ 预测下一 token 的约定。在此语义下,第 $t$ 步'新揭示'的因果训练样本是前缀特征与新目标的配对 $(\phi(k_{t-1}), v_t)$(等价于标准索引下的 $(\phi(k_i), v_{i+1})$),而非常见的同步对 $(\phi(k_t), v_t)$;后者仍是因果的,但优化的是另一个内部目标。
这一对齐偏移是论文'下一潜在预测'视角的来源,也是与 DeltaNet、线性注意力等已有规则唯一的结构性差异,是理解本文定位的关键。
研究动机
标准 Transformer 的自注意力对长度 $N$ 的计算量是 $O(N^2)$,且推理时需维护不断增长的 KV 缓存,长上下文下注意力矩阵与显存读取均成为瓶颈。更深层的问题是:长上下文建模本质上也是持续学习——模型必须在不发生灾难性干扰的前提下在线绑定新证据。Transformer 把这份快速记忆外置为不断增长的 KV 缓存,而线性注意力、RWKV、Mamba、DeltaNet 等循环模型把它压缩进固定尺寸循环状态,但论文指出这些架构的状态更新规则通常只在架构层面给出,其隐含的局部目标函数与时间对齐方式是模糊的:多数方法沿用同步配对 $(\phi(k_t), v_t)$,而真正符合'预测下一潜在'语义的因果样本对应是前缀对 $(\phi(k_{t-1}), v_t)$——预测 $v_t$ 时只有 $k_{t-1}$ 可用。此外,固定学习率对回归型快速权重更新存在尺度失配,训练稳定性和混合精度下的数值鲁棒性都缺乏目标层面的依据。
本文的目标是本文的目标是把基于状态的序列建模显式重述为'自回归下一潜在预测'(autoregressive next-latent prediction)形式的在线持续学习:循环状态 $S_t$ 就是一个从上一时刻前缀写特征 $x_t=\phi(k_{t-1})$ 到新揭示目标 $v_t$ 的快速线性预测器,其更新规则即在线学习规则。在此视角下,作者希望:(1) 明确读后写(RAW)语义下因果训练对的正确对齐;(2) 导出一族与局部目标曲率/尺度匹配的归一化一阶更新,统一覆盖回归(DeltaNet 型)与内积(线性注意力/Mamba-2 型)两类目标,以及标量、逐通道、滑窗三种动力学;(3) 保持 $O(N)$ 训练、$O(1)$ 推理状态与 SSD 式块并行训练的兼容性;(4) 在 1.24亿–1.30亿参数、约 492亿 token 的语言建模和变长多位加法外推任务上验证家族的竞争力。
与已有工作不同的是,本文的独特切入在于显式分离此前被混为一谈的四个维度:时间对齐、可塑性($\beta_t$)、遗忘($\lambda_t$)与有界排练(窗口 $B$)。与 Titans/ATLAS/MesaNet 等内部记忆工作相比,本文约束更严格——被适配对象是固定尺寸快速记忆状态,且在因果的下一潜在对齐下在线更新。与 Schlag 等的经典 DeltaNet 相比,作者指出读后写语义下正确的因果对是 $(\phi(k_{t-1}), v_t)$ 而非 $(\phi(k_t), v_t)$——后者虽仍因果,但优化的是不同的内部快速记忆目标。更进一步,论文把自适应滤波领域的 NLMS 稳定化原则引入快速权重注意力:回归损失对 Frobenius 范数的光滑常数恰为 $L_t = \|x_t\|_2^2 + \lambda_t$,因此归一化步长 $\eta_t = \beta_t/L_t$ 是目标匹配的曲率归一化而非任意超参;对内积目标,同样的分母则被重新解释为写幅度控制而非曲率要求。
核心方法
直觉上,固定尺寸矩阵状态 $S_t\in\mathbb{R}^{d_x\times d_v}$ 就是前向传播内部的'学生':每来一个 token 就暴露一个局部训练对 $(x_t, y_t) = (\phi(k_{t-1}), v_t)$,状态对其做一步在线梯度下降。技术上,作者对瞬时岭回归损失 $\ell_t(S) = \frac12\|S^\top x_t - y_t\|_2^2 + \frac{\lambda_t}{2}\|S\|_F^2$ 做一步在线梯度下降,得到 $S_t = (1-\eta_t\lambda_t)S_{t-1} + \eta_t x_t r_t^\top$($r_t = y_t - S_{t-1}^\top x_t$ 为残差),并用 NLMS 式归一化步长 $\eta_t = \beta_t/(\|x_t\|_2^2+\lambda_t+\varepsilon)$、$\beta_t\in(0,2)$,由 $L$-光滑性保证每步下降。另一族内积目标 $\ell^{ip}_t(S) = -\langle S^\top x_t, y_t\rangle + \frac{\lambda_t}{2}\|S\|_F^2$ 产生无残差的加性写入。数字 1/2/3 分别表示标量、逐通道、滑窗动力学,后缀 A 表示内积目标,共六个变体;配套 WY 表示+TriSolve、批量逐通道 TriSolve、ParallelFlow tensorInv 等块并行核,以及 log 空间正衰减重整化($\alpha_t \le 1-\varepsilon_\gamma$ 截断),保证与 SSD 式训练管线精确兼容。
核心创新有二。其一是读后写对齐:在自回归'预测下一潜在'语义下,第 $t$ 步揭示的局部训练样本是前缀特征 $\phi(k_{t-1})$ 与新目标 $v_t$ 构成的对,而非传统快速权重规则的同步对 $(\phi(k_t), v_t)$;这一步偏移改变了快速记忆被训练的信息——只在预测时真正可获得的前缀信息上训练,使'状态更新'与'在线学习'在语义上严格一致。其二是目标匹配的 NLMS 归一化:回归损失的光滑常数恰为 $L_t = \|x_t\|_2^2 + \lambda_t$,取 $\eta_t = \beta_t/L_t$($\beta_t\in(0,2)$)可由引理 3.1 保证每步瞬时目标下降;滑窗版用窗口协方差谱范数 $\mu^{(B)}_t = \lambda_{\max}(\bar C^{(B)}_t)$ 作分母,使注入量与衰减比例不随窗口大小 $B$ 系统性放大。这与 DeltaNet、线性注意力使用固定或与目标无关的学习率有本质区别:步长不再是为稳定性拍脑袋的超参,而是由局部目标的几何性质决定。
方法步骤详情
方法步骤如下。(0) 边界约定:$x_1 := 0$、$\eta_1 := 0$(无因果对时不更新),写流整体后移一位,读取为写后读 $o_t = S_t^\top\phi(q_t)$。(1) Falcon-1(标量回归):$S_t = (1-\eta_t\lambda_t)S_{t-1} + \eta_t x_t r_t^\top$,残差 $r_t = v_t - S_{t-1}^\top x_t$,$\eta_t = \beta_t/(\|x_t\|_2^2+\lambda_t+\varepsilon)$。(2) Falcon-2(逐通道回归):每列独立步长 $\eta_{j,t} = \beta_{j,t}/(\|x_t\|_2^2+\lambda_t+\varepsilon)$,更新 $S_t = S_{t-1}(I - \lambda_t\mathrm{Diag}(\eta_t)) + x_t(\eta_t\odot r_t)^\top$,因平方误差目标逐列分解,这严格等价于 $d_v$ 个独立标量更新。(3) Falcon-3(滑窗回归):在最近 $B$ 个因果对上优化窗口平均损失,$\mu^{(B)}_t = \lambda_{\max}(X_t^\top X_t)/B_t$,$\eta_t = \beta_t/(\mu^{(B)}_t+\lambda_t+\varepsilon)$,注入 $\frac{\eta_t}{B_t}\sum_{j\in I_t}x_j r_{j,t}^\top$,所有残差在更新前状态处计算。(4) 三个 A 变体把残差替换为直接写入目标 $y_t$(Falcon-3A 写窗口平均互协方差 $\bar N^{(B)}_t = \frac{1}{B_t}\sum x_jv_j^\top$),分母改用窗口能量 $\bar E^{(B)}_t$。(5) 实现层面给出循环、掩码注意力并行、SSD 式块并行三种等价形式:Falcon-1 用 WY 表示+单次 TriSolve;Falcon-2 用共享 Gram 矩阵的批量逐通道三角求解;Falcon-3 用 ParallelFlow 张量求逆;全部配合 log 空间正衰减重整化与数值稳定反向传播。
技术新颖性
技术新颖性体现在四点。(1) 统一优化视角:线性注意力/Mamba-2 的加性累积是内积目标的梯度步,DeltaNet 是回归目标的梯度步,二者被同一框架收编,差异只在目标函数与时间对齐,而不是看似无关的架构技巧。(2) 逐通道 NLMS:Falcon-2 把步长提升为向量 $\eta_t\in\mathbb{R}^{d_v}$,由于平方误差损失逐列分解,这有严格的可分离性论证($d_v$ 个独立标量下降步),而非启发式的向量化。(3) 谱范数窗口归一化:Falcon-3 用 $\lambda_{\max}(\bar C^{(B)}_t)$ 而非迹上界作为光滑尺度,理论上保证更新幅度对名义窗口 $B$ 不变,且可通过 $B_t\times B_t$ 小 Gram 矩阵精确或幂迭代计算,无需物化 $d_x\times d_x$ 协方差。(4) 精确块并行工程:WY 表示+TriSolve(Falcon-1)、共享 Gram 的批量 $d_v$ 三角求解(Falcon-2,把注入路径与历史路径合并为单一残差系统,省去每块一次批量 TriSolve,类似 Comba 的单求逆形式)、ParallelFlow 的 block-strict-causal 张量求逆(Falcon-3),均带数值稳定的 log 空间正衰减重整化($\alpha_t = \min(\eta_t\lambda_t, 1-\varepsilon_\gamma)$),把在线学习规则无缝嵌入现代高效训练栈。
实验结果
实验设置:124M–130M 参数模型在 FineWeb-Edu 上以约 492亿 token(10 万步、序列长 1024、全局批 480)训练,基线为 Transformer(RoPE+SwiGLU)、RetNet/LightningAttn、Mamba-2、DeltaNet、Gated DeltaNet。(1) 语言建模困惑度(Table 1):Falcon-1.3 在 FineWeb-Edu 上以 17.10 成为全场最佳(Transformer 17.38、最强循环基线 Gated DeltaNet 17.32、Mamba-2 17.70、DeltaNet 17.84、RetNet 18.79);但 Wiki 上 Gated DeltaNet 仍最优(30.99,Falcon-1.3 为 33.00),LMB 上 GDN 46.70 也低于 Falcon 的 48.70——作者直言这'不是均匀的胜利'。(2) 下游 8 任务(Table 2):零样本平均 Falcon-1A.2 最佳 49.30(Transformer 48.16、Mamba-2 48.80、GDN 48.78),单样本平均 Falcon-1.3 达 49.54 为最佳循环值(Transformer 49.67 仍最高);消融显示 QK-RMSNorm 优于 QK-ℓ2 归一化,上下文条件化 $\eta$ 优于上下文条件化 $\beta$。(3) 变长多位加法长度外推(Table 3,训练宽 1–32 位、教师强制评估 33–48 位 OOD):Falcon-3A.3 平均准确率 87.2 最佳,Falcon-1A.3 85.9 次之,均超 RetNet 82.9、Falcon-1A.1 80.6、Mamba-2 75.2、Transformer 65.8;$d33/d48$ 准确率上 Falcon-1A.3/3A.3 达 100.0/69.0(Transformer 97.0/49.0,Mamba-2 100.0/51.0)。值得注意的是回归版 Falcon-1.3 仅 68.8,几乎与 Transformer 一样差,说明内积(加性)写入对存储主导型任务的外推更友好;作者将此实验定位为支持性证据而非主结果。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| 语言建模(FineWeb-Edu 验证困惑度,124M–130M/50B token) | Perplexity ↓ | 17.10(Falcon-1.3) | 17.32(Gated DeltaNet,最强基线);17.38(Transformer) | 较最强基线低 0.22、较 Transformer 低 0.28 |
| 下游 8 任务零样本平均(PIQA/HellaSwag/Winogrande/ARC/OBQA/SIQA/SciQ) | Zero-shot Avg Acc ↑ | 49.30(Falcon-1A.2) | 48.80(Mamba-2);48.16(Transformer) | 较最强循环基线 +0.5,较 Transformer +1.14 |
| 下游 8 任务单样本平均 | One-shot Avg Acc ↑ | 49.54(Falcon-1.3) | 48.57(Gated DeltaNet);49.67(Transformer) | 较最强循环基线 +0.97,仍略低于 Transformer |
| 变长多位加法长度外推(训练 1–32 位,OOD 33–48 位,教师强制) | Mean Acc ↑ | 87.2(Falcon-3A.3);85.9(Falcon-1A.3) | 82.9(RetNet/LightningAttn);75.2(Mamba-2);65.8(Transformer) | 较最强基线 +4.3,较 Transformer +21.4;Acc@d48 为 69.0 vs Transformer 的 49.0 |
局限与改进
作者承认的局限:主表只评测了标量与滑窗内积变体及一个回归消融(Falcon-2/2A/3 已定义但未单独基准测试),且明确表示实证结论'不是均匀的胜利';加法实验只是受控诊断,仅作支持性证据。我自己的观察:(1) 规模很小(124M–130M、50B token),结论向大模型外推存疑,且 Wiki/LMB 困惑度仍落后 Gated DeltaNet(33.00/48.70 vs 30.99/46.70);(2) 回归族在算术外推上表现崩坏(Falcon-1.3 仅 68.8,与 Transformer 的 65.8 相近),暴露残差修正型写入可能损害纯存储任务的外推,但论文未深入分析机理;(3) 滑窗规则的精确跨段续算需要额外携带最后 $B-1$ 个因果对作为尾部状态,削弱了'纯固定状态'部署的叙事;(4) $\mu^{(B)}_t$ 需对 $B_t\times B_t$ Gram 矩阵做特征值估计(可幂迭代近似),带来实现与计算开销;(5) 内积目标在 $\lambda_t=0$ 时对 $S$ 线性且无下界、无有限极小点,此时归一化只是启发式的写幅度稳定器,理论保证弱于回归族;(6) 训练序列长仅 1024,未在真正长上下文(如 8K+)上验证收益。
独立分析的弱点
独立分析的弱点与改进方向:(1) 尺度天花板——所有结论基于约 130M 模型,Falcon 对 Gated DeltaNet 的困惑度优势仅在 FineWeb-Edu 一列成立,应在 1.4B/7B 规模复验,并对比各自的门控变体;(2) 目标选择与任务类型强耦合——内积写入赢外推、回归写入赢语言困惑度,说明单一目标不完美,可研究在同一状态内混合残差修正与加性写入,或让 $\lambda_t$/$\beta_t$ 由上下文自适应切换两种行为;(3) 滑窗状态携带问题——$B-1$ 尾部样本使跨段推理与状态服务化复杂化,可用状态投影/蒸馏近似尾部信息以恢复纯循环部署;(4) 算力报告缺失——Falcon-3 的 tensorInv 与 Falcon-2 的批量逐通道 TriSolve 未给出吞吐/时延与 Mamba-2 官方核的端到端对比,工程实用性待验证,应补齐硬件效率基准;(5) 理论缺口——引理 3.1 只保证瞬时局部下降,外层自回归损失与内层快速记忆目标的相互作用没有刻画,可补累计遗憾界或在线学习稳定性分析;(6) 长上下文叙事未兑现——论文以长上下文与持续学习为动机,却只在 1024 长度训练、加法任务上测试外推,缺少真实长文本检索/推理任务证据。
未来方向
作者提出的方向:把时间对齐、可塑性($\beta_t$)、遗忘($\lambda_t$)、有界排练(窗口 $B$)作为可分离的旋钮进一步研究,并在保持块并行兼容的前提下探索更强变体。基于其成果可延伸的工作:(1) 系统基准测 Falcon-2/2A/3(逐通道与滑窗回归),论文只给定义未给数据,这一空白本身就是一个完整的消融研究;(2) 与门控机制结合(如 Gated DeltaNet 的数据依赖衰减门)测试叠加收益;(3) 一阶与二阶之间插值——用 RLS/MesaNet 式的精确岭解近似增强 Falcon-3 的滑窗更新;(4) 在 10 亿级以上参数与 8K–128K 上下文上验证长度外推优势是否保持,并测试真实长文档任务;(5) 把下一潜在对齐思想移植到测试时训练(TTT)与 Titans 类神经记忆模块,检验对齐偏移是否同样带来增益;(6) 研究窗口大小 $B$ 的可学习化与输入自适应,以及 $\mu^{(B)}_t$ 的更廉价估计;(7) 把正衰减重整化与 FP8/低精度训练结合,检验 NLMS 归一化在低精度下的稳定性优势。
复现评估
复现评估:论文给出项目页 github.com/yifanzhang-pro/fast-weight-attention,正文附 Algorithm 1–4 的完整伪代码(Falcon-2 块并行前向、Falcon-3 循环形式、Falcon-3 ParallelFlow、Falcon-3A 掩码注意力前向),多份附录记录 WY 代数、正衰减重整化、边界 sentinel 约定与稳定反向传播,工程细节披露充分。数据全部公开:FineWeb-Edu 训练集、标准困惑度评测(Wiki/LMB/FineWeb-Edu)与 lm-evaluation-harness 的 8 个下游任务;变长多位加法是合成数据,零成本可生成。算力需求:124M–130M 模型 × 约 492亿 token 约为 10²² FLOPs 量级,单次运行约数百至一两千 A100/H100 GPU 时,属于中等规模,学术实验室可以负担;加法实验单卡即可复现。总体难度中等偏上:Python 层循环形式容易实现且便于验证正确性,主要门槛在 TriSolve、批量逐通道求解与 tensorInv 的 Triton/CUDA 块并行核,需对照伪代码仔细实现 log 空间衰减、正性截断与边界条件;基线可借助 flash-linear-attention 等现成开源库对齐。
论文图表