中文

注意力机制的I/O复杂度,或Flash Attention的最优性

机器学习 2024-02-13 v1 计算复杂性 数据结构与算法 信息论 math.IT

摘要

自注意力机制是流行的Transformer架构的核心,但其存在二次时间和内存复杂度。突破性的Flash Attention算法揭示了I/O复杂度是扩展Transformer的真正瓶颈。给定两个层级的内存层次结构,即快速缓存(例如GPU片上SRAM)和慢速内存(例如GPU高带宽内存),I/O复杂度衡量的是对内存的访问次数。Flash Attention使用N2d2M\frac{N^2d^2}{M}次I/O操作计算注意力,其中NN是注意力矩阵的维度,dd是头维度,MM是缓存大小。然而,这个I/O复杂度是最优的吗?已知的下界仅排除了当M=Θ(Nd)M=\Theta(Nd)时的o(Nd)o(Nd) I/O复杂度,因为需要写入慢速内存的输出是Ω(Nd)\Omega(Nd)。这引出了我们工作的主要问题:对于所有MM值,Flash Attention的I/O复杂度是否最优?我们通过展示一个I/O复杂度下界来完全一般性地解决上述问题,该下界在任意常数因子内,对于任何Md2M \geq d^2,都与Flash Attention提供的上界匹配。此外,对于M<d2M < d^2,我们给出了一种具有更低I/O复杂度的更好算法,并证明它也是最优的。而且,我们的下界不依赖于使用组合矩阵乘法来计算注意力矩阵。我们证明,即使使用快速矩阵乘法,上述I/O复杂度界限也无法改进。我们通过引入一种新的用于矩阵压缩的通信复杂度协议,并将通信复杂度与I/O复杂度联系起来来实现这一点。据我们所知,这是第一个建立通信复杂度和I/O复杂度之间关系的工作,我们相信这种联系可能具有独立的意义,并将在未来证明I/O复杂度下界方面找到更多应用。

关键词

引用

@article{arxiv.2402.07443,
  title  = {The I/O Complexity of Attention, or How Optimal is Flash Attention?},
  author = {Barna Saha and Christopher Ye},
  journal= {arXiv preprint arXiv:2402.07443},
  year   = {2024}
}

备注

24 pages, 3 figures