FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning

发表时间: 2023-07 · arXiv:2307.08691

原文: https://arxiv.org/abs/2307.08691

文章标题: FlashAttention-2:通过更好的并行性和工作分区实现更快的注意力机制
作者: Tri Dao
机构: 1. 普林斯顿大学计算机科学系 2. 斯坦福大学计算机科学系

速读

一句话结论 FlashAttention-2 通过重构注意力机制在 GPU 上的并行策略与工作分区,大幅减少了非矩阵乘法运算和共享内存的读写,在 A100 GPU 上实现了比初代 FlashAttention 快 2 倍的加速,并将前向传播计算效率提升至理论峰值的 73%。

要解决什么问题 Transformer 处理长序列时,注意力层的运行时间和内存占用随序列长度呈二次方增长。初代 FlashAttention 通过分块计算和在线 softmax 技术,将内存占用降至线性并实现 2 到 4 倍加速,但存在一个致命卡点:计算效率极低,仅能达到 GPU 理论峰值算力的 25% 到 40%。导致该卡点的机制有两方面:第一,GPU 上不同线程块和线程束(warp)之间的工作分区不合理,导致在处理长序列且批次较小时 GPU 占用率低下,并产生大量不必要的共享内存(SRAM)读写与同步开销;第二,算法包含过多非矩阵乘法浮点运算。现代 GPU 配备的 Tensor Cores 使矩阵乘法吞吐量达非矩阵乘法的 16 倍,过多的非矩阵乘法严重拖慢了整体执行速度。

怎么做的 FlashAttention-2 的核心思路是重构算法逻辑、全局并行维度与线程块内部的任务分配,最大化矩阵乘法单元利用率并消除多余内存通信。该方法由三个关键设计构成。首先是算法微调以减少非矩阵乘法运算。在初代在线 softmax 中,每次更新输出块需对历史和当前结果同时缩放。FlashAttention-2 改为在内层循环中仅维护“未缩放”的输出版本,并只保存局部最大值与指数和的对数(logsumexp),直到循环结束才进行最终缩放。其前向传播更新输出的核心公式为:$O_i^{(j)} = \text{diag}(e^{m_i^{(j-1)} - m_i^{(j)}})^{-1} O_i^{(j-1)} + \tilde{P}_{ij}^{(j)} V_j$。这使得计算时间更多让渡给高效的矩阵乘法。其次是增加序列长度维度的并行性。初代仅在批次大小和头数量上分配线程块,处理长序列(此时批次较小)时无法填满 GPU。FlashAttention-2 在前向传播中将外层循环(遍历序列长度的行块)调度到不同线程块上无通信并行;反向传播则为每个列块调度一个线程块,通过原子加法更新梯度,显著提升了长序列下的 GPU 占用率。最后是优化线程块内 warp 间的工作分区。初代采用“split-K”方案,将键 K 和值 V 划分给不同 warp,查询 Q 共享,导致 warp 必须将中间结果写入共享内存进行同步累加。FlashAttention-2 将 Q 划分给不同 warp,而 K 和 V 对所有 warp 可见。每个 warp 计算完负责的 $QK^\top$ 局部矩阵后,直接与共享的 V 相乘得到输出,彻底消除了 warp 间的通信与共享内存读写。

效果如何 实验在单张及 8 张 A100 80GB SXM4 GPU 上搭建,端到端测试训练了 1.3B 和 2.7B 参数的 GPT-style 模型,序列长度覆盖 2k 到 16k。对比基线点名了四种路线:代表标准实现的 PyTorch 原生注意力、代表 IO 感知优化的原始 FlashAttention、代表编译器自动优化的 Triton 版 FlashAttention,以及代表底层算子库优化的 xformers 库 FlashAttention(cutlass 版本)。量化结果显示,在注意力基准测试中,FlashAttention-2 速度是原始 FlashAttention 和 xformers 版的 2 倍,是 Triton 版的 1.3 到 1.5 倍,最高达 PyTorch 标准实现的 10 倍。在 A100 上,前向传播速度最高达 230 TFLOPs/s,触及硬件理论峰值的 73%。端到端训练中,相比不使用 FlashAttention 的基线提速 2.8 倍,单卡吞吐量达 225 TFLOPs/s,模型 FLOPs 利用率达 72%。该方法的局限在于工程代价:为平衡寄存器溢出和共享内存容量,必须针对不同注意力头维度手动调整分块大小(如 64x64 或 128x128),尚未实现自动调优。

A1 主要贡献

本文的核心问题是解决 Transformer 模型在处理长序列时的性能瓶颈。注意力层的运行时间和内存占用随序列长度呈二次方增长,这限制了模型处理长文档、高分辨率图像或长视频的能力。尽管 FlashAttention 通过优化 GPU 内存使用,将内存占用从二次方降低到线性,并实现了2-4倍的加速,但其计算效率(FLOPs/s)仍远低于优化后的矩阵乘法(GEMM)操作,仅达到理论峰值的25-40%。

本文通过分析发现,FlashAttention 效率不高的主要原因是 GPU 上不同线程块和 warp 之间的工作分区不理想,导致了低占用率或不必要的共享内存读写。

为解决上述问题,本文提出了 FlashAttention-2,其研究目标是通过改进并行策略和工作分区来进一步提升注意力计算的效率。主要创新点和贡献如下:

  1. 算法调整以减少非矩阵乘法浮点运算:对 FlashAttention 算法进行了微调,减少了非矩阵乘法(non-matmul)的浮点运算次数。由于 GPU 上的专用计算单元(如 Tensor Cores)使得矩阵乘法的吞吐量远高于非矩阵乘法(可达16倍),这一优化能让计算时间更多地用于高效的矩阵乘法上。
  2. 增强并行性以提升GPU占用率:除了在批次大小和头数量维度上进行并行化,FlashAttention-2 还增加了沿序列长度维度的并行化。这在处理长序列(此时批次大小通常较小)的场景下,能显著提高 GPU 资源的利用率(即占用率)。
  3. 优化线程块内部工作分区:在每个线程块(thread block)内部,重新设计了不同 warp 之间的工作分配方式,以减少它们之间通过共享内存进行的通信和数据读写。

这些改进使得 FlashAttention-2 相比于 FlashAttention 实现了约2倍的速度提升,在 A100 GPU 上的前向传播计算效率达到了理论峰值的50-73%。在端到端的 GPT-style 模型训练中,每块 A100 GPU 的训练速度高达 225 TFLOPs/s,模型 FLOPs 利用率达到72%。

A3 背景知识

2.1 硬件特性

2.2 标准注意力实现

2.3 FlashAttention

A2 方法细节

我们描述了 FlashAttention-2 算法,它包含了对 FlashAttention 的几处调整以减少非矩阵乘法 FLOPs。然后,我们描述了如何在不同的线程块上并行化计算以充分利用 GPU 资源。最后,我们描述了在一个线程块内如何在不同的 warp 之间划分工作以减少共享内存的访问量。这些改进带来了2-3倍的加速,这在第4节中得到了验证。

3.1 算法

3.1.1 前向传播

3.1.2 反向传播

3.2 并行性

3.3 Warp之间的工作分区

A4 实验环境

A4 实验结果

注意力基准测试

端到端性能

A5 结论

FlashAttention-2 比 FlashAttention 快2倍,这意味着现在训练一个16k上下文长度的模型的成本与之前训练一个8k上下文长度的模型相当。这一进步有望推动模型在理解长篇书籍报告、高分辨率图像、音频和视频等领域的应用。同时,FlashAttention-2 也将加速现有模型的训练、微调和推理过程。

未来工作展望
1. 扩展硬件和数据类型支持:计划与研究人员和工程师合作,将 FlashAttention 推广到不同类型的设备(如 H100 GPU、AMD GPU)和新的数据类型(如 FP8)。
2. 针对 H100 的深度优化:下一步计划是优化 FlashAttention-2 以利用 H100 GPU 的新硬件特性(如 TMA、第四代 Tensor Cores、FP8)。
3. 结合高级算法:将 FlashAttention-2 的底层优化与高级算法(如局部注意力、扩张注意力、块稀疏注意力)相结合,可能使我们能够训练上下文更长的 AI 模型。
4. 提升可编程性:与编译器研究者合作,使这些优化技术更易于编程实现。

方法细节中的引用汇总