如果你在 A100 上跑一个 4096 长度的 attention,PyTorch 原生实现和 FlashAttention 的差距,最夸张的时候可以到 10 倍。我第一次看到这个 benchmark 数字时,觉得提交者是不是在 benchmark 里做了手脚。后来我拿着 profiler 一步步追,才发现慢的根源不在“算”,而在“搬”。

GPU 里有两类存储:HBM 和 SRAM。HBM 是显存,容量大但带宽有限;SRAM 是芯片上的高速缓存,容量小但带宽高一个数量级。FlashAttention 能让 attention 变快,靠的就是尽量让数据待在 SRAM 里,别去 HBM 里来回跑。
三行 PyTorch,贵在搬数据
标准 attention 的计算公式很简单:Attention = softmax(QK^T)V。在 PyTorch 里,你可能写过这三行:
s = torch.matmul(q, k.transpose(-2, -1))
p = torch.softmax(s, dim=-1)
o = torch.matmul(p, v)
这三行代码,每一行都是一个独立的 CUDA kernel。第一行算出 S 矩阵后,要把它写回 HBM;第二行 softmax 再从 HBM 把 S 读出来,算完 P 又写回去;第三行再读 P 和 V。一个 N×N 的中间矩阵,被写进去又读出来,白白走两趟。
问题在于,HBM 带宽并不便宜。A100 的 HBM 带宽大约 2TB/s,听起来很高,但 S 和 P 的大小随序列长度平方增长。序列长度 4096 时,S 有 1600 万个元素,FP16 下 32MB。一次写、一次读,来回两次就是一百多兆的流量。这些流量花的时间,反而超过了矩阵乘法本身。
你可能会问:为什么不让这三个算子共享中间结果?因为 PyTorch 的算子彼此独立,只能通过显存交换数据。这是库抽象层付出的代价。
可以想象一下做菜:HBM 是冰箱,SRAM 是灶台。你切一次菜就放回冰箱,再拿出来炒,当然又慢又费劲。FlashAttention 的做法是把食材一次拿到灶台边,切、炒、调味全在台上完成,最后只把成品放回冰箱。
- HBM
- GPU 的“主存”,容量可达几十 GB,延迟高,带宽相对低。
- SRAM
- 计算核心旁边的片上缓存,容量小(A100 约 20MB),但带宽接近 HBM 的十倍。
FlashAttention 的三板斧
论文用了三个关键技巧,我把它们逐个拆开。
第一招:分块(Tiling)。把 Q、K、V 切成小块,每次只把一小块加载到 SRAM,在片上算出这一块的 attention,累加到输出。块内产生的 S 和 P 只存在于 SRAM,整个计算过程不产生 N×N 的中间矩阵。中间矩阵没了,HBM 的读写压力就小了一大半。
分块尺寸的选择很讲究。块太小,SRAM 利用不充分;块太大,放不下且并行度下降。论文里针对 A100 选择了类似 64×64 的块大小,这需要根据硬件 microbenchmark 来调。
第二招:Online Softmax。这是真正的难点。softmax 的归一化需要看完一整行才能知道最大值和总和,但分块时你手上只有一部分,怎么办?FlashAttention 为每个 token 维护一个运行状态:当前最大值 m 和归一化和 l。每处理一个块,就更新 m、l,同时把之前累加的输出 O 按比例缩放,让它们和新的指数值保持匹配。等所有块处理完,再做一次最终归一化。每一步都是精确的,不丢任何数学信息。
第三招:反向传播时重计算。训练时为了算梯度,标准实现会把正向的 S 和 P 保存在内存里,这就是 O(N²) 的内存开销。FlashAttention 选择不存,反向时需要什么就重新算一遍。看起来是拿计算换内存,但省掉的 N×N 矩阵读写,远比重计算更贵。
一张表看清:它比原生快在哪
| 维度 | PyTorch 原生 | FlashAttention |
|---|---|---|
| Kernel 数量 | 至少 3 个 | 1 个 fused kernel |
| 中间矩阵 | N×N 写回 HBM | 留在 SRAM |
| 内存复杂度 | O(N²) | O(N) |
| 数学结果 | 精确 | 精确 |
我最初的误解和验证
我第一次听说 FlashAttention 时,以为它又是一个 sparse attention 的变体,砍掉一些不重要的注意力权重来加速。但读论文时我愣住了:文中反复强调“exact”,即结果和标准 attention 一模一样。为了确认,我把自己模型里的 attention 替换成 FlashAttention 版本,跑了一遍测试集,输出误差在 1e-6 级别。那一刻我才意识到,过去的优化靠“改模型”,而 FlashAttention 靠“改底层实现”——这是完全不同的两条路。
为什么 PyTorch 原生做不到?
你可能觉得,思路这么清晰,PyTorch 为什么不直接改掉?答案在于,PyTorch 是高层库,很难跨算子做这种融合。每个算子都要有独立 kernel,kernel 之间只能通过显存传递数据。虽然 torch.compile 也在尝试自动融合,但对于 FlashAttention 这种需要精细管理 SRAM 的算法,自动编译还达不到手工 CUDA kernel 的精度。
这里我说的“PyTorch 原生实现”特指手动写的多 kernel 版本。最新 PyTorch 的 torch.nn.functional.scaled_dot_product_attention 在支持时会自动调用 FlashAttention kernel,其实已经是 FlashAttention 的受益者了。
长上下文训练成为可能
FlashAttention 最直接的价值,是让长上下文模型真正落地。序列长度翻一倍,标准 attention 的内存消耗就翻四倍。到了 128k 上下文,一个 N×N 的中间矩阵就有几十 GB,单张 A100 都放不下。FlashAttention 把内存复杂度从平方降成线性,128k 甚至更长才变得可训练。
这也可以解释为什么很多长文本模型都依赖这类优化:不是模型结构变了,而是同样的模型结构,现在能塞进显存了。
不是银弹,但方向对了
FlashAttention 的局限也清晰:它针对 NVIDIA GPU 的 SRAM 架构设计,换到别的芯片要重写;它只优化 attention 本身,如果模型瓶颈在 MLP,加速有限;它用了一些内存布局假设,自定义 attention mask 等需求可能要改 kernel。
但它给整个社区带来的启发已经足够大:在 GPU 上,时间不总是花在计算上,很多时候花在搬数据上。FlashAttention 告诉我们,不用发明新的数学,只需重新审视数据流动的路径,就能让训练快 5-10 倍。
如果你还想验证这些数字,建议把 FlashAttention 论文的算法 1 和算法 2 对着代码读一遍,再跑一遍官方 benchmark。看完你会明白,同样的数学,换一种数据流动方式,能快到什么程度。
原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/445.html