将ONNX转TensorRT时ScatterND Plugin未找到的解决求助
我之前在TensorRT 7.1.3.4和ONNX 1.6.0的组合下也踩过一模一样的ScatterND插件坑,给你几个亲测有效的解决思路:
方法一:自定义并注册ScatterND插件到TensorRT
TensorRT 7.1.x版本默认没有内置ScatterND的官方插件,所以需要手动添加插件实现并完成注册:
编译ScatterND插件动态库:
你可以参考TensorRT官方samples里的插件模板,或者社区开源的适配实现,编写ScatterND的插件代码(要确保版本是1,和ONNX节点的版本完全匹配),然后编译成动态链接库(Linux下是.so,Windows下是.dll)。在Python中注册插件:
先初始化TensorRT插件库,再注册ScatterND的插件creator,代码示例如下:import tensorrt as trt # 初始化日志器 TRT_LOGGER = trt.Logger(trt.Logger.WARNING) # 加载所有可用的插件库 trt.init_libnvinfer_plugins(TRT_LOGGER, "") def register_scatternd_plugin(): # 获取ScatterND插件的creator(版本1) plugin_creator = trt.get_plugin_creator("ScatterND", "1", "") if not plugin_creator: raise RuntimeError("Failed to find ScatterND plugin creator (version 1)") # 注册插件 trt.register_plugin_creator(plugin_creator, "") # 执行注册 register_scatternd_plugin()注意要把编译好的插件库路径添加到系统环境变量中(Linux的
LD_LIBRARY_PATH,Windows的PATH),确保Python能正确加载到这个库。
方法二:修改PyTorch代码,避免导出ScatterND算子
如果不想折腾插件,可以直接在PyTorch层面重构ScatterND的逻辑,用TensorRT原生支持的算子组合来替代,这样导出的ONNX模型就不会包含ScatterND节点:
比如原来的ScatterND操作:
output = torch.scatternd(indices, updates, shape)
可以改成手动索引赋值的方式:
# 初始化输出张量 output = torch.zeros(shape, dtype=updates.dtype, device=updates.device) # 将indices转成tuple形式的索引,直接完成赋值 output[tuple(indices.t())] = updates
这种写法导出ONNX后,会生成TensorRT能直接解析的节点,完美避开插件问题。
方法三:升级TensorRT版本
如果项目允许版本升级,最省心的方式是把TensorRT升到8.0及以上版本——TensorRT 8.x已经内置了ScatterND算子的支持,不需要额外插件,直接就能解析包含ScatterND的ONNX模型。不过要注意同步升级ONNX版本到1.8+,兼容性会更好。
额外注意事项
- 如果你选择插件方案,一定要保证插件的版本(这里是version 1)和ONNX模型中ScatterND节点的版本完全一致,否则还是会报错。
- 测试时可以先单独验证插件是否能被正确加载,再进行ONNX转TensorRT引擎的操作。
内容的提问来源于stack exchange,提问作者sp_713

