在介绍 FlashAttention 之前,先统一本文使用的符号。
| 符号 | 含义 |
|---|
| N | Sequence Length,序列长度 |
| D | Head Dimension,每个 Attention Head 的特征维度 |
| Q | Query 矩阵 |
| K | Key 矩阵 |
| V | Value 矩阵 |
为什么 Attention 很慢?#
Transformer 中最重要的计算之一就是 Self-Attention, 从公式上看,Attention 主要由两个矩阵乘法和一个 Softmax 组成:
SPO=QKT=Softmax(S)=PVQ,K,V∈RN×D而:
S,P∈RN×N最终:
O∈RN×D而问题在于这个 N×N 的中间矩阵,当 Sequence Length N 增大时,Attention Matrix 的大小按照 N2 增长,当序列长度从 1024 增长到 8192 时,这个中间矩阵的元素数量会增加 64 倍。
FlashAttention 的核心思想并不是减少 Attention 的 FLOPs 计算量,而是:
不要把巨大的 Attention Matrix 写回 GPU 的 HBM,而是把计算拆成小块,让数据尽可能留在 GPU 的片上高速存储中。
在 Standard Attention 实现中,中间矩阵 S,P 只是为了下一个计算阶段使用,却被反复写入和读取 HBM,后面的 FlashAttention 就是围绕如何减少这些 HBM IO 展开的。
1 / 4
Standard Attention 实现
IO-Aware Algorithm:Tiling#
FlashAttention 所指的 IO-Aware 思想实际上就是 Tiling。
在矩阵乘法的优化中,我们介绍了 Tiling,比如矩阵 A 乘以矩阵 B,我们将 A 矩阵和 B 矩阵进行分块,逐块加载到 shared memory 中,读取 shared memory 完成计算,而不是直接读取 global memory (HBM)。
对于 Attention 计算 S=QKT 我们也可以使用 Tiling,一块一块的计算 Score Matrix。
Online Softmax#
先回顾普通的 Safe Softmax,三次遍历分别求最大值、求指数和以及计算结果。
FlashAttention 不希望保存完整的 Score Matrix,因为它很占空间,而标准的 Softmax 实现需要完整的 Score Matrix 的一行,因此我们需要优化 Softmax 的实现让它也可以一块一块地处理。
Safe Softmax 标准实现假设已经处理:
x1,⋯,xk定义两个变量 running max m 以及 running sum l:
mk=1≤i≤kmaxxilk=i=1∑kexi−mklN 就是我们想要的 Softmax 的分母。
现在来了新的 block:
xk+1,⋯,xk+b新的最大值:
mnew=mk+b=max(mk,max(xk+1,⋯,xk+b))=max(mk,mblock)由于旧的 lk 是按照 mk 计算的,因此需要引入纠正系数:
correction factor=emk−mnew不难得到:
lnew=lk+b=lkemk−mnew+j=k+1∑k+bexj−mnew=lkemk−mnew+lblock这样,我们只需要维护:
m,l就可以将 Softmax 的前两次循环,求最大值循环和求和循环,压缩为一次循环。
我们实际上把三次循环缩减到了两次循环,而这就是 Online Softmax。
Flash Attention#
Online Softmax 还不够,我们真正需要的是 O=PV,我们希望 O 也可以在线更新。
假设我们将 Q∈RN×D 切分为一个一个的 Qblock∈Rtile_q×D,同理将 K∈RN×D 和 V∈RN×D 分别切分为一个一个的 Kblock∈Rtile_k×D 和 Vblock∈Rtile_k×D。
Flash Attention我们先考虑 tile_q=1,tile_k=b 的情况,这个时候我们的 Oblock∈RD 和 Qblock∈RD 变成了一维向量,记作 o 与 q。
在 CUDA 实现中,每一个线程块处理到一个 q,我们后续推导专注一个线程块上的计算,所以默认 q 是不变的。
假设我们处理到某个 Sblock 时已经处理了:
x1,⋯,xk其中 xi=qkiT,其中 kiT 是 KT 矩阵中的一个列向量。
我们在 online softmax 中定义了 running max m,running row sum l,也就是:
mk=1≤i≤kmaxxilk=i=1∑kexi−mk当处理完 Sblock 时会引入:
xk+1,⋯,xk+b我们可以得到新的 m 和 l:
mnew=mk+b=max(mk,mblock)lnew=lk+b=lkemk−mnew+lblock现在考虑 o,它是:
o=j=1∑NlNexj−mNvj其中 vj∈RD 是 V 中的一个行向量。
我们想让 o 也能在线更新,于是我们定义新的变量 oi:
oi=j=1∑iexj−mivj那么自然就有
o=oN/lN现在我们推导 ok 和 ok+b 的关系:
ok+b=j=1∑k+bexj−mk+bvj=j=1∑kexj−mk+bvj+j=k+1∑k+bexj−mk+bvj=(j=1∑kexj−mkvj)emk−mk+b+j=k+1∑k+bexj−mk+bvj=ok×emk−mk+b+j=k+1∑k+bexj−mk+bvj=ok×emk−mnew+[exk+1−mnewexk+2−mnew⋯exk+b−mnew]vk+1vk+2⋮vk+b太好了,我们将 oold 和 onew 联系到一起了,现在我们可以在线更新 o 了!
而且我们可以很轻松地扩展到 tile_q=1 的情况,Flash Attention 的计算过程如下:
Flash Attention 计算过程演示
完整的 Flash Attention 实现#
见 我的实现。