中文

Token Statistics Transformer:基于变分率减少的线性时间注意力

机器学习 2024-12-24 v1

摘要

注意力算子是Transformer架构的关键区分因素,这类架构在多种任务上展现了最先进性能。然而,Transformer注意力算子常带来巨大的计算负担,其计算复杂度随token数量呈二次方增长。本文提出一种注意力算子,其计算复杂度随token数量呈线性增长。我们通过扩展先前工作,指出Transformer结构天然源于“白盒”架构设计,即网络各层旨在实现对最大编码率减少目标(MCR2^2)的递增优化步骤。具体而言,我们推导了MCR2^2目标的新型变分形式,证明其无限梯度下降所得的架构导致一种新的注意力模块,即Token Statistics Self-Attention(TSSA)。TSSA具有线性计算和内存复杂度,与典型注意力架构形成鲜明区别——后者计算token之间的两两相似性。在视觉、语言和长序列任务上的实验表明,仅替换标准自注意力为TSSA(即我们称的Token Statistics Transformer, ToST),即可在计算效率和可解释性方面显著优于传统Transformer,同时保持竞争性能。我们的结果也在一定程度上质疑了成对相似性风格注意力机制对Transformer架构成功的关键性。代码将在https://github.com/RobinWu218/ToST公开。

关键词

引用

@article{arxiv.2412.17810,
  title  = {Token Statistics Transformer: Linear-Time Attention via Variational Rate Reduction},
  author = {Ziyang Wu and Tianjiao Ding and Yifu Lu and Druv Pai and Jingyuan Zhang and Weida Wang and Yaodong Yu and Yi Ma and Benjamin D. Haeffele},
  journal= {arXiv preprint arXiv:2412.17810},
  year   = {2024}
}

备注

24 pages, 11 figures