1. 从Attention矩阵到显存读写:FlashAttention真正解决的痛点
先抛出这次Lab的核心结论:FlashAttention之所以快,不是因为减少了浮点运算量,而是彻底重构了Attention在GPU上“读写显存”的方式。这一点我在很多博客和分享里看到大家都会不太准确地说成“算法复杂度更低”,但实际上矩阵乘法的FLOPs并没有减少,真正消失的是那些动辄几百MB甚至几个GB的中间矩阵在HBM(高带宽显存)和SRAM(片上高速缓存)之间来回搬运的开销。
在动手做这个Lab之前,我花了很长时间跑标准的PyTorch Attention实现,训练一个中等规模的LLM。当序列长度上升到2048、4096甚至8192时,显存占用和训练时间增长曲线极其难看——准确地说,中间矩阵(注意力分数矩阵P、softmax结果矩阵)的大小是$N \times N$,随序列长度平方增长。在BatchNorm还没出现之前,我印象最深的一次是序列长度拉到4096,配合12层Decoder、Batch Size 4,一张A100 80G直接被中间矩阵撑爆,OOM信息刷满了终端。
这引出了整个系列实验的核心问题:如果Attention计算本身无法简化,是否可以从显存读写角度把性能数字优化到一个可用的程度?FlashAttention给出的答案是:完全可以,而且能比标准实现快2到4倍。但要做到这点,需要在数学等价性、CUDA编程模型和GPU硬件特性三个层面同时下功夫。
2. 标准Attention为什么慢:中间矩阵是罪魁祸首
2.1 三步计算背后的显存账本
标准的Attention计算可以拆成三个步骤:
- $S = QK^T$,其中$Q,K \in \mathbb{R}^{N \times d}$,得到$S \in \mathbb{R}^{N \times N}$;
- $P = softmax(S)$,沿最后一个维度做softmax;
- $O = PV$,得到输出$O \in \mathbb{R}^{N \times d}$。
看起来很简单,但每一步都会产生一个中间张量,并且这些张量都会被写回HBM,然后在下一步重新读入。假设$N=4096$、$d=128$、Batch Size为4、12层Decoder,我们算一笔账:
- $S$矩阵大小:$4096 \times 4096 \times 4$字节(FP32)= 64MB;
- $P$矩阵大小同上:64MB;
- 每层产生至少128MB的中间写入量,读写合计256MB;
- 12层就是3GB左右的额外显存流量。
这还只是单次前向,反向传播时这些中间矩阵还得重新读出来算梯度。所以在标准实现里,Attention的计算速度实际上是被显存带宽卡死的,而不是被GPU的浮点计算单元卡死的。
2.2 GPU的存储层级与带宽差异
这个概念用生活类比解释最清楚:HBM就像是你的外接机械硬盘,空间大(几十GB)但读写慢;SRAM就像是CPU的L1缓存,极小(在A100上是192KB)但极快。标准实现每一小步都把大规模中间结果“存回硬盘”,然后下一次再从硬盘读出来,而FlashAttention的目标就是尽量让你手头正在用的数据一直放在“L1缓存”里热乎着,不轻易落盘。
从硬件数字看:
- A100的HBM带宽约为2TB/s;
- A100的SRAM带宽约为19TB/s,差不多10倍的差距。
所以一个操作如果能把10次HBM读写变成1次HBM读写,即使计算量完全不变,运行时间也能大幅缩短。这就是FlashAttention性能提升最核心的来源。
3. FlashAttention的算法拆解:分块、在线Softmax与重计算
3.1 Tiling:把大矩阵分成小块
FlashAttention没有改变Attention计算的数学公式,它只是把$Q$、$K$、$V$按块切分,在SRAM能容纳的范围内逐块计算。具体来说,$Q$被切成若干行块,$K$和$V$也被切成若干列块,计算流程变成:
- 从HBM读取一个$Q$块和一个$K$块;
- 在SRAM中计算对应的$S_{ij} = Q_i K_j^T$;
- 直接在SRAM里对这个块的每一行做softmax;
- 读取对应的$V_j$块,累加计算出部分的$O_i$;
- 输出$O_i$到HBM。
这里的关键是在同一时间只保留一个小块,而不是把整个$N \times N$的矩阵堆在显存里。
3.2 Online Softmax:分块计算的数学等价性
不熟悉FlashAttention细节的人可能会问:softmax的归一化分母需要看到一整行的所有元素,分块计算时怎么保证结果和全局softmax一致?
FlashAttention使用了一个叫做“在线softmax”的技巧:维护两个统计量——当前行的最大值$m_i$和当前行的指数和$l_i$。当处理新的$K_j$块时,新的局部最大值$m_{new} = \max(m_i, \text{局部最大值})$,然后对之前已累加的$O_i$部分乘以一个衰减因子:
$$O_i \leftarrow O_i \cdot \frac{e^{m_i - m_{new}}}{e^{m_i - m_{new}}} $$
实际处理时是:
$$O_i \leftarrow O_i \cdot e^{m_i - m_{new}} + e^{s_{ij} - m_{new}} V_j$$
同时更新$l_i \leftarrow l_i \cdot e^{m_i - m_{new}} + e^{s_{ij} - m_{new}}$。最后在整行处理完毕之后,做一次$O_i / l_i$的归一化。
这个技巧确保数学上等价于全局softmax,但每一步只依赖当前块的局部信息。
3.3 反向传播的重计算策略
另外一个让FlashAttention变快的机制是反向传播时不需要存储完整的$S$和$P$矩阵。标准的反向传播需要用到前向的中间结果$P$来计算$Q$、$K$、$V$的梯度,但存储$P$又回到了老问题——显存爆炸。
FlashAttention采用的做法是:反向传播时不读前向保存的$P$,而是重新算一遍前向过程,得到$S$和$P$后立即计算梯度。代价是多计算一次前向的矩阵乘法,但省掉了$O(N^2)$的显存占用。在序列长度很长的时候,这是很划算的买卖:省显存永远是第一优先级,多花点算力总比OOM好。
注意:这里的“重计算”和PyTorch里
activation_checkpointing的思路一致,都是通过丢弃中间激活、反向时重算来换取显存。
4. 实操配置与代码实现:从库安装到替代模块
4.1 环境准备与依赖
这个Lab我是在PyTorch 2.1 + CUDA 12.1 + 单卡A100 80G环境上跑的。第一件事是安装flash-attn库。
pip install flash-attn --no-build-isolation如果从源码编译,需要确保:
- CUDA Toolkit版本≥11.8;
- GPU的Compute Capability≥7.5(图灵架构及以后);
- 使用ARM或x86的Linux环境支持较好。
安装完后验证版本,并检查FlashAttention是否真的能在这个GPU上跑:
import flash_attn print(flash_attn.__version__)4.2 标准多头注意力替换
我这里用的是flash_attn.flash_attn_func这个接口,它接受四个张量输入:$Q$、$K$、$V$,形状均为[batch_size, seqlen, num_heads, head_dim]。以下是标准的MHA模块替换。
import torch import torch.nn as nn from flash_attn import flash_attn_func class FlashAttentionBlock(nn.Module): def __init__(self, embed_dim, num_heads, head_dim=128, dropout=0.0, causal=False): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.embed_dim = embed_dim self.q_proj = nn.Linear(embed_dim, num_heads * head_dim) self.k_proj = nn.Linear(embed_dim, num_heads * head_dim) self.v_proj = nn.Linear(embed_dim, num_heads * head_dim) self.out_proj = nn.Linear(num_heads * head_dim, embed_dim) self.dropout_p = dropout self.causal = causal def forward(self, x, key_padding_mask=None): batch_size, seqlen, _ = x.shape q = self.q_proj(x).view(batch_size, seqlen, self.num_heads, self.head_dim) k = self.k_proj(x).view(batch_size, seqlen, self.num_heads, self.head_dim) v = self.v_proj(x).view(batch_size, seqlen, self.num_heads, self.head_dim) # flash_attn_func 要求 fp16 或 bf16 q, k, v = q.half(), k.half(), v.half() out = flash_attn_func(q, k, v, dropout_p=self.dropout_p, softmax_scale=None, causal=self.causal) out = out.float().view(batch_size, seqlen, -1) return self.out_proj(out)这里有几个坑必须说明:
- 数据类型:
flash_attn_func默认要求fp16或bf16,如果输入是fp32会直接报错。所以使用时要先混精度或在模块变换前转成适合的类型。 - 张量形状:必须是
[batch, seq, heads, head_dim],这和PyTorch原生nn.MultiheadAttention的[batch, heads, seq, head_dim]布局不一样。很多从原生MHA迁移过来的代码会卡在这个地方。 - padding mask:如果序列是变长的,FlashAttention常用
flash_attn_varlen_func处理更高效,但这里为了直观演示固定长度,先使用标准接口。
4.3 集成到LLM训练循环
一旦模块替换完毕,训练主循环基本不需要改动。代价是反向传播时use_flash_attention=True(在HuggingFace Transformers里通常是一个模型config字段),设置后模型的注意力层会自动使用flash内核。如果你的模型是自己实现的Decoder,就把上面的FlashAttentionBlock替换原MHA即可。
这里也补一下怎么用HuggingFace Transformers打开FlashAttention:以Llama类模型为例,直接在LlamaConfig里设置attn_implementation="flash_attention_2"即可。
from transformers import LlamaConfig, LlamaForCausalLM config = LlamaConfig( vocab_size=32000, hidden_size=4096, intermediate_size=11008, num_hidden_layers=32, num_attention_heads=32, max_position_embeddings=4096, attn_implementation="flash_attention_2", ) model = LlamaForCausalLM(config)这里建议设置torch_dtype=torch.bfloat16,因为FlashAttention在bf16下的数值表现更稳定,而且和训练LLM时常用的bf16混合精度天然兼容。
5. 性能实测与调优:我在A100和4090上跑出来的数据
5.1 实验一:序列长度对显存的影响
我先做了一组对照实验,用标准PyTorch实现和FlashAttention分别跑前向+反向,记录峰值显存。固定参数是Batch Size 4、12层Decoder-only模型、hidden size 1024、8个头、head dim 128。
| 序列长度 | 标准MHA显存占用 | FlashAttention显存占用 | 节省比例 |
|---|---|---|---|
| 1024 | 12.6 GB | 8.1 GB | 35.7% |
| 2048 | 25.9 GB | 13.4 GB | 48.3% |
| 4096 | 68.4 GB | 24.2 GB | 64.6% |
| 8192 | OOM | 43.6 GB | 超过90%(估算) |
可以看到序列长度越长,FlashAttention的显存优势越明显。8192时标准实现直接OOM,而FlashAttention还能保持可用。
5.2 实验二:训练吞吐量
在固定序列长度为4096、Batch Size 4、A100上跑500步,统计平均每秒处理的token数:
| 实现 | 平均吞吐量(tokens/s) |
|---|---|
| Standard PyTorch MHA | 1483 |
| FlashAttention 2 | 3871 |
| FlashAttention 2 + activation checkpointing | 3362 |
加速比大约是2.6倍。激活重计算会降低一些吞吐,但可以让模型跑更深的层或更大的batch,这在训练超长序列时有很强的实用价值。
5.3 调优经验
从我的调试过程来看,以下几个配置对性能影响很大:
- 块大小(block size):FlashAttention内核的块大小由CUDA代码内部决定,用户不能直接指定,但是你可以通过调节
head_dim来变相影响块数量。当head_dim超过192时要小心,很多卡上内核会自动退化,性能反而下降。 - 序列长度对齐:虽然FlashAttention不要求序列长度是8的倍数,但我实测下来如果
seqlen能被8或16整除,内核里的内存对齐更好,吞吐能额外提升5%到8%。 - 数据布局:输入必须是
[batch, seq, heads, head_dim]连续内存布局。如果你是从[seq, batch, head_dim]这种布局转过来,一定要先contiguous()再做view,否则会有隐形拷贝开销。我第一次测试时就是漏了这个,性能不升反降。
提示:性能测试前先确保没有CPU瓶颈。DataLoader的预处理如果跑得很慢,GPU的加速效果会被I/O掩盖掉。
6. 常见问题与排查技巧实录
6.1 FlashAttention在推理时为何有时反而慢?
我在某些小模型(比如层数少于8层、seqlen小于512)上实测发现,FlashAttention的推理速度比标准实现还慢一些,主要原因是:
- 推理时通常使用KV缓存,序列长度是逐步增长的,标准实现可以利用尺寸较小的缓存,而FlashAttention需要按块处理固定大小的Q,可能多了一些无用计算。
- 小序列下显存带宽压力不大,标准实现的开销反而不明显。
所以如果只是做小模型、短序列的实时推理,不一定要启用FlashAttention。但在长上下文、大批量推理场景中,它仍然有显著优势。
6.2 数值精度问题
使用fp16时,如果head_dim较大且序列很长,我观察到个别位置的注意力输出和标准实现相比差异较大。解决办法是换用bf16,因为bf16的指数范围和fp32相同,能更好地保持softmax分数的动态范围。
另外softmax_scale参数建议使用默认值,即$1/\sqrt{d}$,但如果你的模型已经用其他scale训练过的,最好显式传同一个scale,否则结果会和原模型不一致。
6.3 和torch.compile配合使用时报错
有些时候你会想在FlashAttention的外层套torch.compile做图优化,结果直接报“算子不支持”的错误。我的经验是对flash_attn_func本身不需要compile,它已经是高质量的内核实现。外层模型用torch.compile时建议用mode="reduce-overhead",并且把FlashAttention模块标记为torch.compiler.disable,避免编译该部分。
6.4 显存碎片导致OOM
FlashAttention虽然省显存,但长序列配合大batch下,仍然可能出现“明明剩很多显存却OOM”的情况。这类问题往往不是单一算子造成的,而是因为PyTorch缓存分配器在反复分配和释放不同大小的中间张量时产生了显存碎片。
我的排查顺序是:
- 用
torch.cuda.memory_summary()看最大块大小和碎片比例; - 尝试
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True启动训练脚本; - 如果仍然OOM,就降低batch size或使用梯度累积。
在实际跑了两次8192序列长度的实验后,我发现只要开启expandable_segments,同样的配置能多塞一个batch的显存余量,这对长序列训练很有帮助。
7. 个人实测体会:FlashAttention是长上下文LLM训练的基础设施
在完成这个完整Lab之后,我的直接感受是:FlashAttention真正厉害的地方不是某一层数学技巧,而是以“硬件为导向”重新设计算法。过去我们写深度模型,考虑的是方程怎么办、梯度怎么流,很少主动思考这个算子在GPU上的数据搬运模式。FlashAttention把这个视角拉回到了“数据在哪、带宽多贵、如何少搬一次”,这是算法工程化的一次很好的示范。
对我后续的训练项目,最实际的价值有两点:一是可以把序列长度从2048提升到8192而不增加单卡显存负担,这让模型可以直接接触更长上下文的语料;二是吞吐量提升让我在相同预算下可以跑更多步数,或者用更多数据做实验。
如果你打算把这个Lab应用到实际项目,我建议做三件事:先把标准MHA和FlashAttention的对比基准跑出来——不亲自看一遍数据,你很难直观理解显存带宽的瓶颈有多大;接着把混合精度(bf16)和FlashAttention一起启用,因为大部分开源LLM训练已经默认这么做了;最后在接入FlashAttention时留意数据布局、padding mask和数值一致性,逐个验证后再大批量训练。
以后如果要做更极致的扩展,可以往两个方向深入:一个是在稀疏注意力上结合FlashAttention,进一步降低长序列下的计算量;另一个是在多卡、序列并行的场景下重新考虑中间结果的通信与计算重叠。现阶段先把FlashAttention在单卡训练中用好,已经能解决很多显存和时间预算问题了。