导读

在大型语言模型(LLM)处理超长序列时,全上下文自注意力机制的二次复杂度成为可扩展性的核心瓶颈。传统做法是在推理时采用窗口或稀疏注意力等受限执行策略来降低计算和内存开销,但训练阶段仍依赖全上下文注意力——这种训练与推理语义的脱节,导致模型在推理时无法依赖训练时用到的完整上下文信息,进而损害泛化稳定性和长上下文任务表现。内蒙古大学和香港科技大学(广州)的研究团队注意到这个被广泛忽视的“训练-推理一致性”问题,并提出了全新的解决方案。 该工作的核心创新在于:将片段级执行从推理优化提升为训练和推理共用的建模假设。通过设计一个固定大小的KV尾部作为跨片段可微分接口,并配合截断反向传播(TBPTT)限制梯度传播范围,作者使得模型在训练阶段和推理阶段遵循完全相同的前向执行语义。同时,为了不损失远程信息访问能力,框架额外支持只读的检索前缀机制。这一设计不仅理论上保证了目标函数的精确梯度计算(而非近似),还在实践中取得了与全上下文注意力相当的困惑度性能,并将128K上下文预填充的峰值内存降低至FlashAttention的约1/6。 对于关注长上下文LLM训练和推理优化的研究人员,这篇论文提供了一种突破性的视角——不再是简单地升级算力或压缩注意力,而是从执行语义一致性出发,从根本上消除训练-推理鸿沟。它提出了一个简单但有效的框架,并经过多个长上下文基准的验证,是2026年ICML上值得细读的工作。

论文基本信息

Figure 1. Peak GPU memory consumption during long-context prefill. 来源:原论文 PDF 第 1 页。 **

**

摘要

基于Transformer的大语言模型在长上下文生成中面临严峻的可扩展性挑战,根源在于全上下文自注意力的计算和内存成本呈二次增长。在计算和内存资源有限的情况下,许多推理高效的长上下文方法仅在推理时采用受限上下文或片段级执行,但训练仍依赖全上下文注意力,导致训练和推理的执行及状态转换语义不匹配。针对这一不足,作者提出一种训练-推理一致的片段级生成框架,训练和推理遵循完全相同的片段级前向执行语义。 具体地,训练阶段通过将梯度传播严格限制在从紧邻前一个片段携带过来的KV状态上,来强制执行与推理的一致性;同时,前向传播允许针对特定注意力头访问过去的KV状态,但这些访问不参与梯度计算。这一设计的核心是,只有固定大小的KV尾部作为可微分的跨片段接口状态,该状态在训练和推理中完全一致。在长上下文基准测试上,该方法取得了与全上下文注意力相当的性能;与强大的推理高效基线相比,在延迟和内存的权衡上具有竞争力;并且大幅提升了非常长上下文场景下的可扩展性——例如,在128K上下文长度下,预填充阶段的峰值内存比采用FlashAttention的全上下文注意力降低约6倍。

引言:论文要解决什么问题

随着大语言模型在文档理解、持续对话和复杂推理等长上下文应用中的普及,如何高效地扩展Transformer到极长序列成为核心难题。全上下文自注意力的计算复杂度为O(L²),其中L是上下文长度,这使得当L达到数万甚至十几万个token时,GPU内存和计算时间都难以承受。 现有解决方案分为两类:一类是语义保持的执行级优化,如FlashAttention和Chunked Prefill,这类方法在数学上精确等价于全上下文注意力,因此能在不改变模型输出的前提下降低内存和延迟;另一类是受限执行策略,如窗口注意力、稀疏注意力等,它们主动丢弃部分上下文信息以换取效率。然而,图1(原论文Figure 1)清晰地展示了:即使使用FlashAttention,128K上下文下的预填充峰值内存仍接近80GB,而Chunked Prefill虽有所降低但仍达约60GB,远远超出单GPU容量。这意味着,仅依赖执行级优化带来的资源节省,已不足以支撑更长上下文场景的实用部署。 更关键的问题是训练与推理的语义鸿沟。大多数现有方法只在推理时采用受限执行(如窗口注意力或片段级执行),而训练阶段仍保留全上下文注意力。这种不匹配导致:模型在训练中可以自由使用当前片段之外的完整历史信息来更新参数,但在推理时却无法获取这些信息,因为推理时的执行规则被限制了。这会造成模型对训练时可用但推理时不可用的信息产生过度依赖,从而削弱长上下文设定下的稳定性和泛化性。作者指出,该问题的本质是执行语义(execution semantics)和状态转换语义(state-transition semantics)在训练和推理之间存在根本性不一致。 针对上述痛点,本文的工作目标是:设计一种训练-推理一致的片段级执行框架,使得模型从训练开始就以受限的片段式前向语义进行学习,从而消除推理时的不适应。同时,框架必须保留足够的远程信息访问能力,以防性能过度下降。

方法:核心思路与技术路线

1. 片段级执行与跨片段状态接口

作者将输入序列切分为固定大小的片段(segment)。处理每个片段时,模型仅对当前片段内的token执行因果自注意力。为了跨片段传递信息,每个片段计算完成后,保留其KV缓存的一个固定大小的尾部(tail),称为“携带的KV尾部”(carried KV tail),记为C_{i-1}。这个尾部作为唯一的可微分接口,从前一片段传递到当前片段。携带的KV尾部大小是固定的(例如128个token的KV),因此跨片段信息流被约束在一个有限、可控的通道中。 在训练和推理中,携带的KV尾部以完全相同的方式参与当前片段的前向计算:当前片段内的注意力头可以访问该尾部中的KV,从而获取前一片段的部分上下文信息。该尾部是唯一允许参与梯度传播的跨片段状态,所有更早的历史KV都不直接携带梯度。

2. 训练时的梯度限制:截断反向传播(TBPTT)

为了确保训练阶段的梯度计算与推理阶段的一致性,作者引入截断反向传播(Truncated Backpropagation Through Time, TBPTT)。其核心思想是:沿着跨片段的状态链(由携带的KV尾部串联而成),梯度传播只允许向后回溯K步。也就是说,当计算当前片段损失的梯度时,只有最近K次携带的KV尾部状态转换参与梯度计算,更早的历史片段的状态转换产生的梯度被截断(不计算)。 这种截断不是对真实梯度的近似,而是在“训练-推理一致的有限递归”假定下的精确梯度计算,因为推理时模型也只会用到最近K步的携带KV尾部(如果K=1,则只依赖紧邻前一段)。具体来说:

  • 假设当前片段为第i段,其前向计算依赖于携带的KV尾部C_{i-1},而C_{i-1}又依赖于前一片段的输出和其携带的C_{i-2},以此类推。
  • 在反向传播中,只有连续的K段之间的梯度链被保留;对于K步之前的状态,计算图被截断。
  • 这么做的好处是:训练时模型无法通过梯度学习到长期历史信息(因为截断了),从而被迫只依赖最近K段携带的KV状态来更新参数,这与推理时仅拥有有限历史上下文的情形完全一致。

消融实验表明,K=1(即仅依赖紧邻前一段)就是最优设置。换句话说,仅仅依靠前一个片段携带的KV尾部就足以实现与全上下文注意力接近的性能。

3. 前向只读的检索前缀

仅仅携带最近一个片段的KV尾部,显然不足以让模型访问更早的历史信息。为了在不引入额外梯度依赖的情况下扩展长程证据访问,框架增加了一个只读的检索前缀机制(retrieved prefix)。具体来说,模型维护一个“过去KV池”(past-only KV pool),其中存储了所有历史片段的KV缓存(但只用于前向读取,不参与梯度)。 在处理当前片段时,除了使用携带的KV尾部C_{i-1}外,模型还可以选择性地读取该池中的一个KV前缀R_{i-1}(例如检索与当前片段相关性最高的若干历史token的KV)。该检索前缀在整个前向传播中只读(forward-only),即:在计算当前片段损失时,梯度不会通过检索路径反向传播到更早的片段。这意味着,检索行为本身不参与参数更新,仅作为模型前向推理时获得额外上下文的一种手段。 这种设计保证了核心跨片段学习完全发生在携带的KV尾部及其梯度链路中,检索前缀只是辅助,不会破坏训练-推理一致性。

4. 头-层稀疏长程架构

为了高效地实现上述两种跨片段信息流(携带KV尾部 + 检索前缀),作者设计了头-层稀疏的长程头机制(Head- and Layer-Sparse Long-Range Heads)。其基本思想是:并非所有注意力头都需要处理远程信息,大多数头只负责局部计算。 具体实现分为两步:

  • 非长程层(non-long-range layer) 在该类层中,注意力头分为两种:
  • 局部头(local heads):对当前片段内的token执行因果注意力,同时可以访问携带的KV尾部(即前一段尾部)。
  • 长程头(long-range heads):仅对当前片段内的token执行因果注意力,不访问任何跨片段KV。这类头不承担跨片段信息传递任务。
  • 长程启用层(long-range-enabled layer) 在此类层中,局部头的行为不变;长程头则额外被允许访问从过去KV池中检索到的前缀R_{i-1}(蓝色路径),以及携带的KV尾部C_{i-1}(绿色路径)。但注意,检索路径仅在前向使用,不参与梯度。

这种设计带来的好处:只有少数层中的少数头(即长程头)需要访问检索前缀,而大多数层和大多数头只进行局部状态携带计算。因此,训练和推理时的前向语义完全一致,且计算开销可控。图3(原论文Figure 3)直观展示了这种架构:在非长程层中,局部头访问当前片段+前一段尾部,长程头只做局部;在长程启用层中,长程头额外访问检索前缀。

5. 训练-推理一致的总体流程

综上,每个片段的前向计算接收两个跨片段输入:

  • 携带的KV尾部C_{i-1}(可微分,参与梯度传播限K步)
  • 可选的检索前缀R_{i-1}(只读,不参与梯度)

更新片段后输出新的携带尾部C_i(其大小固定)。在训练阶段,TBPTT确保梯度沿携带尾部链最多传播K步;检索路径和历史路径均被截断梯度。在推理阶段,完全等价的前向执行被复用,没有训练时额外的梯度计算。由此,训练和推理的执行语义严格对齐。

配图:方法结构

Figure 1. Peak GPU memory consumption during long-context prefill. 来源:原论文 PDF 第 1 页。 Figure 2. Training–inference consistent segmented execution. A sequence is processed segment by segment with two cross-segment inputs: a carried KV tail Ci−1 (the only differentiable state that propagates across segments) and an optional retrieved prefix Ri−1 read from a past-only KV pool. During training, TBPTT with depth K truncates credit assignment along the state chain (red cross), so gradients flow through C for at most K segment transitions (blue brace), while the retrieval path and earlier history are forward-only (no gradient). 来源:原论文 PDF 第 3 页。 Figure 3. Head- and layer-sparse long-range retrieval. (a) In a non-long-range layer ℓ/∈Llong, local heads attend to within-segment tokens and the carried KV state from the previous segment (green), while long-range heads use within-segment causal attention only (orange). (b) In a long-range-enabled layer ℓ∈Llong, local heads remain unchanged, while long-range heads additionally attend to a retrieved prefix from a past-only KV pool (blue). In all cases, attention within the current segment remains causal. 来源:原论文 PDF 第 4 页。

实验:设置、指标与结果

数据集与模型

实验主要在PG19数据集上进行评估。PG19是由arXiv书籍组成的测试集,常用于长上下文语言建模。作者使用两种基础模型:LLaMA2-32K(上下文长度32K)和LLaMA2-80K(上下文长度80K)。训练时,模型采用上述训练-推理一致片段级框架进行微调;推理时,使用相同的片段设置。上下文长度从4K变化到64K(LLaMA2-32K最大32K,LLaMA2-80K可达64K以上)。

基线

主要的基线是全上下文注意力,即使用FlashAttention-2执行标准自注意力。此外,也包括Chunked Prefill(将预填充分块但保持完整上下文)等执行级基线。对于推理高效基线,论文对比了TVT等仅推理时片段级执行的方法(但原文未在数值细节中列出所有对比,仅提到“与强推理高效基线相比具有竞争性的延迟-内存权衡”)。

指标

主要评价指标包括:

  • 困惑度(Perplexity, PPL):衡量模型对测试数据的预测能力,越低越好。
  • 峰值GPU内存(Peak GPU Memory):在预填充(prefill)阶段,即处理全部上下文生成第一个token时的最大显存占用,单位为GB。
  • 延迟(Latency):生成token的平均时间。但论文主要报告困惑度和内存结果。

主要结果

困惑度比较:在原论文Figure 4中,展示了LLaMA2-32K和LLaMA2-80K在PG19测试集上,不同评估上下文长度(4K、8K、16K、32K、64K)下的困惑度。结果显示,所提出的方法(图中标记为Ours)在每种长度下的困惑度曲线与全上下文注意力(Full Context)几乎重合。例如,在LLaMA2-32K上,32K上下文时两者困惑度均在约5.0左右;在LLaMA2-80K上,64K上下文时困惑度均在约4.8左右。这表明,尽管执行被限制为片段级且仅携带尾部状态,模型的语言建模能力并未显著下降。 峰值内存比较:原论文Figure 1展示了在LLaMA2架构下,不同上下文长度(4K到64K)预填充阶段的峰值GPU内存。FlashAttention-2在64K上下文时内存接近80GB;Chunked Prefill略低但仍超过60GB;而本文方法在64K下不到15GB。外推到128K,摘要中明确报告:本文方法在128K上下文预填充时峰值内存比FlashAttention低约6倍。由于FlashAttention在128K下内存需求约80GB(根据图形趋势推断),本文方法仅需约13-14GB,这使得单GPU就可以处理128K上下文,而全上下文注意力需要约80GB显存(至少需要A100 80GB或更大)。 延迟-内存权衡:虽然未给出具体数值,但论文指出该方法与强推理高效基线(如TVT)相比,在延迟和内存的权衡上具有竞争力。

消融与分析

TBPTT深度K的选择:作者进行了TBP深度K的消融实验,测试K=1,2,3等值对困惑度的影响。结果表明,K=1(即只从紧邻前一段携带梯度)已经足够,且是最优选择。增大K并未带来困惑度的改善,反而可能增加计算开销。这说明,跨片段依赖只需要最新一个片段的KV尾部即可捕获足够的局部连续性,更长范围的依赖可以通过前向只读检索来弥补。该结果验证了框架设计的合理性:更严格的梯度截断反而有助于防止过拟合训练时可用但推理时不可用的远距离信息。 检索前缀的影响:虽然论文未在消融部分详细展开,但从框架设计可以推断,检索前缀机制对于处理需要远程记忆的任务(如文档摘要中的开头关键信息)至关重要。实验设置中,PG19是一个连续文本,可能不需要强检索,因此主要表现困扰适度接近全上下文注意力。但论文未明确给出是否使用检索的对比实验。

配图:实验结果

Figure 4. Perplexity on PG19 test under varying evaluation context lengths. Results are reported for LLaMA2-32K and LLaMA2-80K. 来源:原论文 PDF 第 7 页。

结论:贡献、局限与启发

主要贡献

  • 提出了训练-推理一致的片段级执行框架,将片段级执行从推理优化提升为建模假设,从而从根本上消除训练-推理语义不匹配问题。
  • 理论证明了在严格受控的跨片段接口状态下,截断反向传播可计算推理一致目标的精确梯度,而非近似。
  • 通过头-层稀疏长程架构,高效实现了局部连续通道(携带KV尾部)和远程只读通道(检索前缀)的解耦,保持了训练-推理一致。
  • 在PG19等长上下文基准上,获得与全上下文注意力相当的困惑度,同时大幅降低内存消耗,128K上下文中预填充峰值内存比FlashAttention低约6倍。

局限性

原文未明确说明局限性。可能的局限性包括:检索前缀的构建需要额外存储和检索机制,其效率可能影响整体延迟;该框架目前仅在语言建模任务上评估,对需要复杂推理或精确撤销上下文的长上下文应用(如长文档QA)的效果有待进一步验证;TBPTT K=1的设置可能在某些需要跨长距离细粒度对齐的任务中不足,尽管困惑度指标未显示。

启发

本工作提醒研究社区:在长上下文LLM的优化中,“训练-推理一致性”是一个不容忽视的设计原则。以往仅仅将受限执行看作推理加速的手段,可能导致训练和推理之间的经验差异。本文展示了一种优雅的解法:通过固定大小的KV接口和受控的梯度传播,使得模型天然适应受限推理环境。未来工作可以将此思路扩展到更大的模型(如70B以上)、更长的上下文(如1M token)以及更多的下游任务。此外,检索前缀机制的设计可以进一步与近期流行的RAG或记忆增强方法结合,在保持训练-推理一致的前提下提升远程记忆能力。

原文信息

原文链接:https://arxiv.org/abs/2605.11744v1 (论文PDF可在该链接获取,包含详细公式、图表及分析。)

成为VIP会员查看完整内容
8

相关内容

ICML 2026 | 理解上下文持续学习中的泛化与遗忘
专知会员服务
13+阅读 · 5月28日
Llama-3-SynE:实现有效且高效的大语言模型持续预训练
专知会员服务
36+阅读 · 2024年7月30日
ICML2020 图神经网络的预训练
图与推荐
12+阅读 · 2020年4月4日
干货|从LSTM到Seq2Seq
全球人工智能
15+阅读 · 2018年1月9日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
1+阅读 · 2015年12月31日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
2+阅读 · 2015年12月31日
国家自然科学基金
1+阅读 · 2015年12月31日
国家自然科学基金
4+阅读 · 2014年12月31日
国家自然科学基金
18+阅读 · 2012年12月31日
VIP会员
最新内容
《履带式无人地面战车技术发展现状》
专知会员服务
3+阅读 · 8月2日
《无人机脆弱性利用:网络空间力量的新域》
专知会员服务
3+阅读 · 8月1日
美空军如何将人工智能从战场部署至后方机关
专知会员服务
12+阅读 · 7月31日
《史诗怒火行动:多域前瞻评估》49页报告
专知会员服务
9+阅读 · 7月31日
《英国防部:未来空战系统数字化战略》33页
专知会员服务
6+阅读 · 7月31日
《面向自主飞行网络的智能体人工智能架构》
专知会员服务
9+阅读 · 7月31日
相关基金
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
1+阅读 · 2015年12月31日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
0+阅读 · 2015年12月31日
国家自然科学基金
2+阅读 · 2015年12月31日
国家自然科学基金
1+阅读 · 2015年12月31日
国家自然科学基金
4+阅读 · 2014年12月31日
国家自然科学基金
18+阅读 · 2012年12月31日
微信扫码咨询专知VIP会员