FlashAttention
一、Introduction
为了加快LLM的训练和推理速度,针对transformer注意力机制的特点和GPU等硬件结构,通过分块和重计算来减少HBM读写次数,进而加快注意力计算的一种优化算法。
二、 Background

2.1 硬件特性
GPU的内存结构如上图左侧所示:
SRAM是GPU的片上内存,GPU计算时必须先把数据搬运到SRAM上才能进行计算。内存最小,但是IO速度最快
HBM是高带宽内存,存放模型参数权重、KVcache以及其它计算中间变量。内存较大,IO速度较快。(是否正确)
DRAM是主存,cpu进程存放数据的地方。内存大,IO速度慢。
模型运行时,GPUs拥有海量线程来执行一个操作/函数(kernel),每个kernel将输入从HBM加载到寄存器和SRAM上,然后计算,最后将输出写回HBM。
2.2 操作特性
依据计算和内存访问的平衡,操作可以分为计算密集型操作和内存密集型操作
计算密集型:操作所需时间主要取决于算术运算次数,访问内存只占很少的时间。典型示例为维度很大的矩阵乘法、通道数很多的卷积。
内存密集型:操作所需时间主要取决于内存访问次数,计算花费的时间很少。典型示例为逐元素操作(激活函数,dropout)、reduce(求和、Softmax、Batch Normalization、layer Normalization)。
2.3 标准Attention实现(简化的核心部分)
输入Q,K,V∈RN×d
先计算注意力分数矩阵S(每个token的Q与所有token的K进行计算)
S=QKT∈RN×N
再计算注意力权重矩阵P(对S按行进行softmax)
P=softmax(S)∈RN×N
最后计算输出矩阵O
O=PV∈RN×d
标准的Attention实现需要将S和P存储到HBM中,占用O(N2)的内存
部分操作例如softmax操作是内存密集型操作,大量内存访问导致较慢的运行时间。
还有其它的逐元素操作如掩码(按行)和P的丢弃也加剧了这一现象。
三、 FlashAttention算法的设计与分析
3.1 设计目标
给定 HBM 中的输入Q,K,V∈RN×d,计算注意力输出O∈RN×d 并将其写入 HBM。目标是减少 HBM 访问量,使访问量少于Θ(N2)。
3.2 分块
使用了分块(tiling)和重计算(recomputation)两种方法来实现。
分块计算的详细流程:

- 如图1右侧所示,先从行维度将Q分成Tr个块,将K和V都分成Tc个块。
关于块大小的选取,原则是使SRAM能同时放下Qi、Kj、Vj、Oi、ℓi、mi的前提下,让分块的数量尽可能少。按上述分块大小,所需的总空间为(2Bc×d+2Br×d+2Br)。上述理论推导时忽略了ℓi、mi等低阶空间复杂度变量,实际工程时会采用保守缩小块大小的策略。
- 遍历Q的分块,从HBM取一个块Qi
- 遍历K、V的分块,从HBM取[Kj,Vj],注意二者分块的行序列要一致。
- 在SRAM上计算注意力分数矩阵Sij=QiKjT∈RBr×Bc
- 在SRAM上计算Qi中选中的行的每行最大值,注意力矩阵逐元素取e指数P~ij=exp(Sij−m~ij)∈RBr×Bc,进行softmax操作的分母ℓ~ij=rowsum(P~ij)∈RBr 。
- 更新Qi中选中行的每行行最大值mi,每行各元素的e指数之和ℓi。
- 增量式计算更新Qi中选中的行的注意力输出Oi←diag(ℓinew)−1(diag(ℓi)emi−minewOi+em~ij−minewP~ijVj),写入HBM。
- 将更新后的mi和ℓi写回HBM。
与标准Attention访问次数对比分析:
二者Q,K,V读取次数相同,区别在于FlashAttention不用存S和P,为什么传统attention需要存呢?——因为一般N很大,并且N>>d,SRAM存Nd空间没问题,存不下S和P这种N2空间的。但是每次多了ℓi、mi的读写,总共需要4×Br×Tr×Tc=4NTc次读写,同时多了2×Tr×(Tc−1)×Br×d=2Nd(Tc−1)次读写,复杂度为Θ(N2d2/M),而标准的Attention的HBM读写次数复杂度为Θ(N2)。
对于典型的d(64-128)和M(约100KB)值,d2远小于M,因此FlashAttention所需的HBM访问次数比标准Attention实现的少很多倍。
3.3 重计算
通过分块,无需存储S和P,但是训练时反向传播需要使用S和P来计算梯度。解决办法是重新计算。
通过存储输出O和 softmax 归一化统计量(m,ℓ),我们可以在SRAM中轻松地重新计算注意力矩阵S和P。
相比之下,即使FlashAttention有更多的 FLOPs,重新计算通过减少 HBM 访问来加速反向传播。
3.4 扩展:Block-Sparse FlashAttention
用一个掩码矩阵M∈{0,1}N/Br×N/Bc来表示块与块之间是否需要计算,