news 2026/9/15 14:29:53

TensorRT Python 插件实战:基于 IPluginV3 与 size tensor 实现数据依赖输出形状的 NonZero 算子

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
TensorRT Python 插件实战:基于 IPluginV3 与 size tensor 实现数据依赖输出形状的 NonZero 算子

TensorRT Python 插件实战:基于 IPluginV3 与 size tensor 实现数据依赖输出形状的 NonZero 算子

【免费下载链接】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

本指南以 NVIDIA TensorRT 开源仓库中的non_zero_plugin示例(samples/python/non_zero_plugin/README.md)为核心,讲解如何使用 Python 编写一个输出形状依赖于输入数据内容的自定义插件:NonZero(查找张量中非零元素的下标)。文中将完整还原示例的工程实现思路,并结合 non_zero_plugin.py、plugin_utils.py 等源码深入剖析 IPluginV3 插件体系、size tensor 声明机制、CUDA Python/PyTorch 双后端执行路径,最终帮助你掌握在 TensorRT 中落地「输出形状依赖输入数值」的自定义层的完整方案。

NonZero 插件示例概述

non_zero_plugin是一个完全用 Python 实现的 TensorRT 插件示例,用于执行 NonZero 运算——找出输入张量中所有非零元素的下标。它构建并运行一个只包含单个NonZeroPlugin节点的 TensorRT 引擎,展示了两个关键能力:

  1. 用 Python 编写自定义插件:插件类直接继承trt.IPluginV3及其能力接口,无需编写 C++ 代码;
  2. 支持数据依赖的输出形状:输出张量的维度(非零元素的个数 K)无法仅由输入形状推导,必须由运行时的输入数值决定。

插件支持两种执行后端(通过--backend参数切换):

  • cuda_python:使用 CUDA Python 绑定(driver API + NVRTC 运行时编译 CUDA kernel);
  • torch:直接调用 PyTorch 的torch.nonzero()完成计算。

为什么需要 IPluginV3:数据依赖输出形状的突破口

IPluginV3及其关联接口出现之前,TensorRT 插件的输出形状只能依赖输入形状,无法依赖输入数值。也就是说,像 NonZero 这样「输出大小取决于输入里有多少个非零元素」的算子,用老式插件接口(如 V2 插件)是难以表达的。

IPluginV3OneBuild作为IPluginV3的 build 阶段能力接口,提供了解决这一问题的机制:插件通过get_output_shapes()方法向 TensorRT builder 描述输出形状表达式,其中数据依赖的维度必须用size tensor来表达。

关于该接口的官方语义,可参见仓库中的 Python 绑定文档 pyPluginDoc.h:get_output_shapes()返回用于从输入张量形状计算输出张量形状的表达式,由IBuilder在网络分析阶段调用。

插件类设计:组合四个接口

示例中的NonZeroPlugin一次性继承了四个接口(non_zero_plugin.py):

class NonZeroPlugin(trt.IPluginV3, trt.IPluginV3OneCore, trt.IPluginV3OneBuild, trt.IPluginV3OneRuntime): def __init__(self, backend=None): trt.IPluginV3.__init__(self) trt.IPluginV3OneCore.__init__(self) trt.IPluginV3OneBuild.__init__(self) trt.IPluginV3OneRuntime.__init__(self) self.num_outputs = 2 self.plugin_namespace = "" self.plugin_name = "NonZeroPlugin" self.plugin_version = "1" ...

各接口的职责:

接口阶段职责
trt.IPluginV3基础插件对象基类
trt.IPluginV3OneCore核心提供插件名称、版本、命名空间,以及get_capability_interface()能力分发
trt.IPluginV3OneBuild构建描述输出形状/数据类型、格式支持、构建期配置
trt.IPluginV3OneRuntime运行提供enqueue()在 GPU 上执行实际计算

其中get_capability_interface()把自身返回给 TensorRT,使其能够同时获得 build 与 runtime 能力:

def get_capability_interface(self, type): return self

attach_to_context()返回插件的克隆副本供执行上下文使用,clone()通过复制__dict__实现深拷贝(non_zero_plugin.py):

def attach_to_context(self, context): return self.clone() def clone(self): cloned_plugin = NonZeroPlugin() cloned_plugin.__dict__.update(self.__dict__) return cloned_plugin

声明 size tensor:让 builder 知道输出有多大

NonZero 插件处理形状为 R × C 的二维输入张量。假设其中有 K 个非零元素,且要求按行序输出(每组下标占一行),那么输出形状为 K × 2。

第一步:描述不依赖数据的维度

输出第二维恒为 2,可以直接用IExprBuilder构造常量表达式:

# output_dims[0] = trt.DimsExprs(2) output_dims[0][1] = exprBuilder.constant(2)

第二步:构造 upper-bound 与 opt 表达式

数据依赖维度的范围无法静态确定,因此必须为 size tensor 提供**上界(upper-bound)最优值(opt)**两个IDimensionExpr,TensorRT 会据此做自动调优并为输出张量分配内存:

  • 对于输入规模未知的情况,上界取输入元素总数(R × C,即所有元素都可能非零):
upper_bound = exprBuilder.operation(trt.DimensionOperation.PROD, inputs[0][0], inputs[0][1])
  • 一个合理的估计是「一半元素非零」,因此最优值取上界整除 2:
opt_value = exprBuilder.operation(trt.DimensionOperation.FLOOR_DIV, upper_bound, exprBuilder.constant(2))

第三步:声明 size tensor 并回填输出维度

size tensor 是类型为trt.int32trt.int64标量输出,必须作为插件的一个输出存在。IExprBuilder.declare_size_tensor()需要指定它位于哪个输出索引,示例将它放在非零下标输出之后(索引 1):

num_non_zero_size_tensor = exprBuilder.declare_size_tensor(1, opt_value, upper_bound)

随后把非零下标输出(索引 0)的第一维指向该 size tensor:

# output_dims[0] = trt.DimsExprs(0) output_dims[0][0] = num_non_zero_size_tensor

完整实现位于 non_zero_plugin.py:

def get_output_shapes(self, inputs, shape_inputs, exprBuilder): # First output is 2-D # Second output is a size tensor, which must be declared a scalar (0-D) output_dims = [trt.DimsExprs(2), trt.DimsExprs(0)] upper_bound = exprBuilder.operation(trt.DimensionOperation.PROD, inputs[0][0], inputs[0][1]) opt_value = exprBuilder.operation(trt.DimensionOperation.FLOOR_DIV, upper_bound, exprBuilder.constant(2)) num_non_zero_size_tensor = exprBuilder.declare_size_tensor(1, opt_value, upper_bound) output_dims[0][0] = num_non_zero_size_tensor output_dims[0][1] = exprBuilder.constant(2) return output_dims

注意两个关键约束:

  1. size tensor 必须声明为 0 维(标量),因此output_dims[1]trt.DimsExprs(0)构造;
  2. size tensor 的数据类型必须是trt.int32trt.int64,示例中为trt.int64

仓库中 pyPluginDoc.h 对declare_size_tensor的语义给出了与示例完全一致的官方解释:插件写 K 到第二个输出,TensorRT 用constant()declare_size_tensor(1, ...)分别构造表示 2 与 K 的IDimensionExpr;auto-tuning 使用 opt 值,内存分配使用 upper-bound 值。

输出数据类型与格式组合

get_output_data_types()声明两个输出的类型——非零下标为trt.int32,size tensor 为trt.int64

def get_output_data_types(self, input_types): return [trt.DataType.INT32, trt.DataType.INT64]

supports_format_combination()对每个位置(pos)做格式与类型校验(non_zero_plugin.py):

def supports_format_combination(self, pos, in_out, num_inputs): assert num_inputs == 1 assert pos < len(in_out) type_ok = False # first input should be float16 or float32 if pos == 0: type_ok = in_out[0].desc.type == trt.DataType.FLOAT or in_out[0].desc.type == trt.DataType.HALF elif pos == 1: type_ok = in_out[1].desc.type == trt.DataType.INT32 else: # pos == 2 # size tensor outputs must be NCHW INT64 type_ok = in_out[2].desc.type == trt.DataType.INT64 return in_out[pos].desc.format == trt.TensorFormat.LINEAR and type_ok

可以看出:

  • 输入支持FLOAT(fp32)与HALF(fp16),与--precision参数对应;
  • 下标输出必须是INT32
  • size tensor 输出必须为INT64且使用LINEAR格式,这与 C++ 版 sampleNonZeroPlugin.cpp 中的supportsFormatCombination()实现完全一致。

注册 Plugin Creator:V3 插件的必备伴侣

与 V2 插件需要IPluginCreator类似,V3 插件必须注册一个实现trt.IPluginCreatorV3One接口的 creator(non_zero_plugin.py):

class NonZeroPluginCreator(trt.IPluginCreatorV3One): def __init__(self): trt.IPluginCreatorV3One.__init__(self) self.name = "NonZeroPlugin" self.plugin_namespace = "" self.plugin_version = "1" self.field_names = trt.PluginFieldCollection( [trt.PluginField("backend", np.array([]), trt.PluginFieldType.CHAR)] ) def create_plugin(self, name, fc, phase): backend = None for f in fc: if f.name == "backend": backend = f.data[:-1] if f.data[-1] == 0 else f.data return NonZeroPlugin(backend)

creator 通过PluginField接收backend字段(CHAR类型),在create_plugin()中解析后构造插件实例。主程序里通过插件注册表注册 creator:

plg_registry = trt.get_plugin_registry() my_plugin_creator = NonZeroPluginCreator() plg_registry.register_creator(my_plugin_creator, "")

插件序列化字段由get_fields_to_serialize()提供,backend字符串被编码为字节后放入PluginField

def get_fields_to_serialize(self): return trt.PluginFieldCollection( [trt.PluginField("backend", self.backend.encode(), trt.PluginFieldType.CHAR)] )

构建引擎:ONNX 与 Network API 两种路径

示例支持--net_type {onnx,inetdef}两种建网方式,都使用**强类型(strongly typed)**网络。

ONNX 路径(onnx)

用 onnx-graphsurgeon 构造一个仅含NonZeroPlugin节点的 ONNX 图,节点输出为Y(下标)与Y_num(非零个数),并通过attrs传入backend(non_zero_plugin.py):

onnx_path = "test_NonZeroPlugin.onnx" inputX = gs.Variable(name="X", shape=inp_shape, dtype=precision) Y = gs.Variable(name="Y", dtype=np.int32) Y_num = gs.Variable(name="Y_num", dtype=np.int64) nonZeroPluginNode = gs.Node( name="NonZeroPlugin", op="NonZeroPlugin", inputs=[inputX], outputs=[Y, Y_num], attrs={"backend": args.backend.encode()}, ) graph = gs.Graph(nodes=[nonZeroPluginNode], inputs=[inputX], outputs=[Y], opset=16) onnx.save(gs.export_onnx(graph), onnx_path) # build engine build_engine = EngineFromNetwork( NetworkFromOnnxPath(onnx_path, strongly_typed=True), CreateConfig() )

Network API 路径(inetdef)

先从注册表取回 creator,用PluginField构造字段集合并创建插件对象,再调用network.add_plugin_v3()将插件节点加入网络(non_zero_plugin.py):

builder, network = create_network(strongly_typed=True) plg_creator = plg_registry.get_creator("NonZeroPlugin", "1", "") plugin_fields_list = [ trt.PluginField("backend", args.backend.encode(), trt.PluginFieldType.CHAR) ] pfc = trt.PluginFieldCollection(plugin_fields_list) plugin = plg_creator.create_plugin("NonZeroPlugin", pfc, trt.TensorRTPhase.BUILD) # Populate network inputX = network.add_input(name="X", dtype=trt.float32 if precision==np.float32 else trt.float16, shape=inp_shape) out = network.add_plugin_v3([inputX], [], plugin) out.get_output(0).name = "Y" network.mark_output(tensor=out.get_output(0)) build_engine = engine_from_network((builder, network), CreateConfig())

add_plugin_v3()的签名可从 Python 绑定源码 pyGraph.cpp 确认:接收输入张量列表、shape 输入列表与插件对象。在强类型网络中,输入张量需显式指定 dtype(trt.float32trt.float16)。

两种路径的等价实现可以参考 C++ 版 sampleNonZeroPlugin.cpp:其中同样通过getPluginRegistry()->getCreator("NonZeroPlugin", "0", "")获取 creator,用addPluginV3()建层并markOutput()标记两个输出。

运行期计算:CUDA Python 与 PyTorch 双后端

核心执行逻辑在enqueue()中,按 backend 分两条路径(non_zero_plugin.py)。

CUDA Python 后端:NVRTC 即时编译内核

cuda_python路径读取输入描述中的 R、C 维度,计算 grid 规模(block 256 线程),并把输入/输出设备指针打包成 kernel 参数后通过cuLaunchKernel启动:

R = input_desc[0].dims[0] C = input_desc[0].dims[1] blockSize = 256 numBlocks = int((C + blockSize - 1) // blockSize) d_in = np.array([inputs[0]], dtype=np.uint64) d_out_0 = np.array([outputs[0]], dtype=np.uint64) d_out_1 = np.array([outputs[1]], dtype=np.uint64) args = [d_in, d_out_0, d_out_1, np.array(R, dtype=np.uint32), np.array(C, dtype=np.uint32)] kernelArgs = np.array([arg.ctypes.data for arg in args], dtype=np.uint64) ...

CUDA 内核以字符串常量内嵌在 Python 文件中,例如 fp32 版本(non_zero_plugin.py):

extern "C" __global__ void find_non_zero_indices_float( float const* X, int* indices, unsigned long long* count, int R, int C) { static_assert(sizeof(unsigned long long) == 8U, "unsigned long long must be 8 bytes in NVCC"); int row = blockIdx.x * blockDim.x + threadIdx.x; // Check if the row index is within bounds if (row < R) { for (int col = 0; col < C; ++col) { if (X[col + C * row] != 0.F) { // Increment count atomically and get the previous value unsigned long long index = atomicAdd(count, 1ULL); indices[2 * index] = row; indices[2 * index + 1] = col; } } } }

内核的要点:

  • 每个线程负责一行,遍历该行所有列;
  • 命中非零元素时通过atomicAdd原子递增计数器并取得该元素在输出中的写入位置(行序排列:indices[2*index]=rowindices[2*index+1]=col);
  • 计数器count正是写入 size tensor(输出 1)的值,从而使运行期输出形状与写入数据一致。

KernelHelper(plugin_utils.py)封装了 NVRTC 的完整流程:创建 Program、按设备计算能力(通过cudaDeviceGetAttribute查询 major/minor)生成--gpu-architecture=smXX参数、编译并加载 cubin/PTX,最终通过cuModuleGetFunction取得 kernel 函数句柄。configure_plugin()on_shape_change()中通过cuDeviceGet(0)获取设备句柄供 KernelHelper 使用。

fp16 内核与 fp32 内核结构相同,只是用half const z = static_cast<half>(0.F)比较。从代码结构看,CUDA Python 路径的优势是不依赖 PyTorch 等框架,仅需 cuda-python 与 NVRTC 即可运行。

PyTorch 后端:torch.nonzero

torch路径利用UnownedMemory(plugin_utils.py)把 TensorRT 分配的设备指针包装成 CuPy 数组视图,再转为torch.Tensor,直接调用torch.nonzero()

inp_mem = UnownedMemory(inputs[0], input_desc[0].dims, inp_dtype) out_mem = UnownedMemory( outputs[0], 2 * volume(input_desc[0].dims), np.int32 ) out_1_mem = UnownedMemory(outputs[1], 1, np.int64) a_t = torch.as_tensor(inp_mem.d, device="cuda") out = torch.nonzero(a_t) out_mem.d[: volume(out.shape)] = cp.reshape(cp.asarray(out), (-1,)) cp.copyto(out_1_mem.d, cp.reshape(cp.asarray([out.shape[0]]), (-1,)))

torch.nonzero()默认按行序返回下标(shape K × 2),正好与插件声明的输出布局一致;out.shape[0]即 K,写入 size tensor 输出。主程序在torch后端启动时会先调用torch.cuda.init()初始化 CUDA(cuda_python后端则调用cudaFree(0))。

运行示例与结果验证

环境依赖

示例依赖见 requirements.txt,核心包括:

  • cuda-python==12.9.0(CUDA Python 后端所需);
  • cupy-cuda12x(设备指针包装与数据搬运);
  • torch(PyTorch 后端所需);
  • polygraphy>=0.50.1(建网、构建引擎与运行推理的封装);
  • onnx-graphsurgeononnx(ONNX 建图);
  • numpy==1.26.4
  • 部分包通过--extra-index-url https://pypi.ngc.nvidia.com从 NGC PyPI 索引安装。

运行命令

在满足依赖的 Python 3.10+ 环境中执行:

python3 non_zero_plugin.py [-h] [--precision {fp32,fp16}] [--backend {cuda_python,torch}] [--net_type {onnx,inetdef}]

命令行参数(可在non_zero_plugin.py主程序部分找到默认值):

参数默认值可选值说明
--precisionfp32fp32,fp16输入张量精度,fp16 时输入与内核均使用 HALF 类型
--backendtorchcuda_python,torch运行期计算后端
--net_typeonnxonnx,inetdef建网方式:ONNX 图或 Network API
-h/--help显示完整选项说明

输入数据与结果校验

示例以np.random.normal生成 128 × 128 的随机输入,并随机将部分元素置零;随后用 NumPy 的np.nonzero()计算参考结果Y_ref,通过 Polygraphy 的TrtRunner执行推理(non_zero_plugin.py):

Y_ref = np.transpose(np.nonzero(X)) with TrtRunner(build_engine, "trt_runner") as runner: outputs = runner.infer({"X": X}) Y = outputs["Y"] Y = Y[np.lexsort(np.fliplr(Y).T)] if np.allclose(Y, Y_ref): print("Inference result correct!") else: print("Inference result incorrect!")

由于插件内核以行序写入下标(每行一个非零元素下标对),而np.nonzero(X)的转置结果在行上并非必然有序,示例用np.lexsort(np.fliplr(Y).T)先对结果按行、列排序再比较,以保证与参考结果一致。

运行成功时的标志性输出:

Inference result correct!

与 C++ 版 NonZero 插件的对照

仓库同时提供了 C++ 版实现 samples/sampleNonZeroPlugin,两者互为印证:

维度Python 版C++ 版
主文件non_zero_plugin.pysampleNonZeroPlugin.cpp
内核实现内嵌字符串 + NVRTC 运行时编译nonZeroKernel.cu 编译期生成
接口trt.IPluginV3系列nvinfer1::IPluginV3系列
插件属性backend(cuda_python/torch)rowOrder(行序/列序)
size tensor 声明declare_size_tensor(1, opt_value, upper_bound)exprBuilder.declareSizeTensor(1, *optValue, *upperBound)(sampleNonZeroPlugin.cpp)

C++ 版还展示了rowOrder=false的列序输出场景(输出形状为 2 × K),此时内核需要分两次启动:第一次仅统计非零总数(存入 workspace),第二次依据总数在正确位置写入下标(sampleNonZeroPlugin.cpp),并申请sizeof(int64_t)的 workspace。此外 C++ 版要求输入必须为二维(inputs[0].nbDims != 2getOutputShapes()返回 -1)。这些细节说明了数据依赖输出形状插件在工程上需要注意的边界与资源问题。

关键注意事项与常见陷阱

  1. size tensor 必须是 0 维标量输出trt.DimsExprs(0),并放在声明的输出索引处,不能省略;
  2. size tensor 的类型约束:文档要求INT32INT64,示例与supports_format_combination()均强制 size tensor 输出为INT64+LINEAR格式;
  3. 强类型网络NetworkFromOnnxPath(..., strongly_typed=True)create_network(strongly_typed=True)缺一不可,输入 dtype 需显式指定;
  4. creator 生命周期:creator 对象必须在 build 与 infer 全程存活,注册表条目才能保持有效(C++ 版注释也强调了这一点,见 sampleNonZeroPlugin.cpp);
  5. 内存上限:TensorRT 依据 upper-bound 分配输出缓冲,因此上界必须真实覆盖所有可能输出规模(本示例为 R × C),opt 值则影响自动调优质量;
  6. 结果有序性:内核按行序写入、参考结果需排序后比较,避免因遍历顺序不同导致校验误报。

总结

通过non_zero_plugin示例可以看到,TensorRT 的IPluginV3系列接口(IPluginV3OneBuildIPluginV3OneRuntimeIPluginV3OneCore)为 Python 开发者提供了完整的自定义层能力,而 size tensor 机制则让「输出形状依赖输入数值」的算子(NonZero、稀疏类算子、动态 ROI 类算子等)能够被正确建网、调优与分配内存。示例中 CUDA Python(NVRTC 运行时编译)与 PyTorch 双后端的切换方式,也为在真实项目中权衡「零框架依赖的裸 CUDA 执行」与「复用成熟框架算子」提供了可直接复用的工程范式。

进一步的扩展学习资源位于仓库内:Python 插件官方文档说明(含declare_size_tensorget_output_shapessupports_format_combination等接口语义)、C++ 版 NonZero 插件示例,以及 samples/python/plugin_utils.py 中可复用的 CUDA 上下文、UnownedMemory、KernelHelper 等工具类。

【免费下载链接】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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 14:28:42

前端摇一摇红包实现:设备方向传感器+CSS3动画+jQuery调度

简介&#xff1a;这是一份面向前端初学者与网页交互开发者的学习型代码资源&#xff0c;聚焦HTML5CSS3动画与jQuery事件驱动的趣味抽奖交互实现。资源通过模拟手机摇一摇触发红包开启的完整流程&#xff0c;帮助开发者掌握Canvas绘图、CSS3关键帧动画、transform变换及jQuery D…

作者头像 李华
网站建设 2026/9/15 14:25:20

CAXA在非标自动化设计中的高效应用与实战技巧

1. 非标自动化设计师的日常挑战非标自动化设计这个行当&#xff0c;最让人头疼的就是客户那些"千奇百怪"的需求。上周刚遇到个案例&#xff1a;某食品厂要求设计一条能自动给月饼"盖章"的生产线&#xff0c;但特殊之处在于——每个月饼要盖三个不同图案的章…

作者头像 李华
网站建设 2026/9/15 14:24:00

工控安全网关技术解析与信创适配实践

我无法根据您提供的输入内容生成符合要求的博文。原因如下&#xff1a;输入中缺少必要字段&#xff1a;项目正文、关键词、摘要描述三项均为空&#xff08;仅包含空行或未提供实质内容&#xff09;&#xff0c;而根据任务定义&#xff0c;这三项是进行专业拆解与延展的基础原料…

作者头像 李华
网站建设 2026/9/15 14:21:55

梆梆加固脱壳实战:从内存Dump到DEX修复全流程

搞安全的兄弟&#xff0c;对“梆梆加固”这四个字应该都不陌生。这几年做Android样本分析、App隐私合规检测、漏洞挖掘&#xff0c;碰到梆梆加固的频率极高。这玩意儿在移动应用加固市场里占了不少份额&#xff0c;很多银行类、政企类、头部互联网App都在用。它的保护能力也确实…

作者头像 李华