含NonZero操作的PyTorch模型转换为TensorRT的可行方案咨询
PyTorch含NonZero算子转TensorRT的解决方案
算子替换方案(推荐,无需自定义开发)
你提到的用Where或是Equal+LogicalNot组合替换NonZero的方案完全可行,是目前成本最低的适配方案,操作步骤如下:
- 提前在PyTorch代码中完成算子替换,不要等到ONNX导出后再修改,避免出现算子映射问题
- 等价实现逻辑如下,可直接对齐NonZero的输出格式:
# 原调用NonZero的代码(默认返回[N, dims]格式的坐标张量) nonzero_coords = input_tensor.nonzero() # 替换写法1:仅用Where实现,适配所有非零判断场景 nonzero_coords = torch.stack(torch.where(input_tensor != 0), dim=1) # 替换写法2:Equal+LogicalNot组合实现,逻辑和写法1完全等效 nonzero_coords = torch.stack(torch.where(torch.logical_not(torch.eq(input_tensor, 0))), dim=1)
- 上述写法用到的
Where、Equal、LogicalNot、Stack都是ONNX和TensorRT原生支持的算子,导出ONNX后可直接转换,无需额外开发。
兜底适配方案(无法替换算子时选择)
如果业务逻辑限制不能修改原有NonZero调用,可选择以下两种方案:
- 自定义TensorRT插件:自行实现NonZero对应的TensorRT算子插件,在ONNX转TensorRT阶段注册对应算子映射即可,适配你的输入输出维度即可正常运行
- 推理流程拆分:将NonZero之前的计算逻辑导出为TensorRT模型运行,NonZero及后续计算量较小的逻辑放在CPU侧用原生PyTorch/ONNXRuntime运行,不会产生明显的性能损耗。
验证注意事项
- 算子替换后先在PyTorch侧对比替换前后的输出是否完全一致,避免逻辑错误
- 导出ONNX后先做一次推理验证输出正确性,再转TensorRT进行精度校验,降低后续调试成本
内容的提问来源于stack exchange,提问作者Alexender Iotko
相关产品推荐
相关产品推荐

