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 原生不支持的算子,例如 BiDAF(Bidirectional 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.cpp | Hardmax 插件实现(基于 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.so(Linux)或customHardmaxPlugin.dll(Windows)。动态库加载的副作用是执行插件实现中的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 → Einsum:
Compress会根据第二个张量中为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 中锁定的关键版本包括:
onnx==1.18.0、onnx-graphsurgeon>=0.3.20、numpy==1.26.4、cuda-python==12.9.0、nltk==3.9.1、wget>=3.2、requests==2.32.4、tqdm==4.66.4、pyyaml==6.0.3;Windows 平台额外安装pywin32,并指定了--extra-index-url https://pypi.ngc.nvidia.com这一额外包源。安装 CMake(构建插件动态库需要);
安装 cuBLAS(插件实现依赖 cuBLAS 库;本样例自 2024 年 1 月起将其列为显式前置条件,因为插件改用
cublasCreate自行创建 handle);Windows 构建需要 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 构建(PowerShell)
mkdir build; pushd build cmake .. -G "Visual Studio 15 Win64" / -DTRT_LIB=C:\path\to\tensorrt\lib / -DTRT_INCLUDE=C:\path\to\tensorrt\lib / -DCUDA_INC_DIR="C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v<CUDA_VERSION>\include" / -DCUDA_LIB_DIR="C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v<CUDA_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定义控制导出符号; - 要求 C++17 标准(
target_compile_features(..., cxx_std_17))。
构建产物为build/libcustomHardmaxPlugin.so(Linux)或build/Debug|customHardmaxPlugin.dll(Windows),与 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:最小形状 batch=1、最优形状 batch=8、最大形状由
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_LENGTH(64),否则会触发断言并提示增大该常量。
第四步:深入插件实现(源码级剖析)
IPluginV3 插件架构
本样例的插件基于IPluginV3接口体系实现(README 的 Changelog 显示 2026 年 3 月已从 IPluginV2DynamicExt 迁移到 IPluginV3)。插件类同时继承三个能力接口(见 plugin/customHardmaxPlugin.h 第 28-31 行):
IPluginV3OneCore:提供getPluginName()/getPluginVersion()/getPluginNamespace()等身份信息;IPluginV3OneBuild:负责构建期形状推导、格式校验与 workspace 大小计算;IPluginV3OneRuntime:负责运行期的enqueue()执行。
getCapabilityInterface()按PluginCapabilityType(kCORE/kBUILD/kRUNTIME)返回对应的能力接口指针,TensorRT 在构建与执行阶段分别通过这些接口与插件交互。
关键方法逐一解读
getNbOutputs():返回 1,即单输出插件;getOutputDataTypes():输出数据类型与输入保持一致(透传);getOutputShapes():输出形状与输入完全相同(Hardmax 不改变形状);supportsFormatCombination():仅支持DataType::kFLOAT且格式为PluginFormat::kLINEAR的线性排布,且不允许输入输出类型不同;configurePlugin():将负数 axis 归一化为正数(mAxis += inDims.nbDims);onShapeChange():运行时形状变化时,用samplesCommon::volume()计算mDimProductOuter(axis 之前各维乘积)、mAxisSize(axis 维大小)与mDimProductInner(axis 之后各维乘积);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()检查异步错误。
需要说明的代价与局限:插件用同步的cudaMemcpy(Device→Host)读取最大值,会阻塞流水线;且该并行策略在"axis 维很大、其余维很小"时高效(例如形状(1, 512, 3)、axis=1),反之若 axis 维很小(例如 axis=2)则串行开销明显。源码注释也明确指出:一个更聪明的插件应当识别这种不对称性,把最耗时的维并行化。这也是用 cuBLAS 原语"拼装"算子的典型取舍——若改用自定义 CUDA kernel 可完全规避这些瓶颈。
Creator 与注册机制
HardmaxPluginCreator继承IPluginCreatorV3One,在构造函数中声明唯一的axis字段(PluginFieldType::kINT32),createPlugin()从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, type=trt.PluginFieldType.INT32) field_collection = trt.PluginFieldCollection([axis_attr]) plugin = plugin_creator.create_plugin( name="CustomHardmax", field_collection=field_collection, phase=trt.TensorRTPhase.BUILD )随后通过network.add_plugin_v3(inputs=[input_layer], shape_inputs=[], plugin=plugin)把插件挂到强类型网络上。这证明了一条重要事实:同一个 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 的argmax+put_along_axis构造 one-hot 结果; - 用插件构建引擎执行推理,断言与参考实现逐元素完全相等。
测试覆盖了负数 axis、多维张量、正负值混合输入等多种边界情况,可作为后续开发自定义插件的通用测试范式。
版本演进记录(Changelog)
README 记录了该样例的演进历程,能帮助读者理解当前代码形态的来由:
- 2026 年 3 月:Hardmax 插件从 IPluginV2DynamicExt 迁移到 IPluginV3(当前源码即 IPluginV3 形态);
- 2025 年 10 月:迁移到强类型(strongly typed)API,对应
sample.py中STRONGLY_TYPED标志的使用; - 2025 年 8 月:不再支持 Python < 3.10;
- 2024 年 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),仅供参考