torch.compile深度剖析:碎算子融合收益显著,GEMM场景基本白搭
先讲个我前几天调过的实际例子。同事递过来一个推荐模型说训练速度上不去怀疑是数据加载的问题。我上去没先看数据管线直接nsys profile跑了一轮对着 kernel 列表扫了一眼模型里占比最大的几个 GEMM kernel 吃掉了超过 60% 的 GPU 时间剩下的几十个 elementwise、mask、gather 之类的碎算子数量倒是很多单个却不到 20 微秒。他之前已经试过torch.compile看完结果直接说“没用”。我说你把编译日志贴出来。他开了TORCH_LOGScompilation发现大部分计算图确实被捕捉了但 Triton 生成的新 kernel 基本都是那堆碎算子GEMM 仍然走的是 cuBLAS。这个现象其实已经把结论写在脸上了torch.compile 对碎算子会认真做融合和代码生成对 GEMM 基本只能“让路”——因为后者早就被 cuBLAS、cutlass 优化到了接近理论峰值。这个例子基本就是这篇文章的全部结论torch.compile 不是一开就灵的开关它是一个编译管线。收益多少取决于你的计算图里到底有多少值得融合的碎片化算子。如果你的模型里到处都是 LayerNorm、Softmax、逐元素乘加、mask 操作的组合torch.compile 帮你省掉的是显存搬运和 kernel 启动时间这部分快一倍完全可能如果你的模型核心就是一堆大 GEMM那它绕不过 cuBLAS能给你留点面子的边际收益都有限——所以叫“基本白搭”。这篇文章面向 AI Infra、算法工程和 GPU 性能调优方向的读者。我会先把 torch.compile 的编译机制讲清楚再用一个能直接跑的 benchmark 模板量化碎算子收益然后解释为什么 GEMM 这条赛道没有编译器发挥空间最后给一套判断模型该不该开 torch.compile 的实操方法。1. 为什么先要弄清torch.compile 是编译管线不是一行魔法1.1 编译管线做了什么很多人对 torch.compile 的第一印象来自官方快速上手教程三行代码把你的 model 包一层然后看 loss 下降。这个演示太顺滑导致不少人把它当成模型层面的“性能开关”——开了就快不开就慢。但实际用下来经常发现根本不是这么回事有人甚至遇到开了以后变慢、爆显存、结果对不上。问题的根源在于torch.compile 不是一个统一的开关而是一套“编译管线”。它做的事情分三段。第一段是图捕获。PyTorch 是动态图框架Python 代码执行到哪算子就调用到哪。编译的第一步是用 torch._dynamo 把你传入的函数或模块重新追踪成一张静态计算图FX graph。追踪过程中遇到 Python 控制流、动态分支、不支持的第三方算子时_dynamo 会停下把这个位置标记为 graph break然后从断点继续追下一段图。最终一个模型可能被切成多段子图。第二段是后端编译。默认后端是 TorchInductor。它拿到子图后会做算子融合、buffer 复用、死代码消除等图优化然后为 GPU 生成 Triton kernel为 CPU 生成 C kernel。第三段才是执行。编译好的子图替换掉原来的 eager 执行路径后续同一段代码直接跑编译产物。理解了这个流程你就能理解为什么收益和模型结构强相关编译器优化的是“计算图本身”如果图里是大块大块的 GEMM 调用它能做的事非常有限如果图里是几十个小算子交织在一起融合空间才大。这也是我为什么坚持说讨论 torch.compile 之前先讨论你的模型长什么样。1.2 三种 mode 到底差在哪torch.compile 的 mode 参数经常被忽略但它直接影响收益方向和副作用。PyTorch 官方提供三种常用模式default、reduce-overhead、max-autotune。default 模式编译时间较短做基本的算子融合和代码生成适合先跑通验证收益。reduce-overhead 模式在 default 基础上引入 CUDA Graph 技术把一系列 kernel 启动录制成一个图减少 CPU 侧启动开销和 GPU 端的 launch 间隙。它最大的收益点就在碎算子密集的场景。因为碎算子本身执行时间短kernel 启动开销占比高CUDA Graph 能一次性灌入几十个 kernel 的启动信息把启动排队时间压下来。代价是显存占用上升且输入 shape、设备状态等必须固定否则会重新捕获。max-autotune 模式则会对算子做穷举式调优GEMM 会尝试 Triton 手写模板、cuBLAS、cutlass 等多种实现选出最快的一个。编译时间肉眼可见变长大模型上可能从几分钟变成几十分钟甚至更久。它并不是默认推荐的起点而是在确认收益后进一步“压榨”的手段。三种 mode 的取舍可以看这个表格mode编译耗时运行时优化重点典型副作用适用判断default较短算子融合、代码生成基本无大多数情况先跑这个reduce-overhead中等CUDA Graph压 launch overhead显存占用升高shape 固定碎算子多、launch 占比高max-autotune很长对算子做 autotuneGEMM 也会试 Triton/cutlass编译时间爆炸、显存上升、OOM 风险算子 shape 稳定且收益明确时再开我见过不少人用默认模式跑了一遍没加速就宣布 torch.compile 没用。实际上默认模式对很多碎算子的融合收益已经不错但对 launch overhead 的削减不如 reduce-overhead如果模型里碎算子太多建议直接对比 reduce-overhead 和 default而不是只用一个模式就下结论。2. 碎算子为什么能快一倍2.1 碎算子的性能瓶颈是搬运和排队“碎算子”是我在调优时常用的叫法泛指数量大、单个计算量小、以显存访问为主的算子典型代表是 elementwise 运算add、mul、gelu、silu、reduction 类softmax、layernorm、rmsnorm、mask 和 concat/split 操作。它们的名字你可能天天见但很少把它们当成一个整体来分析。碎算子性能有问题根源有两个。第一个是 kernel 启动开销。GPU 上一个算子对应一个 kernelCPU 要把 kernel 的所有参数、网格配置、寄存器分配等信息通过驱动下发到 GPU。这个开销大约在 3-10 微秒取决于驱动和上下文状态。碎算子的 kernel 本身执行时间往往只有几微秒到几十微秒启动开销占比相当可观。你做一次 100 次碎算子的 egel 循环光启动就吃掉一大块时间。第二个是显存带宽瓶颈。碎算子的算术强度很低——它没多少计算可做大部分时间花在把数据从显存读进来、算个简单结果、再写回显存。我常用一个生活类比特来解释碎算子像一辆卡车只装了几件货却要在多个仓库之间来回跑路上时间占大头GEMM 则是满载集装箱的干线运输能跑多快取决于发动机本身。拿最经典的y sigmoid(x)举例。eager 模式下PyTorch 实际会先调用一个 kernel 算出临时结果 tmp再调用另一个 kernel 把 tmp 转成 y。假设数据量是 D每个 kernel 都要经历一次读和一次写总访存量是 2D 2D 4D。如果编译器把它融合成一个 kernel只读一次 x、写一次 y总访存量变成 2D。在显存带宽就是瓶颈的前提下理论加速就是 2 倍左右再算上启动开销被省掉实测经常能到 1.5-2 倍。这就是“碎算子能快一倍”的第一层含义。2.2 融合是核心编译器为碎算子写一站式内核融合和代码生成是编译器在碎算子上真正发力的地方。它不负责提高单个算子内部的效率而是把多个算子的计算逻辑合并到一个 kernel 里。举个例子y tanh(w * gelu(x) b)。eager 执行需要四个 kernelgelu、mul、add、tanh中间张量在显存里被反复写读。torch.compile 会把它变成一个 Triton kernel一次性读出 x在寄存器里依次完成 gelu、乘 w、加 b、tanh最后写回 y。中间过程完全不用碰显存省掉的不只是访存量还有 kernel 启动的次数。reduction 类算子也能融合。LayerNorm、RMSNorm 这类算子需要先算均值/方差再做归一化。朴素实现需要至少两个 kernel一个做统计一个做归一化。编译器可以生成一个 kernel用 split 处理或缓存部分结果把统计和归一化合并。更极端的情况如果 LayerNorm 后面还接一个 elementwise 激活从一个 kernel 变成三个 kernel 再变回一个 kernel 都有可能。我最近在一台 A100 上对一个普通的 transformer block 做测试不开 flash attention、纯粹 PyTorch eager 实现时attention 内部的计算路径里有十几个小 kernelQK^T、scale、mask_fill、softmax、dropout、V每个都很短。torch.compile 会把 mask、scale、softmax 甚至 dropout 融合进 attention 计算中Kernel 数量从十几个降到四五个wall time 下降非常明显。这种收益在低 batch、小 hidden size 的模型上尤其明显因为这时候 GEMM 本身就不算大碎算子的占比反而高。2.3 一个可以直接跑的 benchmark 模板为了量化收益我建议每个团队都维护一个固定的 benchmark 脚本对比 eager、default、reduce-overhead、max-autotune 四种配置在同一批算子上的表现。下面是一个简化的模板可以跑碎算子和 GEMM 两类典型负载。import time import torch import torch.nn as nn def bench(fn, *args, warmup20, iters200): # 预热 for _ in range(warmup): fn(*args) torch.cuda.synchronize() start time.perf_counter() for _ in range(iters): fn(*args) torch.cuda.synchronize() return (time.perf_counter() - start) / iters * 1000 # ms # 碎算子elementwise 链 layer norm 混合 def scattered_ops(x, w, b): h torch.nn.functional.gelu(x) h h * w b h torch.nn.functional.layer_norm(h, h.shape[-1:]) return torch.tanh(h) # GEMM普通大矩阵乘 def gemm_only(a, b): return a b x torch.randn(4096, 4096, devicecuda, dtypetorch.float16) w torch.randn(4096, devicecuda, dtypetorch.float16) b torch.randn(4096, devicecuda, dtypetorch.float16) a torch.randn(4096, 4096, devicecuda, dtypetorch.float16) c torch.randn(4096, 4096, devicecuda, dtypetorch.float16) compiled_default torch.compile(scattered_ops, modedefault) compiled_reduce torch.compile(scattered_ops, modereduce-overhead) compiled_max torch.compile(scattered_ops, modemax-autotune) compiled_gemm_def torch.compile(gemm_only, modedefault) compiled_gemm_max torch.compile(gemm_only, modemax-autotune) print(scattered eager:, bench(scattered_ops, x, w, b)) print(scattered default:, bench(compiled_default, x, w, b)) print(scattered reduce:, bench(compiled_reduce, x, w, b)) print(scattered max:, bench(compiled_max, x, w, b)) print(gemm eager:, bench(gemm_only, a, c)) print(gemm default:, bench(compiled_gemm_def, a, c)) print(gemm max:, bench(compiled_gemm_max, a, c))这个脚本在我的测试环境里碎算子部分 default 能到 1.4-1.7 倍收益reduce-overhead 能到 1.6-2.0 倍左右max-autotune 相比 reduce-overhead 的额外提升有限因为瓶颈已经从计算变成了访存。GEMM 部分则明显不同三种模式基本都在 0.95-1.05 倍之间晃悠有些 shape 上甚至比 eager 慢 5%。需要注意具体的加速倍数和 GPU 型号、数据 shape、库版本强相关不要拿着固定数字到处套用。你真正要关注的是趋势碎算子的收益来自访存次数减少和启动开销削减规律稳定GEMM 的收益非常有限规律也稳定。3. GEMM 为什么基本白搭3.1 先看 GEMM 的优化天花板cuBLAS 有多卷GEMM 指的是通用矩阵乘法C A B。它是深度学习计算量的绝对大头也是 GPU 厂商性能调优的核心目标。NVIDIA 的 cuBLAS、cuBLASLtAMD 的 rocBLAS以及开源的 cutlass都在 GEMM 上倾注了巨大的工程资源。cuBLAS 的 GEMM kernel 做了什么它把输出矩阵切成多个 tile分发到不同的线程块每个线程块把数据分块搬进 shared memory使用 register blocking 让寄存器里的数据被反复复用加载下一步数据的同时当前 step 的计算已经在流水线上执行也就是 double buffering / pipeline对于某些特定维度还会用 split-K 来并行化 K 方向。这些优化产物的效果是对于能容纳进 L2 缓存的大矩阵 GEMMcuBLAS 的实现可以达到设备峰值浮点能力的 90% 甚至更高。在这种接近物理极限的水平下留给编译器的优化空间几乎为零。你可以把 cuBLAS 理解成钛合金做的高速公路torch.compile 再聪明也不太可能凭空再造一条更快的路。3.2 Inductor 对 GEMM 的真实策略默认调库autotune 才碰 TritonTorchInductor 对 GEMM 的处理策略和它对碎算子的处理策略完全不一样。在 default 模式或 reduce-overhead 模式里Inductor 对torch.mm、torch.bmm、nn.Linear这类大算子绝大多数情况下会直接调用底层库也就是继续走 cuBLAS。它在这些模式下对 GEMM 几乎不做内核级改造只做上层的图优化。这就解释了为什么你在默认模式下跑一个大 GEMM 的模型torch.compile 的收益非常小——因为它根本没去写新的 GEMM kernel。到了 max-autotune 模式Inductor 会尝试用 Triton 手写 GEMM kernel。但这里有个很现实的问题Triton 生成的 GEMM 要超过 cuBLAS需要 autotune 找到特别合适的 tile 尺寸、block 配置和流水线深度。对标准 shape 的方阵乘法Triton 很难赢 cuBLAS对非标准 shape、宽而扁的矩阵、某些特殊数据类型偶尔能领先一点点但幅度通常很小而且编译时间明显增加。所以结论很直接如果你模型的核心是几个大矩阵乘法torch.compile 大概率白搭。它不是没干活是这条路上前人已经修得太好了。3.3 GEMM 在什么场景还能“捡点漏”说“基本白搭”不等于完全无用GEMM 集成进编译图之后还有一些边角收益值得注意。第一种是 epilogue 融合。GEMM 后面往往跟着 bias add、激活函数、LayerNorm 之类。常规做法是 GEMM 输出写到显存下一个算子再读出来算。torch.compile 可以把这些 epilogue 操作融进 GEMM kernel减少中间张量的写读。cuBLASLt 本身也提供类似接口但对用户来说直接用 torch.compile 就能享受到这是实打实的好处。第二种是非常规 shape 或特殊数据类型的 GEMM。比如 M 很小例如只有 1-32 行、K 很短的线性层cuBLAS 的表现未必最优Triton 生成的定制 kernel 反而有一定空间。FP8、稀疏 GEMM 等新特性也给了编译器一点发挥余地。但这些都是“场景性捡漏”不能作为通用结论推广到所有模型。判断的关键还是那一条你的模型里 GEMM 的时间占比有多高。如果 60% 以上的 GPU 时间都在 GEMM 上那 torch.compile 就算帮你在剩余 40% 里省一点整体加速也不会超过 10%。这就是我在开头同事那个例子里看到的情况。4. 动手前先做“算力画像”判断模型该不该开 torch.compile4.1 两个便宜有效的画像指标我在给任何模型做优化前都要先做一次“算力画像”核心就是搞明白两类算子的时间占比。方法很简单跑一轮 profiling。第一种方式是用 PyTorch 自带的 profilerimport torch from torch.profiler import profile, ProfilerActivity with profile(activities[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: model(inputs) torch.cuda.synchronize() print(prof.key_averages().table( sort_bycuda_time_total, row_limit30))关注两个指标。第一个是 GEMM/conv 等大算子的cuda_time_total在总时间中的占比。你可以把mm、bmm、conv2d、addmm这类算子列出来它们通常和 cuBLAS/cuDNN 关联。如果这些算子加起来超过总 GPU 时间的一半torch.compile 的全局收益天花板就非常低这就是 Amdahl 定律在 GPU 优化里的体现。第二个指标是 kernel 数量和单 kernel 平均耗时。如果一个模型有几百上千个 kernel单个 kernel 平均只有十几微秒那它大概率是碎算子主导torch.compile、CUDA Graph 这类技术会非常有用。反之如果 kernel 数量少、每个 kernel 动辄几十上百微秒说明计算主体已经是大算子编译收益有限。我建议把 profiling 的结果做成一张表在模型上线前就存档。遇到“torch.compile 没用”的反馈时回头翻一下这张表很快就能定位问题。4.2 动态 shape、控制流、CUDA graph 的暗坑就算模型是碎算子主导torch.compile 也还有几个隐藏陷阱最容易踩的是动态 shape。torch.compile 编译出来的 kernel 通常针对特定 shape 做了 tile 选择和循环展开。如果模型在推理或训练过程中经常改变输入 shape编译器每遇到一个新 shape 就可能重新编译一次导致性能忽好忽坏。PyTorch 提供dynamicTrue参数来尽量复用已编译的 shape bucket但并不能完全避免重新编译。经验是如果你的业务数据 shape 非常不稳定先评估一下编译收益能不能覆盖重编译的代价。第二个坑是 graph break。模型里只要有 Python 控制流、动态索引、某些不支持的三方算子_dynamo 就会把图打断成多段。每一段单独编译段与段之间要回到 eager 模式优化效果大打折扣。排查方法是设置torch._dynamo.config.log_level或直接看TORCH_LOGScompilation日志里的GraphBreak信息。如果模型里 graph break 太多建议先修代码把控制流改成静态可追踪的形式。第三个坑是 CUDA graph 的副作用。reduce-overhead 模式会使用 CUDA Graph图形显存占用会明显上升而且在图捕获期间不允许改变 shape。有些模型开了 reduce-overhead 后直接 OOM不是编译本身出错而是 CUDA graph 对显存做了额外预留。这时候要么退回 default要么缩小编译范围。4.3 实操建议先子模块再全模型我的建议是不要一上来就对整个大模型开 torch.compile。大模型编译时间久、显存开销大、排查困难出了问题时很难定位是哪一个子图拖了后腿。更稳妥的做法是先做子模块级实验。把模型拆成几个部分比如 transformer 的 attention 块、FFN 块、embedding 处理等先对单个子模块做 torch.compile对比收益和编译时间。通常你会发现收益最大的子模块往往是那些包含大量 elementwise 和 reduction 的模块而不是纯 GEMM 的部分。确定收益后再逐步扩展编译范围直到整体达到目标。还要学会“放过”某些地方如果某个子模块 graph break 很多、收益很小可以直接用torch.compiler.disable给它关掉让它继续走 eager逼着编译器把精力留在更有价值的地方。我自己在实际项目里的标准流程是先开 default 模式跑通用 profiler 确认 kernel 数量是否有明显下降如果仍有大量碎算子再换 reduce-overhead 模式压 launch overhead最后才考虑 max-autotune并且只对收益明显的子图开启避免全局编译时间爆炸。5. 常见问题与排查技巧实录5.1 一张速查表解决大部分 torch.compile “疑难杂症”这几年下来我在各种项目里攒了不少 torch.compile 使用中的典型问题整理成一张速查表方便你排查时对照。现象可能原因排查方向编译一次特别久autotune 范围太大、图太复杂先退回 default限制编译范围设置缓存目录开 compile 后显存涨、OOMCUDA graph 预留显存 / autotune buffer 过多换 mode缩小编译范围关 reduce-overhead首次调用很慢编译冷启动 / CUDA graph capture用真实输入预热一次持久化编译缓存动态 shape 导致反复变慢每个新 shape 重新编译设置 dynamicTrue用定长 padding 输入精度和 eager 结果有差异Triton kernel 调度/reduce 顺序不同量化误差关键算子跳过 compilegraph break 太多收益低控制流、不支持算子打断图打印 TORCH_LOGS重写控制流保留 eager 路径DDP 多卡训练没提升all-reduce 无法融入编译图确认通信占比考虑打通通信计算重叠优化这里面有两个我特别想强调的点。第一是关于编译缓存。torch.compile 的编译过程是可以缓存的默认路径是$TORCHINDUCTOR_CACHE_DIR或系统临时目录。如果你在做超参实验同一个模型结构反复编译一定要把缓存目录落到持久化存储上否则每次重启进程都重新编译一遍白白浪费几十分钟。第二是关于精度。Triton 生成的 kernel 在很多情况下会改变归约顺序尤其 fp16/bf16 的情况下LayerNorm、Softmax 的结果会有微小差异。大多数场景下不影响训练收敛但如果你在意严格一致可以在关键算子处不开编译用原始 eager 路径跑。我在部署推理模型时曾经遇到过 fp16 下某个 layer 的误差积累最后排查到就是 torch.compile 对某个 reduction 的改写导致后来手动给那个 layer 单独关掉才解决。5.2 排障时我常用的调试手段排查 torch.compile 问题我常用的工具不多但都很管用。第一个是TORCH_LOGScompilation。这个环境变量会打印编译过程中的图捕获、图优化、代码生成信息是定位 graph break 和编译流程问题的首选。你可以在日志里看到每个子图为什么被打断哪个算子不支持编译一目了然。第二个是torch._inductor.config.trace.enabled True。打开它之后Inductor 会把编译过程中的中间产物写到指定目录包括生成的 Triton 代码。你完全可以打开生成的.py文件看看编译器到底为你的算子写了一份什么样的代码。我经常这样确认某个算子到底是被融合了还是直接调库了。前面提到 GEMM 默认走 cuBLAS也是通过这种途径确认的。第三个是nsight computencu配合生成的 Triton kernel。如果你想深入看某个编译后的 kernel 的性能指标可以用 ncu 对单个 kernel 测 occupancy、memory throughput、achieved peak FLOPs。这不是每个调优场景都需要但碰到那种“融合了但没变快”的疑难问题时能帮你确认瓶颈到底是不是还卡在带宽上。第四个技巧是缩小问题范围。遇到编译后行为异常我会把torch.compile从整个模型缩小到一个 block、一个 layer、甚至一个算子逐步二分定位。很多时候问题不在编译器而在你的代码里某个异常分支被编译器“优化”掉了或者某个自定义 autograd.Function 和编译器发生了冲突。我还习惯在代码里加一个环境变量开关方便随时切换 eager 和编译模式。比如在训练脚本里用一个--torch-compile参数默认是关闭的这样每次实验都能快速对比开启前后的收益而不需要改代码重新部署。写在最后的一点个人体会做 GPU 性能调优这几年我越来越觉得工具本身的“光环”会误导人。torch.compile 刚出来的时候大家都喊着“以后不用手动融合算子了”真用一段时间才发现它是一把好用的刀但用在哪、怎么用仍然需要你对手里的工件足够熟悉。我现在的做法已经固定成一套流程拿到一个模型先做算力画像看碎算子和 GEMM 的占比再开 default 模式跑一遍确认编译产物里有足够多的新 kernel然后决定要不要上 reduce-overhead 压 launch overhead以及要不要对特定子图开 max-autotune。最后永远保留一个可切换回 eager 的开关防止编译链路在某个版本升级后突然出问题。还有一个心得是编译优化不是银弹但它也不是毫无用处的装饰。它真正厉害的战场是那些“乱糟糟”的碎算子——把十几次显存搬运压缩成一次把几十次 kernel 启动折叠成几个。理解了这一点你就能在任何新模型上快速判断torch.compile 值不值得为它折腾一场。