多层Transformer的梯度可在近线性时间内近似
机器学习
2024-10-16 v2 人工智能
计算与语言
摘要
流行Transformer架构中自注意力机制的计算复杂度给训练和推理带来了重大挑战,并成为长输入的性能瓶颈。是否有可能显著降低多层Transformer模型中梯度计算的二次时间复杂度?本文证明了一种新颖的快速近似方法可以在几乎线性的时间内计算梯度,其中是输入序列长度,同时在整个模型上保持多项式小的近似误差。我们的理论适用于一般损失函数,并且当多层Transformer模型包含许多实际子模块时,如残差连接、因果掩码和多头注意力。通过提高梯度计算的效率,我们希望这项工作将基于我们的理论结果促进长上下文语言模型更有效的训练和部署。
引用
@article{arxiv.2408.13233,
title = {Multi-Layer Transformers Gradient Can be Approximated in Almost Linear Time},
author = {Yingyu Liang and Zhizhou Sha and Zhenmei Shi and Zhao Song and Yufa Zhou},
journal= {arXiv preprint arXiv:2408.13233},
year = {2024}
}