Skip to content

FlashAttention

5.2.1 Flash Attention

FlashAttention 是一种精确 Attention 算法,不是近似 Attention。 它的目标不是改变 Transformer 的数学定义,而是改变 Attention 在 GPU 上的计算方式:

普通 Attention 会显式构造巨大的注意力矩阵

S=QKRn×n,

然后再计算

P=softmax(S),O=PV.

FlashAttention 的关键是:

不显式保存完整的 (QK^\top \in \mathbb{R}^{n\times n}) 注意力矩阵,而是分块计算,并在 GPU SRAM 中完成局部计算,只把最终结果写回 HBM。

原始 FlashAttention 论文将其定义为一种 IO-aware exact attention algorithm,核心是通过 tiling 减少 GPU 高带宽显存 HBM 与片上 SRAM 之间的数据读写。

1. 普通 Attention 的主要瓶颈
Attention(Q,K,V)=softmax(QKd)V.

普通 Attention 的理论计算复杂度是:O(n2d),显存复杂度主要来自注意力矩阵:S,PRn×n

所以显存复杂度是:

O(n2)

当序列长度 (n) 很大时,(n^2) 会非常夸张。

例如:n=8192,则注意力矩阵大小为 81922=67,108,864,如果使用 FP16,每个元素 2 bytes,那么一个矩阵大约需要:

67,108,864×2134 MB.

但训练时通常不止保存一个矩阵,还要保存 softmax 结果、dropout mask、中间激活,用于反向传播。因此显存消耗会远大于这个数。

更本质的问题是:

GPU 很快,但显存读写相对慢。普通 Attention 的大量时间花在 HBM 与 SRAM 之间搬运 (n\times n) 矩阵。

FlashAttention 要优化的正是这个 IO 瓶颈。


2. GPU 内存层级:为什么 IO 很重要?

可以粗略理解为 GPU 有两类重要存储:

存储位置说明容量速度特点用途
HBM(GPU 显存)离计算核心较远,访问延迟更高大(几十 GB 到上百 GB)带宽高,比 CPU 内存快很多,但相对片上缓存较慢深度学习大张量主要存放处参数、激活值、梯度等大规模数据
SRAM / Shared Memory / Registers(片上缓存)离计算单元非常近,访问延迟低小(远小于 HBM)极快适合频繁重复访问的数据常配合矩阵乘法、卷积、Attention 等 block/tile 级计算

HBM 全称是 High Bandwidth Memory,可以理解为 GPU 上的大容量显存。在深度学习中,HBM 主要存放大规模张量:

模型参数、输入数据、激活值、梯度、优化器状态、中间计算结果

SRAM 全称 Static Random Access Memory ,它是 GPU 芯片内部的高速存储。在深度学习中,它主要用于:

临时存放正在计算的一小块数据

可以总结概括为:

  • HBM 是 GPU 的大仓库,但每次从 HBW 取数据到计算核心都有代价
  • SRAM 是计算核心旁边的小工作台

FlashAttention 的核心就是:

尽量在工作台 SRAM 上连续完成更多操作,少访问大仓库 HBM

普通 Attention 的流程大概是:

Q,KHBM读入Q,K片上 SRAM计算核心:QKSRN×N片上 SRAM写回SHBMSHBM读入S片上 SRAM计算核心:softmax(S)PRN×N片上 SRAM写回PHBMP,VHBM读入P,V片上 SRAM计算核心:PVORN×d片上 SRAM写回OHBM

问题是 SP 都是 N×N,非常大,反复在 HBM 和 SRAM 之间搬运代价极高。

FlashAttention 的想法是:

Q,K,V 分块搬进 SRAM,在计算核心中完成局部矩阵乘法和在线 softmax,将输出 O 逐块累积在 SRAM 中,完整的 SP 始终不写回 HBM,最终只将 O 写回 HBM。

3. FlashAttention
(1)分块矩阵乘法

Attention 中的 QKV

Q=(q1q2q3qn)K=(k1k2k3kn)V=(v1v2v3vn),QKVRn×d

QKV 分块:

Q=[Q1Q2QTQ],K=[K1K2KTK],V=[V1V2VTK]QiRBq×d,Kj,VjRBk×d.

其中,BqBk 分别是每个分块 QiKjVj 的词向量的个数。

对于每一个 query block (Q_i),FlashAttention 逐个遍历 key/value block。这里以 Bq=2,Bk=3 为例:

Qi=[qi1qi2],Kj=[kj1kj2kj3],Vj=[vj1vj2vj3]

计算小块 attention score:

Sij=QiKjTd=1d(qi1kj1Tqi1kj2Tqi1kj3Tqi2kj1Tqi2kj2Tqi2kj3T)RBq×Bk

这个小矩阵 (S_{ij}) 会放在 SRAM / shared memory 中,用完就丢,不写回 HBM。

接下来会经过局部 Softmax 更新,得到 softmax(Sij) ,具体如何处理后面再讨论。

最后,局部 Value 加权:

oij=softmax(Sij)VjRBq×d

对于 query block Qi,这样逐个遍历 key/value block (K1,V1),(K2,V2),,(KTk,VTk)

oi={oi1,oi2,,oiTk}

将这些结果相加,得到最终 attention 对应 Qi 位置上的输出:

Oi=(softmax(QKd)V)i=j=1TkoijRBq×d

将每个 query block 遍历的输出拼接起来,就得到了最终的 Attention 模块的输出:

O=Concat(O1,O2,,OTQ)Rn×d
(2)Online Softmax

回到之前留下的问题,在得到小块 attention score Sij 之后,我们需要将其进行局部 softmax 更新。分块计算的难点是:虽然每次只看到一个 KV block,但最终必须得到全局 softmax 的结果。

对第 (i) 个 query token,标准 Attention 输出是:

oi=j=1nexp(sij)t=1nexp(sit)vj,sij=qikjTd

注意 softmax 的分母是整行求和 t=1nexp(sit),但分块运算时我们只拿到一个 block,例如只看到 si1,si2,,siBk,无法直接计算出 softmax 的分母。因此,我们就无法直接对每个 block 单独做 softmax。

Online Softmax 的核心思想是:维护运行统计量(running statistics),每来一个新的 block,就更新这些统计量,最终得到与全局 softmax 完全一致的结果。


Online Softmax 的核心

先考虑一个单独的 query token:

qiR1×d,KRn×d,VRn×d

现在受限与 SRAM 内存,只能先进入

K1=[k1,,kN1],V1=[v1,,vN1]

对于 K1,V1 ,计算局部 attention 输出:

sij=qikjdR1,softmax(sij)=exp(sij)r=1N1exp(sir)R1,oiN1=softmax(si)V1=j=1N1softmax(sij)vjR1×d

为避免数值溢出,使用稳定 softmax:

oiN1=j=1N1exp(sijmiN1)vjr=1N1exp(sirmiN1),miN1=max1jN1sij.

但上面的 softmax 方法是不对的,我们还有 K2,V2 没有纳入计算。

对于 qi,K1,V1,维护三个运行统计量:

miN1=max1jN1sijiN1=r=1N1exp(sirmiN1)AiN1=r=1N1exp(sirmiN1)vr

现在处理 K2,V2

miN2=max(miN1,maxN1+1jN2sij)liN2=(miN1miN2)liN1+r=N1+1N2exp(sirmiN2)AiN2=(miN1miN2)AiN1+r=N1+1N2exp(sirmiN2)vr

处理完所有 K,V 之后,最终输出为:

oi=AiN2liN2

现在推广到 query block Qi

QKV 进行分块:

Q=[Q1Q2QTQ],K=[K1K2KTK],V=[V1V2VTK]QiRBq×d,Kj,VjRBk×d

局部 attention score 为:

Sij=QiKjdRBq×Bk.

此时 (Q_i) 中有 (B_q) 个 query token,处理第 j 个 KV block 时,需要对每一行分别维护:

mi(j)RBq,i(j)RBq,Ai(j)RBq×d.

初始化:

mi(0)=,i(0)=0,Ai(0)=0.

处理第 (j) 个 KV block 时,我们已经有 mi(j1),i(j1),Ai(j1)

先计算:

Sij=QiKjTdRBq×Bk

对每一行求最大值:

miblockj=rowmax(Sij)RBq×1

更新全局最大值:

mi(j)=rowmax(mi(j1),miblockj)

分母更新:

i(j)=exp(mi(j1)mi(j))i(j1)+rowsum(exp(Sijmi(j)))RBq×1

加权分子更新:

Ai(j)=exp(mi(j1)mi(j))Ai(j1)+exp(Sijmi(j))VjRBq×d

处理完所有 KV block 后,得到最后的 attention 输出 :

Oi=AiTkiTkRBq×d

最后把所有 query block 的输出拼接起来:

O=Concat(O1,O2,,OTq)Rn×d