SUS backprop:面向长序列输入的变换器线性反向传播算法
机器学习
2025-06-06 v2 人工智能
计算与语言
摘要
对计算图中任意位置进行随机切断以设计无偏梯度估计器是直接的。在切断对计算影响不大的部分时,可为最小的随机梯度方差增加换来显著的反向传播计算节省,这在某些情况下确实如此。这种情况发生在变换器架构的注意力机制中。对于长序列,注意力成为限制因素,因为其计算需求随序列长度 呈二次增长。与此同时,大多数注意力权重变得很小,因为大多数注意力头倾向于将给定 token 只与序列中少数几个其他 token 相连接。这些权重成为切断反向传播的有力目标。我们提出一种由单个参数 控制的简单概率规则,切断大多数注意力权重的反向传播,使每个 token 每个注意力头至多保留 个相互作用。这将注意力反向传播所需计算量降低至 倍,使其从二次复杂度 降为线性复杂度 。我们通过实验验证,对典型变换器模型,切断约 的注意力梯度流(即取 )仅导致 时的梯度方差增加约 ,且随 增大而降低。这一方法可用于高效稀疏矩阵实现,因而有望在处理长序列时,使反向传递的成本相对于前向传递的成本变得微乎其微。
引用
@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