论文
SAS:通过上下文排序的端到端优化实现简单的注意力稀疏化
SAS: Simple Attention Sparsification via End-to-End Optimization of Context Ranking
摘要
训练后注意力稀疏化通过为每个查询选择一小组上下文单元(词元或块)来减少预训练 Transformer 的二次累积注意力成本。现有的可训练方法通常使用轻量级选择器对上下文单元进行评分,然后进行硬性 Top-K 选择,以阻止语言建模损失的梯度。因此,这些方法通常会提取分层的密集注意力分布。尽管这鼓励选择器通过原始模型中的密集注意力权重对上下文单元进行排名,但排名与它们在固定注意力预算(即每个查询的参与上下文单元的数量)下对预测的影响并不直接一致,可能会在不太有用的单元上浪费有限的预算。为了解决这种不一致问题,我们提出了简单注意力稀疏化(SAS),这是一种门控稀疏注意力机制,可以通过语言建模损失来优化端到端的上下文排名。关键思想是在训练期间将选择器的连续分数注入注意力logits中,允许损失通过标准反向传播来更新选择器。我们确定了这个简单设计在实践中良好运行的关键选择:将门以对数形式放置在注意力 Softmax 内部,使用归一化 SoftMax 门根据始终保留的当前块来校准历史上下文,并保留连续选择器分数,以便模型学习相对优先级而不仅仅是硬选择。为了支持长序列训练,我们实现了一个内存高效的 Triton 内核,它将 SAS 集成到 FlashAttention 式计算中。在推理、长上下文理解和代理任务中,SAS 在注意力预算上始终优于可训练的稀疏注意力基线,在预算紧张的情况下收益尤其大,展示了下游任务的更有效的上下文排序。