TensorRT Python 样例详解:为 ONNX 网络添加自定义插件层(onnx_custom_plugin 实战指南)
TensorRT Python 样例详解为 ONNX 网络添加自定义插件层onnx_custom_plugin 实战指南【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT本篇技术指南围绕 TensorRT 开源仓库中的onnx_custom_plugin样例展开完整演示了为 ONNX 网络添加自定义插件层的端到端流程先用 C 基于 cuBLAS 实现一个 Hardmax 算子并封装为 TensorRT 插件IPluginV3编译为共享库后由 Python 侧动态加载注册再通过 ONNX Parser 在解析模型时自动匹配到该插件。读完本文你将掌握插件工程的组织方式、ONNX 图的外科手术式改造方法、插件动态加载与 Python 推理的关键调用链以及一套可复用的插件正确性验证方案。样例概览解决什么问题很多来自生态的 ONNX 模型会包含 TensorRT 原生不支持的算子例如 BiDAFBidirectional Attention Flow双向注意力流问答模型中的Hardmax、Compress、CategoryMapper等节点。onnx_custom_plugin样例给出了一个标准应对思路用 C 实现缺失算子并封装成带 Plugin Creator 的 TensorRT 插件将插件源码编译为动态库.so/.dll在 Python 中加载该动态库使插件注册进 TensorRT 的 PluginRegistry用 ONNX GraphSurgeon 把 ONNX 图中的不支持的算子改写为插件对应的算子名如CustomHardmax用trt.OnnxParser正常解析、构建引擎并执行推理。整个流程覆盖了插件开发 → 动态注册 → 图改写 → 构建推理 → 数值验证的完整闭环是学习 TensorRT 自定义层开发的经典样例。样例目录结构样例位于仓库 samples/python/onnx_custom_plugin 目录各文件职责如下文件/目录作用plugin/customHardmaxPlugin.cppHardmax 插件实现基于 cuBLAS使用 IPluginV3 接口plugin/customHardmaxPlugin.h插件类与 Creator 类的头文件声明model.py下载 BiDAF ONNX 模型并用 ONNX GraphSurgeon 改写不支持的算子sample.py加载插件库、构建引擎并执行问答推理load_plugin_lib.py封装了在 Python 中动态加载libcustomHardmaxPlugin.so的辅助函数test_custom_hardmax_plugin.py用 NumPy 参考实现逐维度、逐轴验证插件数值正确性CMakeLists.txt插件动态库的构建脚本requirements.txt运行本样例所需的 Python 依赖工作原理解析整体数据流本样例的核心思路是用 cuBLAS 实现一个 Hardmax 层即沿指定 axis 取 argmax将最大值位置置 1、其余位置置 0的 one-hot 化算子把它包装成 TensorRT 插件并配套实现一个插件 Creator然后编译成共享库。在 Python 端sample.py启动时首先调用load_plugin_lib()。这个函数通过ctypes.CDLL加载libcustomHardmaxPlugin.soLinux或customHardmaxPlugin.dllWindows。动态库加载的副作用是执行插件实现中的REGISTER_TENSORRT_PLUGIN(HardmaxPluginCreator)宏见 plugin/customHardmaxPlugin.cpp 第 61 行从而把CustomHardmax插件注册进 TensorRT 的 PluginRegistry。此后ONNX Parser 在解析模型时遇到名为CustomHardmax的算子就会从 PluginRegistry 中查到对应的 Creator 并实例化插件。三个不支持的算子如何被处理原版 BiDAF 模型有三个 TensorRT 无法直接解析的节点model.py逐一处理Hardmax → CustomHardmax直接把节点的op改名为CustomHardmax与插件名对齐由插件在推理时接管计算Compress → EinsumCompress会根据第二个张量中为True的索引从第一个张量取值。由于这里的第二个张量恰好是 Hardmax 的输出只有一个位置为 1等价于对两个二维张量做点积。因此样例用Einsum节点equation 为ij,ij-i替换了Compress(Transpose_29, Cast(Reshape(Hardmax)))子图CategoryMapper删除模型输入本来是字符串 token经CategoryMapper转成整数 token。样例直接移除这些节点让网络输入改为整数 token同时把 String→Int 的映射以 JSON 文件保存下来备用。这一套改写逻辑在 model.py 的_do_graph_surgery()中实现最终通过graph.cleanup().toposort()清理孤立节点并输出新的 ONNX 文件bidaf-9-trt.onnx。环境准备Prerequisites安装 Python 依赖pip3 install -r requirements.txtrequirements.txt 中锁定的关键版本包括onnx1.18.0、onnx-graphsurgeon0.3.20、numpy1.26.4、cuda-python12.9.0、nltk3.9.1、wget3.2、requests2.32.4、tqdm4.66.4、pyyaml6.0.3Windows 平台额外安装pywin32并指定了--extra-index-url https://pypi.ngc.nvidia.com这一额外包源。安装 CMake构建插件动态库需要安装 cuBLAS插件实现依赖 cuBLAS 库本样例自 2024 年 1 月起将其列为显式前置条件因为插件改用cublasCreate自行创建 handleWindows 构建需要 Visual Studio 2017 Community 或 Enterprise 版本。具体软件版本要求以 TensorRT 官方安装指南为准。第一步下载并预处理 ONNX 模型python3 model.py脚本会从 ONNX Model Zoo 下载 BiDAF 模型bidaf-9.onnx到models/目录若已存在bidaf-9-trt.onnx则跳过改写流程。改写完成后会生成 TensorRT 可解析的models/bidaf-9-trt.onnx同时导出CategoryMapper_*.json映射文件。第二步构建插件动态库Linux 构建mkdir build pushd build cmake .. make -j popd如果依赖不在默认位置可以手动指定关键变量例如cmake .. -DCMAKE_CUDA_COMPILER/usr/local/cuda-x.x/bin/nvcc # 或把 /path/to/nvcc 加入 $PATH -DCUDA_INC_DIR/usr/local/cuda-x.x/include/ # 或把 /path/to/cuda/include 加入 $CPLUS_INCLUDE_PATH -DTRT_LIB/path/to/tensorrt/lib/ -DTRT_INCLUDE/path/to/tensorrt/include/cmake ..会打印全部可配置变量。如果某个变量被显示为VARIABLE_NAME-NOTFOUND就需要手动指定它或修正其派生来源变量。Windows 构建PowerShellmkdir build; pushd build cmake .. -G Visual Studio 15 Win64 / -DTRT_LIBC:\path\to\tensorrt\lib / -DTRT_INCLUDEC:\path\to\tensorrt\lib / -DCUDA_INC_DIRC:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\vCUDA_VERSION\include / -DCUDA_LIB_DIRC:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\vCUDA_VERSION\lib\x64 # 注意msbuild 通常位于 C:\Program Files (x86)\Microsoft Visual Studio\2017\EDITION\MSBuild\VERSION\Bin # 需要把该路径加入 PATH 环境变量。 msbuild ALL_BUILD.vcxproj popdCMake 脚本要点从 CMakeLists.txt 可以看到构建细节使用add_library(customHardmaxPlugin MODULE ...)生成模块型动态库编译单元包含插件源码以及samples/common/logger.cpp、shared/utils/fileLock.cpp默认TRT_LIB/usr/lib/x86_64-linux-gnu、TRT_INCLUDE/usr/include/x86_64-linux-gnu非 MSVC 时通过find_library查找nvinfer链接nvinfer、CUDA::cudart_static、CUDA::cublas与CUDA::cuda_driver并通过-DTENSORRT_BUILD_LIB定义控制导出符号要求 C17 标准target_compile_features(..., cxx_std_17)。构建产物为build/libcustomHardmaxPlugin.soLinux或build/Debug|customHardmaxPlugin.dllWindows与 load_plugin_lib.py 中的查找路径一一对应。第三步运行推理python3 sample.pysample.py的执行流程如下见 sample.py调用load_plugin_lib()注册CustomHardmax插件若当前目录存在已保存的bidaf.trt引擎则直接反序列化并设置runtime.max_threads 10否则调用build_engine()从bidaf-9-trt.onnx构建新引擎构建时使用STRONGLY_TYPED强类型网络定义并把工作空间workspace内存池上限设为 1 GiB由于输入文本长度可变为每个输入创建优化 profile最小形状 batch1、最优形状 batch8、最大形状由MAX_TEXT_LENGTH 64决定通过common.CudaStreamContext()管理 CUDA stream 生命周期对每个测试用例调用common.do_inference()执行推理模型输出是答案在 context 中的起始、结束位置sample.py据此从分词结果中切出答案文本。成功运行后输出示例 Testing Input context: Garry the lion is 5 years old. He lives in the savanna. Input query: Where does the lion live? Model prediction: savanna Input context: A quick brown fox jumps over the lazy dog. Input query: What color is the fox? Model prediction: brown交互式模式python3 sample.py --interactive可以自己输入上下文和问题 Testing Enter context: Waldo wears a striped shirt. He also wears glasses. Enter query: Who wears glasses? Model prediction: waldo注意交互模式下输入文本分词后的长度不能超过MAX_TEXT_LENGTH64否则会触发断言并提示增大该常量。第四步深入插件实现源码级剖析IPluginV3 插件架构本样例的插件基于IPluginV3接口体系实现README 的 Changelog 显示 2026 年 3 月已从 IPluginV2DynamicExt 迁移到 IPluginV3。插件类同时继承三个能力接口见 plugin/customHardmaxPlugin.h 第 28-31 行IPluginV3OneCore提供getPluginName()/getPluginVersion()/getPluginNamespace()等身份信息IPluginV3OneBuild负责构建期形状推导、格式校验与 workspace 大小计算IPluginV3OneRuntime负责运行期的enqueue()执行。getCapabilityInterface()按PluginCapabilityTypekCORE/kBUILD/kRUNTIME返回对应的能力接口指针TensorRT 在构建与执行阶段分别通过这些接口与插件交互。关键方法逐一解读getNbOutputs()返回 1即单输出插件getOutputDataTypes()输出数据类型与输入保持一致透传getOutputShapes()输出形状与输入完全相同Hardmax 不改变形状supportsFormatCombination()仅支持DataType::kFLOAT且格式为PluginFormat::kLINEAR的线性排布且不允许输入输出类型不同configurePlugin()将负数 axis 归一化为正数mAxis inDims.nbDimsonShapeChange()运行时形状变化时用samplesCommon::volume()计算mDimProductOuteraxis 之前各维乘积、mAxisSizeaxis 维大小与mDimProductInneraxis 之后各维乘积getWorkspaceSize()需要两块 float 数组——一块缓存当前 axis 切片一块放全 1 常量故返回2 * inputs[0].max.d[mAxis] * sizeof(float)attachToContext()克隆插件并用cublasCreate()为克隆体创建独立的cublasHandle_t这是 2024 年 1 月的变更不再依赖attachToContext传入的 cuBLAS context析构函数中cublasDestroy()释放getFieldsToSerialize()把axis属性以PluginFieldType::kINT32形式序列化供引擎序列化/反序列化时保存插件参数。enqueue() 的计算逻辑Hardmax 的数学定义是沿指定 axis 找到最大值所在下标将其置 1其余置 0。插件在 enqueue() 第 225-305 行中用 cuBLAS 原语组合实现了这一功能思路如下cudaMemsetAsync把输出整体清零外层双重循环遍历mDimProductOuter × mDimProductInner个axis 切片对每个切片调用cublasIsamax找最大绝对值元素的下标注意返回值是 1 基索引需要减 1由于cublasIsamax找的是绝对值最大而非数值最大若该元素为负则先把切片拷入 workspace用cublasSaxpy减去最小值等价于平移为全非负再调用一次cublasIsamax得到真正的最大值下标通过cudaMemcpyAsync把该下标对应输出位置写为 1.0最后返回cudaPeekAtLastError()检查异步错误。需要说明的代价与局限插件用同步的cudaMemcpyDevice→Host读取最大值会阻塞流水线且该并行策略在axis 维很大、其余维很小时高效例如形状(1, 512, 3)、axis1反之若 axis 维很小例如 axis2则串行开销明显。源码注释也明确指出一个更聪明的插件应当识别这种不对称性把最耗时的维并行化。这也是用 cuBLAS 原语拼装算子的典型取舍——若改用自定义 CUDA kernel 可完全规避这些瓶颈。Creator 与注册机制HardmaxPluginCreator继承IPluginCreatorV3One在构造函数中声明唯一的axis字段PluginFieldType::kINT32createPlugin()从PluginFieldCollection解析axis值默认 -1后构造插件实例。文件末尾的宏REGISTER_TENSORRT_PLUGIN(HardmaxPluginCreator);在库被加载时自动把 Creator 注册进 PluginRegistry插件名与版本分别为CustomHardmax与1见 plugin/customHardmaxPlugin.cpp 第 61-67 行。Python 侧如何按名字取插件在 test_custom_hardmax_plugin.py 中可以看到脱离 ONNX Parser、纯 Python API 直接使用插件的路径registry trt.get_plugin_registry() plugin_creator registry.get_creator(CustomHardmax, 1, ) axis_attr trt.PluginField(axis, axis_buffer, typetrt.PluginFieldType.INT32) field_collection trt.PluginFieldCollection([axis_attr]) plugin plugin_creator.create_plugin( nameCustomHardmax, field_collectionfield_collection, phasetrt.TensorRTPhase.BUILD )随后通过network.add_plugin_v3(inputs[input_layer], shape_inputs[], pluginplugin)把插件挂到强类型网络上。这证明了一条重要事实同一个 Creator 既可以被 ONNX Parser 隐式调用按算子名匹配也可以被 Python API 显式调用按名字查注册表两种入口共用同一份实现。第五步正确性验证单元测试python3 test_custom_hardmax_plugin.py该脚本对插件做穷举式数值验证遍历维度数1..7对每个维度数遍历所有合法 axis-num_dims .. num_dims-1生成形状随机的输入各维大小 1~3数值范围为(rand - 0.5) * 200覆盖正负混合场景专门考验cublasIsamax的绝对值陷阱分支参考实现hardmax_reference_impl()用 NumPy 的argmaxput_along_axis构造 one-hot 结果用插件构建引擎执行推理断言与参考实现逐元素完全相等。测试覆盖了负数 axis、多维张量、正负值混合输入等多种边界情况可作为后续开发自定义插件的通用测试范式。版本演进记录ChangelogREADME 记录了该样例的演进历程能帮助读者理解当前代码形态的来由2026 年 3 月Hardmax 插件从 IPluginV2DynamicExt 迁移到 IPluginV3当前源码即 IPluginV3 形态2025 年 10 月迁移到强类型strongly typedAPI对应sample.py中STRONGLY_TYPED标志的使用2025 年 8 月不再支持 Python 3.102024 年 1 月改用cublasCreate自行创建 cuBLAS handle不再使用attachToContext传入的 cuBLAS context并把 cuBLAS 列为首要依赖2023 年 8 月ONNX 支持版本更新到 1.14.0移除 Python 3.8 支持。已知问题README 明确声明当前样例没有已知问题。需要再次强调的是插件实现的性能取舍同步cudaMemcpy、对 axis 维大小敏感属于设计权衡而非缺陷作者在源码注释中已明确说明。小结与扩展阅读通过本样例你已经掌握了一套完整的自定义 ONNX 插件层开发范式C 实现与注册REGISTER_TENSORRT_PLUGIN→ CMake 构建动态库 → Python ctypes 加载 → ONNX GraphSurgeon 图改写 → Parser 隐式匹配或 API 显式创建 → NumPy 参考实现数值验证。如需继续深入仓库内还有更多可对照学习的素材Python 侧插件开发的另一范例samples/python/python_plugin展示完全用 Python 编写插件的路径插件生态与源码plugin 目录收录了大量官方插件实现efficientNMS、bertQKVToContext、scatterElements 等是学习各版本插件接口与 kernel 实现的高质量参考快速部署插件模板samples/python/quickly_deployable_plugins构建 Python 绑定的流程可参考 scripts/build_python_wheel.sh。将本样例的算子改写 插件注册方法论应用到自己的模型上即可把 TensorRT 不支持的自定义算子平滑纳入现有 ONNX 推理管线。【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考