FlashAttention

FlashAttention

一、Introduction

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

二、 Background

GPU 内存层级结构:SRAM / HBM / DRAM

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,VRN×d\mathbf{Q},\mathbf{K},\mathbf{V} \in \mathbb{R}^{N \times d}

先计算注意力分数矩阵S(每个token的Q与所有token的K进行计算)

S=QKTRN×N\mathbf{S} = \mathbf{Q}\mathbf{K}^T \in \mathbb{R}^{N \times N}

再计算注意力权重矩阵P(对S按行进行softmax)

P=softmax(S)RN×N\quad \mathbf{P} = \mathrm{softmax}(\mathbf{S}) \in \mathbb{R}^{N \times N}

最后计算输出矩阵O

O=PVRN×d\quad \mathbf{O} = \mathbf{P}\mathbf{V} \in \mathbb{R}^{N \times d}

标准的Attention实现需要将SP存储到HBM中,占用O(N2)\mathbf{O(N^2)}的内存

部分操作例如softmax操作是内存密集型操作,大量内存访问导致较慢的运行时间。

还有其它的逐元素操作如掩码(按行)和P的丢弃也加剧了这一现象。

三、 FlashAttention算法的设计与分析

3.1 设计目标

给定 HBM 中的输入Q,K,VRN×d\mathbf{Q},\mathbf{K},\mathbf{V} \in \mathbb{R}^{N \times d},计算注意力输出ORN×d \mathbf{O} \in \mathbb{R}^{N \times d} 并将其写入 HBM。目标是减少 HBM 访问量使访问量少于Θ(N2)\mathbf{\Theta(N^2)}

3.2 分块

使用了分块(tiling)和重计算(recomputation)两种方法来实现。

分块计算的详细流程:

FlashAttention 分块计算流程

  1. 如图1右侧所示,先从行维度将Q分成TrT_r个块,将K和V都分成TcT_c个块。

关于块大小的选取,原则是使SRAM能同时放下Qi\mathbf{Q_i}Kj\mathbf{K_j}Vj\mathbf{V_j}Oi\mathbf{O_i}i\ell_imim_i的前提下,让分块的数量尽可能少。按上述分块大小,所需的总空间为(2Bc×d+2Br×d+2Br)(2B_c \times d + 2B_r \times d +2B_r)。上述理论推导时忽略了i\ell_imim_i等低阶空间复杂度变量,实际工程时会采用保守缩小块大小的策略。

  1. 遍历Q\mathbf{Q}的分块,从HBM取一个块Qi\mathbf{Q_i}
  2. 遍历K\mathbf{K}V\mathbf{V}的分块,从HBM取[Kj\mathbf{K_j},Vj\mathbf{V_j}],注意二者分块的行序列要一致。
  3. 在SRAM上计算注意力分数矩阵Sij=QiKjTRBr×Bc\mathbf{S}_{ij} = \mathbf{Q}_i \mathbf{K}_j^T \in \mathbb{R}^{B_r \times B_c}
  4. 在SRAM上计算Qi\mathbf{Q_i}中选中的行的每行最大值,注意力矩阵逐元素取e指数P~ij=exp(Sijm~ij)RBr×Bc\tilde{\mathbf{P}}_{ij} = \exp(\mathbf{S}_{ij} - \tilde{m}_{ij}) \in \mathbb{R}^{B_r \times B_c},进行softmax操作的分母~ij=rowsum(P~ij)RBr\tilde{\ell}_{ij} = \mathrm{rowsum}(\tilde{\mathbf{P}}_{ij}) \in \mathbb{R}^{B_r}
  5. 更新Qi\mathbf{Q_i}中选中行的每行行最大值mim_i,每行各元素的e指数之和i\ell_i
  6. 增量式计算更新Qi\mathbf{Q_i}中选中的行的注意力输出Oidiag(inew)1(diag(i)emiminewOi+em~ijminewP~ijVj)\mathbf{O}_i \leftarrow \mathrm{diag}(\ell_i^{\mathrm{new}})^{-1}\big(\mathrm{diag}(\ell_i)e^{m_i - m_i^{\mathrm{new}}}\mathbf{O}_i + e^{\tilde{m}_{ij} - m_i^{\mathrm{new}}}\tilde{\mathbf{P}}_{ij}\mathbf{V}_j\big),写入HBM。
  7. 将更新后的mim_ii\ell_i写回HBM。

与标准Attention访问次数对比分析:

二者Q,K,V\mathbf{Q},\mathbf{K},\mathbf{V}读取次数相同,区别在于FlashAttention不用存S\mathbf{S}P\mathbf{P},为什么传统attention需要存呢?——因为一般N很大,并且N>>d,SRAM存Nd\mathbf{Nd}空间没问题,存不下S和P这种N2\mathbf{N^2}空间的。但是每次多了i\ell_imim_i的读写,总共需要4×Br×Tr×Tc=4NTc4 \times B_r \times T_r \times T_c = 4NT_c次读写,同时多了2×Tr×(Tc1)×Br×d=2Nd(Tc1)2 \times T_r \times (T_c-1) \times B_r \times d = 2Nd(T_c-1)次读写,复杂度为Θ(N2d2/M)\mathbf{\Theta(N^2d^2/M)},而标准的Attention的HBM读写次数复杂度为Θ(N2)\mathbf{\Theta(N^2)}

对于典型的d(64-128)和M(约100KB)值,d2d^2远小于M,因此FlashAttention所需的HBM访问次数比标准Attention实现的少很多倍

3.3 重计算

通过分块,无需存储S和P,但是训练时反向传播需要使用S\mathbf{S}P\mathbf{P}来计算梯度。解决办法是重新计算

通过存储输出O\mathbf{O}和 softmax 归一化统计量(m,)(m,\ell),我们可以在SRAM中轻松地重新计算注意力矩阵S\mathbf{S}P\mathbf{P}

相比之下,即使FlashAttention有更多的 FLOPs,重新计算通过减少 HBM 访问来加速反向传播。

3.4 扩展:Block-Sparse FlashAttention

用一个掩码矩阵M{0,1}N/Br×N/BcM \in \{0, 1\}^{N / B_r \times N / B_c}来表示块与块之间是否需要计算,


FlashAttention
https://jaspery.top/2026/08/04/FlashAttention/
作者
Jaspery
发布于
2026年8月4日