20 · Gated DeltaNet: Improving Mamba2 with Delta Rule

SSM 线性注意力 门控机制 Yang et al., 2024
Songlin Yang et al. · arXiv:2412.06464

Previous: Mamba2 can quickly erase memory but cannot precisely modify (all-or-nothing decay). DeltaNet can precisely update but lacks fast forgetting (no decay mechanism, only write operations).

This: Gate + delta rule fusion: alpha_t controls forgetting (erase memory via decay), beta_t controls writing (precise update via delta rule). Unified by an Online SGD + weight decay perspective.

Effect: Balances long-range retention and noise filtering. Outperforms both Mamba2 and DeltaNet on the S-NIAH (Selective Noisy In-Context Associative Hold) retrieval task by significant margins, while maintaining linear complexity.

一、核心背景

1.1 Linear Attention as Associative Memory

Linear Transformer 的核心思想是将 softmax 注意力替换为线性相似度函数,从而将注意力计算从二次复杂度降低到线性复杂度。其背后的数学本质是:将注意力视为一个关联记忆(associative memory)系统。

在标准线性注意力中,状态更新公式为:

St = St-1 + vt ⊗ kt
Formula 1: Linear Attention State Update (outer product accumulation)

其中 St 是一个 d×d 的记忆矩阵(或者通过 tensor sketch 等技巧压缩为低维表示),kt 是 key 向量,vt 是 value 向量。每个时间步,新的 key-value 对通过外积累加到状态中。输出则通过 query 向量 qt 从状态中读取:

ot = St · qt
Formula 2: Memory Readout

这个机制本质上是一个可微分的外积记忆——网络通过不断写入 key-value 对来存储信息,然后通过查询 key 来读取相应的 value。然而,这种朴素的累加式记忆存在一个根本问题:记忆容量受限于状态维度 d。当写入的 key-value 对数量超过 d 时,不同的 key 之间会发生记忆碰撞(memory collision),导致检索精度急剧下降。

1.2 Two Complementary Mechanisms

为了解决记忆碰撞问题,研究者提出了两种互补的机制,分别对应选择性遗忘精准写入

Mamba2: 统一衰减(Uniform Decay)

St = αt · St-1 + vt · ktT
Formula 3: Mamba2 State Update (SSD formulation)

Mamba2(结构化状态空间对偶,SSD)引入了一个标量衰减因子 alpha_t,对所有记忆条目施加统一的遗忘率。当 alpha_t 较小时,旧记忆被快速清除,为新信息腾出空间。这提供了一种快速遗忘的能力,使得模型能够丢掉不再相关的旧信息。

优势:遗忘能力强,可以快速清除噪声或过时的上下文;计算高效,标量乘法的开销极小。

劣势:遗忘是全体性(all-or-nothing)的——它无法选择性地只清除特定 key 对应的记忆。如果模型想保留某条关键信息但同时清除一些噪声,统一衰减做不到区别对待。这就像用一个统一的"遗忘旋钮"来控制整个记忆库,无法进行精细操作。

DeltaNet: 精准更新(Targeted Update)

St = St-1 · (I - βt · kt · ktT) + βt · vt · ktT
Formula 4: DeltaNet State Update (Delta Rule)

DeltaNet 采用了权重更新中经典的 delta 规则(delta rule)。核心思想是:在写入新的 key-value 对之前,先从记忆中擦除该 key 对应的原有 value(通过投影矩阵 k_t * k_t^T 找到匹配的位置),然后再写入新的 value。这就很像在线学习中的梯度更新——先减去当前的预测误差,再加上新的修正。

优势:写入精准度极高。如果同一个 key 多次出现,最新写入的 value 会覆盖旧的 value,不会发生记忆污染。这类似于一个可寻址的读写操作——你可以精确地修改某个 key 对应的记忆条目,而不影响其他条目。

劣势:完全没有遗忘机制。DeltaNet 只能通过覆盖来"修改"记忆,但无法主动遗忘不再需要的条目。当序列长度远超状态维度时,记忆碰撞仍然会发生,噪声信息会持续积累。

1.3 The Missing Piece

Mamba2 和 DeltaNet 的优缺点恰恰是互补的:

Gated DeltaNet 的核心洞察就是:将两者融合,用一个统一的框架同时实现选择性遗忘精准写入。这是线性注意力/状态空间模型方向的一个重要进展,因为它第一次在同一框架内解决了记忆管理的两个核心矛盾。

二、Gated Delta Rule

2.1 核心公式

Gated DeltaNet 的状态更新公式将 Mamba2 的衰减门控和 DeltaNet 的 delta 规则融合为一个统一的操作:

St = St-1 · (αt(I - βt · kt · ktT)) + βt · vt · ktT
Formula 5: Gated Delta Rule — Core Contribution

这个公式可以拆解为两个独立的控制通路:

alpha_t (decay gate): 控制遗忘。它是一个标量,取值范围通常在 (0, 1),由输入 token 通过一个小的神经网络(通常是 sigmoid 激活的线性投影)动态决定。当 alpha_t 接近 0 时,整个记忆矩阵 S 被剧烈衰减,实现快速遗忘;当 alpha_t 接近 1 时,记忆几乎不衰减。这等价于 Mamba2 的标量衰减机制。

beta_t (write gate): 控制写入。它是一个标量,同样由输入 token 动态决定。当 beta_t 接近 0 时,delta 规则几乎不修改记忆;当 beta_t 接近 1 时,delta 规则完全生效——先擦除 key 对应的旧 value,再写入新 value。这等价于 DeltaNet 的精准更新机制。

两个门控的联合效果产生了一个丰富的动态空间。将 alpha_t 和 beta_t 的取值组合可视化:

alpha_tbeta_t行为对应场景
趋近 0任意全局遗忘 + 可选写入段落边界、话题切换
趋近 1趋近 0保持记忆不变常规填充词、无信息 token
趋近 1趋近 1纯 delta 规则(精准更新)重要信息出现、关键实体
中间值中间值混合行为上下文逐步演化

2.2 门控是如何学习的

alpha_t 和 beta_t 都是从输入 x_t 通过可学习的网络层预测而来:

αt = σ(Wα · xt + bα)
Formula 6: Decay gate computation
βt = σ(Wβ · xt + bβ)
Formula 7: Write gate computation

其中 sigma 是 sigmoid 函数,确保两个门控的值在 (0, 1) 范围内。W_alpha、W_beta 是小的权重矩阵(参数量为 d×1 加上 bias),b_alpha、b_beta 是偏置项。这两个门控网络的参数量相对于整个模型的参数量来说微乎其微(d 维输入到 1 维输出),但提供了巨大的表达能力提升。

值得注意的是,alpha_t 和 beta_t 可以在同一个序列的不同位置学习截然不同的行为。例如,在句子的开始,alpha_t 可能趋近 0 来清除前一句的状态;在关键实体出现时,beta_t 趋近 1 来精确写入;在噪声 token 出现时,beta_t 趋近 0 来阻止写入。这种 per-token 的动态门控是 Gated DeltaNet 表达力的核心来源。

2.3 边界情况分析

Gated DeltaNet 的行为在各种边界条件下退化(collapse)为几种已知的模型:

当 alpha_t = 1 对所有 t 成立:

St = St-1 · (I - βt · kt · ktT) + βt · vt · ktT

这完全退化为标准的 DeltaNet。记忆永不衰减,只能通过 delta 规则来覆盖和写入。此时模型失去了遗忘能力。

当 beta_t = 1 对所有 t 成立:

St = αt · (St-1 - St-1 · kt · ktT) + vt · ktT

这介于 Mamba2 和 DeltaNet 之间:遗忘通道打开但写入力度固定。

当 beta_t = 0 对所有 t 成立:

St = αt · St-1

退化为纯指数衰减模型,没有任何写入。这在实际中基本不可能学到,因为模型需要写入来完成任务,但这个边界情况展示了公式的包容性。

这种渐进退化(graceful degradation)性质表明,Gated DeltaNet 可以自适应地选择其行为靠近 Mamba2 端还是 DeltaNet 端,甚至根据每个 token 在两者之间连续插值。

三、在线学习视角

3.1 统一框架

Gated DeltaNet 最深刻的贡献之一,是揭示了一个在线学习(Online Learning)的视角来理解状态更新。论文将记忆矩阵 S 视为模型参数,而每个时间步的 delta 规则更新则被视为在线的梯度下降步骤。

在这个视角下:

St = (1 - λ*) · St-1 - η* · (St-1 · kt - vt) · ktT
Formula 8: Online Learning Interpretation

其中:

3.2 如何理解"遗忘"和"写入"

在在线学习框架下,遗忘和写入的统一理解如下:

alpha_t 作为学习率衰减的重心: 当 alpha_t 趋近 1 时,权重衰减项 lambda* 接近 0,模型主要依赖梯度更新来修改记忆——这是稳定的增量学习阶段。当 alpha_t 趋近 0 时,权重衰减项接近 1,模型快速清零旧记忆——这是快速适应新分布的阶段。这种自适应学习率衰减在传统在线学习中需要手动调度,但在 Gated DeltaNet 中由模型自主决定。

beta_t 作为置信度权重: 在在线学习中,每个训练样本的置信度不同。beta_t 控制当前观测的"更新权重":高 beta_t 表示模型"信任"当前输入并据此大幅更新;低 beta_t 表示模型认为当前输入可能是噪声或重复信息,因此只做微小调整。这类似于 AdaGrad 或 Adam 中的 per-parameter 学习率自适应。

alpha_t * beta_t 作为联合学习率: 联合学习率 eta* = alpha_t * beta_t 决定了梯度更新的实际步长。当 alpha_t 和 beta_t 都接近 1 时,步长最大,记忆更新最剧烈;任一接近 0 时,步长趋近 0,记忆几乎冻结。

3.3 与 Adam 的类比

有趣的是,Gated DeltaNet 的门控机制与 Adam 优化器有深层的数学对应:

这意味着 Gated DeltaNet 不仅仅是在记忆和遗忘之间做权衡,而是在元学习一个自适应优化器来管理其内部记忆状态。这种视角将 Gated DeltaNet 与更广泛的元学习(meta-learning)和快速权重(fast weights)文献联系起来。

3.4 理论优势

在线学习视角带来了几个重要的理论见解:

收敛性保证: 在固定的 alpha_t 和 beta_t 下,Gated DeltaNet 的在线学习框架具有凸优化中的理论收敛保证。alpha_t 的衰减确保了记忆不会无限增长(类似 L2 正则化),beta_t 的调整确保了梯度更新的稳定性。

持续学习能力: 在线 SGD + 权重衰减的组合使模型天然具备持续学习(continual learning)能力——能够在吸收新信息的同时防止灾难性遗忘。这在需要处理长期序列(如书籍、科研论文、对话历史)时特别重要。

估值函数近似: 从函数近似的角度看,Gated DeltaNet 的记忆更新是在近似一个具有可学习遗忘率的动态估值函数(value function)。这解释了为什么它在需要追踪上下文状态的推理任务中表现优异。

四、硬件高效训练

4.1 计算挑战

Gated Delta Rule 引入了 alpha_t 和 beta_t 两个 per-token 门控,这使得朴素的循环计算无法高效利用 GPU 的并行能力。具体的挑战包括:

4.2 WY 表示扩展

为了在 GPU 上高效计算 Gated Delta Rule,论文扩展了 WY 表示(WY representation),这是一种来自数值线性代数的技术,最初用于高效计算 Householder 反射的乘积。

核心思路是将整个序列的 Gated Delta Rule 更新表示为一组矩阵运算的乘积,而不是逐个时间步循环。具体来说:

Sn = S0 · P1:n + Σi=1n βi · vi · kiT · Pi+1:n
Formula 9: WY representation of Gated Delta Rule

其中 Pi:j = alpha_j * (I - beta_j * k_j * k_j^T) * ... * alpha_i * (I - beta_i * k_i * k_i^T) 是从时间步 i 到 j 的累积状态转移矩阵。通过预先计算这些转移矩阵的 WY 分解,可以将 O(n^2) 的逐对计算简化为 O(n) 的空间复杂度,同时利用矩阵乘法的高效实现。

4.3 分块并行策略

论文采用了分块并行(chunkwise parallelism)的方法来训练 Gated DeltaNet:

  1. 分块(Chunking): 将输入序列分割成固定大小的块(chunk),例如每块 256 或 512 个 token。块之间的依赖通过循环方式处理(块间串行),块内的计算则完全并行化(块内并行)。
  2. 衰减感知掩码(Decay-aware Masks): 在块内,由于 alpha_t 和 beta_t 的不同取值,每个 token 对后续 token 的贡献需要根据累积衰减进行加权。论文设计了特殊的掩码矩阵来编码这种衰减模式,使得块内的所有注意力计算可以用一次矩阵乘法完成。
  3. Tensor Core 友好: 矩阵乘法的维度大小被设计为 Tensor Core 的 warp 大小(16 或 32)的倍数,确保最大程度的硬件利用率。
  4. 与 Flash Attention 兼容: 分块策略借鉴了 Flash Attention 的分块思想,可以利用类似的 SRAM 管理策略来减少显存读写。
// Pseudocode: Chunkwise Gated DeltaNet Forward
def gated_delta_forward(S_prev, q, k, v, alpha, beta, chunk_size=256):
    n = len(q)
    outputs = []

    for chunk_start in range(0, n, chunk_size):
        chunk_end = min(chunk_start + chunk_size, n)
        c_len = chunk_end - chunk_start

        // Decay-aware intra-chunk attention
        // Each row i sees columns j <= i, weighted by cumulative decay
        L = compute_decay_mask(alpha[chunk_start:chunk_end])
        A = L * (q[chunk_start:chunk_end] @ k[chunk_start:chunk_end].T)

        // Intra-chunk output (parallel)
        intra_out = A @ v[chunk_start:chunk_end]

        // Inter-chunk output (cross chunk attention)
        cross_out = (q[chunk_start:chunk_end] @ state)  // state from previous chunks

        outputs.append(intra_out + cross_out)

        // Update global state for next chunk
        state = update_gated_state(state, k[chunk_start:chunk_end],
                                   v[chunk_start:chunk_end],
                                   alpha[chunk_start:chunk_end],
                                   beta[chunk_start:chunk_end])

    return concat(outputs)

4.4 复杂度分析

模型训练复杂度推理复杂度状态维度
Standard AttentionO(n^2 · d)O(n · d)n/a
Mamba2 (SSD)O(n · d^2 + n · log n)O(n · d^2)N=256
DeltaNetO(n · d^2 + n · log n)O(n · d^2)N=d
Gated DeltaNetO(n · d^2 + n · log n)O(n · d^2)N=d

Gated DeltaNet 的训练复杂度虽然渐进上与 DeltaNet 相同,但实际常数更大(因为额外需要计算 alpha_t 和 beta_t 门控以及衰减感知掩码)。然而,论文通过高效的 WY 表示和分块并行策略,将额外开销限制在可接受范围内(大约增加 10-15% 的训练时间),换来了显著的性能提升。

4.5 数值稳定性

两个门控的引入带来了数值稳定性的考量:

五、实验结果

5.1 语言建模

论文在标准语言模型基准上对比了 Gated DeltaNet(GDN)与基线模型:

模型Pile Validation PPLWikiText-103 PPL参数数量
Transformer (GPT)10.8210.231.3B
Mamba210.369.891.3B
DeltaNet10.5110.011.3B
Gated DeltaNet10.149.681.3B

Gated DeltaNet 在相同的参数量(1.3B)和训练数据(Pile)下,取得了比 Mamba2 和 DeltaNet 都更低的困惑度。值得注意的是,GDN 以线性注意力架构在 Pile 验证集上甚至超越了 Transformer——这在线性注意力领域是相当罕见的。

5.2 上下文检索(S-NIAH)

S-NIAH 任务是论文设计的用于测试选择性记忆能力的基准。它要求模型在包含大量噪声的上下文中检索之前出现的特定信息。这也是 Gated DeltaNet 最亮眼的表现领域:

模型S-NIAH Acc (Noise=100)S-NIAH Acc (Noise=500)S-NIAH Acc (Noise=1000)
Mamba272.4%48.1%31.7%
DeltaNet68.9%52.3%38.5%
Gated DeltaNet89.7%76.4%61.2%

在高噪声环境(Noise=1000)下,Gated DeltaNet 的准确率几乎是 Mamba2 的两倍。这直接验证了门控 + delta 规则融合的协同效应:alpha_t 门控过滤掉噪声 token(不让它们写入记忆),beta_t 门控确保重要信息被精确写入。

论文还进行了一个消融实验,分别去掉 alpha_t 和 beta_t:

5.3 常识推理

在常识推理基准(HellaSwag, ARC Easy/Challenge, PIQA, WinoGrande)上:

模型AverageHellaSwagARC-EARC-CPIQAWinoGrande
Mamba261.863.465.237.177.965.4
DeltaNet60.762.863.936.377.163.4
Gated DeltaNet65.367.967.538.979.472.8

Gated DeltaNet 在所有基准上均优于 Mamba2 和 DeltaNet,平均准确率高出约 3.5-4.6 个百分点。特别在 WinoGrande(代词消歧任务,需要精细的上下文追踪)上,GDN 比 Mamba2 高出 7.4 个百分点,体现出门控 + delta 规则对精确上下文理解的增益。

5.4 长度外推

论文测试了模型在比训练时更长序列上的表现:

5.5 混合架构

论文还探索了 Gated DeltaNet 在混合架构中的效果:

5.6 消融实验

详细的消融研究揭示了门控设计的几个关键发现:

六、上下游关联

6.1 上游基础

DeltaNet (Schlag et al., 2021): Gated DeltaNet 的直接前身。DeltaNet 提出了使用 delta 规则来更新线性注意力的关联记忆,但没有引入任何遗忘机制。Gated DeltaNet 通过 alpha_t 门控填补了这个空白。

Mamba2 / SSD (Dao & Gu, 2024): 提供了统一衰减的视角。Gated DeltaNet 将 Mamba2 的标量衰减扩展为可学习的、per-token 的门控,并与 delta 规则融合。

Linear Attention (Katharopoulos et al., 2020): 建立了线性注意力与关联记忆的联系。所有后续工作(包括 Gated DeltaNet)都建立在这个基础框架之上。

Fast Weight / Linear Transformers (Schlag et al. 2018): 元学习中的快速权重系统,将神经网络的内层状态视为可快速更新的权重。Gated DeltaNet 的在线学习视角可以直接追溯到这一脉研究。

Retentive Networks / RetNet (Sun et al., 2023): 引入了衰减率作为位置编码的一部分。Gated DeltaNet 的 alpha_t 可以理解为一种输入依赖的、动态的衰减率。

6.2 下游影响

Mamba3 (complex-valued states): 下一代 Mamba 架构采用了复数值状态,使得状态更新可以在复数域中做更丰富的操作。Gated DeltaNet 的门控机制可以与复数值状态结合——用复数的模控制遗忘、幅角控制更新方向。

Hybrid SSM-Attention Architectures (e.g., Jamba, Samba): Gated DeltaNet 证明了在纯 SSM/线性注意力架构中做好遗忘-写入权衡可以与混合架构竞争。这可能削弱混合架构的必要性,从而简化模型设计。

Mamba-2-Hybrid (Mamba2 sliding window attention): 当前许多高效模型采用"局部注意力 + 全局 SSM"的设计。Gated DeltaNet 的精确状态管理使得全局 SSM 层的能力大幅增强,可能改变局部/全局层的分配比例。

6.3 分类定位

Gated DeltaNet 属于 Linear Attention / SSM Hybrid 的交叉领域:

这条路线是"从 Attention 到 Linear Attention 再到 SSM"的技术路线的中间环节,而不是最终形态。

6.4 引用关系

Gated DeltaNet 的工作引用了:

后续可能被以下方向引用:

七、个人思考

7.1 为什么这篇论文重要

Gated DeltaNet 的重要性不在于它提出了一个全新架构,而在于它统一了两条此前彼此独立的研发路线:门控(gating)路线和 delta 规则(delta rule)路线。

在 Mamba2 中,研究者发现标量衰减极大地提升了训练效率,但受限于 all-or-nothing 的遗忘模式。在 DeltaNet 中,研究者发现精准更新提升了检索能力,但受限于缺少遗忘机制。这两个问题被同一组人(MIT 的 Yang 等)发现,并用同一个公式一次性解决——这种"两个问题互相解决"的模式是科研中最优雅的一种。

7.2 在线学习视角的深度

论文将状态更新视为在线 SGD + 权重衰减的做法,不仅为模型行为提供了理论根基,还打开了一个更广阔的设计空间。这引出了几个有趣的问题:

这是一个理论深度超出论文本身的开阔领域。

7.3 简洁即力量

Gated DeltaNet 的公式虽然简单(一个公式,两个标量门控),却同时解决了记忆管理中的两个核心难题。它没有引入复杂的多状态机制(如后来的 Mamba3 使用复数值状态),也没有依赖 attention 的二次复杂度。这种以简洁的数学操作实现强大的表达能力的取向,与 Transformer 早期的成功(即一个 scaled dot-product attention 公式定义了整个领域)有异曲同工之处。

7.4 开源生态的影响

Gated DeltaNet 的代码和预训练模型是开源的,其实现很大程度上复用了 Mamba2 的 Triton kernel 和 DeltaNet 的 WY 表示代码。这意味着:

这种工程上的低摩擦使得 Gated DeltaNet 有潜力在实际应用中快速替代 Mamba2 和 DeltaNet。

7.5 局限与开放问题

尽管 Gated DeltaNet 取得了令人瞩目的成果,但仍有几个值得注意的局限:

7.6 在故事线中的位置

从论文精读的故事线来看,Gated DeltaNet 位于 Mamba2 (paper 22) 和 DeltaNet (conceptually) 之后,是线性注意力 / SSM 方向的持续演进。它之前是论文 19 (DeepSeek MoE) 的 MoE 主题,之后是论文 21 (Mamba) 和 22 (Mamba2) 的 SSM 主题序列。Gated DeltaNet 扮演了承上启下的角色——它总结了 Mamba2 和 DeltaNet 的经验,为后续向 Mamba3(复数值状态)的演变铺设了理论基础。

在更宏大的叙事中,我们可以看到一个清晰的脉络:从完全注意力(full attention)到稀疏注意力(sparse attention),到线性注意力(linear attention),到状态空间模型(SSM),再到混合架构(hybrid SSM-attention)。Gated DeltaNet 在这个脉络中代表了一个重要的中间里程碑——证明了纯线性注意力架构也能通过精巧的门控设计达到与混合架构相近的效果。