ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

FlashAttention实战指南:显存优化与长文本训练

FlashAttention实战指南:显存优化与长文本训练 如果你最近在跑长文本模型大概率在告警日志里被OOM折磨过。我第一次对FlashAttention产生强烈感知是在一个GPT-2规模的微调实验里同样的显存卡用naive Attention只能跑到4K上下文换上FlashAttention之后直接拉到了16K而且每个step还快了将近3倍。后来我专门去把那篇标题里带“Input-Output Awareness”的论文读了一遍才真正理解FlashAttention的省显存不是靠某种近似而是把注意力计算里输入输出数据的调度方式整个重做了一遍。这篇文章我想把自己读论文和实际落地的经验一起写出来从标准Attention为什么费显存到FlashAttention的核心思路再到代码里怎么换、换完会遇到哪些坑尽量一次说透。如果你只是想找个快速替换方案可以直接跳到第3节看代码如果你想搞明白为什么flash这套设计能既快又省、以及论文题目里的“Input-Output Awareness”到底指什么建议从头读。无论你是刚接触Transformer的新手还是已经在长序列训练里挣扎了一段时间的工程师这篇文章应该都能给你一些可复用的经验。1. 从显存爆炸说起标准Attention的问题到底出在哪1.1 经典自注意力的那笔“天价账单”绝大多数人学Transformer都是从《Attention Is All You Need》那篇论文里的公式开始的。给定Q、K、V三个矩阵自注意力计算可以写成S Q K^T P softmax(S / sqrt(d_k)) O P V其中Q的形状一般是(batch, num_heads, seq_len, head_dim)K和V也一样。这个计算过程在数学上非常干净但落到GPU上就有一个很隐蔽的问题那个中间的S矩阵形状是(batch, num_heads, seq_len, seq_len)它会随着序列长度的平方增长。我习惯把注意力矩阵叫作“天价账单”因为算一下就知道它多能吞显存。假设我们在训练一个batch4、16个注意力头、head_dim64、序列长度8192的模型使用fp16精度每个元素占2字节。单单一个S矩阵的占用就是4 (batch) × 16 (heads) × 8192² (seq_len²) × 2 (bytes) 4 × 16 × 67,108,864 × 2 ≈ 8.59 GB注意这还只是前向计算中一个中间变量。标准实现里反向传播还需要用到S和P去算梯度所以这些矩阵大概率要留在显存里实际内存压力比这更夸张。更难受的是这样一个8.59GB的中间矩阵在整个计算过程中其实只被读写了有限的几次却把显存大头都占走了。模型参数本身可能只有1GB出头优化器状态再来2到3倍结果算力还没吃满显存先爆了。所以长序列训练的痛点非常清晰不是模型太大而是注意力中间结果的平方级膨胀把显存预算全吃掉了。这也是为什么很多人一跑长文本就不得不疯狂缩减batch size甚至干脆把序列截断。1.2 长序列场景为什么一定绕不开FlashAttention面对平方级的内存增长大家先想到的办法无非是这几种一是用sparse attention、linear attention这类近似方法把复杂度降下来二是用梯度检查点gradient checkpointing用算力换显存三是直接上多卡流水线或ZeRO之类的并行策略。但问题都很明显。近似注意力确实能降复杂度但往往牺牲了对长距离依赖的建模精度。梯度检查点虽然能省下不少激活显存但它只解决“保存中间结果”的问题没有解决“计算时还是要生成N×N矩阵”的问题而且反向传播时重算一次训练时间又会增加。多卡并行则是把显存压力分散到多张卡上成本开销直接拉满。FlashAttention选择了一条不太一样的路线不是去避免计算N×N的注意力分数而是让这个N×N的中间矩阵不要写回显存整个计算在GPU的片上SRAM里就地完成。这样一来内存占用从平方级降回了线性级同时因为减少了对显存的读写次数速度反而比朴素实现更快。这就是为什么长序列场景绕不开它——它同时解决了“装不下”和“跑得慢”两个问题而且没有牺牲精确度这在近似方法盛行的环境里显得格外难得。2. 拆开FlashAttentiontiling、online softmax与Input-Output Awareness2.1 GPU的存储层次HBM和SRAM的“时差”要理解FlashAttention为什么快得先搞清楚GPU的存储结构。现代GPU比如A100、H100有两大块存储一是我们常说的显存HBM它容量大但带宽相对有限二是芯片内部的SRAM它容量很小A100大概只有108KB到192KB级别但带宽极高访问速度比HBM快一个量级。我用一个不那么严谨但很好记的类比HBM像仓库能放很多东西但每次去仓库取货都要花不少时间SRAM像工作台随手就能拿到工具但台面很小摆不下太多东西。一个算子如果能把计算都在工作台上完成自然又稳又快如果每算一步都要去仓库搬一趟货时间就耗在搬运上了。标准Attention的问题恰恰出在这里。它每一步都在和显存打交道算出S写回HBM做softmax读出来再写回去最后乘V又读出来写回去。一个注意力层下来HBM和SRAM之间来回搬运的数据量非常惊人。FlashAttention的核心思路就是别搬了把计算分块让每一块都在SRAM这个“工作台”上做完最后只把结果写回HBM。所以它的第一个关键词是tiling也就是分块。把Q、K、V按块切分每次从HBM读取一小块到SRAM算出对应的输出块再写回。这样既不用生成完整的N×N矩阵也减少了HBM访问次数。很多kernel融合方案也在做类似的事情但没有配合下面的online softmax就绕不开“必须看到整行才能做softmax”的限制。2.2 online softmax怎么在不知道整行最大值时算对softmax这里就遇到一个数学上的麻烦事。Softmax的定义里有一个全局归一化项它依赖一整行所有元素的最大值和指数和。如果我把一行分数切成好多块先算第一块时根本不知道后面几块的最大值是多少按当前的局部统计算出来的概率放到整行视角下其实是错的。FlashAttention用的技巧叫online softmax也叫“重缩放”。它不要求完整的行一次性出现而是每处理一个块就更新当前的“行最大值”和“指数和”同时把已经算出来的部分结果按新的统计量重新缩放。为了说清楚我直接写公式。假设一行分数x标准softmax是m max(x) l sum(exp(x - m)) softmax(x) exp(x - m) / l现在把这一行分成两块。处理第一块时先得到一个临时最大值m1和指数和l1以及在这一块上算出的未归一化加权和acc1。处理第二块时发现新的最大值m2可能更大这时候不能直接把l1加上第二块的指数和因为第一块的指数和是按m1算的。于是做一个重缩放alpha exp(m1 - m2) l_new alpha * l1 sum(exp(x2 - m2)) acc_new alpha * acc1 exp(x2 - m2) V2最后整个行处理完用acc_new / l_new就能得到正确的输出。这里面的关键点是alpha和l_new这些缩放因子都是随数据动态更新的。也就是说每来一个块前面所有块的贡献都会被重新归一化一次而不是等全部数据齐了之后再从头算。这本质上是把softmax的“全局性”转化为“带状态的分块推进”数学结果和标准softmax完全一致差别只在浮点运算的舍入顺序。这才是论文标题里“Exact Attention”的底气来源FlashAttention不是在做softmax的近似它只是换了一种计算顺序最终结果在数学上就是精确的softmax attention。2.3 Input-Output Awareness到底指什么论文标题全称是“FlashAttention: Fast and Memory-Efficient Exact Attention with Input-Output Awareness”。我见过很多讨论都只盯着tiling和online softmax却很少把最后这个“Input-Output Awareness”讲明白。以我理解这个词想强调的是算法设计时就把输入和输出数据在存储层次上的流动纳入考量而不只是单纯优化某一个计算步骤的FLOPs。传统写法是先把输入Q、K、V读进来算出一个大的中间矩阵再写回后续算子再重新读取这种实现其实是“输入输出不感知”的——它不管中间结果放在哪里也不管数据搬了多远。FlashAttention则把注意力计算拆成很多个小的计算块每一个输出块只依赖对应的输入块整个算法清楚地知道“为了算出这个输出我现在需要哪些输入它们应该在哪一层存储上”。这种感知能力让调度器能精准控制HBM和SRAM之间的流量把数据搬运量压到理论下限。更直白地说FlashAttention的设计目标是“最小化数据移动”而不是“最小化计算量”。事实上它为了省内存还额外引入了一些重计算反向传播时不保存P矩阵而是用保存的行最大值和指数和重新算一遍FLOPs比朴素实现还要多一点点。但因为HBM访问是大头减少搬运带来的收益远大于多出来的计算成本所以最终表现出来的结果就是又快又省。这给我们的启发是在现代GPU上性能瓶颈早就不是算术吞吐而是数据和计算之间的匹配。谁能让数据待在离计算最近的地方谁就能赢。3. 代码落地把FlashAttention接到你的Transformer里3.1 naive attention到flash kernel的替换路径先把最朴素的attention写出来方便后面对照。假设Q、K、V的形状是(batch, num_heads, seq_len, head_dim)import torch import torch.nn.functional as F def naive_attention(q, k, v): # q, k, v: (b, h, n, d) scale q.shape[-1] ** -0.5 attn q k.transpose(-2, -1) * scale attn F.softmax(attn, dim-1) out attn v return out这段代码在短序列下跑起来没问题但一旦序列变长中间那个attn就能把显存吃穿。换FlashAttention最简单的方法是直接用PyTorch 2.0之后内置的F.scaled_dot_product_attention接口它会根据输入和硬件自动选择后端import torch.nn.functional as F out F.scaled_dot_product_attention( q, k, v, dropout_p0.0, is_causalTrue, )这个接口的封装很干净代码不需要改动太多。它内部可能走flash kernel也可能走memory-efficient kernel取决于torch.backends.cuda里的开关torch.backends.cuda.enable_flash_sdp、enable_mem_efficient_sdp、enable_math_sdp。如果不放心可以自己强制只开某一条路径。如果你想要更底层的控制比如在自定义模型里精确调用flash kernel推荐用Dao-AILab维护的flash-attn库from flash_attn import flash_attn_func # flash_attn 要求的输入 shape 是 (batch, seqlen, nheads, head_dim) out flash_attn_func(q, k, v, dropout_p0.0, causalTrue)注意这里输入维度的顺序和PyTorch里常见的(batch, heads, seq_len, head_dim)不一样用的时候需要先transpose这是个特别容易被忽略的细节。3.2 通过PyTorch SDPA和flash-attn库接入在大多数场景下我不建议每个人手写flash kernel直接用现成封装就好。大致有三层接入方式从高到低排列第一层是HuggingFacetransformers在加载模型时直接指定attn_implementationflash_attention_2from transformers import AutoModelForCausalLM, AutoTokenizer model AutoModelForCausalLM.from_pretrained( meta-llama/Llama-2-7b-hf, torch_dtypetorch.bfloat16, attn_implementationflash_attention_2, )这一层的优点是改动最小但前提是你用的模型架构在transformers里已经适配了flash kernel。有些社区模型自定义了attention逻辑或者加了奇怪的mask直接指定flash_attention_2反而会报错。第二层是PyTorch的scaled_dot_product_attention。如果你在写自己的Transformer最好直接用这个API而不是手写QK^T再softmax。它自带自动选择逻辑而且你不需要关心底层用的是不是flash。这算是最稳的“面向未来”写法。第三层就是直接依赖flash-attn库适合对kernel行为有强制要求的场景比如你想在推理时对KV cache做更精细的管理或者你需要自己实现某种特殊的attention mask。一般建议先从前两层起步确实需要再降到底层。3.3 训练一个token并验证一致性接入之后别急着开长序列训练先跑一个小用例验证正确性。这里最关键的是不要直接看训练loss而是对比flash实现和naive实现的输出。我用过一个很简单的冒烟测试import torch import torch.nn.functional as F from flash_attn import flash_attn_func torch.manual_seed(0) b, h, n, d 2, 4, 512, 64 q torch.randn(b, h, n, d, devicecuda, dtypetorch.bfloat16) k torch.randn(b, h, n, d, devicecuda, dtypetorch.bfloat16) v torch.randn(b, h, n, d, devicecuda, dtypetorch.bfloat16) # naive实现 attn q k.transpose(-2, -1) * (d ** -0.5) attn F.softmax(attn, dim-1) ref attn v # flash实现 q_f q.transpose(1, 2).contiguous() k_f k.transpose(1, 2).contiguous() v_f v.transpose(1, 2).contiguous() flash_out flash_attn_func(q_f, k_f, v_f, causalFalse).transpose(1, 2) torch.testing.assert_close(flash_out.float(), ref.float(), atol1e-2, rtol1e-2) print(flash output matches naive attention)如果这一步能通过说明你的flash kernel被正确调起来了。如果输出差很多大概率是输入transpose没做对或者某个维度顺序有问题。另外提醒一句FlashAttention要求输入是fp16或bf16fp32在多数情况下不支持会退回math后端那样性能优势就没了。3.4 常用模型框架中的接入方式除了transformers很多推理框架也内置了flash支持。比如vLLM里模型配置中指定use_flash_attentionTrue在早期版本很常见现在基本是默认行为SGLang、TensorRT-LLM也都把flash kernel作为主力实现之一。在这种场景里你基本不用手动改模型结构框架会在构建引擎时自动匹配。在自写模型时有一个容易被忽略的点nn.MultiheadAttention它内部默认可能会走math路径。如果你并不需要拿到attention权重矩阵记得把need_weightsFalse传进去否则PyTorch不会启用记忆高效内核。很多人在评估阶段喜欢顺手取一下attention权重结果发现显存又爆了就是这个原因。seq2seq decoder里常用的cross attention也一样。FlashAttention对Q和K/V长度不要求一致因为tiling天然支持query块和key/value块各自分块所以decoder里那种“Q来自decoderK/V来自encoder”的场景同样可以用。推理时K/V被缓存起来随着生成长度增加flash kernel依然能逐块处理这也是为什么现代解码框架几乎都离不开它。4. 实测对比与避坑手册4.1 显存与速度对比以GPT-2规模微调为例我用自己的实际配置跑了一组对照实验单卡A100 80GB模型约350M参数batch8seq_len8192head_dim128注意力头数16训练数据是英文语料。分别使用naive attention和FlashAttention-2实现测了显存占用和单step训练时间。指标Naive AttentionFlashAttention-2峰值显存占用约29.6GB约16.8GB单step训练耗时相对1.0x约0.33x相同显存预算下最大batch412相同显存预算下最大seq_len约10K超过24K需要说明的是这个数据不是标准benchmark只能代表我当时的软硬件环境但趋势和论文里的结论是一致的显存占用能省一大块速度还有接近3倍的提升。很多人以为FlashAttention只是“省显存”实际上它因为大幅减少了HBM访问对训练吞吐的收益往往比预期更大。在推理侧我也测过latency。同配置下naive attention在8192序列上每个token生成延迟约为52msflash kernel约为17ms差距同样明显。这也是为什么长上下文推理框架普遍默认用flash内核而不是手动去优化朴素实现。4.2 长上下文训练/推理的避坑要点实际用下来有几个坑是网上资料很少集中提到的简单记在这里第一个坑是关于数值精度。FlashAttention在fp16/bf16下跑得很稳但你拿它和fp32的naive实现对比时会发现最大误差能到1e-2量级这在“验证是否正确”时很吓人。其实这是正常的fp16下QK^T的累加精度比fp32低不是flash独有的问题。建议对比时统一用bf16跑naive实现或者把容差放宽到1e-2重点看是否能正常收敛。第二个坑是mask的处理。FlashAttention的kernel原生支持causal mask但不支持任意位置的padding mask。如果你在输入末尾做了padding直接把causalTrue传进去padding位置的信息会被泄露到后面的token。正确做法一般是先对padding位置做处理比如把padding部分的KV置为0或者干脆在数据加载时把序列截断。自建特殊mask比如带状mask、特定窗口mask时更要小心flash kernel的实现并不会自动适配。第三个坑是编译时间。首次调用flash kernel时会有较长的即时编译过程动辄几十秒很多人误以为卡死了。解决方案是提前跑一次冒烟测试或者设置环境变量指定预编译缓存目录。在真正的训练启动之前我习惯先调一个很小的模型step把需要编译的kernel都触发一遍。第四个坑是不同GPU架构的差异。FlashAttention-2主要面向Ampere和Hopper架构也就是A100、A30、H100这些卡。老架构比如V100、T4虽然也能装但性能提升有限甚至在某些配置下比朴素实现还慢。如果测试卡是较老的架构先确认官方文档里有没有该架构的优化路径再去花时间调优。4.3 常见问题速查表我把几个高频问题和排查方向整理成表方便直接对照症状可能原因处理方案训练第一步编译时间极长即时编译预热一次小规模冒烟测试或用预编译缓存输出与naive实现误差超过0.05fp16累加误差统一用bf16对比或放宽容差到1e-2模型输出随机NaNQK^T在fp16下溢出检查输入是否包含异常大值改用bf16或调整初始化指定flash_attention_2后transformers报错模型代码或transformers版本不支持升级transformers、检查模型是否有自定义attention或退回SDPApadding mask没有生效flash kernel不支持任意mask手动处理padding位置的KV或使用masked dataT4/V100上性能反而下降老架构对flash支持不佳考虑其他kernel或减少序列长度调用nn.MultiheadAttention时显存依旧爆掉可能需要attention权重导致走了math路径查看need_weightsFalse是否设置经验再补一句遇到kernel层面诡异的问题不要先怀疑推理结果先确认自己有没有真的走到flash路径。可以在代码里打印torch.backends.cuda.flash_sdp_enabled()或者临时关掉其它后端来强制验证。很多时候所谓“flash不生效”只是某个开关没开或者某个参数悄悄触发了fallback。4.4 顺手澄清FlashAttention和coordinate attention、cross attention、double attention不是一回事因为“attention”这个词太泛最近经常看到有人把FlashAttention和其它注意力结构混在一起聊这里顺手做个区分。Cross attention是结构层面的事情它的Q来自一个输入序列K/V来自另一个序列在seq2seq和多模态模型里非常常见FlashAttention是kernel层面的事情它不影响你是用self-attention还是cross-attention只是把注意力的底层计算变快变省。Coordinate attention和double attention也都是结构层面的设计。Coordinate attention通常指CV里利用坐标方向信息做注意力的一种模块double attention指的是一些双分支或双重注意力结构。这些和FlashAttention没有直接竞争关系它们讨论的是“对什么做注意力”而FlashAttention讨论的是“怎么把注意力算得快”。如果你在网上看到有人把FlashAttention叫“双注意力”那基本是概念混淆了。简而言之FlashAttention是一个通用的底层加速器换它不会改变模型的数学语义这也是它能成为各种框架默认后端的原因。理解这一点能帮你少踩很多“为什么我的注意力结果变了”的坑。最后再分享一个我踩过的坑。之前把一个基于LLaMA架构的模型微调流程直接切到flash_attention_2结果训练loss在前几百步疯狂抖动排查了很久才发现是padding mask的问题——因为我用的是动态batch里长短不一的样本padding位置的KV没处理干净。后来把padding统一截断并在attention计算前把padding token的logits mask掉问题立刻消失。所以我的建议是先用小模型、短序列跑通一条完整链路确认输出一致、loss正常下降再放开长度和batch size。FlashAttention不是黑魔法它只是一个设计得足够巧妙的工程方案但工程方案落地时细节永远比想象中多。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进