Transformer GPU性能优化:从算子融合到FlashAttention,吃满显存带宽
Transformer 已经成为大模型时代绕不开的底座从 LLM 到多模态、视频生成主干网络基本全是它。但很多人真到自己上手在 GPU 上跑 Transformer 时才发现明明显卡不差训练却慢得离谱推理更是半天吐不出一个 token。快手团队在公开分享里反复强调过一个观点Transformer 要跑得快不能只停留在改 PyTorch 代码的层面真正要吃满 GPU必须做到底层优化。这话听起来硬核其实说白了就是——你得先把硬件喂饱再谈模型调优。这篇文章我想从算子、显存、kernel、部署这几个角度把 Transformer 在 GPU 上的性能瓶颈和优化路径拆开讲清楚。无论你是做 LLM 微调、推理部署还是自己手写 Attention 实验都有参考价值尤其适合那种“明明用了 GPU 却没感觉变快”的情况。1. Transformer 为什么“喂不饱”GPU瓶颈到底在哪很多人的第一反应是GPU 算力这么强Transformer 怎么还会慢问题往往不在“算力不够”而在“数据搬运不过来”。我见过太多同学把模型跑起来后第一件事就是盯 GPU 利用率发现只有 30%就觉得是代码没写好。但其实 GPU 底层优化要解决的问题比“利用率”这几个字复杂得多。1.1 算力强但显存带宽是短板GPU 的算力这些年涨得非常快。以 A100 为例FP16 的 Tensor Core 算力能做到 312 TFLOPS 左右H100 的峰值还要再翻几个量级。但显存带宽呢A100 的 HBM 带宽大约是 2TB/sH100 在 3.35TB/s 上下。算力和带宽之间的悬殊差距决定了 GPU 并不是所有计算都能跑在峰值。这里有个概念很关键算术强度也就是“每个字节的数据搬运到计算单元后能支撑多少次浮点运算”。如果算术强度太低GPU 的计算单元大部分时间都在等数据从显存搬过来这时候性能瓶颈就是“访存受限”而不是“计算受限”。Transformer 里很多算子恰恰是访存密集型的比如 LayerNorm、Softmax、残差连接、激活函数这些操作本身的计算量不大但要把整个张量读一遍再写一遍非常吃带宽。我用一个生活化的类比CPU 好比研究能力很强的专家但资料存在远处的仓库里。如果专家大部分时间都在等资料送过来而不是在看资料那就算他再聪明产出也上不去。GPU 的计算单元就是专家HBM 就是仓库带宽就是那条运输通道。Transformer 这种模型对运输通道的要求特别高尤其是长序列场景。1.2 Attention 的平方复杂度才是最大麻烦如果说显存带宽是“慢性病”那 Attention 机制就是“急性病”。Transformer 的核心公式大家都熟Attention(Q, K, V) softmax(QK^T / sqrt(d)) · V。问题就在其中的 QK^T 这一步它会产生一个 N×N 的注意力分数矩阵N 是序列长度。序列长度翻一倍这个矩阵的存储量级直接翻四倍。我算过一笔账假设 batch size 是 16head 数量是 32序列长度是 4096head_dim 是 128如果用 FP16 存储 QK^T 的注意力分数矩阵需要多少显存batch * heads 512 个注意力矩阵每个矩阵有 4096 * 4096 16.7M 个元素总元素数就是 512 * 16.7M ≈ 8.59B每个 FP16 元素占 2 字节光这一个中间结果就是 17GB 以上。很多人的显卡总共才 24GB这还没算 Q、K、V 本身、前馈网络和其他层的中间激活。这也是为什么长序列训练时经常 OOM为什么 FlashAttention 这类技术能成为“救世主”。所以做 GPU 底层优化时注意力机制永远是需要优先盯住的点。它不是“优化一下能快 20%”的问题而是“不优化可能根本跑不起来”的问题。1.3 Kernel Launch 与内存拷贝被忽视的隐形成本除了访存另一个容易被忽略的点是 kernel launch 开销。Transformer 的一个层里有多少操作Embedding、LayerNorm、QKV 投影、reshape、transpose、attention、concat、output projection、残差连接、FFN、激活、第二个残差再算上 dropout 和最后的 LayerNorm起码二三十个算子。在 PyTorch 里每个算子基本对应一次 GPU kernel 启动。kernel 启动本身是有开销的尤其当单个 kernel 很小、执行时间极短时CPU 端启动 kernel 的时间甚至可能超过 GPU 端执行的时间。这就像点外卖每道菜都用独立包装、派一次车结果配送成本比菜本身还贵。频繁的小 kernel 会让 GPU 在大部分时间里处于“等活干”的状态。所以在优化 Transformer 时一个很重要的思路就是把零散的算子“攒起来”减少 kernel 启动次数同时减少中间张量在显存里的读写次数。这也就是后面要说的算子融合。2. GPU 底层优化全景图从硬件特性到算子重写了解瓶颈之后再看优化手段就有了方向。GPU 底层优化的范围很广从数据精度、算子融合、到注意力实现方式、再到编译器自动优化每一层都能带来可观的收益。这里我把最关键的几块拆开讲。2.1 Tensor Core 与数据精度把硬件最擅长的用起来现在主流的 NVIDIA 显卡都内置 Tensor Core这是专门为矩阵乘加设计的硬件单元。A100 的 FP16 Tensor Core 算力远远超过普通 FP32 CUDA Core。很多人跑 Transformer 还在用默认的 FP32这等于手持一台跑车却一直挂着低速挡在开。混合精度训练和推理的核心思路就是把模型中占大头的矩阵乘法放到 Tensor Core 上用 FP16 或 BF16 计算同时把精度敏感的算子如 LayerNorm、Softmax保留在 FP32。BF16 和 FP16 相比指数位和 FP32 一样所以动态范围更稳在大模型训练里更常用。还有更新的 FP8、INT8能把吞吐再往上顶但需要更谨慎地处理量化误差。如果你的 PyTorch 版本比较新最简单的做法是使用 torch.autocast 或直接调用model.to(torch.bfloat16)。不过要留意混合精度不是无脑打开就完事。Loss Scaling、梯度的精度保持、某些算子的数值稳定性都需要实践确认。我在实际项目里的习惯是能上 BF16 就不上 FP16LayerNorm 和 attention 里的 softmax 强制走 FP32这样既省显存又快模型效果基本不掉点。2.2 算子融合少几次读写快一倍以上算子融合是 GPU 底层优化里性价比最高的手段之一。它的核心思想很简单把多个连续算子合并成一个或少数几个 kernel减少中间张量的显存分配和读写。数据搬运少了算术强度上去了整体速度自然就快了。举一个最常见的融合场景一个 Transformer block 里经常是 LayerNorm 残差连接 QKV 投影连在一起。如果不做融合计算流程是先读上一层的输出算 LayerNorm写回显存再读回来加残差再写回再读出来做矩阵乘法。每一趟都有完整的一次显存读和写。融合之后可以在一个 kernel 里直接把 LayerNorm 归一化、残差相加和矩阵乘法一次性完成中间数据只留在寄存器或共享内存里不落回 HBM。实现算子融合有几种方式一是手写 CUDA kernel灵活但开发成本高二是用 Triton 这类 DSL 写开发效率更高三是用 PyTorch 2.0 的 torch.compile 自动做图优化和 kernel 融合。对绝大多数人来说先试 torch.compile 是最快路径它能在不改模型代码的情况下把很多小算子自动融合掉并且还能捕获 CUDA Graph降低 kernel launch 开销。2.3 FlashAttention把“全局注意力”变成“分块注意力”FlashAttention 值得单独讲因为它几乎是 Transformer 性能优化里最重要的一次突破。传统 Attention 的问题在于为了算 softmax必须把完整的 QK^T 分数矩阵写到显存里再读回来做归一化这个中间矩阵极其占显存。FlashAttention 的思路是分块计算不保留完整的 N×N 分数矩阵。具体来说它把 Q、K、V 都切成小块在 GPU 的 SRAM 里逐块计算通过 Online Softmax 的技巧维护每个分块的 running max 和 running sum从而在不需要完整分数矩阵的情况下得到正确的 softmax 结果。这样显存占用从 O(N²) 降到了 O(N)而且因为避免了大矩阵的读写速度也有大幅提升。打个比方以前你要把一整仓库的货全部搬到展台上才能开始统计和计算FlashAttention 是直接在仓库里按区域分块盘点每块盘完就把结果汇总。省下来的搬运时间相当可观尤其序列越长越明显。现在 PyTorch 里可以直接用F.scaled_dot_product_attention它会根据输入自动选择后端包括 FlashAttention、Memory-Efficient Attention 等。如果你的场景是长序列、大 batch这个 API 换上去通常立刻能看到显存下降、速度上升。3. 实操方案怎么一步步把 Transformer 加速原理聊完落到实操。我自己的优化习惯是分四步走先 profile 摸清家底再做精度和编译层面的低成本优化然后针对推理场景做 KV Cache 和批处理优化最后才考虑多卡并行。这个顺序基本遵循“成本从低到高、收益从确定到不确定”避免一上来就投入大量精力去做复杂改造。3.1 第一步先 Profile再优化别靠感觉很多人优化性能是靠直觉觉得某个模块应该慢就去优化那个模块。但实际跑出来的热点往往和直觉不一样。做 GPU 底层优化的第一步永远是用工具看清楚时间都花在哪了。我常用的工具有这么几个nvidia-smi 看显存占用和 GPU 利用率只能粗略看PyTorch Profiler 可以看每个算子的耗时和显存分配定位到具体是哪一层、哪个操作在消耗时间Nsight Systems 看 kernel 在整个时间轴上的分布能发现 GPU 空等和 CPU 预处理瓶颈Nsight Compute 则更底层能看某个具体 kernel 的 SM 占用率、访存吞吐、指令流水线效率。拿到 profile 数据后重点关注几个指标GPU 利用率是否长期偏低、每个 kernel 的耗时占比、是否存在大量小 kernel、访存带宽是否打满。举个例子我曾经在优化一个生成模型时profile 后发现 LayerNorm 相关操作竟然占了近 18% 的时间看上去不痛不痒但它访存量极大而且被拆成了好几个小 kernel。后来把 LayerNorm 和前后两个操作融合成一个 kernel这部分时间直接降到了 3% 以内整体推理速度提升了差不多 10%。如果没有 profile这一步根本发现不了。3.2 第二步开启混合精度和编译优化profile 完之后最省事也最有效的“无脑操作”就是混合精度和编译优化。我建议按这个顺序来先确认数据精度。如果是训练尝试torch.autocast(device_typecuda, dtypetorch.bfloat16)如果是推理可以直接把模型转成model.to(torch.bfloat16)。注意不是所有显卡都支持 BF16老一些的架构比如 V100对 BF16 支持有限这时用 FP16 更合适。打开 PyTorch SDPA。把模型里手写的 attention 替换成F.scaled_dot_product_attention后端自动选择 FlashAttention 或 Memory-Efficient Attention。如果能用 FlashAttention长序列场景收益立竿见影。试试 torch.compile。对模型执行model torch.compile(model, modereduce-overhead)它会做算子融合、CUDA Graph 捕获等一系列优化。第一次运行会有编译开销预热后跑起来明显更稳更快。不过也要注意有些自定义算子或者动态 shape 场景下 torch.compile 可能报错或不生效需要用fullgraphTrue或改动态 shape 设置来调试。如果用的是 HuggingFace Transformers检查它的attn_implementation参数直接指定flash_attention_2比自己改模型省事得多。这套组合拳下来多数情况下一个 Transformer 的推理或训练速度能提升 2 倍以上而这几乎不需要改动模型结构纯粹是把硬件能力和最新库的优势吃满。3.3 第三步推理侧的 KV Cache 与连续批处理如果目标是部署推理那只看单次前向延迟是不够的还要看吞吐。LLM 推理分成两个阶段prefill 阶段处理用户输入的所有 tokendecode 阶段逐个生成新 token。decode 阶段最烦人的地方在于每一步生成新 token 时模型都要把前面所有 token 的 K 和 V 重新算一遍。解决方式是 KV Cache把历史 token 的 K、V 向量缓存下来避免重复计算。这是推理优化里必做的动作。KV Cache 的大小有个很直观的公式2 × 层数 × batch size × 序列长度 × KV head 数 × head_dim × 数据类型字节数。上下文一长KV Cache 会非常吃显存这也是为什么长上下文推理成本高。为了把 KV Cache 用得更高效业界搞出了 PagedAttention也就是 vLLM 的核心思路把 KV Cache 分成固定大小的块按需分配减少显存碎片提高利用率。同时配合 continuous batching连续批处理把不同请求的 decode 阶段拼到一起执行GPU 就不再因为单个请求的串行生成而大面积空转。我自己在部署 Qwen 这类模型时模型权重的显存占比往往只有三分之一剩下的基本都被 KV Cache 吃掉不优化这块根本扛不住高并发。3.4 第四步多卡并行与通信优化别急着加卡最后才是多卡并行。很多人觉得一块卡不够快就上八块卡但如果没有优化好单卡多卡可能更慢因为通信开销会拖垮整体收益。多卡训练常用的是数据并行、张量并行、流水线并行和 ZeRO 优化。数据并行简单每个卡一份完整模型但要不停做梯度同步也就是 all-reduce 通信张量并行把单个矩阵切到多张卡通信频率更高一般用在单机多卡且模型很大时。通信这块底层走的是 NCCL 库。优化方向主要有三个一是把通信和计算重叠比如在反向传播算梯度的时候同时异步发起上一层的梯度同步二是减少通信量比如梯度压缩三是选择高效的通信策略比如用 ring all-reduce 而不是简单的 broadcast reduce。实际工程里如果不是做大模型训练通常不建议一上来就弄多卡先保证单卡优化到位再评估通信瓶颈。4. 实战中常见的坑显存爆炸、驱动崩溃与莫名其妙的慢理论和方法都聊了最后必须说说实战里那些“翻车现场”。我做性能优化这些年遇到最多的不是优化不生效而是被各种环境和硬件层面的问题卡住有时候一个问题排查一整天结果发现是驱动或配置的问题。分享几个高频坑。4.1 OOM 不一定是显存真的不够训练 Transformer 时最常见的报错就是 CUDA Out of Memory。很多人第一反应是减 batch size但有些 OOM 其实跟 batch size 关系不大而是某些中间张量占用了不合理的空间。最常见的就是 QK^T 注意力分数矩阵前面算过长序列下它能撑爆显存。还有梯度全量保存、优化器状态Adam 需要两倍于模型参数的显存、中间激活值这些都是显存大户。解决思路有几个能用 FlashAttention 就尽量用它直接把 N×N 的中间矩阵干掉了开 gradient checkpointing用计算换显存在反向传播时不保存所有中间激活而是在需要时重算一次显存实在紧张时可以调低 batch、缩短序列长度、减小微批大小。还有一个小经验torch.cuda.empty_cache()只是把缓存还给 PyTorch 的内存池不是还给操作系统它不能根治显存碎片化别指望靠这个救急。4.2 GPU 驱动崩溃与 D3D 设备移除在 Windows 环境下跑 Transformer 训练或者推理偶尔会碰到“GPU 发生崩溃”或“D3D 设备已移除”之类的报错这个在深度学习圈子里其实很常见很多第一次遇到的人都以为是显卡坏了或者代码写错了。真实原因往往是 Windows 的 TDRTimeout Detection and Recovery超时检测与恢复机制在起作用。Windows 默认会检测 GPU 是否有响应如果某个 kernel 在一个较长周期内一直占着 GPU 不返回系统就认为 GPU 挂了然后重置设备。而 Transformer 训练里有些算子执行时间很长尤其在大模型、长序列场景下很容易触发这个看门狗机制。解决方式分几种一是把 TDR 的超时时间调大在注册表HKLM\SYSTEM\CurrentControlSet\Control\GraphicsDrivers下增加TdrDelay单位是秒可以调到 20 甚至更高二是升级或重装驱动确保 CUDA 版本、PyTorch 版本和驱动互相匹配三是如果跑的是长时间训练更推荐直接用 Linux 环境这也是为什么生产环境几乎都用 Linux 的原因。此类问题更多是环境稳定性的问题不是模型本身的问题判断清楚了就不会慌。4.3 为什么用了 GPU 却还是很慢这是新人最容易疑惑的问题我把它整理成一个速查表遇到了可以对着排查现象可能原因快速验证方式GPU 利用率极低几乎为 0模型没有放到 CUDA 上比如漏了model.to(cuda)打印next(model.parameters()).deviceGPU 利用率波动大CPU 跑满DataLoader 预处理太慢num_workers0调大num_workers开pin_memoryTrue单次推理延迟高但显存不高Kernel launch 开销占比大batch size 太小用更大的 batch 或捕获 CUDA Graph显存占用高、速度一般手动实现的 attention 没走 SDPA或者开了 FP32换 SDPA开 BF16/FP16所有操作都正常但速度上不去硬件过热降频或功耗限制用nvidia-smi看温度和功耗CUDA 不可用回退到 CPUPyTorch 版本和 CUDA 驱动不匹配跑torch.cuda.is_available()确认排查这类问题还有个硬核手段用 gpu-burn 这类压力测试工具先烤机。我之前遇到过一次训练间歇性崩溃模型代码怎么看都没问题跑 gpu-burn 一晚上直接暴露了硬件供电不稳定。先确认硬件没问题再去折腾软件优化能省掉很多无效劳动。5. 从快手视角看底层优化为什么工程团队必须自己下探聊了这么多技术细节再回到最开始那个问题为什么像快手这样的团队会反复强调 GPU 底层优化答案其实很现实因为规模化之后性能直接就是成本。5.1 业务规模越大性能越等于钱短视频、推荐、AIGC 这些业务每天都有海量推理请求。同样的模型如果推理速度能提升 20%意味着同一批 GPU 可以支撑多 20% 的流量或者在同等流量下少买 20% 的卡。这个数字放大到几千张卡、几万张卡的规模时省下的成本是千万级的。所以大厂会专门养团队做算子优化、做推理引擎、做显存调度不是炫技是算过账的。个人开发者虽然不需要自研一套底层设施但思路可以借鉴与其一上来就租八张 A100不如先把一张卡吃透。很多时候模型性能上不去不是卡不够好而是 PyTorch 默认路径没有把硬件优势发挥出来。学会看 profile、学会用 SDPA、torch.compile、混合精度这些现成优化往往比加卡更划算。5.2 底层优化不是单点技巧而是体系化的能力GPU 底层优化从来不是某一个技巧的胜利。它是一套从模型结构、算子实现、kernel 融合、显存管理、通信调度到硬件特性适配的完整链路。快手强调“需要 GPU 底层优化”本质上是提醒大家模型结构和数据 pipeline 只是其中一环剩下的性能空间藏在更底层的地方。对普通工程师来说我的建议是不一定每个人都非要会手写 CUDA kernel但至少要理解 GPU 的执行模型比如线程怎么组织、共享内存为什么快、访存带宽为什么是稀缺资源、kernel launch 为什么有开销。有了这个底层认知你再去看那些高级优化工具就会觉得它们每一步都顺理成章而不是黑魔法。最后再分享一点我个人的习惯。每次拿到一个新的 Transformer 项目我不会急着上复杂方案而是先做最基础的三件事把显存带宽和算力比算一遍把主要算子的耗时 profile 一遍再把 BF16、SDPA、torch.compile 这三件套开起来看效果。绝大多数场景下这三步做完已经有非常可观的提升。之后再根据 profile 结果决定要不要深入写 Triton kernel、要不要调显存管理策略。底层优化不是一上来就写代码而是先搞清楚你的 GPU 到底在等什么、缺什么然后再去解决那个最关键的短板。