Command Palette
Search for a command to run...
MARCH:利用内容路由状态锚扩展循环记忆
MARCH:利用内容路由状态锚扩展循环记忆
Ming Zhang Kaisen Yang Shu Yu Ermo Hua Ning Ding Xia Hu Bowen Zhou Chaochao Lu Youbang Sun
摘要
Transformer 强大的长上下文检索能力在很大程度上归功于随上下文长度增长的 token 级记忆。然而,这种灵活性在训练中带来了二次方的计算复杂度,并在自回归推理中导致键值缓存随上下文线性增长。循环式替代方案通过将整个历史压缩为固定大小的状态来实现高效解码,但在依赖回忆的任务上往往表现不佳,因为较早的关联通常会被后续更新覆盖,仅保留最近的上下文信息。本文提出了跨上下文历史的记忆锚路由(MARCH),这是一种网络架构,能够在保持长序列计算效率的同时,将状态空间模型扩展到固定大小维度之外。MARCH 定期将累积的循环状态检查点缓存为状态锚,并将每个锚与一个紧凑的、以内容为条件的锚键相关联。这使得 MARCH 能够维护一个可随上下文长度增长而扩展的记忆库,从而在历史分辨率与记忆成本之间提供可控的权衡。在每个 token 处,MARCH 生成一个锚查询,以关注所有因果可用的状态锚,输出则通过对所有历史锚以及当前状态进行注意力式聚合来计算。实验表明,经过标准预训练后,MARCH 在常识推理、LongBench 和上下文检索任务上一致优于多种线性注意力变体。这些结果表明,内容路由的状态缓存能够在保持其原生计算路径的同时,显著增强循环长程记忆。
一句话总结
上海人工智能实验室、清华大学和复旦大学的研究人员提出 MARCH(Memory-Anchor Routing across Context History),该方法通过定期将循环状态检查点缓存为按内容路由的状态锚点并对其进行聚合,使状态空间模型超越固定大小状态的限制;经过标准预训练后,MARCH 改进了长上下文检索,并在常识推理、LongBench 和上下文检索上优于线性注意力变体。
核心贡献
- 提出 MARCH(Memory-Anchor Routing across Context History),作为一种网络架构,通过周期性缓存累积循环状态检查点作为状态锚点,并将其与紧凑的、以内容为条件的锚点键相关联,增强循环状态空间模型。
- MARCH 维护一个随上下文长度增长的内存库,通过对因果可见的锚点和当前状态进行注意力式聚合来计算输出,将内存容量与密集的逐 token 更新解耦,同时保留高效循环。
- 经过标准预训练后,MARCH 在常识推理、LongBench、上下文检索和 NIAH 上持续优于多个线性注意力变体,消融研究表明通过检查点密度和稀疏路由可实现可控的检索效率权衡。
引言
大型语言模型越来越需要跨长文档、多轮交互和上下文示例整合信息。标准自注意力支持细粒度回忆,但会产生二次训练成本,以及与序列长度线性增长的 key-value 缓存。线性注意力和现代循环模型提供恒定内存解码,但将整个历史压缩到固定大小的状态中,因此较早的关联可能被覆盖或削弱,无法从最新状态中恢复。先前工作扩展了内存容量或保留多个压缩状态,但对历史状态的查询相关检索仍未充分探索。作者引入 MARCH(Memory-Anchor Routing across Context History),该方法定期保存循环状态快照作为锚点,并使用紧凑的已学习描述符将每个 token 路由到相关的早期状态,同时保留高效循环。
方法
作者引入 MARCH,这是一种按内容路由的循环内存框架,支持从不断演化的循环状态的早期版本中进行选择性检索。MARCH 不是基于预定义状态索引或时间尺度进行路由,而是根据内容将每个查询与各个历史状态进行匹配。
如框架图所示,该方法在处理 token 时周期性地对累积循环状态设置检查点,形成状态锚点库。每个检查点与一个共享的已学习锚点 token 的出现配对,其隐藏表示产生紧凑的路由键。对每个文本 token,路由查询对所有因果可见锚点以及一个已学习的空选项进行评分,使模型仅在有用时使用历史内存。所得路由概率定义可见锚点状态的加权组合,并使用该 token 的标准循环查询进行读取。该历史读出被加到当前状态读出上,保留原生循环路径,同时引入一条按内容依赖到达早期内存的路由。状态锚定和按内容路由检索共同将原本瞬态的状态轨迹转变为长期内存的持久来源。
锚定过程首先定义一组有序文本边界。在每个边界之后使用共享的已学习锚点嵌入插入一个锚点位置。文本位置应用基础循环更新,使矩阵值状态在锚点边界之间持续演化。在处理边界处的文本 token 后,MARCH 立即对所得累积状态设置检查点以形成状态锚点。后续锚点位置不修改循环状态;其隐藏表示提供与该检查点关联的路由元数据。由于循环在锚点边界之间不会被重置,每个状态锚点编码到该位置为止的累积前缀,追踪单个循环内存的时间演化,并在后续衰减和 delta 更新削弱其内容之前保留较早版本。
为使路由显式依赖于每个状态锚点的内容,锚点位置仅读取与其对齐的状态检查点。同一输入表示被投影为紧凑的路由键。对齐读出通过标准输出投影和残差路径纳入锚点位置。因此,下一层的表示依赖于状态锚点,产生的路由键也以对齐状态锚点所保留的内容为条件。尽管所有锚点位置共享相同的已学习输入嵌入,但在第一层之后它们获得不同的、依赖状态的表示。
对于基于内容的路由,模型将文本 token 的归一化隐藏状态投影为路由查询,并将其与每个可见锚点的键进行评分。为使模型能够绕过历史内存,可见锚点集合会加入一个负载固定为零的空选项。路由概率通过对可见锚点和空选项的 logits 进行 softmax 计算。由于所选路由概率直接对历史状态读出加权,其得分由语言建模目标联合优化。该聚合形式还支持稀疏变体,通过将聚合限制在得分最高的可见锚点上,以最小的性能下降降低聚合成本。
给定路由概率后,将因果可见状态锚点聚合为查询相关的历史状态,并使用与当前状态相同的状态读取查询进行读取。所得历史读出随后加到当前状态读出上。这种加法形式保留原始循环路径,并将历史检索作为辅助残差分支引入,而不修改底层循环更新。由于路由概率直接影响层输出,路由查询和来自锚点的键被端到端优化。
作者将 MARCH 实现为两阶段生产者-读取器计算。按照硬件高效的分块形式,生产者以适合张量核心加速的块处理循环更新,计算每个 token 的当前状态输出,并在每个锚点边界对循环状态设置检查点。所得状态锚点由历史读取器消费。受 I/O 感知原则启发,读取器联合分块处理查询 token 和状态锚点,在查询块中复用每个锚点分块,并将路由得分计算、在线 softmax 更新和加权状态读出累积融合为流式归约。这种融合调度避免物化密集的 token 到锚点路由矩阵,也避免物化大得多的每个锚点候选读出张量,从而减少中间存储和相关内存流量。如效率分析所示,尽管存在历史检索成本,融合的密集实现在更长序列长度下的吞吐量超过标准注意力基线,并且核心运行时间更低。
实验
论文通过从头开始在 50B tokens、16K 上下文条件下预训练模型,并将 MARCH 与 Gated DeltaNet、一种对数线性变体和 Transformer 基线进行比较来评估 MARCH。测试包括零样本常识推理、长上下文理解、大海捞针检索和上下文检索,并对块大小和路由设计进行消融。MARCH 持续优于循环基线,尤其在检索密集型和长上下文任务上,同时在短上下文理解上仍与全注意力模型具有竞争力。消融研究表明块大小 512 能平衡准确性和成本,稀疏路由提高效率,已学习的空选项有益。
在八个零样本常识推理基准上,MARCH 在每个任务上都优于普通和对数线性 Gated DeltaNet 变体,并将平均准确率提高到两个 Transformer 基线之上。相对普通循环基线,在 OpenBookQA 上的提升最大。MARCH 在八个任务中的六个上超过标准 Transformer,在四个上超过 24 层 Transformer,表明其短上下文语言理解具有竞争力。MARCH 在所有八个常识推理基准上持续优于普通和对数线性 Gated DeltaNet,单一任务最大增益来自 OpenBookQA 上相对普通骨干网络。MARCH 在八个任务中的六个上超过标准 Transformer,在四个上超过更深的 Transformer,同时平均得分高于两个全注意力基线。
在十二个 LongBench 任务上,MARCH 在每个任务上都优于普通 Gated DeltaNet 及其对数线性变体,相对 Gated DeltaNet 的平均相对增益约为 25%。改进在多文档问答和摘要中尤为显著,平均得分超过标准 Transformer,并接近 24 层 Transformer。MARCH 在所有十二个 LongBench 任务上都优于两个 Gated DeltaNet 变体,任务涵盖单文档问答、多文档问答、摘要和少样本学习。最大相对增益出现在多文档问答中,MARCH 相对更强的循环基线大幅提高了 2WikiMultihopQA 和 MuSiQue 得分。摘要也显著受益,QMSum 相对更强的循环基线提高 32%。MARCH 的平均 LongBench 得分高于标准 Transformer,并接近 24 层 Transformer。
MARCH 在所有六个基准上均改善了相对于两个 Gated DeltaNet 变体的上下文检索,在每个任务上都相对更强的循环基线取得相对增益。平均检索准确率从 20.5 上升到 23.3,相对提升 14%。Transformer 基线在平均值上仍然领先,但 MARCH 缩小了差距,同时持续增强循环骨干网络。MARCH 在所有六个检索任务上都优于更强的 Gated DeltaNet 基线,相对增益从 SQuAD 的 8% 到 TriviaQA 的 23%。相对循环基线的最大相对改进出现在 TriviaQA、NQ 和 FDA 上。六个基准的平均检索准确率相对更强的 Gated DeltaNet 基线提高 14%。MARCH 在平均检索准确率上落后于 Transformer 基线,但相比 Gated DeltaNet 变体有显著改进。
匹配训练和推理块大小表明,块大小 512 在检索质量和锚点数量之间提供最佳整体平衡。较小的块会改善一些长上下文结果,但增加内存和路由成本;较大的块通常因检查点变稀疏而降低检索性能。推理时改变块大小提供灵活的准确率与内存权衡,分层 Fenwick 树组织仍具有竞争力。512 的块大小相对锚点数量提供最佳整体检索质量。较小的块产生更密集的锚点,改善一些长上下文 NIAH 结果,但需要更多内存和路由开销。较大的块大小降低锚点密度,通常会降低检索质量,尤其在较长上下文下。推理时,更密集的锚点倾向于以更高成本改善检索,而过于稀疏的锚点导致显著退化。状态库的 Fenwick 树组织性能接近参考方案,表明锚点组织具有灵活性。
默认密集路由器使用查询键维度 64 和已学习的空选项,提供最佳整体权衡,在常识推理、LongBench 和 NIAH 平均值上最高。更大的路由器维度改善检索但降低其他聚合得分,而 top-4 路由在常识推理和检索上接近,但在 NIAH 上较弱。移除空选项会降低所有聚合指标,表明学习忽略不相关历史状态是有益的。默认密集路由配置配合空选项在常识推理、LongBench 和 NIAH 平均值上领先,但在检索上略低于更宽的路由器。Top-4 路由在常识推理和检索上几乎与密集路由相当,但在 NIAH 上落后,移除空选项会持续降低聚合性能。
在零样本常识推理和 LongBench 任务中,MARCH 持续优于普通和对数线性 Gated DeltaNet 变体,平均超过标准 Transformer,并在若干任务上接近或超过更深的 Transformer。它还在循环基线上增强了上下文检索,尽管全注意力 Transformer 在平均值上仍领先。消融研究表明,块大小 512 能最佳平衡检索质量和开销,而默认密集路由器配合已学习的空选项在常识推理、长上下文和检索基准上提供最强的整体权衡。