最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条int8掉点5个确实挺常见的,校准集最好覆盖各种光照和物体分布,别只用训练集里抽的图,我试过用验证集里挑几百张多样性的图能拉回来不少。动态shape的话,建议直接固定输入尺寸,Jetson上跑分割一般分辨率不会变,省掉一堆麻烦。F.interpolate这种算子,onnx里用resize模式导出,再在trt里用plugin或者干脆替换成卷积上采样,比硬刚报错省时间。Layer Fusion别太指望自动,试下trt的onnx-parser版本换新点,或者手动把bn和conv合并了,效果比默认强。
int8掉点大概率是校准集太单一,换点纹理丰富的图试试,另外自定义算子能合到plugin里就别硬刚onnx。
动态shape建议固定尺寸跑,Jetson上省心太多,融合效果得看onnx版本和trt版本匹配。
int8掉5个点太正常了,先别急着怪校准集,检查下onnx导出时opset版本和动态轴设置,F.interpolate建议换成固定尺寸输入或者用TRT的resize层替代。至于layer fusion,官方那个对自定义结构很保守,你可以试试用torch2trt的legacy模式,或者手动把几个常用模块(比如conv+bn+relu)封装成plugin。顺便问下你用的TRT版本是8.x吗?有些算子融合对版本很敏感。
int8掉点大概率是校准集太单一,试试用验证集随机抽500张图跑一圈,效果立竿见影。
动态shape建议直接固定分辨率,省事不说,trt的优化空间还能大不少。
说实话你这两天的经历我太懂了,当初我转一个分割模型到Xavier上也是差点把头发薅光。动态shape这块我建议你直接固定输入尺寸,除非业务必须,否则trt的优化空间会被动态维度和显存分配策略拖累不少,实在要动态就试试用优化profile分几个档位。F.interpolate那个警告我遇到过,最后是换成了resize+pad的组合才消掉,或者你干脆在onnx里用Resize算子重写一下导出逻辑。int8掉5个点其实还算正常,校准集别用原图,最好是推理时真实场景的crop,而且每类像素比例要均衡,我上次用500张验证集做校准才把掉点压到2个以内。Layer fusion说实话别抱太大希望,trt的自动融合对卷积+bn+relu这种常规结构还行,但你的分割网络里那种skip connection和上采样它经常识别不了,我后来是手动把几个小算子合并成自定义plugin才提速的。还有个小坑,onnx导出时opset版本最好选11以上,不然有些aten算子直接崩。你要是能把报错log贴出来,说不定能帮你看看具体是哪一步的问题。
int8掉点先别急着怪校准集,先看看是不是没设torch.nn.functional.interpolate的recompute_scale_factor,或者导出onnx时把opset版本拉高到13以上能省不少事。动态shape的话,我建议直接固定一个batch和输入尺寸,Jetson上跑实际部署基本用不到动态,省心很多。Layer Fusion其实要靠onnx-graphsurgeon手动改图,纯靠trt自动融合确实有限,特别是那些带resize的分支。另外量化校准我试过用1000张验证集子集比500张随机图效果好得多,你试试看是不是数据分布差太多。
int8掉点大概率是校准集太单一,试试用验证集随机抽500张图跑一遍calibration,效果立竿见影。
动态shape别硬刚,直接固定尺寸输入,Jetson上性能还能再提一档。
int8掉点先别急着怪校准集,检查下onnx导出时有没有把BN层折叠掉,我上次就是这里没处理好直接掉4个点。F.interpolate建议换成固定尺寸输入或者用trt的resize层替代,动态shape的话干脆固定batch=1和输入分辨率,Jetson上部署没必要搞太花哨。层融合那个得看算子类型,我实测过一些transpose+conv能合,但很多自定义结构确实合不动,实在不行就手动写plugin,别指望全自动。
int8掉点5个确实有点多,校准集尽量选跟实际场景分布一致的图,而且试试用熵校准而不是默认的minmax,改完可能能救回来一点。F.interpolate这个坑我也踩过,建议导出onnx时把opset版本拉到13以上,或者用reshape+nearest的替代写法避免动态size。Layer Fusion别太指望自动,可以先转个fp16看看速度提升,再手工把conv+bn+relu这类结构用torch.fusion过一遍。另外动态shape如果不需要特别灵活,干脆固定输入分辨率,能省掉很多麻烦。
动态shape这块建议直接固定尺寸,Jetson上跑分割模型一般输入分辨率都是定死的,省掉一堆麻烦。F.interpolate先在onnx里用Resize算子替换掉,别让它自动选模式。int8掉5个点大概率是校准集太单一,多塞点不同场景的图,或者试试用per-channel量化。Layer fusion这功能别抱太大期望,很多时候得手动改网络结构去凑它的融合模式,比如把bn和conv写一起。
F.interpolate这个坑我太熟了,导出onnx时警告多半是因为动态尺寸导致opset版本里插值节点输出shape不确定,建议要么固定输入尺寸,要么在onnx里用resize算子替代,别偷懒直接跑默认转换。动态shape的话,我一般先把输入固定成训练时的尺寸,等验证没问题了再折腾dynamic axes,不然一堆维度推导报错根本分不清是网络问题还是trt的优化问题。int8掉5个点其实不算夸张,校准集最好覆盖各种光照和物体分布,别只用训练集的前几张图,试试用验证集随机抽500张做熵校准,或者换成per-channel量化,有时候能救回1-2个点。Layer fusion这东西别抱太大希望,trt自动融合主要是针对卷积+bn+relu这种标准结构,你的分割网络里大量skip connection和上采样层它不一定认,实在不行就手动把几个小操作并成自定义plugin,或者干脆用torch2trt的legacy模式,虽然老但有些算子反而支持得好。对了,你Jetson上用的是jetpack版本吗?不同版本的trt对opset的支持差异挺大,有时候升级下jetpack比改代码省事多了。
同款链路踩过一遍,你提的这几个点基本是必经之路。动态shape这块,我建议直接用trtexec的minShapes和optShapes参数先跑通再说,别在代码里硬调,等模型能转了再回头优化。F.interpolate那个警告,我后来是把onnx的opset版本升到13以上,然后导出前把插值模式固定成nearest或bilinear,别让它走动态计算,实测能消掉不少坑。int8掉点5个确实偏多,校准集我试过500张和2000张,差别挺大,但更关键的是要用验证集的真实分布,别拿训练集随机抽,还有校准时的batch size和算法(entropy vs percentile)也值得多试几组。Layer fusion那个,官方文档说得挺美,但实际对分割网络效果有限,我后来发现把一些batchnorm和激活函数在PyTorch里提前fold掉,反而比等TensorRT自动融合更有效。另外,你如果Jetson上跑,建议直接看下TensorRT的plugin库,有些自定义算子在官方sample里能找到现成实现,比自己写plugin省事得多。最后问下,你转trt之后有没有对比过输出张量的数值误差?我遇到过fp32下误差在1e-3以内,但int8某些层会飙到1e-1,后来发现是某些卷积的权重分布太宽,得手动加clip才行。
跟你情况差不多,也是部署到Jetson上,动态shape这块我建议直接固定输入尺寸,或者用trtexec的minShapes和optShapes配合好,别指望导出时候一次搞定。F.interpolate那个警告我后来换成ONNX的Resize算子,虽然还有warning但至少能跑通,精度掉的话校准集最好覆盖所有类别,尤其小目标,不然int8掉点真没法看。至于Layer Fusion,感觉还是得靠onnx-graphsurgeon手动改图,自动融合在分割模型上确实收益有限,不如试试用TRT的network API直接搭个简化版。
F.interpolate这个坑我太熟了,当时转完onnx也是警告一堆,后来直接在导出前把resize操作改成固定尺寸的nn.Upsample,或者干脆在trt里用resize层单独实现,虽然麻烦点但能彻底避开算子不兼容的问题。动态shape的话,如果Jetson上部署场景固定,建议直接设成固定尺寸,能省掉很多优化上的麻烦,比如显存分配和kernel选择都会更激进。int8掉5个点确实偏多,校准集最好从训练集里随机抽500张以上,覆盖各种光照和物体分布,另外校准算法也可以试试entropy和percentile的混合策略,有时候minmax反而更稳。至于layer fusion,其实trt的融合主要看算子类型和内存布局,像conv+bn+relu这种经典组合基本都能自动融合,但如果你用了大量elementwise操作或者自定义插件,融合效果就会差很多,这时候可以试试torch2trt或者onnx-tensorrt的优化级别调高,有些版本对特定模型有额外pass。还有个隐蔽的问题是onnx导出时opset版本,有时候默认的11会让某些节点拆分得很碎,导致trt优化空间变小,手动换到13或17能好不少。最后建议用trtexec先跑一遍看每层耗时,定位到底是哪几个节点拉低了整体性能,再针对性地做算子替换,比盲目调参有效率得多。
interpolate那个警告我建议你直接换成ONNX支持的resize,或者用grid_sample替代,虽然麻烦点但一劳永逸。int8掉5个点确实校准集问题比较大,试试用验证集里采样多样性的图,batch size稍微大点,另外记得开onnx的QDQ模式导出。layer fusion效果不明显的话,大概率是前处理里有些小算子卡住了融合,比如clamp或者transpose,可以在onnx里先手动合并一下。动态shape建议固定到几个常用尺寸,用optimization profile做多档,别直接全动态,Jetson上容易爆显存。
int8掉点这事我太有同感了,校准集最好直接从训练集里抽个几百张带标签的,别用纯背景图,我上次就是校准图太单一结果直接崩了。动态shape建议先固定一个常用分辨率跑通,后面再用trtexec的minShapes和optShapes慢慢调,别一开始就追求全动态。F.interpolate的话试试在onnx里用Resize算子替换,或者干脆在导出前把上采样层换成转置卷积,虽然参数多了点但稳得多。Layer fusion别太指望自动,我一般手动把bn和conv合并了再导出,效果比啥都明显。
动态shape确实是最先要解决的,我建议直接用trtexec转的时候固定一个batch,或者用onnx-simplifier把F.interpolate这类先拆成基础算子再导出,能省掉不少警告。int8掉5个点的话,校准集最好选和实际场景分布一致的图,数量别低于500张,另外试试看用熵校准而不是默认的最小化。Layer fusion别指望全自动,我一般会在onnx里手动合并一些conv+bn+relu,转出来效率提升还挺明显的。
int8掉点先查校准集,别用默认的,挑几百张覆盖多样性的图效果差很多。
F.interpolate建议直接用resize模式导出,别用bilinear,能省不少事。
int8掉点5个其实挺常见的,先别急着怪校准集,试试把校准数据换成训练集里随机抽的几百张,最好覆盖各种光照和类别分布,有时候是校准样本太单一了。另外F.interpolate这种算子建议在onnx导出前手动替换成resize模式,或者直接把上采样倍数固定成常量,动态shape配合trt的优化profile能省不少事。Layer fusion那个确实别抱太大期望,我试过用trt的onnx-parser直接解析比先转onnx再转trt效果要好一丢丢,但自定义层还是得自己写plugin,绕不开的。
int8掉点这个事,校准集确实很关键,建议用训练集的子集但得覆盖各种类别分布,能拿验证集跑一下对比。F.interpolate的话,ONNX导出时把mode换成nearest或者用resize算子配合坐标变换能绕开警告。动态shape建议直接固定到训练时的输入尺寸,除非必须多分辨率,否则省心太多。Layer Fusion可以试试TRT8以后的版本,旧版本对分割网络支持确实一般。你Jetson用的哪个型号?orin和老xavier的优化策略差异挺大的。