首页
学习
活动
专区
圈层
工具
发布
社区首页 >专栏 >从Online Softmax 到 FlashAttention

从Online Softmax 到 FlashAttention

作者头像
AI老马
发布2026-01-13 15:11:27
发布2026-01-13 15:11:27
8810
举报
文章被收录于专栏:AI前沿技术AI前沿技术

本文通过梳理online softmax公式的推导过程,逐步的理清 FlashAttention的优化思路,以及结合分块策略是怎样进行高效实现的。主要内容为:

1)矩阵分块的进一步思考,让attention计算也一步到位。 2)通过引入代理值使得softmax从3-pass算法到2-pass算法的实现 3)flash attention 计算如何进化到1-pass

1,矩阵分块的进一步思考

矩阵分块考虑内存限制的同时,充分利用了显卡并行计算的能力,极大提高了大矩阵计算的效率。

两个分块矩阵A和B相乘得到结果C。中间需要从左到右的遍历对应的A,从上到下的遍历B,此时的中间结果暂存在内存中,直到得到完整的计算结果C,将结果C写回到HBM中。以上过程的关键是:一次读取和一次写入,中间结果不必缓存,进而减少了和HBM内存的交互次数。且多个分块并行运算,最后将结果累加后得到整个矩阵运算的结果,“高,实在是高!”。

那对于attention的计算,是否也可以一步到位的计算,从输入X直接得到最终的输出O?

答案是可行的。对于矩阵的乘法利用分块矩阵可以高效的完成计算,对于attention 依然可以利用分块一次性计算。

图(一),分块flash attention示意图。

对矩阵Q、K、V和O进行分块,然后一次载入一个对应的分块,一站式的完成attention的操作后,直接输出对应分块的最终attention结果。然后写入HBM内存中,这样就极大减少了中间结果的保存和读取,节省了大量IO访问时间。

FlashAttention 要做到一步到位计算注意力,主要的障碍在于softmax。

2,self-attention 矩阵分块的阻碍

2.1,如何一步到位计算注意力

标准计算attenton 的方式是通过以下的公式逐步计算。

其中X为前softmax逻辑矩阵,A为注意力分数,O为最终输出。全局来看,标准的注意力需要存储中间的计算结果,与HBM多次交互进行读写,效率不高。

优化的目标是利用矩阵分块一次性的从输入直接得到输出,一步到位,即公式:

其中最大的阻碍就是在计算softmax时候,需要依赖所有的元素最大值和所有元素的exp和。

公式中 保证了数值不会溢出。其中的 , 为N个元素中的最大值,需要遍历一次所有的元素才可以得到,即传统的softmax 计算需要多次的遍历所有元素。

2.2,softmax算法从3-pass 到 2-pass 的优化

传统的softmax算法为3趟算法(3-pass)即需要遍历三次。从三趟到两趟算法之间,仅需要一个代理!

初始化: 第1个元素到第 i 个元素中的最大值。

第1个元素到第 i 个元素的exp指数和,每个元素需要减去所有元素的最大值。

为最终softmax值。

图(二),softmax从3-pass到2-pass公式变化。

从3-pass 到2-pass 的关键点是,在计算所有元素的指数和时,引入了中间代理值 ,代理值不再需要遍历全部元素得到最大值,只需要截止到当前时刻的最大值即可。引入代理值有两个优势:

  • • 代理值的当前元素和前一元素存在递推关系,具体如公式 (8‘)所示,详细证明过程见参考论文。
  • • 当i=N即,最后一个元素时,原始的指数和,和代理值是相同的。

通过引入代理值,可以将3-pass中的第2趟求指数和合并到第一趟,即求最大值和指数和一块进行。具体如上图中 Algorithm 2-pass 所示。

3,1-pass FlashAttention

通过引入代理值将softmax的3-pass减少到了2-pass,对于softmax已经无法再进行压缩到1-pass。但是对于attention似乎还可以使用代理值技巧,重复昨天的故事,登上“一趟算法”的快船。

  • • 先定义变量,

把softmax的2-pass算法带入到attentin计算中。 其中 表示Q、O、V矩阵的第k行。矩阵的每一行操作都是相同,可以并行计算,故以其中的第k行进行介绍。同理 转置矩阵K的第i行。对于输出结果中的第i个元素

  • • 2-pass 的flashattention 算法。

在第一趟时,已经求的所有元素的最大值和相应的指数和

在第二趟时,先计算第i个注意力分值 (标量值),再与V矩阵的第i行相乘得到一个向量,与前一个历史向量相加后即得到输出。

遍历结束后最终的输出 即为输出矩阵的第K行结果。

图(三),flashattention从2-pass到1-pass公式推导

如何将attention 的2-pass 算法压缩的1-pass?如法炮制,同样对 引入中间代理值 , 两者之间也存在一定关系:

  • • 代理值的当前元素和前一元素存在递推关系,具体如公式 (12‘)所示,详细证明过程见参考论文。
  • • 当i=N时原始值和代理值是相同的。

此时计算输出值仅仅依赖于标量值 和输入元素以及矩阵V元素。可以将其放在一次for循环中完成。如上图算法所示。

4, flashattention 如何分片

利用矩阵分块提高计算的效率。以下算法对于Q矩阵还是单行,即第K行,对于K和V矩阵进行分块。

假设:分块的大小为b,总共有 #tiles块,

此时的 为一个向量,维度大小为[1,b]

图(四),分片后的flashattention算法

相比于1-pass 的flashattention算法,需要注意 表示在序列长度L方向上,按分块大小b,一步一步的滑动。同理对于V矩阵 在外层也是在L方向上滑动,但内存还需要在分块对每个行遍历。

另外一个区别是求局部最大值,先在分块b个元素内求的最大值 ,然后再与上一个历史最大值进行比较。

在以上的算法中关键的一点是:整个流程中仅和分块的大小 b 和单头维度d有关,和输入序列长度L完全的解藕了。

总结

flashattention1算法同时在Q和KV 方向进行了分块计算attention,基本的思路是一样的。再进一步,flashattention2做了些工程上的优化,效率进一步提升主要集中在三点。

  • • 减少大量非matmul的冗余计算,增加Tensor Cores运算比例
  • • 在原先的batch和heand并行的基础上,增加seqlen维度的并行,充分的利用多个 SM 并行计算,主要解决单batch长序列的问题。
  • • Warp Partitioning策略优化。主要解决线程块内部并行计算和共享问题。

更多精彩:

历史文章:

显卡知识-算力开挂的GPU

大模型量化-roofline性能分析工具

大模型推理-Flash attention 访问内存优化

大模型推理-page attention 内存分页术

大模型推理-极致化的批处理策略介绍

大模型推理- PD分离部署,势在必行!

大模型推理-高效推理必备KV cache

大模型训练-混合专家系统MoE

大模型训练-Nvidia GPU 互联技术全景图

大模型训练-流水线并行PP

大模型训练-张量并行TP

本文参与 腾讯云自媒体同步曝光计划,分享自微信公众号。
原始发表:2025-07-05,如有侵权请联系 cloudcommunity@tencent.com 删除
目录
  • 1,矩阵分块的进一步思考
  • 2,self-attention 矩阵分块的阻碍
    • 2.1,如何一步到位计算注意力
    • 2.2,softmax算法从3-pass 到 2-pass 的优化
  • 3,1-pass FlashAttention
    • 4, flashattention 如何分片
    • 总结
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档