
本文通过梳理online softmax公式的推导过程,逐步的理清 FlashAttention的优化思路,以及结合分块策略是怎样进行高效实现的。主要内容为:
1)矩阵分块的进一步思考,让attention计算也一步到位。 2)通过引入代理值使得softmax从3-pass算法到2-pass算法的实现 3)flash attention 计算如何进化到1-pass
矩阵分块考虑内存限制的同时,充分利用了显卡并行计算的能力,极大提高了大矩阵计算的效率。
两个分块矩阵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。
标准计算attenton 的方式是通过以下的公式逐步计算。
其中X为前softmax逻辑矩阵,A为注意力分数,O为最终输出。全局来看,标准的注意力需要存储中间的计算结果,与HBM多次交互进行读写,效率不高。
优化的目标是利用矩阵分块一次性的从输入直接得到输出,一步到位,即公式:
其中最大的阻碍就是在计算softmax时候,需要依赖所有的元素最大值和所有元素的exp和。
公式中 保证了数值不会溢出。其中的 , 为N个元素中的最大值,需要遍历一次所有的元素才可以得到,即传统的softmax 计算需要多次的遍历所有元素。
传统的softmax算法为3趟算法(3-pass)即需要遍历三次。从三趟到两趟算法之间,仅需要一个代理!
初始化: 第1个元素到第 i 个元素中的最大值。
第1个元素到第 i 个元素的exp指数和,每个元素需要减去所有元素的最大值。
为最终softmax值。

图(二),softmax从3-pass到2-pass公式变化。
从3-pass 到2-pass 的关键点是,在计算所有元素的指数和时,引入了中间代理值 ,代理值不再需要遍历全部元素得到最大值,只需要截止到当前时刻的最大值即可。引入代理值有两个优势:
通过引入代理值,可以将3-pass中的第2趟求指数和合并到第一趟,即求最大值和指数和一块进行。具体如上图中 Algorithm 2-pass 所示。
通过引入代理值将softmax的3-pass减少到了2-pass,对于softmax已经无法再进行压缩到1-pass。但是对于attention似乎还可以使用代理值技巧,重复昨天的故事,登上“一趟算法”的快船。
把softmax的2-pass算法带入到attentin计算中。 其中 表示Q、O、V矩阵的第k行。矩阵的每一行操作都是相同,可以并行计算,故以其中的第k行进行介绍。同理 转置矩阵K的第i行。对于输出结果中的第i个元素
在第一趟时,已经求的所有元素的最大值和相应的指数和
在第二趟时,先计算第i个注意力分值 (标量值),再与V矩阵的第i行相乘得到一个向量,与前一个历史向量相加后即得到输出。
遍历结束后最终的输出 即为输出矩阵的第K行结果。

图(三),flashattention从2-pass到1-pass公式推导
如何将attention 的2-pass 算法压缩的1-pass?如法炮制,同样对 引入中间代理值 , 两者之间也存在一定关系:
此时计算输出值仅仅依赖于标量值 和输入元素以及矩阵V元素。可以将其放在一次for循环中完成。如上图算法所示。
利用矩阵分块提高计算的效率。以下算法对于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做了些工程上的优化,效率进一步提升主要集中在三点。
更多精彩:
历史文章: