FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
复制标题

DOI:
--
复制
发表时间:
2022-05
期刊:
ArXiv
影响因子:
--
通讯作者:
Tri Dao;Daniel Y. Fu;Stefano Ermon;A. Rudra;Christopher R'e
Tri Dao;Daniel Y. Fu;Stefano Ermon;A. Rudra;Christopher R'e
中科院分区:
其他
文献类型:
--
作者:
Tri Dao;Daniel Y. Fu;Stefano Ermon;A. Rudra;Christopher R'e

文献摘要

相似文献

Transformers在长序列上速度慢,内存消耗大,因为自注意的时间和内存复杂度是序列长度的二次方。近似注意力方法试图通过权衡模型质量来解决这个问题,以降低计算复杂度,但通常不能实现挂钟加速。我们认为,一个缺失的原则是使注意力算法IO感知-占GPU内存级别之间的读取和写入。我们提出了FlashAttention,一种IO感知的精确注意算法,该算法使用平铺来减少GPU高带宽内存(HBM)和GPU片上SRAM之间的内存读/写次数。我们分析了FlashAttention的IO复杂性,表明它需要比标准注意更少的HBM访问,并且对于一系列SRAM大小是最佳的。我们还将FlashAttention扩展到块稀疏注意力,产生一个近似注意力算法,比任何现有的近似注意力方法都快。FlashAttention训练Transformers的速度比现有的基准更快:BERT-large上的端到端挂钟加速15%(以下简称为BERT-large)。长度512)与MLPerf 1.1训练速度记录相比,GPT-2(序列号:长度为1 K),以及在远程竞技场上的2.4$\times$加速(以下长度1 K-4K)。FlashAttention和块稀疏FlashAttention使Transformers中的上下文更长,产生更高质量的模型(GPT-2上的困惑度提高了0.7,长文档分类上的提升点提高了6.4)和全新的功能:第一个在Path-X挑战中实现优于机会的性能的Transformers。长度16 K,61.4%准确度)和Path-256(seq.长度64 K,准确率63.1%)。
Transformers are slow and memory-hungry on long sequences, since the time and memory complexity of self-attention are quadratic in sequence length. Approximate attention methods have attempted to address this problem by trading off model quality to reduce the compute complexity, but often do not achieve wall-clock speedup. We argue that a missing principle is making attention algorithms IO-aware -- accounting for reads and writes between levels of GPU memory. We propose FlashAttention, an IO-aware exact attention algorithm that uses tiling to reduce the number of memory reads/writes between GPU high bandwidth memory (HBM) and GPU on-chip SRAM. We analyze the IO complexity of FlashAttention, showing that it requires fewer HBM accesses than standard attention, and is optimal for a range of SRAM sizes. We also extend FlashAttention to block-sparse attention, yielding an approximate attention algorithm that is faster than any existing approximate attention method. FlashAttention trains Transformers faster than existing baselines: 15% end-to-end wall-clock speedup on BERT-large (seq. length 512) compared to the MLPerf 1.1 training speed record, 3$\times$ speedup on GPT-2 (seq. length 1K), and 2.4$\times$ speedup on long-range arena (seq. length 1K-4K). FlashAttention and block-sparse FlashAttention enable longer context in Transformers, yielding higher quality models (0.7 better perplexity on GPT-2 and 6.4 points of lift on long-document classification) and entirely new capabilities: the first Transformers to achieve better-than-chance performance on the Path-X challenge (seq. length 16K, 61.4% accuracy) and Path-256 (seq. length 64K, 63.1% accuracy).