最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条int8掉点5个确实挺常见的,我建议先别急着动校准集,检查一下onnx导出时有没有把BN层和conv融合掉,这个影响特别大。还有F.interpolate建议换成onnx支持的resize模式,或者干脆在trt里用plugin,虽然麻烦但稳。动态shape的话,如果场景允许,固定到训练时的尺寸能省不少事,我上次搞了个动态batch,结果显存占用直接翻倍。层融合这事别太指望,trt的自动优化对分割网络效果有限,不如手动把一些重复的卷积结构简化下。
之前做检测模型转trt也卡在interpolate上,后来直接改用pytorch的grid_sample配合onnx opset 11才稳下来,建议你查下是不是opset版本太老。int8掉点那个,校准集最好覆盖各种光照和遮挡情况,我试过用500张验证集子集比200张训练集效果好很多。关于layer fusion,可以试试trtexec加--fp16和--builderOptimizationLevel=3,有些小算子确实能并进去,但别指望太多,关键还是减少slice/concat这类低效节点。你动态shape是用min/max/opt profile配的么?如果固定分辨率能接受,直接定死省一堆麻烦。
动态shape这块建议固定输入尺寸或者用min/max/opt profile,能省掉很多麻烦,F.interpolate的话试试把onnx opset调到13以上,或者用torch.nn.functional.interpolate的recompute_scale_factor参数关掉。int8掉点的话,校准集最好覆盖各种光照和物体分布,别只用训练集随便抽几张,另外可以试试per-channel量化或者用TRT的entropy_calibration。算子融合效果不明显可能是你的网络里有些层不支持,可以用trtexec的--dump-profile看看哪些层没被融合,然后手动改写这些层试试。还有个小技巧,导出onnx前把模型设为eval模式,并且用torch.jit.trace,能减少不少警告。
这链路我趟过一遍,F.interpolate那个警告基本可以无视,但转trt前最好把scale_factor改成固定尺寸,不然动态shape下必炸。int8掉5个点大概率是校准集太单一,试试用训练集里随机抽500张带真实分布的数据,calibrator用entropy_v2会好很多。Layer fusion对自定义结构确实没辙,可以试下trtexec的--saveEngine加--tacticSources,偶尔能挤出点性能。还有一个坑,Jetson上记得用JetPack自带的trt版本,pip装的经常不匹配cuda。
说到动态shape这个坑我太有共鸣了,当时我转一个检测模型也是卡在这,后来干脆把输入固定成几个常用分辨率,用多个profile去覆盖,虽然内存占用多点但省心。F.interpolate那个警告其实可以绕过去,导出onnx时把size改成显式数值而不是变量,或者干脆用onnx的resize算子替换掉,trt这边就稳了。int8掉5个点的话,校准集确实很关键,我试过用500张训练图混着验证集去校准,比单用验证集好不少,另外校准时要关掉batch norm的统计更新,不然精度崩得更厉害。Layer fusion这个别抱太大期望,trt自带的融合对常见卷积+bn+relu有效,但像分割网络里的跳连和上采样结构基本不会动,我现在都是手动把一些小的操作并进相邻卷积里,比如把flatten和softmax合成一个plugin。还有个小技巧是转trt前用torch.jit.script把模型过一遍,有些算子能被trace成更标准的格式,onnx导出干净很多。最后建议你装一下trtexec命令行工具,它能输出每层耗时,比单纯看总延迟好定位瓶颈,有时候瓶颈不在算子而在内存拷贝上。
int8掉5个点算正常范围,校准集别光用原图,最好从训练数据里抽几百张带标注的图,用softmax后的概率分布做校准,能稳一点。动态shape的话,建议直接固定一个或两个分辨率,TensorRT对动态尺寸支持还是太费劲,实在要动态就锁住batch维。F.interpolate这种算子,导出ONNX前用torch.nn.functional.upsample替换掉,或者干脆在onnx里用resize节点,能少很多麻烦。Layer Fusion你试试trtexec加--saveEngine和--fp16,然后看下生成的engine日志里有没有显示融合了哪些层,有时候是模型本身结构太碎,融合空间不大。
动态shape这块我建议固定输入尺寸,或者用trtexec的minShapes/optShapes先跑通再调,不然onnx导出时一堆reshape节点特别容易炸。F.interpolate的话试试把onnx opset调到13以上,有时候能少点警告,实在不行就写个plugin,别硬刚。int8掉5个点挺正常的,校准集得选跟实际场景分布一致的图,而且校准批次别太少,我上次用500张比100张效果好不少。层融合的话,你可以先看下trtexec的profiling输出,确认哪些层没合上,有时候是shape不匹配导致的,调下网络里concat的输入顺序可能就有改善。
刚跑完类似的项目,F.interpolate那个坑建议试试在onnx导出时用opset 11以上,然后用nearest或bilinear的固定尺寸版本,动态尺寸的话最好在trt里显式设好profile范围。int8掉点严重大概率是校准集分布和实际场景差异大,试试用验证集里覆盖各种光照和类别比例的样本,或者换成entropy_calibrator_2。层融合不生效的话,检查下是不是有些操作被拆成了多个小算子,可以先用trt的onnx解析器自带优化跑一遍,再手动合并那些未被识别的elementwise操作。另外Jetson上记得开fp16,比int8稳很多,精度损失小还能提速。
我最近也在折腾这个链路,F.interpolate确实是个大坑,建议试试在onnx导出前把它换成torch.nn.functional.upsample里用nearest或bilinear的特定模式,或者直接改用resize算子,能少很多警告。动态shape的话,我一般是在转trt时固定一个batch size和输入分辨率,Jetson上部署通常不需要太灵活,这样能省不少事。int8掉5个点确实有点多,校准集最好用训练集里覆盖各种场景的样本,数量不用太多但要有代表性,还可以试试看用熵校准而不是最小化均方误差,有时候效果差别挺大。Layer fusion那个确实看运气,我试过让模型里多放些卷积+BN+ReLU的结构,trt自己融合得挺好,但要是网络里有很多跳连和特殊层,融合效果就不明显了。另外你有没有试过用trtexec的--saveEngine和--fp16先跑一遍看速度提升?我这边有时fp16比int8还稳,精度损失小不少。要是自定义算子实在绕不过去,可以看看onnx-graphsurgeon手动改图,比重写网络省心多了。
int8掉点大概率是校准集太单一,试试用验证集随机抽500张,另外f.interpolate换成nearest+conv能省不少事。
int8掉点大概率是校准集太单一,换个覆盖全场景的试试,另外动态shape建议固定尺寸或加个padding,能省一堆麻烦。
你这几个坑基本都踩遍了,动态shape最省心的解法是固定输入尺寸或者用trtexec的minShapes反复测几组典型值,别指望完全动态。interpolate导出时把mode换成nearest或者把scale_factor改成具体size能少很多警告,实在不行就自己写个plugin。int8掉点先检查校准集是不是覆盖了所有类别分布,换500张带各类别的图再试,另外记得开strict_type。层融合这事确实玄学,但把batchnorm和激活函数提前fold进conv里,再打开trt的preview模式,能明显看出融合效果。
interpolate换resize算子能省一堆事,校准集记得用验证集抽样覆盖全类别。
int8掉点5个确实有点多,校准集一般得从训练集里挑跟实际场景分布接近的样本,数量别太少,我上次换了个策略直接从验证集抽了500张就好很多。动态shape的话建议固定一个最常用的分辨率,或者用trt的optimization profile设三档,别指望完全动态。F.interpolate那个警告我遇到过,换成onnx的resize算子或者自己写个plugin能绕过去,但比较折腾。Layer fusion别太指望,trt对conv+bn+relu这类常规结构融合还行,复杂点的自定义结构还是得手动改图或者用torch.fx先重写一下。
interpolate这个坑我也踩过,建议导出onnx时把opset版本拉到13以上,或者干脆在导出前把F.interpolate换成torch.nn.functional.upsample_bilinear,至少能少一半警告。int8掉5个点的话,校准集最好覆盖各种光照和物体分布,别只用训练集里抽的,我用500张混合场景图后精度基本拉回来了。至于layer fusion,你可以试试trtexec里开--fp16配合--stronglyTyped,有些卷积+bn+relu的融合其实得靠onnx-simplifier先清理图结构才触发。另外动态shape别用trt的dynamic shape模式,直接固定尺寸+resize输入,Jetson上速度差不了多少还省心。
你这几个点我基本都踩过,F.interpolate建议直接换成ONNX支持的resize模式,或者导出前把size改成固定值,动态shape能省就省,Jetson上能省不少事。INT8掉点的话试试用验证集里挑多样性高的图做校准,别用训练集,另外校准算法换一下可能也有改善。Layer Fusion其实对某些结构有效,但更多时候得配合TensorRT的plugin自己写,比如把BN和激活合并进去。要是实在不想动网络,可以先试下trtexec的精度模式,有时候比代码里直接转稳。
int8掉点先别急着赖校准集,先查下onnx导出时是不是把某些op拆碎了,尤其是F.interpolate建议换成resize模式或者直接用onnx的Resize算子,能少很多幺蛾子。动态shape的话,我一般固定到最常用的几个分辨率,用optimization profile做多档,比完全动态省心太多。Layer Fusion确实别抱太高期待,很多情况得靠onnx-simplifier先捋一遍,再把一些融合不了的子图手动改成trt支持的plugin。你那个语义分割输出头如果是双线性上采样,强烈建议直接换成转置卷积或者反卷积,精度损失小还能让trt自己优化。
我最近也在搞这个,动态shape这块建议直接固定一个或几个常用分辨率,TensorRT对动态shape的支持其实挺折磨人的,特别是配合自定义算子的时候。F.interpolate那个警告大概率是ONNX导出时nearest模式映射的问题,我后来是重写了个简单版的resize算子绕过,或者干脆在预处理阶段统一缩放到固定尺寸。int8掉5个点确实常见,校准集得选跟实际部署场景分布一致的图,别偷懒只用几十张,至少几百张,而且校准方法试试percentile而不是默认的entropy,有时候能救回来一些。关于Layer Fusion,它确实会自动做,但很多融合需要算子满足特定条件,比如conv+bn+relu这种标准的没问题,你要是中间夹了别的操作就废了,建议用trtexec加--dumpProfile看看哪些层没被融合,再手动改网络结构去贴合它。另外Jetson上跑的话,最好直接用TensorRT自带的Python API写推理,别绕onnxruntime那层,省得白折腾。你用的是哪个版本的TensorRT?我这边8.6和8.5的报错信息都不太一样,搞不好是版本兼容性的坑。
动态shape这块我建议直接固定输入尺寸,Jetson上跑语义分割一般分辨率都是定死的,省掉很多麻烦。F.interpolate的话,onnx导出时把mode换成nearest或者用resize算子重写一下能规避不少警告。int8掉点5个确实有点多,校准集最好挑和实际部署场景分布一致的图,采样个500张左右,另外试下trtexec的--calib参数调下量化算法。算子融合其实不用太纠结,trt自己会做,你只要保证onnx导出的图别太碎就行,比如把一些小的elementwise操作合并进卷积里。
这题我熟,上周刚在Orin上踩完一遍。动态shape建议先固定到训练时的分辨率,实在要动态就只用NCHW别碰NHWC,不然TRT的优化直接摆烂。F.interpolate导出时加个torch.onnx.is_in_onnx_export的flag绕过就行,或者干脆用resize算子重写,别让onnx自己猜。INT8掉5个点大概率是校准集太单一,试试多场景混合图片加个熵标定,能救回来2-3个点。层融合这事别太指望官方自动,把常见Conv+BN+ReLU先手动合并了,比啥都强。