ARTICLE DETAIL

资讯详情

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

碎算子快一倍、GEMM白忙活?torch.compile加速原理与实测

碎算子快一倍、GEMM白忙活?torch.compile加速原理与实测 两年前我第一次认真把 torch.compile 接进生产推理管线时脑子里还是那句宣传语——“改一行代码白捡一两倍速度”。结果并不难看但也不惊艳一个以 GEMM 为主体的 7B 模型在 A100 上量出来只有 1.03 倍而同一个仓库里另一个算子很碎、充满逐元素操作的视觉模型直接飙到 1.9 倍。差别不在“用没用 torch.compile”而在于模型里到底以什么类型的 GPU 算子为主。这篇文章想把这件事彻底说透。标题里那句话不是夸张碎算子逐元素类、归约类、形状转换类在 torch.compile 下确实能近乎翻倍而 GEMM通用矩阵乘法基本就是白忙活。为什么因为两者在 GPU 上的瓶颈完全不同前者卡在启动开销和内存搬运后者早就被 cuBLAS/CUTLASS 卷到了硬件极限。下面我会把原理、实测、以及怎么判断自己的模型值不值得用 torch.compile 完整讲一遍。1. 别把 torch.compile 当“开关”它更像一把手术刀1.1 一次典型的“宣传级复现”失败我先说一个很多人都会遇到的场景。团队里看到 PyTorch 官方博客和各类推文都在吹 torch.compile 的加速比于是把一个经典 CV 模型和一个小型 LLM 放进 benchmark 脚本准备“白捡”性能。结果一连跑了三天数据大致如下模型场景eager 耗时torch.compile 耗时加速比ResNet-50推理2.81 ms1.72 ms1.63x小型 7B LLM推理静态形状15.7 ms15.2 ms1.03x小型 7B LLM推理动态形状16.4 ms15.9 ms1.03xResNet 那种每隔几层就有一堆 ReLU、BatchNorm、残差相加的模型收益肉眼可见可一旦模型主体是若干个连续大矩阵乘速度提升就趋近于零。当时旁边一个同事开玩笑说“这哪是开关这是把手术刀只割特定组织。”这个反差并不是 bug而是 torch.compile 的底层逻辑决定的。torch.compile 的核心思路是捕获整张计算图然后对图里的算子做融合fusion和代码生成。但对不同算子融合的空间天差地别。逐元素算子之间融合能把几十次内核启动和多次内存搬运压成一次而 GEMM 本身已经是高度精调过的计算内核你很难再“融合”出额外的带宽节省因为它的调用次数少、每次调用又已经打满了 GPU 的算力。1.2 “加速”不是模型的属性而是算子结构的属性很多人理解 torch.compile 时会搞反因果关系。他们以为“模型 A 跑起来变快了”实际上应该反过来看模型 A 的算子分布里内存受限类算子占了很大比例所以编译器的图优化才能撬动这么多收益。从 GPU 性能工程的角度看任何算子都可以粗分为两类计算受限compute-bound比如大矩阵乘 GEMM、大卷积。一次算子的计算量远大于读写量GPU 的 SM 单元长期处于满载状态。内存受限memory-bound比如逐元素乘加、激活函数、归一化、广播、concat、split。每个元素只需要极少的 FLOPs但必须从 HBM 读进来、算完再写回去整个算子的执行时间基本由访存速度决定。torch.compile 对第二类特别有效因为融合意味着减少中间结果的落盘和重新读取对第一类几乎没有办法因为计算密度已经很高编译器不可能凭空让 Tensor Core 的 FLOPS 翻倍。理解了这一点你就能预测一个模型在 torch.compile 下的大致加速范围而不是跑完才知道。后面第 4 节我会给出一套具体的 profile 判断流程。2. 碎算子快一倍的根本原因内存流量与启动开销被压缩2.1 一个逐元素链路在 Eager 模式下的“搬运账本”我用一个非常常见的残差块来拆输入 x 经过线性层、GELU、另一个线性层然后和残差相加最后乘一个 scale。写成公式是h x scale * gelu(linear2(dropout(linear1(x))))在 eager 模式下PyTorch 会一个算子一个算子地执行每一步都是独立的 kernel launchlinear1GEMM会非常快但输出张量需要写回显存dropout读取 linear1 的输出随机 mask再写回gelu读取 dropout 的输出逐元素计算再写回linear2GEMM读取 gelu 结果再写回multiply读取 linear2 输出逐元素乘 scale写回add读取 x 和上一步结果逐元素相加写回。如果中间张量大小是 1MB那么每一步的 HBM 流量大概是“读 1MB 写 1MB”。6 个算子加起来总共要搬运约 12MB 数据。而如果编译器能把 dropout、gelu、乘 scale、加残差这些逐元素操作融合成一个内核数据只需要从内存读一次、最后写一次HBM 流量直接降到 2MB 左右。光这一项就有 5-6 倍的流量削减空间。实际加速比当然达不到 6 倍因为还有 GEMM 部分的访存和 kernel launch 开销。但方向很明确访存流量下降是碎算子加速的最主要来源。2.2 融合内核到底做了什么“一倍”是怎么来的为什么标题说“能快一倍”而不是三倍、五倍这要从 GPU 执行一个逐元素链路的完整时间构成来看。假设 A100 的显存带宽约 2TB/s那么一个 4MB 的中间张量单次读写的理论时间大概是读 4MB 写 4MB 8MB 8MB / 2000GB/s 4 微秒看到没纯数据搬运其实非常快。但 eager 模式一个 kernel 从 CPU 端发起、进入 GPU 队列、执行、返回整个启动和调度开销通常在 5-20 微秒。于是当张量不大、单个算子执行时间只有几微秒时启动开销反而成了主导。这也是为什么在很多小 batch、短序列场景下eager 模式的 GPU 利用率经常只有 30%-50%大量时间花在排队和等待上。torch.compile 的融合内核把几十次启动压缩成一次同时降低了 HBM 流量。两个收益加在一起实际效果就落在 1.8-2.1 倍这个区间。低于 2 倍是常态因为融合后的内核仍然要读一次、写一次而且中间可能会有寄存器溢出或缓存不友好的访问模式。2.3 除了融合reduce-overhead 还在哪部分出力很多人忽略了一个细节torch.compile(model, modereduce-overhead)不只在做算子融合还会把整个计算图捕获成 CUDA Graph。CUDA Graph 能把 kernel launch 的 CPU 开销从“每个 kernel 几微秒”降到“整个 graph 一次提交”本质上是在消除 CPU 端的 launch 瓶颈。这对碎算子模型的帮助尤其大因为逐元素算子数量多、单次执行时间短CPU launch 开销占比极高。一旦 graph 化启动开销几乎归零。GEMM 模型则没那么敏感——一次 GEMM 本身运行几十上百微秒你省掉那 5 微秒 launch 时间占比只有几个百分点。实测中你会发现对算子密集模型reduce-overhead往往比默认模式多出 5%-10% 的收益对 GEMM 密集模型这个差异几乎看不见。这正是“碎算子能快一倍”的第二个支柱。3. GEMM 白搭不是 bugcuBLAS/CUTLASS 早已逼近硬件极限3.1 GEMM 在 eager 下已经是“精装修”要理解 torch.compile 对 GEMM 无能为力你先得知道 PyTorch eager 底下的 GEMM 是谁在跑。默认情况下PyTorch 的torch.mm、torch.matmul、nn.Linear都会落到 NVIDIA 的 cuBLAS 库在 Ampere 之后很多路径还会走 CUTLASS 生成的模板内核。这些库经过 NVIDIA 十几年的持续调优对每个 shape 都做了一堆 autotune内部可能有几十种不同 tile 大小、流水线深度、寄存器分块策略最后选出最合适的一个。说得直白一点cuBLAS 里的一个 GEMM 内核本身的优化程度已经不亚于甚至高于大部分深度学习团队手工写的 kernel。torch.compile 就算再把图结构扫描一遍也不可能让这个内核再快一倍。它只能做到“和 cuBLAS 差不多快”运气好一点快几个百分点运气差一点反而更慢。3.2 torch.compile 在 GEMM 上实际在做什么你可能会问那 Inductortorch.compile 的代码生成后端看到 GEMM 节点时到底会做什么答案是分情况。小 GEMM 或形状比较奇怪的 GEMMInductor 会用 Triton 生成一个自定义的 GEMM kernel形状规整的大 GEMM它可能直接调 cuBLAS在某些版本和配置下它还会尝试 CUTLASS 的 path。问题是Triton 自动生成的 GEMM 往往只用了很基本的 tile 划分和流水线策略对很多 shape 来说能匹敌 cuBLAS 已经很不错超过的案例少之又少。我在 A100 上测过一组[4096, 4096] x [4096, 4096]的 GEMMeager 和 torch.compile 分别是 0.92 ms 和 0.93 ms基本持平。换到[64, 4096] x [4096, 4096]这种偏瘦的 GEMMtorch.compile 反而有 4% 左右的下降。原因也很简单这种 shape 下 cuBLAS 可以选用更优的 split-k 或者针对小 M 的 kernelTriton 自动调优没那么精细。3.3 少数“GEMM 变快”的例外值得说清楚GEMM 并不是在任何情况下都白搭有几个例外值得注意免得你把结论理解得太绝对。第一个例外是GEMM 大量逐元素 epilogue 融合。比如linear bias gelueager 需要 GEMM 写回结果再读出来做 bias 加和 gelutorch.compile 可以把 epilogue 直接并到 GEMM kernel 里省掉一次读写。这在计算密集型大 GEMM 上优势不明显但在带宽受限的 decode 场景、小 batch 场景收益还是比较可观的。第二个例外是小 GEMM。当矩阵尺寸小到 kernel 本身只跑几微秒时启动开销和调度开销又成了大头torch.compile 通过算子融合和 CUDA Graph 能贴掉这部分成本。所以你会看到许多 LLM 推理框架里torch.compile 对 decode 阶段batch size 小、GEMM 规模小有一些帮助但对 prefill 阶段batch size 大、GEMM 规模大帮助微乎其微。第三个例外是它帮你找到融合边界之外的优化机会。有时候 torch.compile 的自动调优会选择把连续两个小 GEMM 合并成一个大 GEMM或者调整内存布局让后续算子更连续。这种情况不是常规路径但确实存在。我可以负责任地说这类收益不稳定需要逐模型验证不能当成通用预期。4. 拿 profile 说话你的模型是“算子密集”还是“GEMM 密集”4.1 快速分析模式在决定要不要上 torch.compile 之前先花十分钟做一次 profiling比任何经验法则都靠得住。我常用的方法是用 PyTorch 自带的 profilerimport torch from torch.profiler import profile, ProfilerActivity model.eval() x torch.randn(1, 3, 224, 224, devicecuda) with profile(activities[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: for _ in range(10): model(x) torch.cuda.synchronize() print(prof.key_averages().table( sort_bycuda_time_total, row_limit30, top_level_events_onlyTrue ))拿到表格后把 CUDA time 占比最高的算子分个类xxx_elementwise、xxx_reduce、xxx_gemm、convolution。如果 Elementwise 类的时间加起来能占到总 CUDA time 的 30%-50%torch.compile 很值得试如果 80% 以上时间都堆在 GEMM 和卷积上那基本可以预期它在 1.0x-1.1x 之间波动。4.2 推理和训练要分开看推理和训练对算子分布的敏感度很不一样我踩过不少坑这里分开说。推理阶段特别是小 batch 或动态形状场景kernel launch 开销和内存搬运占比远高于大 batch 场景。所以同样的模型torch.compile 在小 batch 推理下更容易看到收益。但要注意如果你在服务端使用了自定义的推理引擎例如 TensorRT、vLLM底层早就做了 CUDA Graph 和算子融合torch.compile 再叠上去几乎没有空间甚至可能干扰已有优化。训练阶段反向传播里同样有大量逐元素梯度算子比如 ReLU 的 mask、LayerNorm 的梯度、Dropout 的 mask 逻辑。这些在 eager 模式下都是一堆细碎内核torch.compile 把它们融合之后收益通常比前向更明显。代价是编译时间更长、显存占用可能更高因为编译器可能为了融合缓存中间结果而保留更多信息。4.3 一个简单的量化判断标准如果你想用更数字化的方式判断可以把 profiling 结果里的算子时间按下面这个权重模型估算总收益估算 0.4 * elementwise占比 0.2 * reduce占比 - 0.1 * gemm占比这里系数是基于我自己的实测经验拍的不是严格数学公式。大体趋势是逐元素占比高收益大归约类算子softmax、LayerNorm有一定收益但不如逐元素夸张GEMM 占比高收益趋近于零甚至为负。算出来的值如果低于 0.1基本可以放弃 torch.compile转去优化 kernel 本身的实现。还有个小技巧直接用torch._dynamo.config.report看编译过程中有哪些图破坏了优化。很多模型因为动态 shape、自定义 autograd.Function、控制流复杂Dynamo 会把图切成很多小块优化效果大打折扣。如果日志显示大量 “graph break”那你根本不用等 benchmark基本可以判断收益有限。5. 实测几个有代表性模型的加速差异5.1 视觉和逐元素密集模型真的能接近两倍我拿一个经典的检测模型和一个小型扩散模型做过测试它们的共同点是包含大量卷积、归一化、激活、上采样、concat。这类模型在 eager 模式下 kernel 数量非常多尤其是上采样和 concat 附近常常会有几十个细小算子排队执行。实测结果如下A100-80GPyTorch 2.3CUDA 12.2batch size 1模型eager 推理时间torch.compile 推理时间加速比ResNet-502.81 ms1.72 ms1.63x小型 U-Net 扩散模型4.64 ms2.51 ms1.85x轻量检测模型6.32 ms3.58 ms1.76x这几个模型里最夸张的是 U-Net 类结构因为它的通道数大、feature map 多逐元素算子的访存量极高融合收益自然就大。我之前说过“能快一倍”在这种模型上是真的能实现甚至有一些框架配合低精度和内存布局优化能到 2 倍以上。值得一提的坑视觉模型如果用了大量动态 shape比如输入尺寸不固定torch.compile 会因为重编译频繁而出现很大的开销。实际使用中最好固定输入分辨率或者给 Dynamo 配置dynamicFalse的 shape 假设避免每来一个新 shape 就重新编译一次。5.2 大语言模型 decode 段GEMM 白搭的主线再来看一个典型的 transformer decoder-only 模型。prefill 阶段输入序列长QKV 投影、attention、FFN 都是大规模的 batch GEMMCUDA time 几乎全在 GEMM 和 FlashAttention 上。torch.compile 在这个阶段的收益非常有限经常只有 1%-3%。有时候因为 Inductor 对 attention 图做了额外剖分反而会引入几个额外的 kernel实测出现 1%-2% 的倒退。decode 阶段更有趣一点每步只生成一个 tokenbatch size 通常很小比如 1 或 8这时候每个 GEMM 的 M 维度很小kernel 本身跑得很快。碎算子比如 RoPE、激活、LayerNorm、残差相加在整步耗时里的占比一下子提高了。我在一个 7B 模型上测得 decode 阶段 torch.compile 能带来 1.06-1.12 倍的收益主要来自 Attention 之外的逐元素融合和 CUDA Graph 消除 launch 开销。但注意这只是相对 eager 而言。如果你的服务端已经用了 vLLM 或 TensorRT-LLM它们已经做了很多同样的事torch.compile 在这里基本没有叠加价值反而可能因为 graph capture 限制了动态 shape 处理增加工程复杂度。5.3 中间地带融合算子已经很多收益会被摊薄还有一种情况是你模型里使用了比较“重型”的融合 kernel比如 FlashAttention、FusedAdam、FusedLayerNorm。这些算子内部已经完成了大量融合外部看起来“碎算子”数量减少了很多torch.compile 能融合的对象自然就少了。我实测过一个用 FlashAttention 的 GPT 模型torch.compile 前后只有 1.02 倍同一个模型把 FlashAttention 换回普通 attentiontorch.compile 能达到 1.15 倍。原因很直白**好 kernel 越多编译器可发挥的余地越少。**这不是坏事说明你的模型本身就处于一个比较优化的状态没必要为了 torch.compile 而 torch.compile。6. 实操时更该做的几件事以及真实效益盘点6.1 跑 benchmark 最容易犯的错误很多人反映 torch.compile 测试结果不稳定这通常不是编译器波动而是 benchmark 方法不对。我说几个最常见的坑没有 warmuptorch.compile 第一次运行时包含编译和 autotune时间可能非常长。必须先把模型跑几轮让 graph capture 和 kernel 选择全部完成再开始计时。没有固定随机种子和输入某些模型包含随机性Dropout 等会导致前后耗时波动。实测时建议把输入固定或者在 eval 模式下禁用随机层。没有同步PyTorch 的算子默认是异步的如果直接time.perf_counter()包住整个推理循环你测的可能只是 CPU 端的入队时间。记得在计时区间末尾加torch.cuda.synchronize()。没有控制动态 shape如果输入 shape 在 benchmark 中变化Dynamo 会不停 recompile把大量开销混进测试结果。要么固定 shape要么显式标记动态维度。6.2 比 torch.compile 更值得做的几件事我在多个项目里的体会是torch.compile 只是性能工程的一部分而且未必是最关键的部分。如果你已经意识到自己的模型是 GEMM 密集与其纠结torch.compile的 1.03 倍不如把时间花在下面这几个方向上把逐元素算子和归约算子手动融合或者在模型设计层面减少碎算子。比如把连续的 LayerNorm、残差、激活合并到自定义 CUDA kernel 里很多时候能拿到和 torch.compile 一样的收益而且行为更可控。处理内存布局和静态 shape。很多时候模型速度上不去瓶颈在 tensor 的 stride 不连续、padding 过多、频繁的 transpose 和 view。让张量保持连续内存布局往往比开编译器更管用。对大 GEMM 使用底层优化库或专用 kernel。比如 FlashAttention 这类融合注意力、业界调好过的 persistent GEMM kernel收益经常远超 torch.compile。用 CUDA Graph 手动捕获推理图。torch.compile 的 reduce-overhead 本质就是帮你做了 CUDA Graph capture但手动做可以更精细地控制哪些部分要 graph、哪些部分要动态处理。我在实际项目里见过一个很有意思的对比同样的视觉模型torch.compile 拿到 1.6 倍后来有人用 TensorRT 做了一遍直接 2.2 倍。这并不意味着 TensorRT 一定比 torch.compile 强而是说不同的优化层次可以叠加选对工具比用一个大而全的开关更重要。6.3 把预期放到正确的位置最后给一张我经验总结的速查表方便你做技术选型时快速判断模型/场景特征torch.compile 预期收益建议动作大量逐元素算子、小张量、动态分支少1.5x - 2.0x直接用优先reduce-overhead大量归约算子、中等张量1.2x - 1.4x可以试但别期待太高大 GEMM、长序列 prefill1.0x - 1.05x考虑底层库和显存优化已用 FlashAttention、FusedLN 等重型融合 kernel1.0x - 1.1x重点查剩余 kernel 的启动开销动态 shape、图 break 频繁可能负优化先修 Dynamo 捕获问题再谈编译做性能工程这几年我最深的体会是任何一个优化工具都有它的适用边界。torch.compile 在当前 PyTorch 生态里确实是一个非常好用的工程化手段但你得先搞清楚它的收益来源才能决定该不该用、怎么用。它的核心贡献在于帮开发者自动完成了大量内核融合和启动开销优化而它对 GEMM 类算子的无能为力恰恰说明那部分优化应该交给更专业的 kernel 库。如果你现在正准备给模型上 torch.compile我建议先花半天时间把 profile 跑清楚按第 4 节的方法算一遍算子分布。如果你的模型是碎算子为主放心上收益大概率让你惊喜如果模型是 GEMM 密集也别失望那说明你真正该关注的是另一个层面的优化。搞 GPU 性能工程本来就是一步步找瓶颈、拆瓶颈的过程torch.compile 只是这一系列工具箱里的一件趁手工具不是终点。
返回列表