最近在用PyTorch训练一个语义分割模型,准备部署到Jetson上,所以想把模型转成TensorRT。但折腾了两天,各种报错快把我整懵了。比如动态shape怎么处理?有些自定义算子(像F.interpolate)在onnx导出时总警告,转trt后直接报错。还有量化精度掉得厉害,int8比fp32掉了5个点,不知道是不是校准集没选好。另外,有没有办法在不重写网络结构的情况下,让TensorRT自动融合一些算子?我看官方文档说支持Layer Fusion,但实际效果好像不明显。希望有经验的大佬能分享下实战经验,特别是从PyTorch->ONNX->TensorRT这条链路下的常见坑和优化技巧,先谢过了!
PyTorch转TensorRT时,有哪些常见的踩坑点和优化思路?
全部回复
共 163 条同感,PyTorch转TensorRT这条链路确实是坑多路滑,尤其是语义分割模型,自定义算子一多就各种翻车。我前阵子搞一个U2Net的变体,也卡在F.interpolate上,ONNX导出时那个警告其实已经暗示了——TRT对双线性插值的支持在动态shape下特别敏感。我后来是把上采样层全换成torch.nn.Upsample,然后在ONNX里把mode设成nearest或bilinear,再配合opset版本12以上,才勉强跑通。不过你这个报错具体是啥?是解析阶段崩还是推理时段错误?如果是后者,多半是shape推导的问题,建议先固定输入尺寸试试。
量化掉点5个确实有点多,校准集大概率是主因。我踩过类似的坑,一开始用验证集随机抽100张,结果int8直接崩。后来换成训练集里均匀采样,并且保证每类像素占比接近训练分布,掉点才压到1-2个点。另外TRT的calibration cache很关键,第一次跑完会把校准结果存下来,后面可以复用,但注意如果换校准集一定要删掉缓存。还有个小技巧:你可以先用fp16跑一下,如果fp16比fp32掉点不多,那int8的校准问题可能性更大。
关于Layer Fusion,说实话TRT的自动融合对语义分割这种多尺度特征网络效果有限,尤其是用了很多skip connection或者concatenate的模型。我试过手动用trtexec加--layerPrecision参数指定部分层用fp16,或者把一些ReLU和卷积手动合并,但收益不大。后来发现真正有效的反而是优化ONNX导出阶段:比如把dropout、identity这些无意义节点去掉,或者用torch.onnx.export时指定dynamic_axes只对batch维度生效,其他维度固定,这样ONNX图干净了,TRT自动融合的成功率会高不少。
另外Jetson平台内存有限,建议你在TRT转完模型后,用trtexec测一下实际延迟和内存占用,有时候模型能转但推理时显存爆了。可以试试开启--workspace参数限制内存,或者对不需要梯度计算的层强制用fp16。别问我怎么知道的,都是泪(手动狗头)。
动态shape这块确实头疼,我后来是固定了输入尺寸才消停的,你可以试试用多组固定shape写个engine,虽然麻烦但稳。自定义算子的话,F.interpolate我一般用trt的plugin重写一下,网上有现成代码能参考。int8精度掉5个点可能是校准集分布和训练集差太多,建议挑个跟实际场景最接近的小数据集做校准。至于layer fusion,实测有些手工融合比如Conv+BN+ReLU写进代码里比全靠自动融合靠谱,trt有时候偷懒不帮你合。
动态shape确实头疼,试试固定输入尺寸或用trt的profile指定范围,量化校准集得覆盖真实分布才行。
动态shape建议固定输入尺寸,F.interpolate用resize模式替代,INT8校准集至少要覆盖500张典型场景图。
同感,动态shape这块确实是PyTorch转TRT的老大难,我试过在onnx里用-1维度配合trt的optimization profile来设动态范围,但不同shape下性能波动挺大的,而且有些层在动态输入时根本没法融合。F.interpolate这个我遇到过,后来发现用torch.nn.functional.upsample配合align_corners=False能减少警告,但转trt后还是得用Plugin重写,或者干脆换成最近邻插值省事。int8精度掉5个点其实挺正常的,尤其是语义分割这种对细节敏感的任务,校准集最好从实际场景中均匀采样,别用训练集里的干净图,我试过用500张不同光照的图片做校准,能把掉点控制在2-3个。关于Layer Fusion,你可以试试trt的onnx parser加一些显式的算子替换,比如把BatchNorm和Conv提前在onnx里fuse掉,或者用trtexec的--saveEngine看哪些层没被融合,手动调整网络结构。另外提个思路,如果模型不复杂,直接用torch2trt这种三方库能省不少onnx导出的麻烦,虽然灵活性差了点。
动态shape建议固定输入尺寸,F.interpolate换成torch.nn.functional.interpolate试试,校准集选和实际场景接近的。
动态shape确实头疼,我一般用固定batch size加padding绕过去,或者用TensorRT的optimization profile配置多档位。F.interpolate在onnx里建议换成torch.nn.functional.upsample,导出时用opset_version=11以上能少点警告。INT8校准集最好挑跟真实场景分布接近的数据,我试过用验证集子集做校准,精度能拉回来1-2个点。Layer Fusion有时候得手动调一下网络结构,比如把连续的卷积+BN+ReLU写进一个block里,TensorRT自动融合效果会好很多。
动态shape这块确实坑多,我一般用固定输入尺寸或者加个pad对齐来绕过,不然ONNX导出时reshape容易炸。F.interpolate建议用onnx官方支持的resize算子替换,或者干脆在TRT里用Plugin手写。INT8掉点5个不算离谱,校准集尽量覆盖真实场景的分布,多试几张图或改用熵校准可能好点。Layer Fusion其实会自动做,但效果跟模型结构有关,你可以用trtexec的--verbose看看哪些层没合上,再手动调整。
玩语义分割转trt确实容易头大,F.interpolate这个坑我当初也踩过,后来发现可以用ONNX的Resize节点替代,或者在导出时把mode固定成nearest或bilinear并指定align_corners,能少很多警告。动态shape的话,建议先在onnx里设成固定尺寸导出,再在trt里用optimization profile指定几个常用分辨率,虽然麻烦但稳定,完全动态的优化空间其实不大。int8掉5个点确实有点多,校准集最好选和真实场景分布接近的图,数量不用太多但要有代表性,我试过分块校准和用KL散度选图,能拉回来1-2个点。至于layer fusion,trt默认会做Conv+Bias+ReLU这种基础融合,但更复杂的结构像残差块里的add+relu它不一定自动处理,可以试试用polygraphy或者trtexec加--best参数看看最终网络里哪些层没被融合。如果实在不想改网络,把PyTorch里的小算子合并成一个大算子也能触发更多融合,比如把几个连续的卷积和激活写成一个自定义plugin,虽然麻烦但效果明显。另外建议用onnx-simplifier先简化模型,很多冗余操作会被剪掉,转trt时少报一堆warning。
动态shape确实头疼,我一般用固定batch加padding来绕过去,自定义算子像F.interpolate建议替换成torch的upsample,或者用onnx的Resize节点。int8精度掉5个点挺正常的,校准集最好选和实际场景分布接近的图片,数量500张起步。层融合方面可以试试trtexec的--best选项,或者手动调一下builder的优化参数,有时候默认策略就是不够猛。
动态shape确实坑多,我一般用固定batch再加padding来绕开。int8校准集选1000张带标签的图效果会好不少。
我之前也踩过动态shape的坑,建议导出onnx时固定一个batch和输入尺寸,或者用trtexec的--minShapes参数显式指定范围,不然trt自己推断很容易崩。F.interpolate这个确实烦,我后来换成了torch.nn.functional.upsample加上align_corners=False才消掉警告。int8掉点的话,校准集尽量选和实际场景分布一致的数据,我用500张随机图比100张专用图效果还差。Layer Fusion在trt8以上版本配合onnx-graphsurgeon手动调一下图结构会好很多,光靠自动融合确实鸡肋。
动态shape这块确实头疼,我一般会在onnx导出时固定一个batch size,然后用TensorRT的optimization profile去适配不同尺寸,能省不少事。F.interpolate这种算子可以试试用torch.nn.functional.upsample替换,或者干脆在导出前把它改成固定尺寸的resize,兼容性好很多。int8精度掉5个点的话,校准集最好挑跟实际部署场景分布接近的数据,数量不用太多但要有代表性,另外可以试试开启trt的strict_type_constraints或者用fp16混合精度过渡下。至于layer fusion,有时候onnx图里一些冗余reshape或transpose会阻碍融合,用onnx-simplifier先清理一下图结构再转trt效果会更明显。
动态shape确实是个大坑,我一般用固定batch+padding来处理,或者直接在trt里设optimization profile,能省不少事。F.interpolate那个警告其实可以在onnx导出时加上opset_version=11来规避,不过trt对双线性插值的支持还是不如原生算子稳定。int8掉点5个的话,校准集最好挑跟实际场景分布接近的数据,别用训练集的随机抽,另外试试开启strict_type_constraints调一下精度敏感层。Layer Fusion其实很多是自动做的,但有些复杂结构得手动写plugin才能发挥效果,你可以先跑个trtexec看下哪些层没融合。
我最近也在折腾这个,动态shape确实头疼,我后来直接固定输入尺寸绕过去了,精度损失用onnx-simplifier加polygraphy调试会好不少。int8校准集得挑跟实际场景分布差不多的图,不然掉点很厉害。层融合的话,可以试试trtexec的--best参数或直接写plugin,实测比默认融合效果好一截。
动态shape建议用trtexec加优化策略,量化校准集最好覆盖所有类别场景,不然掉点太正常了。
动态shape确实坑多,试试用trt的profile显式声明范围。int8掉点严重优先检查校准集分布,跟训练数据越像越好。
跟你情况差不多,之前做道路分割也踩过动态shape的坑。我的经验是尽量固定输入尺寸,实在不行就用trtexec的min/opt/max profile来设范围,但ONNX导出时得把dynamic_axes参数写对,不然转trt就崩。F.interpolate那个警告其实可以试试用torch.nn.functional.upsample替代,或者在onnx导出前用onnx_graphsurgeon手动改图,把插值节点换成TensorRT原生支持的resize层。int8掉点5个确实有点多,校准集最好覆盖真实场景的各种光照和纹理,另外可以试试开启calibrator的entropy模式,或者对敏感层用fp32混跑。至于Layer Fusion,官方文档说得玄乎,实际融合效果跟模型结构关系很大,像一些残差块里的add+relu倒是会自动合并,但复杂算子基本只能靠自定义plugin。还有个小技巧,用TensorRT的Polygraphy工具逐层分析精度差异,能快速定位到是哪一层量化后崩的。
动态shape确实是老生常谈的坑了,我一般直接固定输入尺寸跑TRT,实在需要动态的话可以试试TensorRT的Optimization Profiles,不过配置起来挺烦的。F.interpolate那个警告我遇到过,建议换成torch.nn.functional.interpolate的onnx官方支持版本,或者在导出时加个opset_version=11试试。量化掉点5个还算能接受,校准集建议挑些分布接近真实场景的图,数量500-1000张差不多,另外INT8转完最好跑一遍完整验证集看看哪里掉得最狠。
动态shape确实是个大坑,我一般直接固定输入尺寸或者用trtexec的explicit batch模式绕过去。F.interpolate的警告可以试试在onnx里用Resize算子替代,或者用onnx-simplifier过一遍。int8掉点严重的话,建议多跑几种场景的校准数据,别只拿训练集凑数,另外可以试试开启QAT微调。Layer Fusion在TensorRT8.6之后其实挺给力的,但有些自定义op得自己写插件才能触发融合,网上有现成的插件库可以参考。