算子融合到底是什么?不换硬件不改模型,只改一张“计算图“就能快 43%
动手跑过才敢写。本文所有数字都来自我自己的环境里亲手跑出来的真实输出没有一个是估算的。环境torch 2.10.0cu128/onnx 1.21.0/onnxruntime 1.26.0CPU 推理5060Ti i5 14600k。引言为什么同一张图改两笔就变快了做部署的人常有这种感觉同一个模型在 PyTorch 里跑挺慢扔给 ONNX Runtime / TensorRT 就唰地快起来。很多人以为是格式换得好其实真正的大头是算子融合——推理引擎在加载时偷偷把计算图搓了一遍。这篇文章回答三个问题算子融合到底在融什么—— 把会挨个执行的算子捏成一个。融完之后图变成什么样—— 亲手把融合前后的计算图导出来看。到底快了多少—— 用同一台机器、同一个引擎开关图优化实测对比。第一章、先懂一个道理GPU 最怕来回搬砖类比一下端菜 vs 一锅炖想象食堂后厨切菜、焯水、炒、调味是四个独立岗位。算子不融合时流程是这样的切菜师傅切好 → 端出去放架子上 焯水师傅端进来焯水 → 再端出去放架子上 炒菜师傅端进来炒 → 再端出去放架子上 调味师傅端进来调味 → 端出去装盘每一道中间结果都要从内存里写出去、再读回来。对 GPU 来说这个端进端出就是中间张量在显存里的反复读写是真正的性能黑洞——因为 GPU 算得快但搬数据内存带宽永远比算得慢。算子融合之后改成这样一个师傅切 → 焯 → 炒 → 调一气呵成 中间结果不出灶台最后才端上桌中间结果不落地省掉了大量内存读写还少启动了好几次 kernel。这就是融合的全部秘密。一句话融合 把端进端出的中间张量省掉把多次 kernel 启动合并成一次。第二章、融合前导出图长什么样我搭了一个 10 段Conv BatchNorm ReLU堆叠的小网络3244 万参数输入1×3×224×224导出成 ONNXimporttorch,torch.nnasnnclassDeepFuseNet(nn.Module):def__init__(self,ch64):super().__init__()blocks,cur[],3for_inrange(10):blocks[nn.Conv2d(cur,ch,3,padding1),nn.BatchNorm2d(ch),nn.ReLU()]curch self.bodynn.Sequential(*blocks)self.fcnn.Linear(ch*224*224,10)defforward(self,x):xself.body(x)xx.view(x.size(0),-1)returnself.fc(x)modelDeepFuseNet().eval()torch.onnx.export(model,torch.randn(1,3,224,224),deepfuse.onnx,input_names[input],output_names[output],opset_version17)我环境里真实打印的算子分布参数量 : 32448074 导出图节点总数: 22 算子分布: {Conv: 10, Relu: 10, Reshape: 1, Gemm: 1}注意一个细节模型明明有 10 个 BatchNorm但导出图里一个 BatchNorm 节点都没有。因为 BN 在推理时只是一组 scale/shift早在torch.onnx.export阶段就被预先熔进了 Conv 的权重里。这是导出时融合是融合的第一次发生。剩下 22 个节点里10 组Conv → Relu是最经典的融合对象——它们首尾相接中间的 Relu 结果完全可以不出内存。第三章、融合后图真的瘦了一圈ONNX Runtime 在加载时会做图优化。我用一个optimized_model_filepath把优化后的图导出来和原始图逐节点对比importonnx,onnxruntimeasortfromcollectionsimportCounter soort.SessionOptions()so.graph_optimization_levelort.GraphOptimizationLevel.ORT_ENABLE_ALL so.optimized_model_filepathdeepfuse_optimized.onnx# 让 ORT 把优化结果写出来sessort.InferenceSession(deepfuse.onnx,so,providers[CPUExecutionProvider])raw_opsdict(Counter(n.op_typeforninonnx.load(deepfuse.onnx).graph.node))opt_opsdict(Counter(n.op_typeforninonnx.load(deepfuse_optimized.onnx).graph.node))print(融合前:,len(raw_ops),raw_ops)print(融合后:,len(opt_ops),opt_ops)我环境里的真实输出融合前 节点总数 22 : {Conv: 10, Relu: 10, Reshape: 1, Gemm: 1} 融合后 节点总数 13 : {Conv: 10, ReorderOutput: 1, Reshape: 1, Gemm: 1} 节点减少: 910 个 Relu 全部消失了。它们被熔进了前面的 Conv变成了FusedConvConvRelu 合并算子。节点数从 22 掉到 13少掉 9 个那 10 个 Relu 全没了只新增 1 个布局转换ReorderOutput。换句话说推理时GPU 不再需要单独算 10 次 Relu也不用把每次 Conv 的中间结果写回内存再读出来给 Relu 用。一次算完中间结果留在寄存器/高速缓存里直接给下一步。第四章、实测开图优化到底快多少同一个 ONNX 文件、同一个 ORT 引擎、同一台机器只把graph_optimization_level从关闭切到开启CPU 上跑 100 次取平均先 warmupso_offort.SessionOptions();so_off.graph_optimization_levelort.GraphOptimizationLevel.ORT_DISABLE_ALL sess_offort.InferenceSession(deepfuse.onnx,so_off,providers[CPUExecutionProvider])so_onort.SessionOptions();so_on.graph_optimization_levelort.GraphOptimizationLevel.ORT_ENABLE_ALL sess_onort.InferenceSession(deepfuse.onnx,so_on,providers[CPUExecutionProvider])我环境里的真实耗时PyTorch eager : 75.5183 ms ORT 图优化关闭(原始22节点) : 70.9388 ms ORT 图优化开启(算子融合) : 49.7349 ms 融合加速 (关 → 开) : 1.43x 整体 vs PyTorch eager : 1.52x 最大绝对误差 : 2.887e-08三个要点融合带来的提速是 1.43×70.94 ms → 49.73 ms。这完全是图优化搞的鬼跟模型、跟数据没有任何关系。结果几乎不变融合前后最大绝对误差只有2.9e-08浮点舍入级别精度无损。这个模型还不够大3244 万参数在 CPU 上融合挤掉的是中间张量读写和 kernel 启动。模型越大、算子越碎融合收益越明显——到了 TensorRT 那种把整张图极致重排的引擎收益往往能到几倍甚至十几倍。执行方式推理耗时相对融合后PyTorch eager75.52 ms×1.52ORT 图优化关闭70.94 ms×1.43ORT 图优化开启融合49.73 ms×1.00第五章、融合的几种常见套路除了ConvRelu深度学习里还有一堆天生该融的组合融合套路说明出现场景Conv BatchNormBN 推理时折叠进 Conv 权重导出时已做几乎所有 CNNConv Relu合并成 FusedConv中间结果不出内存CNN 主干Conv Add Relu残差块最经典三个捏成一个ResNet 等残差网络Elementwise 链一串逐元素算子Add/Mul/Relu合并Transformer 的归一化Gemm Add全连接 偏置合并分类头融合的本质遵循一个规律两个首尾相接的算子如果中间没有必须被别的算子也读到的分叉就能安全地合并。合并后少一次中间张量落地、少一次 kernel 启动。总结一张表读懂算子融合问题答案我环境里的验证方式融合在融什么把首尾相接的独立算子合并成一个10 组ConvRelu熔成FusedConv图变什么样节点数变少中间算子消失节点 22 → 13Relu 全消失为什么快省中间张量内存读写 少 kernel 启动实测 Relu 节点清零快了多少仅图优化一项就 1.43×70.94 → 49.73 ms精度有损吗没有误差在浮点舍入级最大绝对误差 2.9e-08什么时候最有用模型越大、算子越碎时收益越明显小模型 CPU 已见效最后一句大白话算子融合不是玄学它就是把每道工序都端进端出的笨办法改成一个师傅从头到尾不离灶台。省下的不是计算量本身而是内存搬砖和 kernel 启动这两笔隐形成本。这也是为什么 ONNX Runtime、TensorRT 这些引擎能白嫖加速——它们什么都没训练只是把计算图搓得更聪明了。