技术摘要:OSDN:在线性注意力中通过可证明的在线预条件改进 Delta 规则
1. 问题陈述
线性注意力和状态空间模型(SSMs)提供了 O(N) 的推理复杂度和恒定内存的 softmax 注意力替代方案,使其适用于长序列处理。然而,它们通常在上下文关联回忆方面表现挣扎,特别是与基于 softmax 的 Transformer 相比。
Delta 规则(如在 DeltaNet 中实现)试图通过将循环状态更新解释为每个 token 回归损失的单步在线梯度下降(OGD)来弥合这一差距:ft(S)=21∥Skt−vt∥F2。更新规则为 St=St−1+βt(vt−St−1kt)kt⊤,其中 βt 是标量门控。
核心局限性:标量门控 βt 充当应用于所有键维度的统一学习率。这忽略了内部目标的特征级曲率。在关联回忆任务中,不同的键(例如,频繁与稀有标识符、稳定与高熵值)表现出不同的经验曲率。单一的标量步长迫使做出妥协:增加它以帮助覆盖一个方向上的陈旧记忆,可能会过度校正另一个已经校准良好的方向。标准自适应优化器(如 Adam)使用对角预条件来解决这个问题,但将此类机制应用于循环层通常需要稠密的二阶状态,或者破坏使 DeltaNet 高效的硬件友好的分块并行化流水线(如 WY 变换)。
2. 方法论:在线缩放 DeltaNet (OSDN)
作者提出了OSDN,它通过对角预条件器 Dt=diag(dt) 增强标量门控,该预条件器通过超梯度反馈在线学习。
2.1. 代数等价性与硬件保留
OSDN 的一个关键洞察是,对梯度步进行右预条件在代数上等价于对写入侧键进行每特征缩放。
标准更新:
St=St−1−βt∇ft(St−1)Dt=St−1+βt(vt−St−1kt)(dt⊙kt)⊤
通过定义预条件键 k~t=dt⊙kt,更新变为:
St=St−1+βt(vt−St−1kt)k~t⊤
这种变换允许 OSDN 严格保留 DeltaNet 的分块 WY 并行流水线。实现仅需在分块 Gram 矩阵和分块内分数中将存储键 K 替换为 K~,除了 O(K) 预条件向量外,不产生任何高维状态开销。
2.2. 解耦的超梯度反馈
预条件器 dt 通过基于在线缩放梯度方法 (OSGM) 框架 [18] 导出的超梯度代理进行在线梯度下降 (OGD) 更新。
预条件器的代理损失为:
ht(dt)=∥∇ft(St−1)∥F2ft(St)−ft(St−1)
关键在于,由于内部回归损失 ft 的精确二次结构,该代理 admit 一个闭式表达式,该表达式仅依赖于键 kt 和标量门控 βt,从而将其与高维状态 St、值 vt 或残差 ut 解耦。
ht(dt)=2∥kt∥22(1−βt⟨dt,kt2⟩)2−1
这种解耦使得两阶段实现成为可能:
- 阶段 1:对键流进行轻量级 O(K) 扫描,以计算预条件键序列 {k~t} 并更新 dt。
- 阶段 2:使用 K~ 代替 K 进行标准分块 DeltaNet 传递。
2.3. 自适应预条件遗忘 (APF)
为了处理键分布发生变化的非平稳语言上下文,作者引入了APF。这增加了一个学习的、token 级、头级的保留门控 rt,h,仅应用于预条件状态 dt(而非记忆状态 St):
dˉt+1=rt,hdt+η∇dht(dt)
这保持了双阶段扫描所需的仿射循环结构,同时允许元优化器“遗忘”陈旧的校准。
3. 理论保证
该论文基于内部损失的精确二次性质建立了两个收敛保证:
- 总体极限超几何收敛(定理 5.1):针对理想的右牛顿比较器(D⋆=Σk†),如果在线学习者实现次线性后悔,则全局目标的次优性以超几何速率 O((C/T)T) 收缩。这显著快于标准的几何(线性)收敛。
- 算法局部残差收缩(定理 5.2):对于实现的对角更新,token 局部残差比率的乘积 ∏t=1Tft(St−1)ft(St) 是有界的。如果对角比较器满足门控牛顿条件(βt⟨d⋆,kt2⟩=1),则收缩速率也是超几何的,在标准 OGD 后悔界限下具体为 O((1/T)T)。
4. 实验结果
实验在340M和1.3B参数规模下的DeltaNet、门控 DeltaNet (GDN) 和Kimi Delta 注意力 (KDA) 骨干网络上进行。
4.1. 上下文回忆(340M 规模)
- 性能:原生 OSDN 将 JRT 风格的上下文回忆提高了32%(从 0.150 提升至 0.198),优于基线 DeltaNet。增益集中在重复上下文任务(例如"JRT-twice"变体)上,在这些任务中键会重复出现,从而使预条件器能够进行校准。
- APF 影响:OSDN-APF 保留了**17%**的增益(0.176),并在 DeltaNet 变体中提供了最佳的 WikiText/LAMBADA 困惑度(30.99 对比 32.00)。
- 泛化能力:在更广泛的基准测试(Commonsense, LongBench)上,OSDN 变体与其基线保持持平,证实该机制是针对性的,而非通用的性能提升器。
4.2. 扩展到 13 亿参数
- 残差收缩:机制层面的信号在规模扩大时得到放大。几何平均残差比率 (qgeo) 从0.432 降至 0.265(减少 39%),几乎是 340M 规模下观察到的改进的两倍。
- 下游持平:尽管内部残差收缩大幅改善,但下游指标(WikiText/LAMBADA 困惑度、Commonsense、LongBench)仍与 13 亿参数的 DeltaNet 基线持平,表明该机制有效转移且未损害通用能力。
4.3. 效率
- 吞吐量:OSDN 变体产生的推理开销极小。在 340M 规模下,吞吐量变化在 ±2.2% 以内。在 13 亿规模下,OSDN-APF 比基线慢 6.8%,归因于预条件器的额外顺序扫描。
- 状态开销:由于 O(K) 预条件向量,持久循环状态仅增加了≤0.05%。
5. 意义与主张
该论文将 OSDN 定位为对 Delta 规则的最小化、针对性扩展,弥合了一阶标量更新与自适应优化之间的差距。
- 机制层面改进:主要贡献是通过在线预条件在关联回忆方面实现了可证明的、机制层面的改进,并通过残差收缩比率的直接诊断得到验证。
- 硬件兼容性:与其他二阶或稠密预条件方法不同,OSDN 保留了分块并行流水线和 O(K) 状态复杂度,使其适用于大规模部署。
- 可扩展性:结果表明,在线预条件的优势在十亿参数规模下能够转移并放大,这表明线性注意力中的“曲率不匹配”是一个根本性瓶颈,可以在不进行架构彻底改革的情况下加以解决。
- 局限性:作者明确指出,OSDN 并非通用的基准改进器;增益集中在涉及重复键的检索任务上。理论保证依赖于内部损失的精确二次假设和单调下降,这些是理想化条件。该方法不声称解决所有上下文学习的局限性,而是专门解决 Delta 规则的写入步几何问题。