← 返回列表

论文综述:SAS——通过端到端优化上下文排序实现简单的注意力稀疏化

SAS: Simple Attention Sparsification via End-to-End Optimization of Context Ranking

原文作者Zhiwei Li, Lei Zhu, Hao Gu, Xiang Hu, Yan Wang, Haitao Mi, Sirui Han, Leo Liang, Zhijiang Guo机构Tencent Hunyuan LLM Frontier; Hong Kong University of Science and Technology论文发布2026-09-11综述日期2026-09-18HF 票数🔺 61
稀疏注意力长上下文大语言模型推理Transformer优化端到端训练
📄 查看原文 →

一、论文是干什么的?

想象你在开一场持续好几个小时的长会,会议记录已经记满了几百页。轮到你现在要发言时,你不可能把几百页笔记从头到尾重新读一遍再组织语言——那样太慢了。更现实的做法是:快速扫一眼笔记的目录或关键词,挑出跟当前话题最相关的那几页,只精读这几页,然后发言。

大语言模型(LLM)在处理长文本时面临的正是类似问题。Transformer 模型里的注意力机制(attention)需要让每一个新词都去”回顾”前面所有已经出现过的词,计算量随着文本长度的平方增长——文本越长,算力开销涨得越猛。这在处理长文档、长对话历史、智能体(agent)多轮任务时会变得非常昂贵。

“注意力稀疏化”(attention sparsification)就是解决这个问题的思路:与其让模型”回顾”全部历史内容,不如训练一个”筛选器”,让它只挑出一小部分真正重要的历史片段来参与计算,其余的直接跳过。这篇论文提出的方法叫 SAS(Simple Attention Sparsification,简单注意力稀疏化),它的核心贡献是让这个”筛选器”能够通过标准的反向传播,跟着语言模型的训练损失一起被优化,而不需要依赖复杂的蒸馏或额外的训练技巧。论文在推理任务、长文本理解任务和智能体任务上做了实验,显示这种简单直接的训练方式效果反而更好。

二、核心方法与创新

2.1 先把历史内容切成”小笔记本”

SAS 首先把输入的上下文切成一个个连续的小块(block),每块固定包含 b=64b=64 个 token(可以理解为 64 个字词)。如果总长度是 nn,那么一共会切出 C=n/bC = n/b 个历史块,记作 {B1,B2,…,BC}\{B_1, B_2, \ldots, B_C\}。当前正在处理的那一小块内容(也就是刚生成或刚读到的最新片段)被称为”当前块”,它是永远保留、不会被丢弃的,因为离当前位置越近的内容通常越关键。

2.2 用一个轻量筛选器给每个笔记块打分

对于每一个历史块,SAS 用一个轻量级的打分器(沿用了此前 SeerAttention-R 论文提出的 AttnGate 结构)计算一个相关性分数:

s=Rθ(q,{KBm}m=1C)\mathbf{s} = \mathcal{R}_\theta(\mathbf{q}, \{\mathbf{K}_{B_m}\}_{m=1}^{C})

这里 q\mathbf{q} 是当前查询向量,KBm\mathbf{K}_{B_m} 是第 mm 个历史块的键向量,θ\theta 是筛选器自己的可学习参数。打完分之后,只保留分数最高的 Top-K 个历史块(再加上当前块)参与真正的注意力计算,其余块直接跳过,从而把原本要平方级增长的计算量降到与”总长度 ×\times 选中长度”成正比。

2.3 关键创新一:把打分”做成”softmax里的对数偏置项,而不是简单硬筛选

以往的可训练稀疏注意力方法(比如 SeerAttention-R)也有这样一个打分器,但它们的做法是:先用打分器的输出去模仿一个”完整版”(稠密)注意力模型每个位置该关注多少,也就是用蒸馏(distillation)的方式训练打分器,跟语言模型本身的训练损失是分开的。问题在于,Top-K 硬筛选是一个”非黑即白”的操作(选中就是1,没选中就是0),这种操作的梯度几乎处处为零,语言模型的损失没法直接告诉打分器”你选错了”或者”你选对了”。

SAS 的做法更直接:把打分器算出来的分数,转换成一个对数形式的”门控值”(gate),直接加到注意力计算内部的 softmax 之前:

oSAS=softmax(qKS⊤+log⁡gS)VS\mathbf{o}_{\mathrm{SAS}} = \mathrm{softmax}\left(\mathbf{q}\mathbf{K}_{\mathcal{S}}^\top + \log \mathbf{g}_{\mathcal{S}}\right)\mathbf{V}_{\mathcal{S}}

其中 S\mathcal{S} 表示被选中的块集合(当前块 + Top-K 历史块)。因为门控值现在是注意力计算这条”可微分链条”上的一环,语言模型的最终损失就可以顺着这条链条一路反向传播,直接更新打分器的参数 θ\theta:

∇θLLM\nabla_\theta \mathcal{L}_{\mathrm{LM}}

打个比方:以前是先让学生(打分器)“模仿”老师(稠密注意力)怎么划重点,划得像不像老师是唯一标准;SAS 则是让学生直接根据”考试成绩好不好”(语言模型损失)来调整自己划重点的方式——目标和最终效果直接挂钩,不再绕一个中间弯子。

论文的消融实验证实了”门放在 softmax 内部”比”门放在 softmax 外部(也就是只对最终的值向量做缩放)“效果好得多:内部门控在训练约1000步时的 GPQA-Diamond 准确率约为 53.7%,而外部门控只有约 43.2%。原因是内部门控能重新分配”注意力质量”本身,而外部门控只是简单缩放最终结果,信号更弱。

2.4 关键创新二:用归一化 softmax 门控校准”历史”与”当前”的权重

历史块的门控值不是随意设定的,而是先对所有历史块的分数做一次 softmax 归一化,再取对数:

log⁡gm=sm−LSE(s),Bm∈H\log g_m = s_m - \mathrm{LSE}(\mathbf{s}), \quad B_m \in \mathcal{H}

这里 H\mathcal{H} 是历史块集合,LSE\mathrm{LSE} 表示 LogSumExp(对数-求和-指数,是 softmax 归一化中常见的数学运算)。当前块则固定门控值 g0=1g_0 = 1(即不加偏置)。虽然 −LSE(s)-\mathrm{LSE}(\mathbf{s}) 这一项对所有历史块是共享的常数,但由于当前块不参与这个归一化,这个常数并不会在后续的注意力 softmax 中被抵消,从而真正起到了”校准历史上下文整体权重相对于当前块”的作用——避免历史块要么被无脑地看得比当前块还重要,要么被压得几乎不起作用。

2.5 关键创新三:保留连续分数,而不是非黑即白的硬选择

有一种直觉的做法是,既然最终要做 Top-K 硬选择,那训练时干脆也用硬选择,配合直通估计器(Straight-Through Estimator,STE)之类的技巧近似求梯度。但论文发现这样做效果明显更差:连续的软门控(soft gate)方案达到约 54.4% 准确率,而硬门控(hard STE)方案只有约 46.0%。

原因在于:软门控的权重天然被 softmax 限制在 0到1 之间,比较稳定;而硬门控在训练早期,打分器还没训练好、经常会把本该重要的块排到 Top-K 之外,这时候硬门控产生的梯度在数值上可能比软门控大出好几个数量级,导致训练不稳定。换句话说,“记笔记时先打个模糊的重要性分数,再挑重点”,比”一上来就非黑即白地划掉不选的内容”更容易学得又稳又好。

2.6 工程实现:把门控融合进 FlashAttention 风格的 Triton kernel

如果直接按公式实现,训练时需要先算出完整的注意力矩阵,再加上门控偏置,这对长序列来说会占用巨大的显存,几乎无法训练。为此,论文在附录中给出了一个 FlashAttention 风格的 Triton kernel 实现:在按”瓦片”(tile)遍历 Key/Value 分块计算 qKS⊤\mathbf{q}\mathbf{K}_{\mathcal{S}}^\top 的过程中,直接把每个被选中历史块的归一化对数门控值融合进去,同时屏蔽掉未被选中的块,当前块则保持不加偏置。反向传播时,则把每个被选中历史块内部所有注意力 logit 的梯度累加起来,得到这个块级门控值的梯度。这样就可以在不生成完整注意力矩阵的前提下,完成稀疏范围内的高效训练。论文提到,实际系统构建在 SGLang 推理框架之上,复用了它的分页 KV 缓存(Paged KV Cache)和 FlashInfer 注意力内核,分布式训练部分则基于字节跳动开源的 VeOmni 框架。

三、使用了哪些模型和计算资源?

基座模型:

  • 主实验(后训练/post-training 稀疏化)使用了 Qwen3 系列的三个规模:Qwen3-4B、Qwen3-8B、Qwen3-14B。
  • 另有一组”继续预训练”(continued pretraining)实验,使用 OLMo3-7B,从其 stage-1 基础检查点continue 训练。

训练数据与配置:

  • 后训练阶段:在 OpenR1-Math-220K 数据集(消融实验用其中约 9.37万条样本)上训练 1 个 epoch,全局批量大小 32,最大序列长度 32,768 个 token,学习率 1e-3,优化器为 AdamW。
  • 继续预训练阶段:批量大小 512,学习率先热身升到 2e-4 再余弦衰减到 2e-5,共训练约 13,000 步,相当于约 500亿(50B)个 token。

计算资源:

  • 论文正文没有直接给出训练所用的 GPU 型号、数量和总训练耗时。
  • 根据论文官方 GitHub 仓库(Tencent-Hunyuan/Simple-Attention-Sparsification)的说明,SAS 对资源需求较低,可以在 8 张 NVIDIA H20 GPU 上完成训练;但仓库中同样没有给出具体训练时长,因此训练时间在论文中未明确说明。
  • 推理速度评测在单张 GPU(张量并行度为1)上进行,开启了 CUDA Graph 优化,测试框架为 SGLang,具体 GPU 型号论文正文未指明。

四、实验结果

SAS 主要对比的基线是 SeerAttention-R——一种同样使用 AttnGate 打分结构、但采用”蒸馏稠密注意力分布”方式训练打分器的可训练稀疏注意力方法;此外还与训练无关(training-free)的方法如 Sliding Window(滑动窗口)、StreamingLLM(固定注意力汇聚点)、Quest(查询感知稀疏)等做了对比,以及一个继续预训练版本的基线 HiLS-Attention。总体结论是:SAS 在推理任务、长文本理解任务和智能体任务上,全面且持续地超过可训练的基线方法,尤其是在注意力预算(可以理解为”允许看多少历史内容”的额度)很紧张的情况下,优势更加明显。

以下是几组代表性数值(均为论文中报告的百分比准确率或分数,budget 指的是被选中参与计算的历史 token 总量上限):

任务/数据集模型budgetSASSeerAttention-R
MATH500(数学推理)Qwen3-4B102490.6584.67
GPQA-Diamond(研究生级问答)Qwen3-8B102453.1739.43
AIME24(数学竞赛)Qwen3-4B204868.8555.83
LongBench(长文本理解,8K以上段落)Qwen3-14B204853.951.5
BFCL多轮(智能体函数调用)Qwen3-4B204832.5029.00
VitaBench配送场景 Pass@4(智能体任务)Qwen3-14B409668.063.0

可以看到,在预算比较紧的 1024 这一档,论文原文指出 SAS 相比 SeerAttention-R 在 MATH500 上领先约6.0到7.7个百分点,在 GPQA-Diamond 上领先约10.6到15.5个百分点(覆盖 Qwen3-4B/8B/14B 三个规模),这正好印证了论文强调的”预算越紧张,端到端优化排序的价值越大”。

消融实验(在 GPQA-Diamond 上,Qwen3-4B、2048 token 预算,训练至1个epoch收敛后测得的准确率)进一步验证了每个设计选择的必要性:

设计维度更优选择数值对比
门控放置位置放在 softmax 内部内部约54.4% 对比 外部约41.6%
门控归一化方式softmax 归一化softmax约54.4% 对比 sigmoid约17.0%
分数保留方式连续软门控软门控约54.4% 对比 硬门控(STE)约46.0%
训练时的注意力范围稀疏范围内训练(效率更高)稀疏约54.8% 对比 全量约54.4%(结果相近甚至略优,且稀疏训练成本更低)

推理加速效果: 在批量大小为1时,随着上下文长度增长,标准稠密注意力的解码延迟线性增加,而 SAS 的延迟几乎保持不变,在 51.2万(512K)token 长度下最高可实现约 5.6 倍的延迟降低(论文给出的具体数值为 64K、256K、512K 下分别约 2.4倍、4.6倍、5.6倍);批量大小为8时,在 6.4万(64K)token 长度下加速比可以达到约 13 倍,且这个加速比对预算大小不敏感。不过论文也指出,当上下文长度极长(如512K)时,Top-K 挑选块本身的计算开销会成为新的瓶颈(约占总耗时的90%),而在较短上下文(如8K)时,真正的注意力计算仍是主要开销(约占58%)。

五、潜在应用与已落地应用

潜在应用方向:

  • 长文档问答、长篇报告摘要、法律/医疗等长文本密集型场景,可以在几乎不损失效果的情况下大幅降低推理成本。
  • 多轮对话或多轮智能体任务(如 BFCL、VitaBench 这类需要不断回顾历史工具调用记录的场景),可以只保留真正相关的历史片段。
  • 由于 SAS 的训练方式是”简单直接的端到端反向传播”,理论上比依赖蒸馏的方法更容易迁移到新的模型架构或新的任务上,减少了工程上”先训练一个稠密教师模型再做蒸馏”的中间步骤。

已落地情况:

六、网络上的讨论与评价

截至综述撰写时(2026年9月18日),该论文在 HuggingFace Papers 页面获得了 61 票,说明在社区中受到了一定关注。通过网络搜索,主要能找到的是论文本身在 arXiv、HuggingFace Papers、Papers with Code、HyperAI 等学术聚合平台上的转载和摘要信息,暂未搜索到 Reddit、Hacker News、X(Twitter)等社交平台上的实质性讨论帖或长评。这可能与论文发布时间较新(2026年9月11日提交)有关,社区讨论可能仍在发酵中,如实说明:未找到具体的社区讨论内容。

七、思维导图

mindmap
  root((SAS:端到端优化上下文排序的稀疏注意力))
    研究背景与问题
      Post-training稀疏选择器梯度阻断
        Top-K硬选择梯度处处为零
      SeerAttention-R依赖注意力蒸馏
        排序目标与预测影响不直接对齐
      本文目标:语言建模损失直接优化选择器
    核心方法SAS
      Block划分与选择器:block size 64,当前块B0恒保留,沿用AttnGate结构
      Log空间Gate注入softmax内部
        公式o等于softmax(qK转置加log g)V
        内部gate优于外部gate 54.4对41.6
      归一化softmax Gate与连续分数保留
        公式log g_m等于s_m减LSE(s),当前块g0恒为1
        软门控54.4对硬门控(STE)46.0
      Triton kernel融合FlashAttention
        避免物化完整注意力矩阵,基于SGLang FlashInfer与VeOmni
    训练与推理配置
      基座模型Qwen3-4B/8B/14B
      继续预训练OLMo3-7B
      训练数据OpenR1-Math-220K 93.7K样本
      8张NVIDIA H20 GPU
    实验设计与结果
      推理任务
        MATH500/GPQA-Diamond/AIME24/25,预算1024下GPQA-Diamond领先10.6到15.5分
      长上下文理解LongBench
        Qwen3-14B预算2048 8K+段落53.9对51.5
      智能体任务
        BFCL多轮预算2048领先3.5分
        VitaBench Pass@4配送场景68对63
      推理加速
        batch1下512K加速5.6倍
        batch8下64K加速约13倍
        长序列Top-K选择成新瓶颈占90%
    理论分析与洞察
      可微分门控使梯度直达选择器
      归一化保证历史与当前块权重可比
      软门控数值稳定性优于硬选择
    影响与展望
      长文档与多轮智能体任务潜在应用
      开源代码GitHub与HuggingFace模型权重
      腾讯混元与港科大工业界降本增效研究