最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条动态shape这块确实头疼,我一般用固定batch size加padding来绕过去,虽然有点浪费但省心。F.interpolate在onnx里建议用torch.nn.functional.upsample替代,能少不少警告。int8精度掉5个点的话,校准集最好覆盖各种光照和场景,我试过分桶校准效果会好一些。至于layer fusion,你可以试试trtexec加--best参数,或者手动调一下onnx的图优化选项,有时候默认策略确实不积极。
int8掉点5个确实挺常见的,校准集最好直接用训练集的子集,而且得覆盖各种难例,别用纯验证集。动态shape建议固定到最大分辨率或者用多档位,TensorRT对动态shape的优化很保守。F.interpolate的话,试试在onnx里用Resize算子替代,或者干脆在trt里用插件,网上有现成的。融合效果不明显可能是你的网络结构本身太规整了,可以试着打开trtexec的--stronglyTyped或者调高工作区大小,有时候会有惊喜。
我之前搞过一阵子这个链路,F.interpolate那个警告基本是ONNX版本和opset的锅,建议直接用onnx-simplifier过一遍,或者把上采样换成pixel shuffle或者conv transpose,虽然不能完全消除但能规避不少坑。动态shape的话,最省事的办法是固定到几个常用分辨率,比如训练时用768x768,部署就锁死这个尺寸,TensorRT的优化能好不少,真要动态的话得用trt的optimization profile,但内存占用和延迟都会变差。int8掉5个点其实挺常见的,校准集最好从训练集里均匀抽,覆盖各类场景,用熵校准试试,另外层级的敏感度分析值得做,有些层对量化特别敏感,可以单独保留fp16。至于Layer Fusion,它确实自动做,但效果不明显的根本原因往往是模型里有些算子不支持融合,比如激活函数用了奇怪的变体,或者有slice/concat这类操作挡了路,你可以用trtexec的--dump-profile看看哪些层耗时大,然后针对性手动改写,不一定重写整个结构,把关键路径上几个瓶颈操作换掉就有效果。还有个容易被忽略的点,Jetson上TensorRT版本和PC上不一样,有些算子支持度有差异,建议直接在板子上跑一下导出流程,别在PC上验证完了再搬,经常能省半天时间。最后,如果实在卡在自定义算子上,试试mmdeploy或者torch2trt这种社区工具,它们对PyTorch算子的兼容性做了很多补丁,有时候比官方流程省心。
我之前也被F.interpolate坑过,后来在onnx导出时把scale_factor换成固定尺寸就好多了,动态shape还是尽量别碰,实在要动态建议直接用trt的python api搭。int8掉点的话,校准集最好覆盖各种光照和物体分布,用500张以上图,然后试试看开启calibration cache会不会稳一点。层融合这块别太指望自动,我都是手动把conv+bn+relu合成一个块再导出,效果立竿见影。
ONNX导出时把f.interpolate换成onnx的resize算子能省不少事,校准集建议挑些边缘纹理丰富的图。
interpolate这个坑太真实了,建议导出onnx时直接把resize模式固定成nearest或者bilinear带align_corners=False,trt对这两种支持比较稳。int8掉点先别急着换校准集,试试看trt的explicit_precision模式,把敏感层单独保fp16,比整体int8强很多。层融合其实要看算子组合,像conv+bn+relu这种大概率能合,但你要先确认下有没有开trt的sparsity和cudnn调优开关,有时候是默认没开所以效果不明显。
另外动态shape我一般是定几个固定档位,比如分辨率1倍和0.5倍各导一个engine,运行时切换,比动态shape省心太多。还有个小技巧,onnx导出前把一些复杂自定义模块用torch.jit.script包一下,能减少很多莫名其妙的warning。
同款折腾路过,上周刚在Jetson上把segnet跑起来,你这几个坑我基本全踩过。动态shape建议直接固定尺寸,除非必须实时变分辨率,不然ONNX导出的dynamic_axes在TRT里经常要配optimization profile,稍微没对齐就报错,我最后是直接resize到固定大小省心。F.interpolate那个警告我也见过,ONNX里会转成Resize节点,但TRT的老版本对Resize的坐标变换模式支持不全,建议先升级到8.5以上,或者干脆在PyTorch里用F.upsample_bilinear试试,导出时算子映射更干净。int8掉5个点我猜校准集要么太少要么分布跟实际场景差异大,试试用验证集里随机抽500张,混合各类别,校准方法用entropy别用minmax,另外记得关掉某些层不量化,像最后几层和softmax前保留fp16会有惊喜。Layer Fusion别指望自动全搞定,TRT更擅长融合conv+bn+relu这类,但跨模块的残差结构它经常懒得动,你可以用trtexec的--dumpProfile看每层耗时,手动把能合并的卷积和激活写进一个block,比如把两个3x3替换成5x5(如果感受野允许)也能提速。还有个冷门坑,Jetson上显存带宽有限,转TRT时尽量把输入改成NHWC格式,某些硬件上能快10%左右。最后建议先跑通fp16版本,int8等模型收敛了再慢慢调,别一上来就追求极致压缩。
ONNX导出时F.interpolate那个警告基本可以无视,但转TRT前最好把size改成scale_factor或者直接用插件,不然容易在动态shape下崩。int8掉5个点的话,校准集最好选和实际场景分布一致的图,别只用训练集里随机抽的,另外试下用熵校准加per-channel量化能救回来一点。层融合其实TRT自己会做,但有些算子得手动用ISetLayerOutput标记才行,建议先跑一下trtexec的profile看看哪些层耗时高再针对性优化。
F.interpolate这个坑我太熟了,试试在onnx导出时把opset版本拉高到13以上,然后给interpolate加个fixed_shape的静态尺寸,或者直接用onnx-simplifier处理一下图结构,能省不少事。int8掉点5个确实有点多,建议校准集尽量贴近实际部署场景的数据分布,而且每类像素占比最好均衡点,另外可以试试trtexec的--fp16和int8混合跑,有时候精度和速度能平衡得更好。Layer fusion这块不用太指望自动,可以先看下onnx导出后的graph里有没有多余的transpose和reshape,手动在pytorch端把通道顺序调对,往往比事后指望trt优化有效得多。
动态shape这块我建议先固定一个常用分辨率导出,实在要动态就只让batch维动态,否则onnx那边一堆reshape节点转trt后特别容易崩。F.interpolate的话试试把onnx opset调到13以上,然后trt这边用8.x版本,对这类算子的支持会好很多。int8掉5个点确实偏多,校准集最好从训练集里随机抽个500-1000张,覆盖各种光照和类别分布,别用验证集。算子融合可以看看trt的profiling结果,有时候是模型里某些小算子阻止了融合,比如clamp或某些elementwise操作,能手动替换掉效果会明显些。
这链路我趟过不少水,动态shape建议直接固定输入尺寸或者用min/max/opt profile,能省掉一半报错。F.interpolate的话试试把onnx opset设到11以上,再不行就手动拆成resize+卷积,别惯着它。int8掉点先检查校准集有没有覆盖全类别,我上次就是背景占比太大导致小目标全废了,换500张带多样性分布的图就直接拉回3个点内。Layer fusion别太指望,trt自己融不了就试着手动改网络结构,比如把bn和relu提前塞进卷积后面,或者用trt的plugin写个融合核,虽然麻烦但效果立竿见影。
看到你说动态shape和f.interpolate,我上周刚踩完这坑,建议导出onnx时把interpolate换成最近的resize算子,或者固定输入尺寸用onnx-simplifier处理下。int8掉点五个其实挺常见的,校准集最好从训练集里均匀抽个几百张,别用验证集,另外试试trtexec的直方图校准方式,有时候比默认的好。层融合的话,可以先看下onnx图里有没有多余的transpose和reshape,手动清理一遍再转trt,效果比指望自动融合明显。
看到你这个帖子我太有共鸣了,上周刚在Jetson Orin上跑通一个检测模型,也是被这些坑折磨得够呛。动态shape这块我建议你直接固定一个batch和输入分辨率,除非业务必须,否则onnx导出时动态轴很容易让TensorRT的优化器犯迷糊,性能反而下降。F.interpolate那个警告我遇到过,最后是在onnx里把它替换成resize节点才消停,或者干脆在导出前用torch.nn.functional.upsample_bilinear2d,虽然麻烦点但稳。int8掉5个点确实有点狠,校准集我建议从训练集里随机抽500张以上,而且要覆盖各种亮度对比度,别用验证集,因为分布可能偏理想化。关于Layer Fusion,实际它更依赖op的排列顺序,你可以在onnx里用graphsurgeon把相邻的conv+bn+relu手动合并成ConvBNReLU,比让TensorRT自动融合靠谱。还有个小技巧,导出onnx时把opset_version设到17以上,有些新算子支持会好很多。另外你提到不重写网络,可以试试torch_tensorrt的dynamic shape模式,但性能可能不如固定shape。最后建议你跑一下trtexec的profiling,看看哪个层耗时最高,有时瓶颈不在算子融合而在内存拷贝。
动态shape这块建议直接固定一个或几个常用分辨率,Jetson上部署没必要追求全动态,能省掉一堆onnx导出的糟心事。F.interpolate的警告一般用onnx的opset11以上能缓解,但转trt前最好把上采样替换成固定kernel的转置卷积或最近邻插值,省得报错。int8掉5个点确实像校准集问题,试试用训练集里覆盖不同亮度/类别的子集,校准方法换成entropy,别用minmax。Layer fusion其实对conv+bn+relu这种常规结构有效,但你的分割模型里如果有很多跳连和小算子,trt不一定会动,可以先用trtexec的--dumpProfile看看哪些层占了时间,再针对性合并或删掉多余reshape。
你这情况我太熟了,去年搞检测模型上Xavier的时候也是这么过来的。动态shape那块建议直接用trtexec的--minShapes和--optShapes参数固定好范围,别指望ONNX导出的dynamic_axes能全自动处理,Jetson上内存带宽本来就紧,动态推理性能波动很大。F.interpolate那个警告基本无解,我后来是写了个自定义plugin包了bilinear采样,虽然代码丑但至少不崩,你可以查下TensorRT的IPooling层能不能直接替换。INT8掉5个点其实算正常范围,校准集建议选跟实际部署场景最接近的500-1000张图,别用训练集,而且试试看开启--calib=entropy或者用percentile=99.9,有时能把差距压到2-3个点。算子融合你可以先跑下trtexec的--dumpProfile看下每层耗时,很多融合没生效是因为模型里有reshape或者transpose打乱了布局,试着用onnx-simplifier把图清理一遍再转,我上次这么搞完融合率直接提升20%。另外注意下TensorRT版本,8.x和7.x对某些算子支持差很多,Jetson上最好刷对应JetPack版本的容器来转,别用PC上的trt硬转再拷过去。
int8掉点先别急着怪校准集,试试看用验证集里熵值最高的那批图做校准,另外记得关掉torch里的batch norm的eval模式,我之前就这么救回来的。动态shape建议固定一个最常用的分辨率,用trt的optimization profile去兜底,别硬刚全动态。F.interpolate这种算子可以在onnx导出前用torch.nn.functional.upsample替换掉,虽然警告还在但至少能出trt。算子融合这事别太指望自动,手动把conv+bn+relu写成一个block,再配合trtexec的--fp16和--tactic-src,效果比默认层融合明显得多。
我也在Jetson上搞过这玩意儿,太懂你说的那种折腾感了。动态shape这块我建议你直接固定输入尺寸,哪怕牺牲一点灵活性,因为TensorRT对动态batch和动态分辨率支持得挺别扭,我之前试过设了优化范围,结果显存占用忽高忽低,反而更麻烦。F.interpolate那个警告其实可以绕过去,导出ONNX前把resize模式换成nearest或者显式指定坐标变换,有些版本对align_corners的处理有bug,转trt后莫名其妙就多出几个层来。int8掉5个点的话,校准集确实很关键,别用训练集直接跑,最好挑些覆盖各种光照和物体分布的图,而且每类像素占比要均衡,我之前用500张图校准比2000张随机图效果还好。至于Layer Fusion,说实话别太指望自动融合,TRT对卷积+BN+ReLU这种常规组合融得好,但像SE模块里的全局池化+全连接就经常拆得稀碎,我后来是手动把几个小算子合并成自定义plugin,虽然麻烦但推理速度能提15%左右。另外你可以在ONNX导出时把opset版本调到13以上,有些新算子映射会更干净,再配合trtexec的--saveEngine和--loadEngine反复测试,省得每次都要从头解析模型。反正这条路就是不断试错,多看看生成的engine日志里每个层的内存占用,比瞎猜管用。
interpolate这个坑我也踩过,建议导出onnx时把坐标变换改成固定尺寸或者用grid_sample替代,能省不少事。动态shape的话可以先固定一个最常用的分辨率,实在要动态就设成-1但记得把profile的优化区间设窄一点。int8掉点大概率是校准集分布和实际场景差太多,试试用训练集里随机抽几百张带标注的图做校准,另外加上entropy calibration和per-channel量化能救回来一点。Layer fusion确实没那么智能,像一些elementwise操作可以手动在onnx里做常量折叠,或者用torch.compile先优化一波再导出。
F.interpolate建议换成固定尺寸输入,动态shape用onnx-simplifier处理下,int8校准试试500张以上多样本。
int8掉点5个确实有点多,校准集最好从训练集里抽,覆盖各种光照和物体分布,别用单独一张大图硬怼。动态shape建议固定到8的倍数,或者干脆转trt时定几个常用分辨率,省得onnx导出时一堆警告。至于layer fusion,可以先试试trtexec加--fp16看速度提升多少,很多时候算子没融合是版本问题,换个TensorRT版本可能就有惊喜。F.interpolate那个建议在onnx里用resize算子替换,pytorch导出时加opset_version=12以上会好点。