如果一个实现多做了一些计算,却运行得更快,我们该怎样解释?单看 FLOPs,似乎很难得出这样的结果;把数据搬运也计入时间,就不奇怪了。事实上,数据访问一度是模型训练和推理最大的性能瓶颈。FlashAttention在此基础上做了一系列优化。
1. GPU 的执行方式
我们可以先从矩阵乘法理解并行计算:一个输出元素尚未算完,并不妨碍另一个元素同时计算。GPU 安排大量计算单元执行这类结构相近的工作,追求总吞吐;CPU 则投入更多控制逻辑和缓存,以适应复杂控制流程并降低单线程延迟。这是一种粗略比较,但足以解释神经网络为什么适合 GPU。
GPU 上的任务按 Grid、Block 和 Thread 组织。一个 Grid 包含多个 Block,Block 被调度到流式多处理器 SM 上;NVIDIA GPU 通常以 32 个线程组成的 Warp 为基本执行单位。同一 Warp 中的线程执行相同指令,只处理不同数据。如果它们进入不同条件分支,硬件需要分批执行各条路径,暂时不在当前路径上的线程会空闲。
2. 存储层次与 memory wall
GPU 的存储从近到远大致包括寄存器、共享内存与 L1 Cache、L2 Cache、全局显存。越靠近计算单元,访问越快,容量也越小。寄存器归单个线程使用,共享内存由同一 Block 的线程共享,跨 Block 的数据交换通常需要经过全局显存。
问题在于,矩阵计算吞吐增长得比显存带宽更快,计算单元可能大部分时间都在等待数据,这就是 memory wall。我们若想提高速度,就不能只问需要做多少运算,还得问每次从全局显存取来的数据能用几次。
3. Roofline 模型
为了把这个问题量化,我们用算术强度表示每搬运一个字节完成的运算量:
$$I=\frac{\text{FLOPs}}{\text{Bytes transferred}}$$
若显存每秒最多送来 $B$ 字节,每字节对应 $I$ 次运算,那么带宽只能支撑每秒 $BI$ 次运算。同时,硬件峰值算力 $P_{peak}$ 又给出了另一个上限。因此
$$P\leq\min(P_{peak},BI)$$

图 1:Roofline 模型中的带宽受限区与计算受限区。来源:JAX Scaling Book。
当 $BI<P_{peak}$ 时,性能主要受带宽限制;当 $BI\geq P_{peak}$ 时,计算单元才成为主要限制。逐元素运算每个元素只做少量计算,却至少需要一次读取和一次写回,通常属于 memory-bound。足够大的矩阵乘法能够复用输入,算术强度更高,更可能接近峰值算力。
例如 ReLU 对 FP32 元素约搬运 8 字节,只做一次比较或选择,算术强度很低。低精度格式减少每个元素的字节数,除了节省显存,也会直接降低传输成本。
4. 减少访存的几种办法
低精度计算让相同带宽一次传输更多元素。实际训练不会让所有步骤都使用同一种格式:权重、激活和矩阵乘法可以采用 BF16 或 FP8,累加、归一化、Softmax 和优化器状态常保留更高精度。
再考虑几个连续的小操作。假如每个操作都单独启动一个 kernel,上一步刚写回显存的结果,下一步又要读出。我们把它们融合起来,中间值便可以留在寄存器或共享内存中。减少了 kernel 启动和中间结果读写,数学上的目标计算保持相同;具体浮点舍入仍可能因执行顺序而有差别。融合也受寄存器、共享内存和算子依赖限制。
还有一种乍看更反常的选择:既然保存某些激活很贵,我们能否先不保存,等反向传播需要时再算?激活重计算就做了这笔交换。它能减少需要保留的激活显存,但未必更快:一般的 checkpointing 常会增加训练时间;只有在节省的传输、融合收益或更大 batch 的收益足够时,额外计算才可能换来净加速。

图 2:激活重计算用额外计算换取较低的激活显存。来源:PyTorch。
5. 合并访存与 Tiling
全局显存按连续数据段传输。同一个 Warp 的线程若访问连续且对齐的地址,硬件可以把请求合并为较少的显存事务;若相邻线程访问的地址相隔整行,即使读取元素数相同,也可能产生更多事务。
回到矩阵乘法,我们还可以让不同输出共享已经取来的输入。Tiling 将矩阵划分为小块,先把子块放进共享内存,再由多个线程反复使用。只看两个 $T\times T$ 子块相乘这一阶段,读入元素数为 $2T^2$,计算量约为 $2T^3$ FLOPs;忽略输出和其他开销后,算术强度随 $T$ 增长。分块的收益由此就能看出来。
分块并非越大越好。共享内存和寄存器有限,过大的块会降低一个 SM 同时驻留的 Block 数;矩阵不能整除分块大小时还会产生无效计算。实际 kernel 要在数据复用、并行度与占用率之间折中。
6. 标准注意力为什么慢
标准注意力写成
$$S=\frac{QK^{\mathsf T}}{\sqrt d},\qquad P=\operatorname{Softmax}(S),\qquad O=PV$$
如果把这三步分别执行,我们就要把 $S$、$P$ 两个 $N\times N$ 中间矩阵写入全局显存,再读回做后续计算。序列一长,这两份中间结果就很昂贵。沿用分块的思路,FlashAttention 将 $Q$、$K$、$V$ 的子块读入片上存储,直接累积输出,避免保存完整的 $S$ 和 $P$。不过,矩阵乘法容易分块,Softmax 的分母依赖整行分数,这个问题还没有解决。
FlashAttention 计算的仍是精确注意力,时间复杂度依然是 $O(N^2d)$。它主要降低了 HBM 与片上存储之间的读写次数。
7. Online Softmax
我们先固定一个 Query,只看一行注意力。假设目前处理了 $s_1,\ldots,s_t$,不保存这整行分数,只维护最大值 $m_t$、分母 $l_t$ 和未归一化加权和 $z_t$:
$$m_t=\max_{1\le i\le t}s_i,$$
$$l_t=\sum_{i=1}^{t}e^{s_i-m_t},\qquad z_t=\sum_{i=1}^{t}e^{s_i-m_t}v_i$$
加入 $s_{t+1}$ 后,新的最大值为
$$m_{t+1}=\max(m_t,s_{t+1})$$
如果新分数更大,之前各项减去的最大值就过时了。但我们不必重读旧分数,因为对所有旧项都有 $e^{s_i-m_{t+1}}=e^{s_i-m_t}e^{m_t-m_{t+1}}$。旧累计量统一乘一个系数,再加上新项即可:
$$l_{t+1}=e^{m_t-m_{t+1}}l_t+e^{s_{t+1}-m_{t+1}},$$
$$z_{t+1}=e^{m_t-m_{t+1}}z_t+e^{s_{t+1}-m_{t+1}}v_{t+1}$$
处理完整行后,我们取 $o=z_N/l_N$,就恢复了原来的加权平均。上式按单个元素写,只是为了把更新看清楚;实现时可以一次加入一个块。这样只保存最大值、分母和部分输出,就能把 Softmax 与矩阵乘法结合起来。反向传播再利用保存的输入和少量统计量重算局部分数,而不读取完整注意力矩阵。