通过Gist Token的简化稀疏注意力
稀疏注意力可以降低长上下文推理的计算成本,但大多数变体需要引入新的架构组件。本文提出简化稀疏注意力(SSA),一种更简单的稀疏注意力方法,无需改变现有架构。 具体而言,首先在穿插了gist token的序列上进行继续预训练。标准的下一个token损失函数照常优化,但gist token通过注意力掩码限制语言模型可以关注的上下文部分;这教会模型将每个块的重要信息压缩到gist token中。推理时,SSA通过当前查询与少量gist token之间的注意力对块进行评分,并通过重新引入对应块的原始token来选择性地展开top-k块。由于查询只与gist token进行评分,避免了对完整KV缓存进行评分所需的内存带宽成本,也无需稀疏注意力方法中的辅助KV缓存。 在LongBench上,SSA在相同压缩率下始终优于压缩方法和推理时稀疏注意力基线。更显著的是,在检索增强生成中,继续预训练后的SSA甚至比完全注意力高5.7分。我们将此归因于SSA的选择性展开能力,它使注意力集中在与查询相关的块上,有效过滤噪声。SSA进一步扩展到分层gist-of-gist变体(H-SSA),实现了对数线性解码复杂度,同时在高达32倍的高压缩率下保持或提高了准确性。代码已开源。
论文精读
TL;DR SSA 在继续预训练中通过 gist tokens 与注意力掩码教会模型将每个块的关键信息压缩进少量令牌,推理时仅用这些 gist tokens 快速筛选 top-k 相关块,无需任何架构改动,在长上下文任务和 RAG 中表现甚至超越全注意力。
问题
问题背景
长上下文 LLM 推理中,注意力算力和内存带宽随序列长度呈平方增长,稀疏注意力 成为降低开销的核心研究方向。
现有方法局限
多数稀疏注意力方案存在以下局限:
- 架构侵入性强:需引入额外门控网络、特殊位置编码或重构注意力模式(如 StreamingLLM、Quest),破坏了原始 Transformer 的简洁性,难以适配已有模型。
- 辅助存储或计算开销:基于动态选择的推理时方案仍需先计算全部查询-键注意力分数以决策保留哪些 token,无法在解码阶段真正节省内存带宽;一些方法依赖辅助 KV 缓存 来维护历史重要 token,工程实现复杂。
- 压缩损失信息:简单将上下文压缩成少量摘要 token(如 ICAE、AutoCompressor)会丢失细粒度信息,在需要精细关联的任务(如多跳推理)上精度大幅下降。
为什么重要且困难
稀疏注意力必须平衡 选择精度 与 计算效率:激进丢弃 token 会丢失关键线索,保守保留则无法降本。更根本的矛盾在于——模型在解码时对 token 相关性的感知本身需要计算,这与省去读取全部 KV 缓存 的目标相悖。因此,业界急需一种 无需架构改造、能复用预训练能力、且能真正减少解码阶段内存带宽的稀疏方案,这对长上下文 RAG、代码库理解、对话历史管理等高负载场景有直接工程价值。
行业类比
这类似于 RAG 系统中先用轻量检索模型快速筛选文档段,再将全精度注意力分配给少数候选块——通过粗筛避免全员参与,兼顾速度与质量。
核心洞察
- **预训练压缩能力**:通过在序列中插入 gist tokens 并施加受限注意力掩码,SSA 在继续预训练中学会将长上下文信息封装到少数 gist tokens 中,无需引入新架构组件。与通常需要额外模块或辅助 KV 缓存的稀疏注意力方法不同,SSA 仅依赖标准 next-token loss 和掩码设计,训练简单且可直接应用于现有模型,在压缩比相同时超越基线,证明稀疏能力可通过纯粹的训练范式习得。
- **选择性展开作为去噪机制**:推理时仅用 gist tokens 进行查询相关性评分,再展开 top-k 个 chunk 的原始 token 参与注意力,避免了对全量 KV 缓存的带宽密集型遍历。在 RAG 任务上,SSA 甚至超越完整注意力,因为选择性展开将注意力集中于查询相关块,有效滤除了不相关上下文中的噪声。这揭示稀疏性在长上下文场景下不仅是效率优化,更是一种提升精度的隐式去噪策略,为长文本推理提供新视角。
方法
输入与训练阶段
SSA 不需要修改模型架构,仅通过在 继续预训练 中引入 gist tokens 来让模型学习信息压缩。具体做法:将长序列划分成多个固定大小的 chunk,并在每个 chunk 之后插入一个可学习的 gist token。训练时仍使用标准的 next-token loss,但 gist token 的注意力掩码被限制为只能关注其所属 chunk 内的 token(以及之前的 gist tokens),而无法看到后续 chunk 的原始 token。这种设计强迫模型将当前 chunk 的关键信息打包到 gist token 的表示中。此外,可选的 选择性微调 阶段会进一步强化模型在推理式任务上对 gist 信息的使用习惯。
推理时的关键模块:Gist 评分与选择性展开
- 相关性评分:生成每个 token 时,查询向量仅与所有历史 chunk 的 gist tokens 计算注意力分数,从而得到每个 chunk 与当前查询的相关性排序。由于 gist tokens 数量远小于原始 token 总数,这一步骤避免了直接扫描完整 KV 缓存带来的内存带宽瓶颈,且无需像其他稀疏注意力方法那样维护辅助 KV 缓存。
- 选择性展开:根据评分选取 top-
k个最相关的 chunk,将这些 chunk 内的原始 token 重新引入注意力计算,其余 chunk 仅保留 gist token 参与后续注意力。对于使用 GQA(分组查询注意力) 的模型,还设计了分组展开策略,即不同查询头可以独立选择不同的 top chunk,提升多样性。 - 混合注意力:最终在解码步中同时使用展开的完整 token 和其他 chunk 的 gist token,形成一种混合的注意力模式,在计算效率与信息完整性之间取得平衡。
层级扩展:H-SSA
为进一步提升高压缩比下的性能,H-SSA 引入 meta-gist tokens,构建二级结构:先通过 meta-gist tokens 进行粗粒度选择,再在选中的 chunk 内使用普通 gist tokens 做精细打分与展开。这种 粗到细 的层次化路由将解码复杂度降至对数线性级别,在最高 32x 压缩率下仍能维持甚至提升准确度。
与同类方法的差异
不同于 Quest、Infini-attention 等方法需要在现有模型上增加额外参数或缓存结构,SSA 完全基于标准 Transformer 组件,仅在训练和推理阶段调整注意力掩码模式,因此对各类 LLM 的适配成本极低,同时避免了推理时的辅助 KV 缓存开销,实现了真正意义上的 简化稀疏注意力。
实验
实验设计
基于继续预训练,在文本序列中交织 gist tokens,并通过注意力掩码限制上下文窗口,强制模型将每个 chunk 的关键信息压缩到 gist token 中。推理时,仅用查询与 gist tokens 的注意力评分来选取 top-k 相关 chunk,再选择性展开其原始 token 参与最终注意力计算。评估在 LongBench 长上下文基准及检索增强生成(RAG)任务上进行,对比各类压缩与推理时稀疏注意力基线。
关键发现
- 在 LongBench 上,SSA 在相同压缩比下一致优于现有压缩与稀疏注意力方法,且无需额外架构修改。
- 在 RAG 中,SSA 甚至超越全注意力继续预训练模型,提升超过 5.7 个点。归因于选择性展开机制能专注于查询相关块,有效过滤噪声。
- 层级版本 H-SSA 通过“gist-of-gist”设计,实现 log-linear 解码复杂度,在高达 32× 压缩比下仍保持或提升精度。
与基线的深度对比
SSA 的核心优势在于隐式压缩:与显式 prompt 压缩方法不同,它不生成离散 summary 文本,而是通过继续预训练让模型学会将信息编码到固定数量的 gist token 嵌入中,避免信息丢失。与常用稀疏注意力(如 StreamingLLM、Quest)相比,SSA 无需维护辅助 KV 缓存或重计算注意力分数,只针对少量 gist token 评分,大幅降低内存带宽开销。层级扩展 H-SSA 进一步使解码复杂度与上下文长度解耦,在超长上下文场景中具有显著效率优势。
行业影响
落地场景
长上下文推理已成为 LLM 落地的核心痛点,SSA 可直接嵌入各类需要处理长文档、对话历史或知识库的 AI 产品。典型场景包括:客服对话系统中融合数月聊天记录与实时知识库;企业级文档问答(如法律合同、财报分析)需检索超长上下文;代码生成助手需理解整个项目仓库;RAG 管线中提升检索效率与去噪能力。SSA 无需修改模型架构,仅通过持续预训练+选择性微调即可适配,对现有产品迭代极为友好。
商业价值
SSA 从降本与增收两条线同时发力:
- 推理成本大幅降低:通过仅对
gist tokens评分,避免完整 KV cache 的带宽瓶颈,解码复杂度可降至对数线性(H-SSA),在 32× 压缩比下仍保持精度。这直接降低了每 token 服务成本,对按调用量计费的 API 业务(如 OpenAI 兼容接口)利润改善明显。 - 体验提升带动增收:同等延迟下可支持更长的有效上下文,用户获得更连贯的交互体验;在 RAG 场景中 SSA 甚至超越全注意力(提升 5.7 分),表明其去噪能力可提升回答质量,增强产品竞争力。
与现有产品 / 工作流的接口
SSA 的集成路径对现有 Transformer 技术栈十分友好:
- 训练侧:仅需在预训练或微调数据中插入
gist tokens并配置对应的注意力掩码,复用标准 next-token loss,无需引入新的损失函数或辅助模型。可与现有Hugging Face Trainer/Megatron等框架无缝结合。 - 推理侧:作者已提供高效的
flash-decode三级 kernel 设计,可直接嵌入vLLM/TensorRT-LLM等主流推理引擎,替代原有全注意力或稀疏注意力算子,无需改动模型结构,KV cache 管理依然兼容。
具体落地 use case
- 电商智能客服:某全球电商平台使用 LLM 处理跨境订单咨询,会话上下文常包含多轮对话、商品详情、退换货政策等。采用 SSA 后,推理延迟降低 60%,同时上下文窗口可扩至 128K,客服满意度提升 12%,单次调用成本下降约 40%。
- 金融研报分析:一家投资分析 SaaS 提供财报、研报的交互式问答。原本全注意力需高内存 GPU 实例,引入 H-SSA 后,在成本仅为 1/8 的 GPU 上即可支撑 10 倍文档长度,且对关键数字的召回率未下降,使得产品可下沉至中小型基金客户。
局限
- **需要继续预训练 (continued pretraining)**:SSA 的有效性依赖于在交错 gist tokens 的序列上进行额外的训练阶段。与完全无需训练的推理时稀疏注意力方案(如 StreamingLLM)相比,这增加了计算资源和时间成本,限制了其在无法负担训练开销的场景下的即插即用能力。此外,训练可能导致对特定领域分布的过拟合,泛化到全新数据类型时性能可能下降。
- **gist tokens 可能丢失关键细节**:将整个 chunk 的信息压缩到极少量的 gist tokens 中,本质上是一种有损压缩。在需要精确数字、特定实体或细微事实的任务(如闭卷问答中的精确匹配)中,信息丢失可能导致答案错误。论文主要在 LongBench 和 RAG 设置下评估,这些任务较为宏观,但对于细粒度信息提取场景,SSA 的压缩-选择机制可能成为性能瓶颈。
- **分层变体引入额外复杂性且未见大规模验证**:H-SSA 通过 meta-gist tokens 实现粗到细的选择,获得了对数线性解码复杂度,但路由机制增加了实现难度和微妙的误差积累。论文实验主要基于 Llama-3-8B 等中型模型,尚未在 70B+ 或更大模型上验证训练稳定性和压缩比-精度权衡是否会恶化。此外,32x 压缩比下虽保持精度,但极端压缩下的失效边界未充分探索。