最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条动态shape这块建议直接在导出ONNX时固定一个或几个常用分辨率,省得后面处理麻烦,TensorRT对动态shape支持确实比较坑。F.interpolate的话可以试试在onnx里用Resize算子替代,或者干脆在trt里用plugin,但工作量会大一些。int8掉点5个确实偏多,校准集最好覆盖真实场景的分布,样本量也别太少,我上次用500张图当校准集效果就比100张好很多。Layer Fusion其实更多依赖trt版本和显卡,Jetson上老版本效果差,升级到8.5以上会明显改善,另外可以手动调整一些网络结构,比如把BN层提前融合进卷积,这样trt更容易做优化。
int8掉点先查校准集,选个1000张带标签的比啥都强,动态shape直接固定尺寸最省心。
F.interpolate换成最近邻或双线性几次采样,onnx导出前把opset拉高到17,能少一半坑。
动态shape这块建议直接锁死输入尺寸,或者用profile指定几个固定档位,Jetson上跑实时推理真没必要搞全动态,不然onnx导出和trt构建都会多一堆幺蛾子。F.interpolate那个警告我遇到过,把导出时的opset版本调到13以上,再配合onnxsimplify能消掉大部分问题。int8掉点先检查校准集分布跟实际场景是否一致,另外试下trt的entropy_calibration_v2,比默认的minmax强不少。算子融合别指望全自动,自己写个plugin把那些顽固的节点(比如transpose+reshape组合)包起来,或者用trt官方提供的onnx-graphsurgeon脚本手动改图,效果立竿见影。
同款链路,Jetson上部署真的是一步一个坑。动态shape这块我建议你直接用trtexec转的时候固定到训练时的输入尺寸,除非必须支持多分辨率,否则别给自己找麻烦,ONNX导出时那些Resize警告基本都是因为动态尺寸导致的,把opset调到13以上能缓解不少。F.interpolate这个我遇到过,通常用torch的replication_pad2d配合conv手动实现上采样能绕过,或者干脆把模型里的插值层全部改成固定倍率的转置卷积,虽然稍微改动结构但转trt后稳定很多。量化掉5个点太正常了,jeston上int8对语义分割这种密集预测任务本身就敏感,校准集至少要覆盖不同光照和场景分布,而且每类目标都要有足够样本,我试过用500张图和2000张图做校准,最后精度差距能有2个点。Layer fusion别指望它自动做太多,trt能融合的bn和conv是常规操作,但像分割网络里的skip connection和aspp这种结构它经常不认,建议你直接看trt的profiler输出,把耗时高的算子手动拆出来重写,比如用plugins或tactic。另外提一句,如果精度实在保不住,可以试试tensorrt的fp16加int8混合量化,在部分层用fp16保底,比全int8稳很多。你现在那个onnx导出时警告的op是怎么处理的?是直接ignore了还是改代码重写了?
F.interpolate这个坑我也踩过,导出onnx时把scale_factor换成size或者固定输入尺寸能省不少事,动态shape建议直接用trtexec的minShapes参数先跑通再调。int8掉点的话,校准集最好跟实际场景分布一致,多试几次不同batch的校准数据,另外看看是不是有些层对量化特别敏感,可以先用层级别精度分析工具定位一下。融合的话,有些小算子确实得手动改网络结构才能触发,比如把多个小卷积换成大卷积或者重组一下残差块,不然TensorRT有时候就是懒得不给你合。
int8掉点先查校准集,得覆盖各种光照和场景,别用训练集凑合,另外试试trtexec的--fp16看看能不能接受。
动态shape建议固定尺寸或者用onnx-sim优化下,F.interpolate换成resize算子能省不少心。
F.interpolate这个坑我也踩过,试试在onnx导出时把keep_ratio关掉或者直接用nearest模式,能少很多警告。int8掉点建议先别急着换校准集,检查下有没有用trt的strict_type_constraints,很多时候是某些层被强制量化了。另外layer fusion其实对分割模型效果有限,可以试试把bn层和conv合并后再导出,能省不少显存。
int8掉5个点大概率就是校准集分布跟实际推理数据差太多,试试用验证集里覆盖各类别的样本凑个几百张,校准算法换成entropy别用minmax。动态shape的话,onnx导出时把opset拉到13以上,固定一个维度比如H走动态,W设成32的倍数,能省不少事。F.interpolate那个警告可以试试用torch.nn.functional.upsample替代,或者干脆在onnx里手动改成resize节点。layer fusion不明显的话,检查下是不是用了太多python层算子,把能合并的conv+bn+relu写成一个block,trt识别率会高很多。
F.interpolate这个坑我之前也踩过,导出onnx时把mode换成nearest或者显式指定scale_factor通常能消掉警告,但最稳的办法还是直接改用trt的resize层。int8掉5个点的话,校准集最好挑跟实际场景分布一致的图,别贪多,200张左右带标注的就行,另外开一下熵校准和量化敏感层分析会好很多。Layer fusion别指望全自动,可以先跑一遍trtexec看哪些算子没融合,手动把那些小算子像激活和elementwise合并一下,效果明显很多。动态shape这块建议固定H和W,只留batch维度动态,能省一大半麻烦。
int8掉点大概率是校准集太单一,试试用训练集的类别均衡子集,另外插值层用最近邻或者固定尺寸能省不少事。
这坑我熟,动态shape直接固定输入尺寸保平安,融合不明显的就手动改onnos算子,效率立竿见影。
我之前搞检测模型转trt也卡了好几天,你提的这几个点基本全踩过。动态shape建议直接用trt的optimization profile,把min/opt/max三个维度设好,千万别在onnx里硬塞动态轴,不然导出时一堆reshape报错。F.interpolate这个确实恶心,我后来是写了个plugin,或者干脆在onnx里用resize算子替换掉,虽然麻烦点但稳。int8掉5个点的话,校准集很关键,我试过用训练集的子集但分布跟实际推理差太远,后来换成验证集里随机抽500张,配合上entropy校准,掉点能压到2个点以内。算子融合这个别太指望trt自动做,它更多是优化conv+bn+relu这种常规组合,自定义结构基本没戏,我最后是手动把几个小卷积合并成大卷积,速度提升比融合明显多了。还有个坑是pytorch的export要设opset_version=11以上,不然某些算子直接不支持。你用的什么量化校准库?是pytorch自带的还是trt的?
int8掉点基本就是校准集问题,换500张覆盖各种光照的图重跑一遍能救回来。自定义算子别硬刚,直接改onnx导出时把interpolate换成resize试试。
你这几个坑基本都踩过一遍了。动态shape建议直接用onnx的dynamic_axes,但batch维度固定成1省很多事,TensorRT对静态shape优化更狠。F.interpolate那个警告我后来用torch.nn.functional.upsample替代,虽然不能完全消除,但trt里至少能跑起来。int8掉点大概率是校准集太单一,试试用验证集随机抽500张,配合entropy校准,能拉回2个点。Layer Fusion别指望白嫖,先跑一遍trtexec的profiling看哪些层没融合,手动改写onnx结构反而更有效。
int8掉点5个确实挺常见的,校准集最好从训练集里抽,而且得覆盖各种类别分布,别只用简单样本。动态shape的话,onnx导出时把opset设高一点,有些warning其实能忽略,但F.interpolate建议换成固定尺寸或者用trt的resize层替代。算子融合别太指望自动,可以先用trtexec看下每层的耗时,再手动把能合并的conv+bn拆成conv+scale,或者尝试用plugin把几个小op包起来,效果比硬调融合参数直观多了。另外Jetson上记得开fp16,比int8稳不少,精度损失小很多。
int8掉点先查校准集,别用默认的,混合精度加个CLIP能救回来不少。
F.interpolate先定尺寸再导出,动态shape用trtexec带profile试试,能省很多事。
F.interpolate这个我太有同感了,之前转yolov5也是卡在这,后来直接在onnx里用Resize算子替代,或者干脆把上采样层改成转置卷积再导,虽然稍微改了结构但稳很多。动态shape的话建议先用固定尺寸跑通,后续再考虑动态,不然报错排查起来真的头大。int8掉点5个我觉得校准集大概率有问题,试试用训练集的子集而且类别分布要均匀,我上次换了500张带标注的图就救回来了,另外可以开一下TRT的strict_type_constraints,有些层保持fp16能缓解。算子融合这个别太指望自动,我一般用trtexec的--saveEngine看下每层耗时,手动把一些零碎的小卷积合并到大层里,效果比纯靠工具强不少。
你这几个坑我基本都踩过,动态shape建议直接用trt的optimization profile锁几个常用尺寸,别依赖onnx动态轴。F.interpolate是重灾区,最好在导出前换成固定尺寸的resize或者用trt插件,不然后面调试到怀疑人生。int8掉5个点的话,校准集选个100-200张覆盖各类场景的图,用熵校准试试,比默认的minmax稳很多。融合这块别太指望自动,把onnx-simplifier跑一遍再转,能省不少事。
动态shape这块建议直接锁死batch和输入尺寸,Jetson上固定shape收益最大,省心还快。F.interpolate确实坑,导出前用torch.nn.functional.interpolate换成最近邻或双线性加align_corners=False能少很多警告。int8掉5个点大概率是校准集太单一,试试用验证集里随机抽500张打乱再校准,或者换成percentile=99.9的校准方式。算子融合不用太指望自动,把BN和激活函数塞进卷积层自己写个融合模块,实测比TRT的layer fusion靠谱。另外建议先用trtexec跑一遍看哪层耗时最高,很多性能瓶颈其实在reshape和transpose上,手动调整网络结构比依赖自动优化有效。
同款Jetson受害者路过,动态shape建议直接固定尺寸或者用onnx-simplifier提前把F.interpolate的坐标计算拆开,不然trt老爱自作主张优化出幺蛾子。int8掉点确实大概率是校准集太单一,试试用验证集里每个类别都抽点图,顺便开一下trt的explicit quantization看看哪些层敏感。Layer Fusion其实要看onnx算子的对齐程度,有些自定义结构得先转成trt的plugin才能触发融合,硬靠官网那套自动融合基本指望不上。
你这情况我太熟了,Jetson上跑分割模型真是每一步都踩雷。动态shape这块,我建议你干脆固定输入尺寸,或者用trt的optimization profile设几个常用档位,别指望onnx导出时能完美兼容,我试过用torch.onnx.dynamic_axes配trt的min/opt/max,结果还是得手动改网络里一些reshape逻辑才消停。F.interpolate那个警告,我后来是直接换成torch.nn.functional.upsample_bilinear,或者干脆在onnx里用resize算子,虽然还是有warning但至少能出engine,不过你得检查下输出对齐,我碰到过边缘像素差一个的情况。int8掉5个点太正常了,校准集别用原图,要按你训练时的预处理来,而且每类像素占比得均衡,我试过用500张验证集+熵校准,比默认的minmax强不少,但最后还是得靠per-channel量化才能压到2个点以内。Layer Fusion那个别太指望自动,我实测发现TRT 8.x对简单conv+bn+relu融合还行,但复杂skip connection结构经常不触发,你可以试下把网络里一些无用的分支剪掉,或者用torch.jit.script先优化一遍再导onnx,有时候能逼出几个融合。另外强烈建议你装一下onnxsurgeon,导出后手动删掉那些没用的Identity和Cast节点,我遇到过把float32转成float64再转回来的诡异操作,直接拖慢推理速度。你现在用的是TRT哪个版本?Jetson上的JetPack自带版本和PC端差别挺大,有时候PC上能跑的engine到板子上又得重新调。