
标准注意力的缺陷

大型语言模型(LLM)依赖于Transformer架构,其中注意力机制至关重要。由于其计算需求巨大,注意力机制的执行(即推理过程)通常在GPU上进行。虽然软件的可复用性很重要,但直接使用现成的GPU内核(如矩阵乘法和softmax函数)来实现注意力机制,往往会牺牲性能。

标准的注意力机制(Transformer 的核心)在处理长序列时,其计算和内存需求会随序列长度呈二次方增长,导致训练和推理速度慢、显存占用巨大。

FlashAttention1
使用现有的内核(Kernels)来实现标准注意力机制,会导致大量数据在GPU与访问速度较慢的全局内存之间往返传输(关于GPU架构的理解,可参阅我之前的文章)。FlashAttention 采用了几种优化技术来减少对全局内存的输入/输出(IO)操作。

FlashAttention2
FlashAttention 虽然已经通过减少对全局内存的访问来提升了内存效率,但它在GPU计算核心的利用率上仍未达到最优,特别是在 A100 这类GPU上,其理论浮点运算性能(FLOPS)的实际利用率仅为 25% 左右
FlashAttention V2 通过消除非矩阵乘法(MatMul)操作和优化线程束(Warps)上的工作划分,提升了 GPU 的利用率。
与 FlashAttention V1 中对每个数据块分别进行归一化不同,V2 版本将所有数据块的归一化操作统一推迟到计算的最后阶段再进行。也就是说,我们首先会尽可能充分地利用 GPU 计算核心,计算出未经归一化的输出结果。
此外,FlashAttention V2 还优化了每个线程块内部的工作划分,以减少线程束之间的通信开销,并将并行性扩展到了多个维度,例如头数、序列长度和批次大小。
FlashAttention V2 达到了理论峰值浮点运算性能(FLOPS)的 50%–70%。得益于 GPU 利用率的提升,FlashAttention V2 的速度比原始版 FlashAttention 快 2 倍。
FlashAttention3
FlashAttention V3 是对 FlashAttention V2 的进一步优化,其驱动力来自 Hopper GPU 架构的先进特性,例如对张量核心(Tensor Cores)低精度矩阵乘法的支持以及多内核的异步执行。FlashAttention V3 实现了 QK 矩阵乘法、Softmax 和 PV 乘法等操作的计算重叠,使它们能够同时进行,并高效地管理内存和计算资源。
在 H100 GPU 核心上,FlashAttention V3 的利用率达到了 75%,而 FlashAttention V2 仅为 35%。同时,它在 FP16 精度下将注意力计算速度提升了 2 倍,并在 FP8 精度下达到了 1.2 PFLOPS 的性能。
进一步利用了新一代 GPU(如 H100)的新特性(如 TMA、WGMMA 指令),进行了底层优化。
支持更长的上下文和更高的计算效率,是训练超长序列大语言模型(如 GPT-4 等)的关键技术之一.
总结
