切换深色模式
FlashAttention
5.2.1 Flash Attention
FlashAttention 是一种精确 Attention 算法,不是近似 Attention。 它的目标不是改变 Transformer 的数学定义,而是改变 Attention 在 GPU 上的计算方式:
普通 Attention 会显式构造巨大的注意力矩阵
然后再计算
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 的理论计算复杂度是:
所以显存复杂度是:
当序列长度 (n) 很大时,(n^2) 会非常夸张。
例如:
但训练时通常不止保存一个矩阵,还要保存 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 的核心就是:
普通 Attention 的流程大概是:
问题是
FlashAttention 的想法是:
把
分块搬进 SRAM,在计算核心中完成局部矩阵乘法和在线 softmax,将输出 逐块累积在 SRAM 中,完整的 和 始终不写回 HBM,最终只将 写回 HBM。
3. FlashAttention
(1)分块矩阵乘法
Attention 中的
将
其中,
对于每一个 query block (Q_i),FlashAttention 逐个遍历 key/value block。这里以
计算小块 attention score:
这个小矩阵 (S_{ij}) 会放在 SRAM / shared memory 中,用完就丢,不写回 HBM。
接下来会经过局部 Softmax 更新,得到
最后,局部 Value 加权:
对于 query block
将这些结果相加,得到最终 attention 对应
将每个 query block 遍历的输出拼接起来,就得到了最终的 Attention 模块的输出:
(2)Online Softmax
回到之前留下的问题,在得到小块 attention score
对第 (i) 个 query token,标准 Attention 输出是:
注意 softmax 的分母是整行求和
Online Softmax 的核心思想是:维护运行统计量(running statistics),每来一个新的 block,就更新这些统计量,最终得到与全局 softmax 完全一致的结果。
Online Softmax 的核心
先考虑一个单独的 query token:
现在受限与 SRAM 内存,只能先进入
对于
为避免数值溢出,使用稳定 softmax:
但上面的
对于
现在处理
处理完所有
现在推广到 query block
对
局部 attention score 为:
此时 (Q_i) 中有 (B_q) 个 query token,处理第
初始化:
处理第 (j) 个 KV block 时,我们已经有
先计算:
对每一行求最大值:
更新全局最大值:
分母更新:
加权分子更新:
处理完所有 KV block 后,得到最后的 attention 输出 :
最后把所有 query block 的输出拼接起来: