首页
学习
活动
专区
圈层
工具
发布
社区首页 >专栏 >FlashAttention | 大语言模型加速的关键组件

FlashAttention | 大语言模型加速的关键组件

作者头像
OpenCV学堂
发布2026-07-24 21:13:32
发布2026-07-24 21:13:32
930
举报

标准注意力的缺陷

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

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

FlashAttention1

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

  • 重计算 —— 不保留前向传播过程中的中间结果供反向传播使用,而是在需要时重新计算它们。
  • 内核融合 —— 将所有子操作合并为单个内核,其中包括矩阵乘法、Softmax操作,以及与值矩阵(V矩阵)的最终乘法运算。
  • 分块处理 —— 查询矩阵(Q)、键矩阵(K)和值矩阵(V)的规模可能相当大。FlashAttention 通过将输入矩阵划分为多个块(Blocks)来管理内核融合,并将融合后的内核应用于这些输入块上,仅需在全局内存中存储少量常量数据。
  • 计算流程 —— 融合后的内核首先读取 Q 和 K 的数据块,并计算用于 Softmax 的 Logits(未归一化的分数)。

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 等)的关键技术之一.

总结

本文参与 腾讯云自媒体同步曝光计划,分享自微信公众号。
原始发表:2026-07-23,如有侵权请联系 cloudcommunity@tencent.com 删除

本文分享自 OpenCV学堂 微信公众号,前往查看

如有侵权,请联系 cloudcommunity@tencent.com 删除。

本文参与 腾讯云自媒体同步曝光计划  ,欢迎热爱写作的你一起参与!

评论
登录后参与评论
0 条评论
热度
最新
推荐阅读
领券
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档