最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条我最近也在搞这个,动态shape建议直接固定尺寸或者用trtexec的min/mid/max优化,不然onnx导出那关就够呛。F.interpolate的话试试把onnx opset版本调高到17以上,或者手动替换成resize层。int8掉点大概率是校准集分布跟实际场景差太多,换一下校准数据或者用熵校准试试。至于layer fusion,有时候得先让onnx simplify清理下冗余节点,再转trt效果会明显些。
我上次int8掉点也是校准集没选对,换了几百张多样本图立马稳了。
我最近也在搞这个链路,PyTorch转ONNX时F.interpolate那个警告基本可以忽略,但转TRT前最好把上采样改成固定尺寸或直接换成PixelShuffle,不然到TensorRT里确实容易炸。动态shape的话,我建议你直接设成固定shape,Jetson上部署如果输入尺寸不变,省掉dynamic axis能少踩一半坑,性能还能提一点。INT8掉5个点确实不正常,校准集得选跟实际部署场景分布一致的图,别用训练集随机抽,最好拿几百张真实推理时的输入图,而且校准算法试试看EntropyCalibrator2,比MinMax稳。Layer Fusion这个你指望它自动搞定不太现实,我实测Conv+BN+ReLU能合,但像残差结构里的add和concat经常要手写plugin,或者先用torch.fx做一下图优化再导出。另外建议你直接上TensorRT 8.6+,对分割模型的优化比老版本好不少,还可以试下onnx2trt那个工具,有些报错信息能看得更明白。
int8掉点5个的话,校准集大概率有问题,试试用验证集里类别分布比较均衡的几百张图,别全用训练集,另外校准算法换成entropy一般比minmax稳一点。动态shape建议固定到2的幂次,比如输入分辨率直接定成544x960,省得onnx导出时一堆reshape警告。F.interpolate可以先试试在onnx里用Resize算子替换,或者干脆把上采样改成转置卷积,虽然会增加点计算量但至少不报错。Layer fusion其实对conv+bn+relu这种经典组合效果还行,但自定义结构就别指望了,实在不行用trtexec的--builderOptimizationLevel=5试试。
int8掉点先别急着怪校准集,先看看是不是用了per-tensor而不是per-channel,换一下能救回来不少。F.interpolate那个警告基本可以无视,但如果你用的是onnx的resize模式,最好在导出前把align_corners和coordinate_transformation_mode对齐,否则trt跑出来的边缘会有一圈鬼影。动态shape的话,我建议你直接固定一个最大分辨率,或者用trt的optimization profile设三个档位,别真让它完全动态,否则很多插件会摆烂。layer fusion这块别抱太大期望,有些融合需要你手动把bn和conv写在一起,或者用torch.fx先做一遍图优化再导出,能少踩一半坑。你用的什么后端版本?tensorrt8.5和9.0的onnx解析器差别挺大的,有时候换个版本比改代码还管用。
int8掉点5个确实偏多,校准集建议直接抽训练集的子集,别用验证集,而且每类像素占比要均衡,不然小目标全废。动态shape这块我建议固定到2的幂次,比如416/832这种,Jetson上内存带宽有限,省下来的时间够你喝杯咖啡。F.interpolate可以先转成ONNX的Resize节点再看情况替换成TRT的插件,我试过用nearest+align_corners=False能绕开不少坑。层融合别指望全自动,像conv+bn+relu这种通常没问题,但遇到残差结构最好手动拆开算一下,有时候显存占用反而降了。
int8掉点确实大概率是校准集的问题,我试过用500张验证集子集做校准比默认的200张好不少,另外校准算法选entropy一般比minmax稳。动态shape建议固定一个batch和分辨率范围,用trtexec的optimization profile限定,别指望全动态。F.interpolate可以试着手动改成resize+padding的组合,或者干脆用ONNX的resize算子替换,很多警告能消掉。Layer Fusion不是所有算子都支持,像一些elementwise+conv的组合得自己写plugin,但尽量先用trtexec测一下,看profiling里哪些层没融合,再针对性处理。
int8掉点先查校准集,试试用验证集随机抽500张,别用训练集。动态shape直接固定尺寸最省事。
interpolate那个警告基本可以无视,但导出时把mode固定成bilinear或者nearest能省掉后面一堆麻烦,动态shape建议直接锁死到部署时的实际输入尺寸,别为了灵活性给自己挖坑。int8掉5个点大概率是校准集分布和真实场景差太多,试试点数加到1000张以上,或者用熵校准加个per-channel,能救回来不少。Layer fusion不用太指望,trt对常规卷积bn relu融合得挺好,但碰到自定义结构就别抱希望了,实在不行就手写plugin,虽然费点时间但可控。
int8掉点这事我调过一阵,校准集最好直接从训练集里抽,覆盖各种光照和类别分布,用500张以上且跑两轮校准,能稳定不少。动态shape的话,onnx导出时把opset版本拉高到13以上,配合trt的优化profile指定几个常用尺寸,别用完全动态的,不然显存和延迟都爆炸。F.interpolate那个警告,导出前手动替换成grid_sample或者把scale_factor改成固定size,基本能绕过去。层融合其实trt自己在做,但你的网络里如果有些小算子比如clamp或者elementwise add,有时反而会打断融合,我试过用torch.fx把相邻的bn和relu先折叠掉,再导出的trt延迟能低个10%。
跟你一样在Jetson上折腾过,interpolate那个警告直接换成固定尺寸或者用roi_align绕开能省不少事。INT8掉点先看看校准集是不是覆盖了所有类别分布,我上次就是背景太多导致小目标全废了。层融合其实得看onnx导出的图结构,把一些reshape和transpose清理干净后效果会明显些,另外建议试试trtexec的--report工具看每层耗时再针对性优化。
int8掉点先查校准集,选个200张带标签的覆盖全场景,比调融合参数管用。
动态shape建议固定分辨率,实在不行用trt的optimization profile分档设,能少踩一半坑。
这链路我跑过一阵子,动态shape确实是最折磨人的,建议能固定就固定,实在不行用trt的optimization profile把min/opt/max设好,不然显存分配和算子选择都会出幺蛾子。F.interpolate那个,onnx导出时最好换成resize算子或者干脆在onnx里手动改成roi align之类的,不然转trt很容易给你拆成一堆小算子,性能直接崩。int8掉5个点其实不算太离谱,但校准集最好从训练集里均匀抽个500-1000张,别用验证集,而且得覆盖各种类别分布,不然某些类别直接废掉。Layer Fusion那个别指望自动做太多,我试过,它主要融合elementwise和conv这类,像残差结构里的add+relu还行,但跨模块的大融合基本得靠onnx graphsurgeon手动调,或者干脆在pytorch里把bn和conv提前融合了,这样导出更干净。还有个坑是trt版本和onnx opset的匹配,我遇到过opset11能过但13报错的情况,建议固定opset=12试试。最后想说,如果项目时间紧,不如直接上torch2trt或者onnx-tensorrt这种工具,虽然灵活度差点,但能省掉大半调试功夫。
int8校准建议多试几组不同场景的图,别只挑训练集,量化掉点太正常了。 动态shape可以先固定一个尺寸跑通,后续再慢慢试优化。
动态shape建议先用固定batch导出,能省一半折腾;F.interpolate换成最近邻或area模式试试,能绕开不少坑。
int8掉点5个不算离谱,校准集最好直接从训练集里抽,覆盖各种类别分布和难例,别用随机图片。F.interpolate建议导出onnx时换成resize算子,或者干脆在trt外面做预处理,省得折腾。融合不明显的可以试试trtexec加--fp16或者--int8参数看下每层耗时,有时候瓶颈在内存拷贝上,跟算子融合关系不大。动态shape建议固定一个batch,高度宽度用32的倍数,能省很多麻烦。
int8掉点大概率是校准集太单一,换那种覆盖边缘和纹理的图试试,能救回来不少。
F.interpolate这个坑我太懂了,导出onnx时把size改成scale_factor能避开大部分警告,但转trt后最好直接用trt的resize层替换掉,省得折腾。int8掉5个点确实不正常,你试试用验证集里覆盖各种类别比例的图做校准,别只用训练集前几百张,另外校准算法换成entropy试试。层融合不明显可能是你模型里有个别算子阻碍了融合,建议先开trt的profiling看看哪些层没被融合,再针对性调整。动态shape建议固定到2的幂次倍,比如512x1024或640x1280,能显著减少优化时间。
动态shape这块我建议你直接固定分辨率,Jetson上跑语义分割没必要搞动态输入,onnx导出时把opset拉到17以上,F.interpolate用nearest或者bilinear的align_corners=False版本能省很多麻烦。int8掉5个点确实太狠了,校准集至少要挑500-1000张跟实际场景分布一致的图,而且别用默认的entropy校准,试试percentile或者minmax,有时候换一下trt的calibration cache缓存就能救回来。Layer Fusion其实不是自动的,你得先打开onnx的graph optimization,再把trt的builder里那个preview的fastmath选项开开,配合layer_norm融合才有点效果,不然就是个摆设。还有个坑是BN层和conv的融合,PyTorch导出的onnx有时候会多出一些reshape节点,你最好先用onnxsimplifier清理一遍,不然trt解析器容易卡死在奇怪的地方。最后你如果不想重写结构,可以试试把自定义算子用plugin包一层,但那个python接口写起来也挺折磨人的,不如直接在onnx里把interpolate换成固定尺寸的upsample。对了,你Jetson上是跑NCNN还是DeepStream?不同pipeline对trt的版本要求差别很大,建议查一下你刷的JetPack对应哪个trt版本,别用了新特性结果板子上不支持。
说实话int8掉5个点挺常见的,校准集最好覆盖各类场景,别用训练集硬怼,我上次换了个贴近部署场景的校准集直接拉回来3个点。onnx导出那块,F.interpolate建议单独替换成resize算子或者固定尺寸输入,省得后面一堆兼容性麻烦。至于layer fusion,你可以先看看trtexec的profiling输出,确认是不是某些op形状导致融合失败,有时候把网络里几个冗余reshape删掉效果立竿见影。动态shape的话,能固定batch就固定,实在要动态就设成opt profile的常用尺寸,别太贪。