中文

自注意力不需要 $O(n^2)$ 内存

机器学习 2022-10-11 v3

摘要

我们提出了一种非常简单的注意力算法,其相对于序列长度所需内存为 O(1)O(1),并给出了自注意力的一种扩展,其所需内存为 O(logn)O(\log n)。这与经常提到的自注意力需要 O(n2)O(n^2) 内存的观点相反。尽管时间复杂度仍为 O(n2)O(n^2),但在现代加速器上,设备内存而非计算能力往往是限制因素。因此,降低注意力的内存需求使得处理比原本可行更长的序列成为可能。我们提供了一种适用于加速器的实用实现,其需要 O(n)O(\sqrt{n}) 内存,数值稳定,且运行时间在标准注意力实现的几个百分点之内。我们还展示了如何在保持内存高效的同时对函数求导。对于序列长度 16384,自注意力的内存开销在推理时降低 59 倍,在求导时降低 32 倍。

关键词

引用

@article{arxiv.2112.05682,
  title  = {Self-attention Does Not Need $O(n^2)$ Memory},
  author = {Markus N. Rabe and Charles Staats},
  journal= {arXiv preprint arXiv:2112.05682},
  year   = {2022}
}