ARTICLE DETAIL

资讯详情

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

算子融合到底是什么?不换硬件不改模型,只改一张“计算图“就能快 43%

算子融合到底是什么?不换硬件不改模型,只改一张“计算图“就能快 43%

动手跑过才敢写。本文所有数字都来自我自己的环境里亲手跑出来的真实输出,没有一个是估算的。环境:torch 2.10.0+cu128/onnx 1.21.0/onnxruntime 1.26.0(CPU 推理)(5060Ti + i5 14600k)。

引言:为什么"同一张图改两笔"就变快了?

做部署的人常有这种感觉:同一个模型,在 PyTorch 里跑挺慢,扔给 ONNX Runtime / TensorRT 就唰地快起来。很多人以为是"格式换得好",其实真正的大头,是算子融合——推理引擎在加载时,偷偷把计算图"搓"了一遍。

这篇文章回答三个问题:

  1. 算子融合到底在融什么?—— 把会挨个执行的算子,捏成一个。
  2. 融完之后图变成什么样?—— 亲手把融合前后的计算图导出来看。
  3. 到底快了多少?—— 用同一台机器、同一个引擎,开关图优化实测对比。

第一章、先懂一个道理:GPU 最怕"来回搬砖"

类比一下:端菜 vs 一锅炖

想象食堂后厨,切菜、焯水、炒、调味是四个独立岗位。算子不融合时,流程是这样的:

切菜师傅:切好 → 端出去放架子上 焯水师傅:端进来焯水 → 再端出去放架子上 炒菜师傅:端进来炒 → 再端出去放架子上 调味师傅:端进来调味 → 端出去装盘

每一道中间结果,都要从内存里写出去、再读回来。对 GPU 来说,这个"端进端出"就是中间张量在显存里的反复读写,是真正的性能黑洞——因为 GPU 算得快,但搬数据(内存带宽)永远比算得慢

算子融合之后,改成这样:

一个师傅:切 → 焯 → 炒 → 调,一气呵成 中间结果不出灶台,最后才端上桌

中间结果不落地,省掉了大量内存读写,还少启动了好几次 kernel。这就是融合的全部秘密。

一句话:融合 = 把"端进端出"的中间张量省掉,把多次 kernel 启动合并成一次。


第二章、融合前:导出图长什么样?

我搭了一个 10 段Conv + BatchNorm + ReLU堆叠的小网络(3244 万参数),输入1×3×224×224,导出成 ONNX:

importtorch,torch.nnasnnclassDeepFuseNet(nn.Module):def__init__(self,ch=64):super().__init__()blocks,cur=[],3for_inrange(10):blocks+=[nn.Conv2d(cur,ch,3,padding=1),nn.BatchNorm2d(ch),nn.ReLU()]cur=ch self.body=nn.Sequential(*blocks)self.fc=nn.Linear(ch*224*224,10)defforward(self,x):x=self.body(x)x=x.view(x.size(0),-1)returnself.fc(x)model=DeepFuseNet().eval()torch.onnx.export(model,torch.randn(1,3,224,224),"deepfuse.onnx",input_names=["input"],output_names=["output"],opset_version=17)

我环境里真实打印的算子分布:

参数量 : 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 so=ort.SessionOptions()so.graph_optimization_level=ort.GraphOptimizationLevel.ORT_ENABLE_ALL so.optimized_model_filepath="deepfuse_optimized.onnx"# 让 ORT 把优化结果写出来sess=ort.InferenceSession("deepfuse.onnx",so,providers=["CPUExecutionProvider"])raw_ops=dict(Counter(n.op_typeforninonnx.load("deepfuse.onnx").graph.node))opt_ops=dict(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} 节点减少: 9

10 个 Relu 全部消失了。它们被熔进了前面的 Conv,变成了FusedConv(Conv+Relu 合并算子)。节点数从 22 掉到 13,少掉 9 个(那 10 个 Relu 全没了,只新增 1 个布局转换ReorderOutput)。

换句话说:推理时,GPU 不再需要单独算 10 次 Relu,也不用把每次 Conv 的中间结果写回内存再读出来给 Relu 用。一次算完,中间结果留在寄存器/高速缓存里直接给下一步。


第四章、实测:开图优化到底快多少?

同一个 ONNX 文件、同一个 ORT 引擎、同一台机器,只把graph_optimization_level关闭切到开启,CPU 上跑 100 次取平均(先 warmup):

so_off=ort.SessionOptions();so_off.graph_optimization_level=ort.GraphOptimizationLevel.ORT_DISABLE_ALL sess_off=ort.InferenceSession("deepfuse.onnx",so_off,providers=["CPUExecutionProvider"])so_on=ort.SessionOptions();so_on.graph_optimization_level=ort.GraphOptimizationLevel.ORT_ENABLE_ALL sess_on=ort.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.52
ORT 图优化关闭70.94 ms×1.43
ORT 图优化开启(融合)49.73 ms×1.00

第五章、融合的几种常见套路

除了Conv+Relu,深度学习里还有一堆"天生该融"的组合:

融合套路说明出现场景
Conv + BatchNormBN 推理时折叠进 Conv 权重(导出时已做)几乎所有 CNN
Conv + Relu合并成 FusedConv,中间结果不出内存CNN 主干
Conv + Add + Relu残差块最经典,三个捏成一个ResNet 等残差网络
Elementwise 链一串逐元素算子(Add/Mul/Relu)合并Transformer 的归一化
Gemm + Add全连接 + 偏置合并分类头

融合的本质遵循一个规律:两个首尾相接的算子,如果中间没有"必须被别的算子也读到"的分叉,就能安全地合并。合并后少一次中间张量落地、少一次 kernel 启动。


总结:一张表读懂算子融合

问题答案我环境里的验证方式
融合在融什么把首尾相接的独立算子合并成一个10 组Conv+Relu熔成FusedConv
图变什么样节点数变少,中间算子消失节点 22 → 13,Relu 全消失
为什么快省中间张量内存读写 + 少 kernel 启动实测 Relu 节点清零
快了多少仅图优化一项就 1.43×70.94 → 49.73 ms
精度有损吗没有,误差在浮点舍入级最大绝对误差 2.9e-08
什么时候最有用模型越大、算子越碎时收益越明显小模型 CPU 已见效

最后一句大白话:算子融合不是玄学,它就是把"每道工序都端进端出"的笨办法,改成"一个师傅从头到尾不离灶台"。省下的不是计算量本身,而是内存搬砖和 kernel 启动这两笔隐形成本。这也是为什么 ONNX Runtime、TensorRT 这些引擎能"白嫖"加速——它们什么都没训练,只是把计算图搓得更聪明了。

返回列表