在介绍FlashAttention之前,先介绍 LLM 推理的两个阶段,一个是 prefill 阶段,一个是 decode 阶段。

假设输入为[Batch_size=1,句子长度len,模型隐藏维度d],数据类型为FP16

标准 Attention 在做什么

首先,计算计算

S=QKT

将结果写入 HBM ,然后应用一个 softmax

P=Softmax(S)

从 HBM 中读取 S ,再将结果 P 写入 HBM

O=PV

然后从 HBM 读取 P ,再将 O 写入 P。整个过程中读入了[Q,K,S,P,V] 写入了 [S,P,O]。每次矩阵乘法是一次加乘,即两次浮点运算,一共执行了 4len2d 次,一共搬运了 4len2+4lend 的数据,即执行 8len2+8lend 次。

prefill与decode阶段的区别

在 prefill 阶段,我们会把整段 prompt 交给 LLM ,此时模型得到 KV cache,输入是一个 nd 维度的向量, 显然,对于每一层都一样,我们通过 Roofline 模型计算计算强度

=(FLOPS)访(Byte)

以A100为例,显存宽带为 2 TB/s ,浮点性能为 312 TFLOPS ,当模型计强度超过 3122=156 的时候,性能受限于浮点性能,当低于这个数的时候,性能受限于宽带,这个值也叫做硬件平衡点。

在 prefill 阶段,每一层都输入一个完整的句子,可以得到计算强度为

=4len2d8len2+8lend

以往的博客告诉我们, prefill 是 computed-bound 的,而 decode 是memory-bound 的,然后会告诉我们下面这些东西。

len 较大的时候,计算强度 = d2 ,此时远超硬件平衡点,故性能受限于浮点性能。

而在 decode 阶段,输入长度变为 1 ,计算强度变成 12 ,受到宽带约束。

实际的具体情况还是要具体分析。

假设 d=4096 ,代入几个值,仅以理论情况分析。

N=256 ,计算强度约为 120 ,仍为 memory bound。

N=512 ,计算强度约为 227 ,成为 computed bound。(在 B200 上面,仍然为 memory bound)

N=1024 ,计算强度为 409,在 B200 上面,仍然为 memory bound

如果我们使用了 chunked_prefill 技术,则也需要考虑这上面的带宽问题。

FlashAttention 问题的背景

看它解决了什么问题,就得以当时的视角看问题是怎么样的。

如果 attention 是一个标注的矩阵乘法,则上面的问题可以忽视

实际上,它还有 softmax ,normalization 以及 scaling 等,同时期还有其他的attention计算方式,例如 sparse attention,linear attention,kernel-based attention 等。它们尝试将算法的时间复杂度降低,但是实际运行起来也没比普通 attention 快多少。

于是尝试在带宽上优化attention。

显然,以上访问出现了中间结果S与P,它被计算并存储,然后重新调用,如果我们一直放在SRAM里面,就能减小访存次数。

Tiling

如果我们想让中间结果一直在SRAM上面,那么我们就需要SRAM一直去存储它们,假设这些数据类型都一样,可以存储 M 个数字。

回到上面的三次计算,我们看看每个时间段需要哪些变量参与

步骤 Q K V S P O
1 1 1 1
2 1 1
3 1 1 1

每个时刻至少需要3份,与此同时,QK其中一个可以一直放在内存块里面,所以简单来看我们可以分成四份。

假设我们取 Qi 的部分大小为 Brd , Kj,Vj 的大小 Bcd ,那么中间变量 S,P 的大小即为 BrBc ,结果 V 大小为 Brd

为了确保能分成四份,首先令 Br=M4d ,这样第一个块就满足了。然后看第二个块 Bc=M4d 也可以了,但是第三个块不一定了,于是选择将 Br 加上一个 Br=min(Br,d) 限制,这样每个块都小于等于 M4d 不会超过限制。

这就是 FlashAttention 分块的思路,将 SRAM 分成四份,然后中间变量不写入 HBM。

此时继续计算,会遇到几个问题:

  • softmax公式为 softmax=exijexj ,在计算attention的时候,需要一整行的元素,但是我们已经将矩阵分块了,S得到的仅仅是一部分结果。
  • 如果仅做推理,则不需要用到 P这个中间结果,但在训练及反向传播的时候,还是需要这些中间变量。

下面我们来解决这些问题

online Softmax

普通的softmax公式为

softmax=exijexj

在 attention 中,这个是逐行计算的,为了防止过大数造成精度损失或者其他情况,根据平移不变性,我们要减去这一行最大值 m(x)

softmax=exim(x)jexjm(x)

问题就出现在这里,我们无法一下取得这一行的最大值。

假设我们这一行的attn被分成两块计算。

x=[x(1),x(2)]

并且已经得到两块的最大值 m(x(1)),m(x(2)) 以及各自分别求的 softmax 值 l(1),l(2)

先合并最大值

m(x)=max(m(x(1)),m(x(2)))

然后合并指数和

l(x)=em(x(1))m(x)l(x(1))+em(x(2))m(x)l(x(2))

这样,在计算 softmax 的时候,公式为

softmax=exim(x)l(x)

现在完成了理论基础,这里还没完,对于前面已经产生的输出 O 还需要重新计算他的贡献。

假设在之前计算中,我们得到了这一行每个块的 mi,li,Oi ,在这一次计算中,得到了新的 mij,lij,Pij

首先会更新全局最大值

mi=max(mi,mij)

然后更新 softmax 分母

li=emimili+emijmilij

接下来更新 O 数组,之前的 O 数组是基于 mi,li 进行缩放的。

当前贡献是 PijVj ,也需要被缩放,这部分贡献是

Oidiag(li)1(diag(li)emimiOi+emijmiPijVj)

于是,我们就解决了这些问题。不考虑drop与mask,我们得到了它的伪代码

FlashAttention伪代码

Recomputation

在反向传播的时候,我们有两个中间变量 P,S 需要被存储,然而这两个东西一直在 SRAM 上面。作为中间变量,可以用 Q,K,V 三个变量进行计算。