最近在部署一个检测模型,PyTorch转ONNX时设了dynamic axes,用onnxruntime测没问题。但转到TensorRT(8.6)就报“Invalid shape”或者干脆构建失败。网上教程都说设minShapes、optShapes、maxShapes就行,但我试了还是不行,感觉是维度顺序或者opset版本的问题?另外模型里有ROIAlign和NMS这类自定义op,是不是必须用plugin?有没有人用torch2trt或者onnx-tensorrt成功部署过带动态batch的检测模型?求分享个能跑通的配置或工作流,感激不尽。
PyTorch转ONNX再转TensorRT,Dynamic shape一直报错,有踩坑的大佬吗?
全部回复
共 8 条之前跑Faster RCNN也踩过一模一样的坑,ONNX Runtime没问题但TRT就是报Invalid shape,最后发现是opset版本太新导致的,TRT 8.6对opset 17以上的某些动态算子支持不完整,降到opset 12或者13试试,ROIAlign和NMS这两个就别指望直接转了,TRT官方plugin里没有现成的,建议用onnx-graphsurgeon把这两个节点抠出来换成自定义plugin,或者干脆在TRT外面用torch处理完再送进去。
torch2trt对动态shape支持其实挺差的,它内部很多层还是静态的,我后来改用onnx-tensorrt那个仓库,但需要自己编译,而且对动态shape的优化也不够好,最终还是老老实实把模型拆成两段,backbone和head用TensorRT,ROIAlign用CUDA手写,虽然麻烦但至少能跑通。你检查下维度顺序,TRT默认是NCHW,如果你的模型是NHWC就要在转换时加个transpose,另外minShapes和optShapes不能随便设,要确保每个维度都合法,比如batch最小不能小于1,而且所有输入输出的维度都要对应上。
还有个细节,动态shape模式下TRT会做kernel autotuning,如果显存不够或者构建时间过长也会莫名失败,建议把构建时workspace调到2G以上,并且设置trt_network_creation的flag为EXPLICIT_BATCH。如果实在搞不定,可以试试TensorRT 8.6的onnxparser直接加载动态ONNX,不用自己手动设置profile,让parser自动推导,但前提是ONNX里的动态维度必须用symbolic name,不能用-1。
另外你用的检测模型是两阶段还是一阶段的?如果是两阶段,那个第二个stage的输入维度是动态变化的,TRT处理起来特别麻烦,建议把整个模型拆成两个engine,第一个输出proposals,第二个接收动态数量的proposals,这样每个engine的shape都能固定下来。我最后是这么跑通的,虽然推理时多一次数据传输,但稳定多了。
ROIAlign和NMS确实得走plugin,TensorRT原生不支持这俩op,onnxruntime能跑是因为它自己实现了。dynamic shape报错八成是opset版本问题,建议ONNX用opset 17以上,另外检查下输入维度是不是NCHW,TensorRT对layout特别敏感。我之前用onnx-tensorrt(就是仓库里那个parser)搞过带动态batch的Faster RCNN,min/opt/max的batch设成1/4/8能过,但得把NMS单独拆出来用plugin,不然构建必挂。torch2trt对动态shape支持更烂,别浪费时间。
动态shape这坑我太熟了,ROIAlign和NMS这种自定义op八成得自己写plugin,ONNX转TRT对带这些的图支持很迷。另外你检查下ONNX的opset版本,建议用17以上,还有维度顺序搞成NHWC有时候能绕开一些bug。torch2trt对动态batch支持也一般,我最后是直接绕开onnx用TensorRT的PythonAPI手动搭网络,虽然麻烦但可控性强。你试试把动态维度只放在batch上,其他维度固定死,说不定能过。
动态shape这块儿确实坑多,尤其ROIAlign和NMS在TensorRT里基本绕不开plugin,建议先确认下ONNX导出时这些op有没有被拆成小算子,不然转TRT肯定炸。我之前用torch2trt跑过带batch的检测,但最后发现还不如固定shape省心,动态batch性能还打折。你试试把opset降到13以下,然后min/opt/max的H和W别设成一样,有时候是维度顺序(NCHW vs NHWC)在搞鬼。
遇到过一模一样的坑,大概率不是opset的事,而是ROIAlign和NMS在转TRT时根本没被映射到原生层,必须走plugin,否则动态shape的shape推理铁定崩。建议先固定batch把流程跑通,再回头搞动态,别一上来就全动态。torch2trt对检测模型支持其实挺看版本的,我用的是trt 8.5+onnx-tensorrt,把NMS换成TRT自带的EfficientNMS插件后,动态shape才勉强过了,min/opt/max设成1/1/4这种倍数关系能少踩不少坑。另外检查下onnx里输出的shape是不是带动态维度,有时候转出来是静态的,TRT那边就报invalid shape了。
如果你坚持要动态batch,可以试试把onnx里所有reshape的0改成实际维度值,或者干脆用onnxsimplify清理一遍再转,很多报错都是图结构里藏着隐式静态shape导致的。
动态shape这块儿TensorRT确实比ONNX Runtime矫情不少,尤其是ROIAlign和NMS这种自定义op,8.6版本基本绕不开plugin。建议你先固定住batch维,把NMS和ROIAlign的输出shape用onnx-simplifier修一下,或者在导出时把这两个op的坐标输入设成静态,只留图像输入动态试试。我上次是被opset11的Cast坑过,换到opset13配合onnx-tensorrt的legacy模式才跑通,torch2trt对动态支持反而更烂,不如直接走onnx-tensorrt的pythonAPI手写profile。你贴一下报错的具体层名吧,八成是某个Resize或Gather的输入shape推断炸了。
ROIAlign和NMS确实是坑,TensorRT原生不支持这俩,得自己写plugin或者用TRT的NMS插件(8.6有内置但限制挺多)。动态shape报Invalid shape大概率是onnx里某些维度推导成了-1或者跟profile对不上,建议先用polygraphy跑一下onnx的shape推断,看哪一层断了。opset版本也关键,ROIAlign最好用opset11以上导出,低版本算子拆解后TRT认不出来。
Dynamic shape这块确实坑多,我去年搞YOLOX部署的时候也卡了好几天。你说的Invalid shape大概率不是min/opt/max设得不对,而是ONNX里某些维度没被正确标记成动态,尤其是reshape或者transpose后面带出来的维度,onnxruntime宽容一点能跑,TRT就直接翻脸。建议用polygraphy或者trtexec加--verbose把每层shape打出来看,通常能定位到具体哪一层炸的。opset版本也有影响,ROIAlign在opset11之后才有原生支持,但TRT对它的动态支持很烂,基本得自己写plugin或者换成TRT自带的插件。NMS更是重灾区,TRT 8.6虽然有EfficientNMS_TRT,但要求输入格式跟torchvision那套对不上,得在导出前把后处理改写成TRT能吃的形式。torch2trt对动态shape支持其实一般,onnx-tensorrt更稳一点但也要看parser版本匹配。我最后是走ONNX+trtexec显式指定 optimization profile 才跑通的,动态batch记得所有输入包括image shape都要一起设进去,别只设batch那一维。