最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条动态shape建议直接固定batch和输入分辨率,Jetson上部署一般不需要太灵活,省掉一堆麻烦。F.interpolate那个警告可以先转成onnx-simplifier试试,很多时候能消掉。int8掉5个点大概率是校准集分布和实际场景差太多,试试用验证集里最接近部署数据的几百张图,或者改用entropy校准。Layer Fusion别指望自动,我一般手动把conv+bn+relu合并,或者用torch.fx重写一下,效果立竿见影。另外你导出onnx时记得把opset版本调到13以上,有些算子兼容性好很多。
int8掉点大概率是校准集太单一,试试换点带纹理的图,f.interpolate建议用onnx的resize算子替代。
动态shape直接用trt的优化profile,设好范围一般能解决,层融合别太指望,先跑通再说。
int8掉点先查校准集,选个几百张覆盖全场景的图,比啥都管用。动态shape建议固定尺寸,省心太多。
int8掉点大概率是校准集太单一,换些难样本试试,融合不如自己重写plugin效率高。
这坑我熟,动态shape直接固定尺寸跑,量化用一万张图校准能救回来一点。
int8校准集确实关键,用500张类分布均衡的图试下,另外f.interpolate换成nearest或trt自带resize层能省不少事。
动态shape这块建议固定一个或者用多档shape先跑通,不然onnx导出后TensorRT的profile能折腾死人。F.interpolate的话试试在onnx里用resize算子替代,或者直接改成upsample+conv,虽然丑但稳。int8掉点5个确实校准集嫌疑大,试试用训练集的validation子集,外加每类像素均衡采样,我上次这么搞直接拉回2个点。Layer fusion别太指望自动,先开着trtexec的--fp16和--sparse试试,有些融合是隐式发生的。另外你Jetson上记得用明确指定tactic source,不然某些层会选到GPU特化kernel导致速度反而慢。
int8掉点大概率是校准集太单一,试试用验证集随机抽500张做校准,别用训练集。
动态shape别硬刚,固定尺寸省心太多,性能还能提一截。
int8掉点先别急着怪校准集,试试看用验证集里覆盖各类别像素比例的图做校准,别用训练集,我上次这么搞直接拉回2个点以内。动态shape的话,如果不想写plugin,建议直接固定尺寸输入,或者用trtexec的minShapes和optShapes多测几组,比在onnx里折腾dynamic_axes省心多了。F.interpolate那个警告我遇到过,后来是把导出时的opset版本调到12以上,再配合torch.onnx.export的opset_version参数,警告就没了,转trt也不报错。Layer fusion确实别抱太大期望,像残差结构这种它自己会合,但很多自定义组合还是得手写plugin,或者试试torch_tensorrt这个库,有时候它的自动优化比onnx-trt路径效果好。
插值算子建议导出时把onnx的opset调到13以上,转trt用静态shape先跑通再谈优化。校准集选个几百张带标注的就行,关键要覆盖各种光照和物体类别。
校准集最好从训练集里均匀抽,别只挑简单样本,int8掉点可以先试试逐层敏感度分析,再针对性用混合精度。
int8校准集最好从训练集里随机抽500张以上,覆盖各种场景,不然掉点很正常。
动态shape建议固定分辨率,实在不行用onnx-simplifier先优化一遍再转trt。
interpolate这个坑太经典了,onnx导出时最好把size改成scale_factor,或者直接换成nearest+pad的组合,能省不少事。int8掉5个点的话,校准集可以试试多塞些不同光照和噪声的图,另外用entropy校准比minmax稳很多。Layer Fusion其实对卷积+BN+ReLU这种常规组合效果比较好,自定义算子基本指望不上,建议先把onnx-simplifier跑一遍,很多冗余节点能消掉。顺便问下你Jetson上是Orin还是Xavier?不同平台的tensorrt版本对op支持差异还挺大的。
int8掉点五个确实挺常见的,校准集最好挑那种覆盖各种光照和物体分布的图,别用训练集硬凑,我上次换了个贴近实际场景的校准集直接拉回三个点。动态shape如果你是固定分辨率部署,直接用静态shape加padding到32的倍数能省很多事,实在要动态就锁死hw的档位。自定义算子像interpolate建议在onnx导出前先重写成交替的resize加conv,或者干脆转trt时用plugin,不然老警告早晚变报错。层融合你试试trtexec加--fp16和--enable-fusion,有时候是onnx图太碎没被优化,先用onnx-simplifier跑一遍再转会好很多。
说实话你这几个坑我全踩过,尤其是F.interpolate那个,onnx导出时警告还算好的,转trt直接给你整个不支持,后来我干脆在onnx里把它替换成resize+conv的组合才勉强跑通。动态shape的话建议先固定一个最常用的输入尺寸,比如720p或者1080p,用trt的optimization profile去设置min/opt/max,别一开始就想着全动态,Jetson上显存本来就紧,动态shape反而容易触发重编译导致延迟抖动。int8掉5个点确实挺多的,校准集我建议用训练集的子集,但一定要覆盖各种光照和物体分布,千万别只用验证集里那些干净图,另外校准算法也试试entropy和percentile的切换,有时候默认的entropy对分割任务就是不太友好。至于Layer Fusion,说实话trt的自动融合对常见CNN结构还行,但你这种带自定义算子的,它经常因为op不支持就跳过一整块,导致融合效果大打折扣,我后来是手动把几个相邻的conv+bn+relu先打包成caffe风格的block再导出,融合率明显上去了。还有个思路你可能没试过,就是直接用trt的onnx parser跑一遍,看它报哪些op不支持,然后反推回pytorch去改,比盲目调参数高效得多。最后想问下你用的tensorrt版本是8.x还是9.x,我这边8.6和9.0对同一个onnx的兼容性差异还挺大的,说不定换个版本就能少报几个错。
int8掉点基本都是校准集的问题,试试用验证集500张以上做校准,动态shape建议固定尺寸再优化。
F.interpolate导出时换成nearest或固定尺寸能省不少事,层融合确实得看算子匹配度。
int8掉点大概率是校准集太单一,换点不同场景的图试试,能救回来不少。
自定义算子别硬刚,能用built-in替换就换,fusion效果真得看运气。
我最近也在搞这个,PyTorch转TensorRT的坑真是踩到麻木。动态shape的话,建议你干脆固定输入分辨率,或者在onnx导出时把opset版本拉到13以上,然后用trtexec的minShapes和optShapes去适配,别指望onnx自动处理好。F.interpolate那个警告我遇到过,记得在onnx里用Resize算子替代,导出前把mode改成linear或nearest,别用bilinear的默认参数,不然转trt必炸。int8掉点5个其实还算常见,校准集别用训练集,最好拿验证集里覆盖各种光照和场景的图,量化敏感层比如最后的卷积可以单独跳过量化。Layer Fusion这事吧,官方文档说得挺美,但实际得看你的算子是不是都映射到了trt的原生层,很多自定义逻辑会被拆成多个小kernel,融合效果自然差。我试过最有效的办法是先把模型导出成onnx,再用onnx-simplifier过一遍,把冗余节点清掉,最后转trt前用trtexec的--saveEngine看下每层耗时,定位瓶颈层再手动改结构。另外,如果你实在不想重写网络,可以试试用TensorRT的plugin接口把那些顽固的自定义op包起来,虽然写起来麻烦,但至少不用改前向逻辑。
int8掉点五个确实大概率是校准集的问题,试试用验证集里覆盖不同光照和物体形态的图,数量拉个500张以上,校准算法换成entropy_2看看。F.interpolate建议在onnx里直接替换成resize算子,或者导出时加opset_version=11以上,能少很多警告。动态shape如果只是batch维度变化,建议固定hw,用trt的optimization profile指定三个维度,能省不少事。层融合那个别太指望,我实测更多是靠torch2trt或者onnx-tensorrt里头的graph surgeon手动合并卷积+bn+relu,效果立竿见影。
int8掉点大概率是校准集太单一,换点复杂场景的图试试,另外f.interpolate建议换成最近邻上采样能省不少事。
动态shape直接锁死固定尺寸跑,Jetson上部署没必要搞动态,稳才是王道。
int8掉点大概率是校准集太单一,试试换多样性高的数据,或者用熵校准,能救回来不少。
动态shape建议直接固定尺寸输入,省心省力,性能还更稳。
int8掉点大概率是校准集分布跟实际场景差太多,试试用验证集里挑几百张覆盖各种光照的图做校准,能拉回来不少。
动态shape先固定成几个常用尺寸,用optimization profile分段处理,别直接设-1,不然TRT优化全乱套。