注意力机制的I/O复杂度,或Flash Attention的最优性
机器学习
2024-02-13 v1 计算复杂性
数据结构与算法
信息论
math.IT
摘要
自注意力机制是流行的Transformer架构的核心,但其存在二次时间和内存复杂度。突破性的Flash Attention算法揭示了I/O复杂度是扩展Transformer的真正瓶颈。给定两个层级的内存层次结构,即快速缓存(例如GPU片上SRAM)和慢速内存(例如GPU高带宽内存),I/O复杂度衡量的是对内存的访问次数。Flash Attention使用次I/O操作计算注意力,其中是注意力矩阵的维度,是头维度,是缓存大小。然而,这个I/O复杂度是最优的吗?已知的下界仅排除了当时的 I/O复杂度,因为需要写入慢速内存的输出是。这引出了我们工作的主要问题:对于所有值,Flash Attention的I/O复杂度是否最优?我们通过展示一个I/O复杂度下界来完全一般性地解决上述问题,该下界在任意常数因子内,对于任何,都与Flash Attention提供的上界匹配。此外,对于,我们给出了一种具有更低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