中文

SUS backprop:面向长序列输入的变换器线性反向传播算法

机器学习 2025-06-06 v2 人工智能 计算与语言

摘要

对计算图中任意位置进行随机切断以设计无偏梯度估计器是直接的。在切断对计算影响不大的部分时,可为最小的随机梯度方差增加换来显著的反向传播计算节省,这在某些情况下确实如此。这种情况发生在变换器架构的注意力机制中。对于长序列,注意力成为限制因素,因为其计算需求随序列长度 nn 呈二次增长。与此同时,大多数注意力权重变得很小,因为大多数注意力头倾向于将给定 token 只与序列中少数几个其他 token 相连接。这些权重成为切断反向传播的有力目标。我们提出一种由单个参数 cc 控制的简单概率规则,切断大多数注意力权重的反向传播,使每个 token 每个注意力头至多保留 cc 个相互作用。这将注意力反向传播所需计算量降低至 c/nc/n 倍,使其从二次复杂度 O(n2)O(n^2) 降为线性复杂度 O(nc)O(nc)。我们通过实验验证,对典型变换器模型,切断约 99%99\% 的注意力梯度流(即取 c2530c \sim 25-30)仅导致 n2000n \sim 2000 时的梯度方差增加约 1%1\%,且随 nn 增大而降低。这一方法可用于高效稀疏矩阵实现,因而有望在处理长序列时,使反向传递的成本相对于前向传递的成本变得微乎其微。

关键词

引用

@article{arxiv.2505.15080,
  title  = {SUS backprop: linear backpropagation algorithm for long inputs in transformers},
  author = {Sergey Pankov and Georges Harik},
  journal= {arXiv preprint arXiv:2505.15080},
  year   = {2025}
}

备注

21 pages, 9 figures; main results unchanged, Fig.5 updated, some text rearranged