文章摘要
让大模型支持长上下文输入需求迫切,但全注意力机制有计算复杂度高等问题,分块稀疏注意力也效果不佳。腾讯混元团队开源的HiLS - Attention,提出新范式,解决两大痛点。实验显示,它能兼顾效率与性能,打破“不可兼得”困境。

让大模型支持更长的上下文输入,已经成为智能体、深度推理以及海量资料整合等AI应用场景的核心需求。但标准全注意力机制的计算复杂度会随着序列长度呈平方级增长,同时还存在长度外推能力差、KV缓存随序列长度线性膨胀导致显存压力过大这三大核心问题。

近期,腾讯混元官方团队正式开源了HiLS-Attention,也就是分层地标稀疏注意力,提出了全新的分块稀疏注意力实现范式,首次从数学层面同时解决了分块重要性估计表达力不足、选择过程端到端不可导这两个核心痛点,真正实现了稀疏注意力的优化落地。

为了突破全注意力的性能瓶颈,业界将目光转向分块稀疏注意力:将完整上下文切分为多个独立的chunk,每个查询只选择与当前内容最相关的Top-K个chunk进行注意力计算,这样可以将计算和显存开销控制在常数级别。但遗憾的是,迄今为止没有任何一种分块稀疏注意力方案能够真正追平标准全注意力的效果。

分块稀疏注意力效果不佳的核心原因在于chunk的选择不够准确。要实现精准的chunk选择,首先需要准确估计每个chunk的重要性,但现有的主流方法都存在先天缺陷。

最常见的方案是使用均值池化来生成chunk的摘要表示,直接将chunk内所有key的平均值作为摘要key,这种方式计算得到的分数本质上是token logit的均值;另一些方案则改用最大logit来近似。但实际上,真正的chunk重要性应该是LogSumExp形式,其行为完全取决于chunk内的logit分布,只有在注意力均匀分布时均值池化才准确,而只有在单个token独占注意力时最大logit才准确。但真实场景中的logit分布会随着查询、注意力头以及输入数据剧烈变化,无论使用均值还是最大logit,都会系统性地错估chunk的重要性,导致真正关键的chunk被遗漏。

比如在最简单的单针大海捞针任务中,使用均值池化的NSA、DashAttention、InfLLM v2等方案在8K长度范围内就已经出现明显的性能下降,因为这类任务中关键信息由少数token集中呈现,均值池化会将这种注意力尖峰稀释掉,导致关键信息被掩盖。

既然非参数化的均值或最大logit表达能力不足,自然的思路是为每个chunk学习一个参数化的摘要表示,以更精准地概括整个chunk的内容。但现有的方法几乎都忽略了一个致命的断点:即使使用了参数化的summary,这些方法也仅用它来打分和选择Top-K的chunk,一旦chunk ID被硬选择出来,summary和打分结果就会被丢弃,不再参与后续的注意力计算。这导致语言建模损失的梯度无法反传到summary和打分过程中,Top-K选择是离散且不可导的操作,模型无法通过梯度反馈来优化summary的生成,使得summary的学习变成盲训,无法真正学会抑制无关chunk、提升关键chunk的权重。

这就引出了两个核心研究问题:一是需要构建数学表达能力足够的chunk重要性估计方法,二是需要让chunk summary能够随着语言建模损失进行端到端的训练。只有同时解决这两个问题,才能真正实现高效且精准的稀疏注意力。

HiLS-Attention的核心设计:让chunk选择变成可微分的分层softmax

HiLS-Attention的核心思路是将上述两个研究问题拆分为两个部分逐一解决。

第一部分:使用一阶泰勒展开构建表达能力足够的chunk算分函数

研究团队通过对LogSumExp进行一阶泰勒展开,发现chunk的对数重要性可以被近似为一个优雅的形式,该形式由两个部分组成:

  • 相关项:$ q^T c^k $,其中$ c^k $是chunk的summary key,本质上是对chunk内所有key进行一次注意力加权求和得到的结果;
  • 偏置项:$ c^b $,也就是该分布的熵,它可以自适应地在两种极端情况之间进行插值:当分布越均匀时,偏置项越接近$ \log S $,而当分布越集中时,偏置项越趋近于0。

这个熵偏置项正好弥补了均值和最大logit各自的缺陷:均值池化丢失了分布的集中度信息,最大logit丢失了分布的分散度信息,而熵偏置项可以统一这两种情况,让代理分数在任意分布下都能贴合真实的chunk重要性。

要生成这个summary key,只需要在每个chunk的末尾添加一个特殊的摘要token(landmark token),使用它来学习所有潜在query的中心表示,再对chunk内的token进行一次注意力计算,即可得到$ c^k $和$ c^b $。每个chunk的计算复杂度仅为$ O(S) $,整个序列的总复杂度为$ O(N) $,彻底摆脱了全注意力的平方级计算开销。

第二部分:打通端到端反传的断点,让summary参与完整的训练过程

HiLS将注意力权重分解为两级softmax:

  • chunk内softmax:在每个被选中的chunk内部,决定各个token之间的相对注意力权重;
  • chunk间softmax:使用代理质量$ \hat{Z}_{i,c} $来决定每个chunk整体能够分配到的注意力比例。

关键在于,代理质量$ \hat{Z}_{i,c} $直接被嵌入到前向传播的注意力权重计算中,这样语言建模损失的梯度就可以顺着前向计算图,反向传播到summary key和地标token上,让模型能够通过梯度反馈来优化summary的生成,真正实现端到端的训练。这一设计彻底解决了之前的断点问题,让chunk选择过程成为可学习的端到端过程,且训练和推理全程都保持原生的稀疏性。

系统实验验证:从345M到7B模型的全面验证

研究团队在345M、1.4B和7B三个不同的模型参数规模上进行了系统的实验验证,得到了高度一致的结论:

  • 短文本场景性能无损失:345M和1.4B模型从零开始训练时,HiLS-Attention在不同上下文长度和训练阶段的语言建模困惑度(PPL)与标准全注意力几乎完全重合,在8K长度下甚至略优于基线;
  • 超长上下文外推能力极强:仅使用8K长度的上下文进行训练,即可实现4M上下文(相当于512倍训练长度)的免训练外推,且保持90%以上的大海捞针任务准确率,远超标准全注意力;
  • 低成本改造现有模型:可以将OLMo3-7B这类标准全注意力模型快速转换为HiLS-Attention模型,仅需要续训50B token即可完成切换,短程任务性能无损失,在LongBench长序列任务的in-domain场景下甚至超越了全注意力基线,无缝继承了HiLS的外推能力,在out-of-domain长度场景下显著碾压YaRN等其他基线方案;
  • 推理速度大幅提升:在512K上下文长度下,prefill阶段加速13.5倍,单步decode阶段加速15.7倍。

这一结果首次打破了稀疏注意力长期以来“效率与性能不可兼得”的困境,实现了两者的同时提升。

研究团队最初的目标是让HiLS-Attention逼近全注意力诱导的chunk选择,但实验结果带来了意外的惊喜:HiLS不仅追平了朴素的分块稀疏注意力,还在长上下文检索任务上反超了标准全注意力。

这一现象的原因可能在于压缩本身具备去噪的效果:标准全注意力的固有问题在于,只要一个token的logit不是负无穷,就会分配到少量的注意力质量,随着上下文长度增加,这些无关token的微小噪声会不断累积,污染检索信号。而HiLS将多个key压缩为一个summary key时,不对齐的噪声会相互抵消,共享的语义信号则被保留下来,使得检索结果更加纯净,这也是HiLS在变量追踪(VT)这类多跳任务上能够比全注意力高出多达50%性能的核心原因。

写在最后

回顾HiLS-Attention的设计逻辑,其实非常清晰:

  1. 稀疏注意力的核心瓶颈在于chunk选择不准确;
  2. chunk选择不准确的根源在于均值或最大logit的系统性失准;
  3. 尝试使用参数化summary进行优化时,又遇到了端到端反传的断点问题;
  4. HiLS通过一阶泰勒展开构建了足够表达能力的算分函数,结合分层softmax打通了梯度反传的路径,一举解决了表达力和可微分两个核心问题。

HiLS-Attention证明了稀疏注意力可以同时提升效率和效果,其效果提升的根源可能来自于压缩带来的去噪效果,得到了更纯净的检索表征。这才是真正将分层稀疏注意力做到“Done Right”的解决方案。

以上内容不代表本平台立场,仅供读者参考