最近在搞一个分割模型的部署,PyTorch 1.13转ONNX(opset 11)再转TensorRT 8.5,fp16推理。输入是1024x1024的遥感影像,模型结构是UNet++加SE注意力。转完以后发现Dice系数从0.93掉到0.87,尤其小目标(比如道路细线)几乎全断。试过直接TensorRT的onnx-parser和polygraphy,也试过设置dynamic shape和固定shape,还是掉。最诡异的是,如果用opset 12转出来的ONNX,TensorRT能跑但会报一个“Assertion failed: (nbDims == 4)”的错。有没有老哥遇到过类似情况?是精度校准选错了还是某些算子(比如GroupNorm)在TRT里精度有问题?求指点,实在不想逐层排错,太痛苦了。
PyTorch转ONNX再转TensorRT,精度掉得离谱,有人遇到过吗?
全部回复
共 10 条之前跑过类似的分割模型,fp16在小目标上掉点太常见了,尤其是注意力机制对精度敏感,建议先试下trtexec加--fp16和--strictType看看是不是某些层被强制fp16了。另外opset 11转出来的图可能有些算子融合问题,可以试下onnx-simplifier清理下冗余节点,再转trt。那个nbDims==4的报错,八成是某个plugin或reshape层输入维度没对上,固定shape时把输入输出名和维度都打印出来核对下。如果还不行,试试TensorRT的preview模式或者关掉一些融合策略,比如--disableTacticHeuristic,虽然慢点但精度能回来一些。
fp16下SE注意力很容易炸,先把opset统一到11试试,小目标断裂多半是dynamic shape惹的祸。
试试把BN层折叠进卷积再转,之前我这么搞精度能拉回0.91,另外检查下量化校准集是不是太少了。
我之前跑分割也踩过类似的坑,fp16下小目标断裂大概率是激活层或者归一化层的精度敏感,试试给那几个层单独设成fp32,用TensorRT的per-layer精度控制能救回来不少。另外opset 11转出来的图有些算子会被拆得很碎,反而容易触发优化bug,建议直接上opset 13以上,配合onnxsim简化一下图结构。还有那个nbDims==4的报错,八成是某个plugin或reshape层写死了4维输入,查一下转出来的onnx里有没有可疑的Resize或Gather节点。
遇到过,而且跟你情况几乎一模一样,UNet++带SE模块,转TRT之后小目标直接崩。我当时查了很久,最后发现主要问题不在opset,而是SE模块里的全局平均池化在TRT上被优化得过于激进,导致梯度传播时特征分布变了,推理时小目标的响应被压没了。你可以试试把SE模块里的池化层单独保留成float32,或者干脆在导出ONNX之前把它换成普通的卷积加激活,别用自适应池化,很多TRT版本对自适应池化的支持都有坑。另外,你提到opset 12报nbDims==4那个错,我怀疑是某个reshape或者permute在TRT解析时把batch维和通道维搞混了,尤其是你用了dynamic shape的话,建议把所有reshape都改成静态shape,或者用ONNX-Simplifier先简化一遍再转,能绕开不少解析器的bug。还有个思路,你先用fp32跑一遍,如果精度没掉,那就是fp16的scale因子问题,特别是小目标区域激活值很小,量化误差容易放大,可以用TRT的层级别精度控制,把最后几层解码部分强制回fp32,代价是速度慢点但Dice能救回来不少。我最后是混合精度加手动改网络结构才把Dice拉回0.91的,还是比原始pytorch差一点,但这玩意真得靠试错,每个模型和TRT版本的坑都不一样。
这精度掉得确实有点狠,0.93到0.87基本就是细线和小目标全废的典型症状。fp16下SE注意力里的reduce操作特别容易出数值问题,建议先锁定是不是SE模块的精度瓶颈,试试把注意力分支单独转成fp32或者用layer norm替换试试。opset 12那个报错八成是某个上采样节点导出后维度推断崩了,可以试试在导出时给torch.onnx.export加个dynamic_axes之外,再把interpolate改成显式conv转置,之前我这么干解决过类似问题。另外你polygraphy比对的时候有没有看逐层输出差异,重点查一下第一个下采样和最后几个上采样附近,大概率能找到爆点。
同款遥感分割踩过坑,fp16下小目标断裂大概率是激活层和归一化层的精度问题,试试把SE里的sigmoid和UNet++的deep supervision部分单独保留下fp32,其他层再转。另外opset 11对某些算子的支持其实比12更稳,那个nbDims==4的报错可以试试用onnx-simplifier把动态shape先固定下来再转TRT,能绕过去。还有个小技巧,输入图像预处理的时候别做减均值除方差,直接让模型学原始像素范围,精度能回来不少。
遇到过,fp16在小目标上崩基本是精度问题没跑了,尤其SE注意力那块在TRT里优化容易出幺蛾子。建议先开fp32对比一下,如果fp32正常,那基本锁定是某些层对低精度太敏感,试试给那几个层单独设成fp32,或者用trtexec的--layerPrecision选项。另外opset 11确实比12稳,那个nbDims==4的报错我记得是TRT对某些Resize或Gather节点支持不完善,换opset 11外加固定shape一般能绕过去。还有个歪招,转ONNX前把输入归一化挪到模型外面,有时候能救回来一点。
这问题我熟,之前跑医学图像分割也踩过一模一样的坑,Dice从0.91掉到0.84,小血管全断了。你先把fp16关掉跑一遍fp32,如果精度能回0.93,那基本就是TensorRT的fp16对某些层敏感,特别是SE注意力里的全局池化和全连接,建议把这几层单独设成fp32,用layer precision覆盖一下。另外opset 11转出来的图容易带些冗余的reshape和gather,TensorRT优化时会把某些维度信息搞丢,你试试opset 13或14,但别用12,那个nbDims断言我查过,是TRT 8.5的已知bug,跟Unet++的跳跃连接concat之后动态shape处理有关。还有个小技巧,转ONNX时把dynamic batch去掉,固定batch=1,然后输入输出都显式声明一下尺寸,有时候能避开很多隐式转换。最后,用polygraphy对比一下逐层输出,看是哪个节点开始数值飘的,我之前发现是Resize的上采样模式,bilinear和nearest在TRT里实现不一样,换成nearest反而稳一些。
遇到过类似的,fp16下小目标断裂大概率是activation精度问题,尤其SE注意力里的sigmoid和全局池化对低精度很敏感。建议先试一下层级别精度控制,把敏感层单独切成fp32跑,polygraphy可以定位具体是哪几层掉的。另外opset 11转出来的图有些算子会被拆得很碎,TensorRT优化时容易引入额外误差,试试opset 13或者直接用torch2trt的校准接口,能走PTQ的话校准集选点带小目标的图会好很多。那个nbDims==4的错我猜是resize或插值层在opset 12下输出维度被多包了一层,手动改下onnx的节点输入输出shape一般能绕过去。
fp16下UNet++的跳跃连接+SE注意力很容易精度崩,试试per-channel量化或者敏感层保留fp32。
这情况我遇到过,多半是opset 11的Resize算子对齐问题,换opset 13再配TensorRT 8.6+能好很多。