讲解 FlashAttention 如何突破注意力的显存访问瓶颈:用 Tiling 分块把计算搬到 SRAM 中完成、反向传播以重计算换显存、通过 kernel 融合减少显存读写,并延伸到块稀疏 FlashAttention。
https://zhuanlan.zhihu.com/p/626079753
https://zhuanlan.zhihu.com/p/639228219
简介
当前Transformer是LLM中大量使用的基础模型构件。如图1,Transformer的核心组件是多头注意力,其计算复杂度和空间复杂度是序列长度的二次方,已经有许多近似注意力的方法尝试减少注意力的计算和内存开销。例如,稀疏近似和低秩近似的方法,将计算复杂度降低到了序列长度的线性或亚线性,但这些方法过于关注FLOPs(浮点数计算次数)的减少,而忽略了IO读写的内存访问开销。在现代GPU中,计算速度已经远超过了显存访问速度,Transformer中的大部分计算操作的瓶颈是显存访问。

表1展示了多头注意力的标准实现伪代码,针对序列长度N和多头注意力维度d,其公式为:

上面的公式按计算步骤可以写为下面三步:

常规实现中S、P等中间结果均需要写回到HBM内存中,这就导致了两个问题:
1)HBM内存开销大:除输入输出外这些中间结果也需占用HBM;

Flash Attention技术
针对上一节的分析结果,作者提出了如下措施来优化IO:
1)Tiling切片:利用更高速的SRAM代替HBM;


Tiling切片
如图2左边子图所示,观察GPU内存层级,SRAM的读写速度比HBM高一个数量级,但其容量要小很多。通过kernel融合的方式,将多个操作融合为一个操作,将数据放到高速的SRAM进行读写,可以代替HBM的角色,从而有效提高访存瓶颈。但SRAM的内存大小十分有限,无法一次性完成所有数据的完整注意力计算,因此必须进行分块计算,使得每块计算需要的内存不超过SRAM的大小。
分块计算的难点在于softmax的计算。由于计算softmax的归一化因子(分母)时,需要获取到完整的输入数据,进行分块计算的难度比较大。论文中也是重点对softmax的分块计算进行了阐述。如表2所示,第3行对Q、K、V进行了Tiling切片操作,第10行中,作者针对每块计算各自的softmax和相关的额外统计量m(x),l(x),然后在第12行中,使用这些额外统计量完成标准的softmax计算。
重计算
作者的优化思路之一是不存储注意力计算过程中的中间结果,即基于trade-off的思想,在IO瓶颈的性能优化中,将IO压力转移到计算中,以突破原有的性能水平。在标准注意力实现中,需要用到多个中间矩阵,但作者并没有保存这些矩阵。而是采用重计算的方式,仅保存了两个统计量,需要使用这些中间结果时再基于高速的SRAM上快速地重新计算,相比于标准注意力中从HBM中读取很大的中间注意力矩阵的方式,重计算尽管增加了额外的计算量FLOPs,但仍然能够减少运行时间。如表3所示,由于重计算导致FlashAttention中计算次数GFLOPs有一定程度增加,但HBM读写量大幅下降,最终的执行时间上来看性能收益明显。

Kernel 融合
对于性能受限于内存带宽的操作,进行加速的常用方式就是kernel融合。kernel融合的基本思路是将多个操作融合成一个操作,减少读写HBM的次数。Tiling分块计算使得我们可以用一个kernel来完成注意力的所有操作。从HBM中加载输入数据,在SRAM中执行所有的计算操作(矩阵乘法,mask,softmax,dropout,矩阵乘法等计算步骤),再将计算结果写回到HBM中。如表2,只在第2行对HBM进行读操作,在第12、13行对HBM进行写操作,除此之外没有额外的HBM访问。如图3所示,巨量的HBM读写量下降中Kernel融合的贡献是最大的。
Block-Sparse Flash Attention
作者将Flash Attention扩展到近似注意力,提出了块稀疏的Flash Attention,其IO复杂度比Flash Attention小,其系数与稀疏度成正比。作者通过调整表2中的FlashAttention算法,大部分步骤不变,只计算注意力矩阵的非零块,跳过零块。应用块状稀疏性可以通过稀疏性直接改善IO复杂性,证明能够近似于任意稀疏。如图3(左)所示,作者测试了块大小的影响,右图的实验结果则展示了,作者验证了随着稀疏度的增加,块稀疏FlashAttention的运行时间成比例地提高。


