SAS:让稀疏注意力直接为预测结果学习上下文排序
长上下文模型的注意力计算会随着序列长度呈二次增长,因此,许多研究尝试只保留少量上下文 token 或块,以降低累计注意力成本。问题在于:如果可用的上下文预算固定,模型真正需要的并不是“最像原始稠密注意力分布”的内容,而是那些在被保留后最能改善最终预测的内容。
现有方法的错位
常见的可训练稀疏注意力方案会加入一个轻量选择器,为上下文单元打分,然后使用硬 Top-K 规则挑选内容。硬选择虽然便于控制计算量,却会切断语言建模损失传回选择器的梯度。因此,训练过程往往转向蒸馏原始模型的逐层稠密注意力分布,让选择器学习“原模型更关注什么”。
这种目标和有限预算下的实际需求并不完全一致。一个上下文单元可能获得较高的稠密注意力权重,却未必是删减其他内容后最有价值的候选。预算越紧,这种排序误差越容易造成性能损失。
SAS 的核心设计
SAS,即 Simple Attention Sparsification,试图用较简单的结构解决这一问题。它不在训练时立即把选择器分数硬化为离散选择,而是将连续分数作为门控信号注入注意力 logits,并放在 Softmax 内部。这样,语言建模损失可以通过标准反向传播直接影响选择器,使其围绕最终预测结果学习上下文优先级。
论文强调了三个实现细节:
- 在 Softmax 内使用对数形式的门控。 这种处理能让门控与注意力 logits 保持兼容,同时避免把选择机制简单地变成不可导的截断操作。
- 采用归一化 Softmax 门控。 历史上下文通常需要经过筛选,而当前块始终保留。归一化门控用于校准两类信息的相对贡献,减少因尺度不一致带来的偏差。
- 保留连续选择器分数。 模型学习的是候选之间的相对优先级,而不只是“入选”或“落选”两个标签,这有助于在预算收紧时形成更细致的排序。
为了支持长序列训练,作者还实现了一个面向 Triton 的内存高效内核,将 SAS 融入 FlashAttention 风格的计算流程。这样,方法不仅停留在选择策略层面,也考虑了稀疏机制在实际训练系统中的执行开销。
影响与局限
根据论文介绍,SAS 在推理、长上下文理解和智能体任务上,相比可训练稀疏注意力基线取得了更稳定的表现,且在注意力预算较紧时优势更明显。这说明,稀疏注意力的关键不只是减少被计算的上下文数量,还在于让排序目标与下游预测真正对齐。
这项工作对长上下文推理、上下文压缩和高效推理具有启发意义:未来的稀疏化方法可能不必过度依赖对稠密注意力的模仿,而应直接优化有限上下文预算下的任务损失。不过,素材未提供具体模型规模、数据集、压缩比例或绝对性能数字,因此其收益仍需结合完整论文中的实验设置进一步判断。
评论
正在确认登录状态……
正在加载评论……