ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

FlashAttention实战:重构显存读写,让长序列LLM训练不再OOM

FlashAttention实战:重构显存读写,让长序列LLM训练不再OOM 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然后在下一步重新读入。假设$N4096$、$d128$、Batch Size为4、12层Decoder我们算一笔账$S$矩阵大小$4096 \times 4096 \times 4$字节FP32 64MB$P$矩阵大小同上64MB每层产生至少128MB的中间写入量读写合计256MB12层就是3GB左右的额外显存流量。这还只是单次前向反向传播时这些中间矩阵还得重新读出来算梯度。所以在标准实现里Attention的计算速度实际上是被显存带宽卡死的而不是被GPU的浮点计算单元卡死的。2.2 GPU的存储层级与带宽差异这个概念用生活类比解释最清楚HBM就像是你的外接机械硬盘空间大几十GB但读写慢SRAM就像是CPU的L1缓存极小在A100上是192KB但极快。标准实现每一小步都把大规模中间结果“存回硬盘”然后下一次再从硬盘读出来而FlashAttention的目标就是尽量让你手头正在用的数据一直放在“L1缓存”里热乎着不轻易落盘。从硬件数字看A100的HBM带宽约为2TB/sA100的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.8GPU的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_dim128, dropout0.0, causalFalse): 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_maskNone): 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_pself.dropout_p, softmax_scaleNone, causalself.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_attentionTrue在HuggingFace Transformers里通常是一个模型config字段设置后模型的注意力层会自动使用flash内核。如果你的模型是自己实现的Decoder就把上面的FlashAttentionBlock替换原MHA即可。这里也补一下怎么用HuggingFace Transformers打开FlashAttention以Llama类模型为例直接在LlamaConfig里设置attn_implementationflash_attention_2即可。from transformers import LlamaConfig, LlamaForCausalLM config LlamaConfig( vocab_size32000, hidden_size4096, intermediate_size11008, num_hidden_layers32, num_attention_heads32, max_position_embeddings4096, attn_implementationflash_attention_2, ) model LlamaForCausalLM(config)这里建议设置torch_dtypetorch.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显存占用节省比例102412.6 GB8.1 GB35.7%204825.9 GB13.4 GB48.3%409668.4 GB24.2 GB64.6%8192OOM43.6 GB超过90%估算可以看到序列长度越长FlashAttention的显存优势越明显。8192时标准实现直接OOM而FlashAttention还能保持可用。5.2 实验二训练吞吐量在固定序列长度为4096、Batch Size 4、A100上跑500步统计平均每秒处理的token数实现平均吞吐量tokens/sStandard PyTorch MHA1483FlashAttention 23871FlashAttention 2 activation checkpointing3362加速比大约是2.6倍。激活重计算会降低一些吞吐但可以让模型跑更深的层或更大的batch这在训练超长序列时有很强的实用价值。5.3 调优经验从我的调试过程来看以下几个配置对性能影响很大块大小block sizeFlashAttention内核的块大小由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时建议用modereduce-overhead并且把FlashAttention模块标记为torch.compiler.disable避免编译该部分。6.4 显存碎片导致OOMFlashAttention虽然省显存但长序列配合大batch下仍然可能出现“明明剩很多显存却OOM”的情况。这类问题往往不是单一算子造成的而是因为PyTorch缓存分配器在反复分配和释放不同大小的中间张量时产生了显存碎片。我的排查顺序是用torch.cuda.memory_summary()看最大块大小和碎片比例尝试PYTORCH_CUDA_ALLOC_CONFexpandable_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在单卡训练中用好已经能解决很多显存和时间预算问题了。
返回列表