学习如何遗忘:针对长上下文稀疏注意力的微调
摘要
本文提出了一种新方法,用于对具有稀疏注意力的transformer语言模型进行微调,以实现高效的长上下文推理,通常优于使用精确注意力训练的模型,并介绍了高效的实现和一个新的开源库。
arXiv:2608.19920v1 公告类型:新
摘要:许多先前工作通过稀疏注意力解决了键值(KV)缓存选择和压缩问题,以实现transformer语言模型的长上下文推理,而无需过多硬件预算。我们提供了一种新方法,用于对具有稀疏注意力的模型进行微调。它适用于任何KV缓存策略,可在中等硬件预算下运行(例如,单个40 GB RAM的Nvidia A100 GPU),并允许模型与策略共同适应,通常优于使用精确注意力(序列并行)训练的模型。我们还提供了H2O稀疏注意力的高效实现(我们实验中的领先策略),并支持专用的缩放点积注意力内核。KeysAndValues(https://github.com/awslabs/keys_values),一个用于长上下文推理和微调的新开源库,提供了易于使用且高性能的代码,涵盖本文讨论的所有方法。
查看缓存全文
缓存时间: 2026/08/21 10:14
# 学习如何遗忘:面向长上下文稀疏注意力的微调方法 来源:https://arxiv.org/html/2608.19920 **作者** Matthias Seeger 注:通讯作者 邮箱:[email protected] 单位:Amazon Web Services Vihang Patil 单位:Amazon 邮箱:[email protected] Konstantinos Benidis 单位:Amazon Web Services 邮箱:[email protected] Sebastian Schelter 单位:柏林工业大学 邮箱:[email protected] ###### 摘要 大量先前研究通过稀疏注意力机制解决键值缓存的选择与压缩问题,旨在使Transformer语言模型能够在有限的硬件预算内实现长上下文推理。本文提出一种新的模型微调方法,该方法适用于任意键值缓存策略,可在中等硬件配置下运行(例如单块配备40GB显存的NVIDIA A100显卡),并允许模型与缓存策略协同优化。实验表明,该方法常能超越采用精确注意力机制(序列并行)训练的模型。我们还提供了H2O稀疏注意力(本实验中的最优策略)的高效实现方案,包含专用的缩放点积注意力内核支持。开源库 KeysAndValues (https://github.com/awslabs/keys_values) 为长上下文推理与微调提供了易用且高性能的代码实现。 ## 1 引言 现代大型语言模型需要处理极长的上下文(即大量标记),以调用具有较大输出的工具、运行思维链推理或维持多轮对话。虽然朴素Transformer实现的计算复杂度随上下文长度呈二次方增长,内存消耗呈线性增长,但在实现近似线性时间与恒定内存缩放方面已取得显著进展。其中,稀疏注意力是一个特别富有成效的方向:将键值信息存储在固定大小的键值缓存中,缓存满时按策略淘汰插槽,由此提出了多种淘汰策略。 本文研究如何在中等硬件预算下(实验使用单块配备40GB显存的NVIDIA A100显卡为4B参数模型计算梯度)对采用稀疏注意力的Transformer语言模型进行后训练。我们提出的新方法适用于任意键值缓存策略,且除稀疏注意力外无需额外近似处理。长上下文基准测试表明,训练算法能使模型与键值缓存策略协同优化,常优于采用精确注意力(序列并行)训练的模型。此外,该微调方法的运行资源需求与稀疏注意力推理相当,可与分组查询注意力、量化等键值缓存压缩策略正交结合使用。 重击者预言机(H2O)是最知名的稀疏注意力策略之一。本文提出多项H2O优化方案,使其在专用缩放点积注意力内核支持下实现更高效的执行。H2O的多个变体在我们的实验中超越其他键值缓存策略,而快速实现方案为达到与vLLM等顶尖推理库(主要依赖上下文或序列并行)相竞争的延迟水平迈出了重要一步。 主要贡献包括: - • 提出在保留任意键值缓存策略的前提下微调Transformer语言模型的新方法。该方法运行资源需求与稀疏注意力推理相当,通过嵌套激活检查点、CPU卸载,结合自动求导保存张量打包技术利用键值缓存缓冲区的线性递归特性,以恒定资源处理任意长度序列。 - • 对重击者预言机(H2O)缓存策略进行方法论与实现改进,包括提供可返回求和注意力权重的Triton代码及FlashInfer缩放点积注意力内核。 - • 在多个长上下文基准测试中进行综合评估,证实当使用稀疏注意力推理时,训练算法常优于采用序列并行训练的模型。 - • 开源库 KeysAndValues (https://github.com/awslabs/keys_values),专为长上下文推理与微调设计。 ## 2 相关工作 长上下文推理的键值缓存压缩研究已有大量成果。简单思路包括分组注意力头以减少需存储的键值向量,或在查询-键矩阵中引入低秩结构。这类方法需在预训练阶段采用。缓存缓冲区可被量化至8位、4位甚至更低精度。 稀疏注意力是通用性强且具多种实现形式的方案。Big Bird规定固定的注意力稀疏模式。重击者预言机(H2O)将在第3.1.1节详述。Q-Hitter结合H2O与量化技术。SnapKV采用类似H2O的求和注意力权重,但仅在生成过程中某个节点做决策。预期注意力试图在强假设下估计键值缓存信息的未来相关性。FlexGen展示如何利用缓存层级最大化吞吐量。FastGen提供在不同缓存策略间投票的元策略。CAKE使用类H2O评分进行稀疏推理,并在层间分配总内存预算。其他稀疏注意力技术还包括多种改进方案。KVPop通过FlexAttention高效计算的未来注意力目标学习缓存策略。qTTT在测试时通过少量梯度更新改善推理结果。MInference与KVPress是提供多种稀疏注意力方法的开源库。ShadowKV是包含键值缓存选择的高吞吐长上下文推理系统。SCBench对长上下文推理方法进行综合实证分析。 优化的缩放点积注意力内核对快速推理与训练至关重要,FlashAttention开创了这一领域。FlashInfer针对1≪N_q≪N_k的推理场景优化。FlexAttention允许指定掩码与评分修改代码。 我们的主要贡献在于长上下文微调。现有研究可分为两类:第一类通过修改多头注意力解决第3.2节详述的困难。LongLoRA通过键值排列与重塑加速多头注意力,但未减少键值内存且要求缓存长度比例较小。原生稀疏注意力将多种固定策略的稀疏注意力内核集成到模型架构中。深度求索稀疏注意力是其改进变体。索引缓存通过层间共享索引器加速原生稀疏注意力。但这类方法需在预训练阶段采用,显著增加成本,且灵活性较低。长序列生成对顶部与底部三分之一层采用静态稀疏模式,中间层使用全注意力,仅略微加速训练而未降低内存需求。DMC采用键值缓存压缩形式,将新信息追加或累积到最近插槽,并提出相应训练启发式方法。从Mamba兴起的长短期记忆网络复兴尝试均未达到可支撑预训练成本的竞争力。YOCO提出单键值缓存块服务所有层的非Transformer架构。 第二类方法不压缩键值缓存,而是将存储与计算分布至多设备,可通过环形注意力或序列并行实现。尽管环形注意力观察到长序列可分块且计算图沿块轴分解,但最终块的计算依赖于所有层的键值缓存缓冲区,无法在单一设备表示。OOMB采用分块处理、激活检查点等技术,支持原生稀疏注意力与深度求索稀疏注意力。其通过自动求导隐藏键值缓存缓冲区的方式需定制CUDA内核,而我们通过差分编码管理缓冲区大小,实现策略无关的实现。 长旋转位置编码结合序列并行与非均匀旋转位置编码搜索。其他调优位置编码、数据混合与微调方案但不压缩键值缓存的研究包括多项工作。长上下文强化学习系统LongStraw强调不同采样间共享提示图与键值缓存,使用OOMB计算梯度但也可配置我们的方法。高度优化的上下文/序列并行实现是长上下文推理与微调的当前最优方案。 在第3.3节我们将讨论稀疏注意力方法在实践中应用较少的原因及改进方向。 ## 3 长上下文微调 本节首先介绍稀疏注意力与键值缓存,对H2O缓存策略提出多项改进,使其在专用内核支持下实现更高效执行。随后详述主要贡献:采用稀缺资源微调嵌入稀疏注意力的模型的新方法,其资源需求与稀疏注意力推理相当。与当前顶尖的上下文/序列并行技术不同,本方法中的GPU可用于提升吞吐量或处理更大模型。
相似文章
使用稀疏Transformer进行生成建模
OpenAI推出了稀疏Transformer,一种深度神经网络,将注意力机制的复杂度从O(N²)优化到O(N√N),使得能够对长度超过以前30倍的序列进行建模,适用于文本、图像和音频领域。该模型采用稀疏注意力模式和基于检查点的内存优化技术,可以训练深达128层的网络,在多个领域实现了最先进的性能。
语法引导的稀疏注意力机制:实现高效可解释的Transformer
本文介绍了一种针对Transformer的语法引导稀疏注意力机制,旨在通过利用语言结构来提高效率和可解释性。
混合大语言模型中的注意力遗忘:思维链微调如何破坏长程记忆及其修复方法
本文发现,在混合线性注意力模型中,思维链监督微调通过将注意力梯度偏向短程模式,从而降低长上下文召回能力,并提出QK-Restore——一种无需训练的方法,可在保留推理性能的同时恢复长上下文召回。
分层稀疏注意力机制的正确实现:迈向无限上下文建模
提出HiLS注意力机制,一种基于块的稀疏注意力方法,通过语言模型损失端到端学习块选择,性能可与全注意力媲美,同时支持超长上下文外推和更快的推理速度。
@VukRosic99: 长上下文Transformer面临两大瓶颈:二次注意力计算和KV缓存(在1M tokens时可达数百GB)…
MiniCPM-SALA是一款9B参数的混合注意力模型,通过在稀疏注意力和线性注意力之间交替插入(每3个线性层插入1个稀疏层)来克服长上下文Transformer的二次计算和KV缓存瓶颈。在256K tokens下,其推理速度比Qwen3-8B快3.5倍,并能在消费级GPU上支持高达1M tokens。该模型采用经济高效的持续训练方法,训练成本降低约75%。