资讯动态

TensorRT 自定义算子插件实战(二):双输入融合算子 customGatedTanh

发布时间:2026/10/8 22:38:25 来源:尧图企业网站定制
承接上一篇TensorRT 自定义算子插件实战一从零手写 customScaledTanh。上一篇我们用简单的单输入、逐元素算子走通了插件的完整骨架本篇把难度往前推一步实现一个双输入、融合算子 customGatedTanhout tanh(a·x) × sigmoid(b·y)重点不是重讲一遍插件的外壳而是对比出从单输入到双输入插件代码到底要改哪些地方。目录一两个算子到底差在哪二、本案例的算子与网络三、Python 端导出双输入自定义 op四、C 端头文件——两个类的声明五、CUDA 核函数六、C 端Plugin 类的实现——只看多输入改了什么6.1 差异一enqueue —— 取两个输入指针6.2 差异二supportsFormatCombination —— 多一个 case差异总结七、C 端PluginCreator 类的实现差异八、构建与验证九、参数的传递过程十、实践经验十一、总结与展望一两个算子到底差在哪我做的两个算子 customScaledTanh上一篇与 customGatedTanh本篇90% 的插件代码完全一样两个类、三个构造函数、serialize、clone、Creator 那套真正不同的只有四处。先给一张对照表全文就围绕它展开维度customScaledTanh单输入customGatedTanh双输入融合算子公式k·tanh(a·x)tanh(a·x) · sigmoid(b·y)网络拓扑Input→1 个 Conv→customInput→2 个独立 Conv→customPython symbolicg.op(…, x, k_f, a_f) 单输入g.op(…, x, y, a_f, b_f)双输入核函数签名(input, output, k, a, …) 单指针(input1, input2, output, a, b, …)双指针enqueue 取数只取 inputs[0]取 inputs[0] 和 inputs[1]supportsFormatCombination2 个 case1入1出3 个 case2入1出底层实现用 tanh手写 1/(1expf(-z))FP16 kernel有实现并未实现这里埋了个坑见第十节一句话概括本篇要传达的核心插件从单输入升级到双输入本质是让 TRT 知道你多了个输入——它体现在 Python 端多传一个张量、enqueue多拿一个指针、supportsFormatCombination多一个 case、kernel多读一个张量四处。而融合的价值是把 tanh、sigmoid、mul 这三个原本要分三个 kernel 执行的运算合并成一个 kernel省启动开销、少一次中间张量的内存往返。二、本案例的算子与网络算子定义双输入、逐元素融合输出形状仍等于输入形状o u t p u t t a n h ( a ⋅ x ) ⋅ s i g m o i d ( b ⋅ y ) output tanh(a·x) · sigmoid(b·y)outputtanh(a⋅x)⋅sigmoid(b⋅y)网络结构关键变化两个独立卷积并联各自输出一个分支汇入自定义算子两个卷积的输入都是同一个 input0各自用独立权重把单通道升到 2 通道卷积的输出 x1、x2 作为 customGatedTanh 的两个输入因为 x1、x2 形状相同都是 [1,2,5,5]输出 output0 与它们形状也相同所以 getOutputDimensions 仍然直接返回 inputs[0]逐元素融合算子输出 任一输入形状。其输入与输出元素是一一对应的如下图所示这就是融合算子的典型形态两个分支汇合进一个节点在注意力、门控GLU、LSTM这类结构里非常常见。三、Python 端导出双输入自定义 opsymbolic 方法接收两个张量 x、y并一起传进 g.op。python 部分内容不多完整贴出importtorchimporttorch.onnximporttorch.nnasnnimportonnximportonnxsimclassCustomGatedTanhImpl(torch.autograd.Function):staticmethoddefsymbolic(g,x,y,a,b):# 双输入x、y 两个张量都作为节点输入returng.op(custom::customGatedTanh,x,y,a_fa,b_fb)staticmethoddefforward(ctx,x,y,a,b):returntorch.tanh(a*x)*torch.sigmoid(b*y)classCustomGatedTanh(nn.Module):def__init__(self,a,b):super().__init__()self.aa self.bbdefforward(self,x,y):returnCustomGatedTanhImpl.apply(x,y,self.a,self.b)classModel(torch.nn.Module):def__init__(self):super().__init__()self.conv1nn.Conv2d(1,2,(3,3),padding1)# 分支 1self.conv2nn.Conv2d(1,2,(3,3),padding1)# 分支 2self.actCustomGatedTanh(2,1.5)forminself.modules():ifisinstance(m,nn.Conv2d):nn.init.kaiming_normal_(m.weight,modefan_out,nonlinearityrelu)defforward(self,x):x1self.conv1(x)# 分支 1x2self.conv2(x)# 分支 2xself.act(x1,x2)# 双输入融合returnxdefexport_norm_onnx(input,model):file./sample_customGatedTanh.onnxtorch.onnx.export(modelmodel,args(input,),ffile,input_names[input0],output_names[output0],opset_version11)model_onnxonnx.load(file)model_onnx,checkonnxsim.simplify(model_onnx)assertcheck onnx.save(model_onnx,file)print(onnx exported simplified)if__name____main__:torch.manual_seed(1)inputtorch.rand(1,1,5,5)modelModel().eval()export_norm_onnx(input,model)与上篇算子的差异symbolic(g, x, y, a, b)参数从单输入 x变成双输入 x, yg.op(“custom::customGatedTanh”, x, y, …) 把两个张量都作为节点输入。forward(ctx, x, y, a, b)计算从 k·tanh(a·x) 变成 tanh(a·x)·sigmoid(b·y)融合了 x 和 y。Model 里是两个独立卷积conv1、conv2各自 (1→2) 输出分支最后 self.act(x1, x2) 汇合。导出的 ONNX 里customGatedTanh 节点会有两条输入边——这正是后面 TRT 侧多输入信息的来源。四、C 端头文件——两个类的声明头文件和上一篇一模一样两个类、三个构造函数、同一套接口唯一的差别是把 mParams 里的两个成员从 {k, a} 换成了 {a, b}。这里不再重复贴全文只给差异片段相关类名称最好都换掉不然不方便维护classCustomGatedTanhPlugin:publicIPluginV2DynamicExt{public:CustomGatedTanhPlugin(conststd::stringname,floata,floatb);// parse、clone 用CustomGatedTanhPlugin(conststd::stringname,constvoid*buffer,size_t length);// 反序列化用// ... 其余接口声明和上一篇完全相同 ...private:conststd::string mName;std::string mNamespace;struct{floata;// 参数从 {k, a} 换成 {a, b}floatb;}mParams;};classCustomGatedTanhPluginCreator:publicIPluginCreator{// ... 和上一篇完全相同 ...};五、CUDA 核函数核函数是本篇差异最直观的地方kernel 从读一个张量变成读两个张量并且 sigmoid 是手写的。cu 内容不多完整贴出// custom-gatedTanh.cu#includecuda_runtime.h#includemath.h#includecuda_fp16.h// 双输入 kernelinput1、input2 两个指针__global__voidcustomGatedTanhKernel(constfloat*input1,constfloat*input2,float*output,constfloata,constfloatb,constintnElements){constintindexblockIdx.x*blockDim.xthreadIdx.x;if(indexnElements)return;// 融合tanh(a·x) * sigmoid(b·y)sigmoid 手写为 1/(1expf(-z))output[index]tanh(a*input1[index])*1.0f/(1.0fexpf(-b*input2[index]));}voidcustomGatedTanhImpl(constfloat*input1,constfloat*input2,float*outputs,constfloata,constfloatb,constintnElements,cudaStream_t stream){dim3blockSize(256,1,1);dim3gridSize(ceil(float(nElements)/256),1,1);customGatedTanhKernelgridSize,blockSize,0,stream(input1,input2,outputs,a,b,nElements);}与上篇算子的差异双输入指针kernel 签名从 const float* input 变成 const float* input1、const float* input2host 封装 customGatedTanhImpl 也多了 input2。索引仍一一对应虽然有两个输入但都是逐元素、且两个输入和输出形状完全相同所以线程 index 仍然同时作为两个输入和输出的下标——一个线程同时读两个输入的同一位置、算一个输出。这是逐元素融合和单输入逐元素在索引上本质一致的原因。sigmoid 手写CUDA 里没有现成的 sigmoid()用 1.0f / (1.0f expf(-z)) 实现expf 是 float 版指数函数。相比 Python 里直接用 torch.sigmoid这是从 PyTorch 到 CUDA 的翻译——库函数要自己用基础数学函数拼出来。一个差异本篇的隐患这个 kernel只有 FP32 版本没有像第一篇那样的 FP16 分支。但下面 supportsFormatCombination 却声明支持 FP16——这会导致一个潜在问题详见第十节。六、C 端Plugin 类的实现——只看多输入改了什么大部分接口三个构造、getPluginType、serialize、clone、destroy……和上一篇是一样的不再展示。真正的差异集中在两个方法enqueue取两个输入和 supportsFormatCombination多一个 case。6.1 差异一enqueue —— 取两个输入指针int32_tCustomGatedTanhPlugin::enqueue(...)noexcept{intnElements1;for(inti0;iinputDesc[0].dims.nbDims;i){nElements*inputDesc[0].dims.d[i];}customGatedTanhImpl(static_castconstfloat*(inputs[0]),// 第一个输入 x1static_castconstfloat*(inputs[1]),// 第二个输入 x2 ← 多出来的static_castfloat*(outputs[0]),mParams.a,mParams.b,nElements,stream);return0;}对比上一篇上一篇 enqueue 只取 inputs[0]单输入本篇多取一个inputs[1]。inputs 是一个 const void* const* 指针数组inputs[0] 指向第一个输入、inputs[1] 指向第二个输入、outputs[0] 指向输出——双输入就是这个数组里多一个元素。6.2 差异二supportsFormatCombination —— 多一个 caseboolCustomGatedTanhPlugin::supportsFormatCombination(int32_tpos,constPluginTensorDesc*inOut,int32_tnbInputs,int32_tnbOutputs)noexcept{switch(pos){case0:// 输入 0return(inOut[0].typeDataType::kFLOAT||inOut[0].typeDataType::kHALF)inOut[0].formatTensorFormat::kLINEAR;case1:// 输入 1 ← 多出来的return(inOut[1].typeDataType::kFLOAT||inOut[1].typeDataType::kHALF)inOut[1].formatTensorFormat::kLINEAR;case2:// 输出return(inOut[2].typeDataType::kFLOAT||inOut[2].typeDataType::kHALF)inOut[2].formatTensorFormat::kLINEAR;default:returnfalse;}}对比上一篇的 2 个 case输入 0 输出本篇变成3 个 case输入 0 输入 1 输出。inOut 数组的长度 输入数 输出数 2 1 3pos 从 0 到 2。每多一个输入这里就多一个 case。差异总结方法单输入 scaledTanh双输入 gatedTanhenqueue 取数inputs[0]inputs[0] inputs[1]supportsFormatCombinationcase 0、1case 0、1、2除此之外插件的其他实现都不用改。七、C 端PluginCreator 类的实现差异Creator 和上一篇几乎完全一致唯一差异是参数名从 k/a 变成 a/bCustomGatedTanhPluginCreator::CustomGatedTanhPluginCreator(){mAttrs.emplace_back(PluginField(a,nullptr,PluginFieldType::kFLOAT32,1));mAttrs.emplace_back(PluginField(b,nullptr,PluginFieldType::kFLOAT32,1));mFC.nbFieldsmAttrs.size();mFC.fieldsmAttrs.data();}IPluginV2*CustomGatedTanhPluginCreator::createPlugin(constchar*name,constPluginFieldCollection*fc)noexcept{floata0,b0;std::mapstd::string,float*paramMap{{a,a},{b,b}};for(inti0;ifc-nbFields;i){if(paramMap.find(fc-fields[i].name)!paramMap.end()){*paramMap[fc-fields[i].name]*reinterpret_castconstfloat*(fc-fields[i].data);}}returnnewCustomGatedTanhPlugin(name,a,b);}deserializePlugin、getPluginName、REGISTER_TENSORRT_PLUGIN宏等都和上一篇逐字相同。注意Creator 完全不涉及输入个数——它只负责参数a、b输入边是 TRT 从 ONNX 图里解析后自动接上的。八、构建与验证main函数比较简单直接读取相关的onnx文件进行本地引擎构建再推理即可。与上篇的唯一区别就是推理时候读取的onnx文件不一样。#include iostream#include memory#include utils.hpp#include model.hppusing namespace std;int main(int argc, char const *argv[]){Model model(models/onnx/sample_customGatedTanh.onnx, Model::precision::FP16);if(!model.build()){LOGE(fail in building model);return0;}if(!model.infer()){LOGE(fail in infering model);return0;}return0;}验证方式与上一篇相同分别用 PyTorch 跑 ONNX 模型、用 C 加载 TRT 引擎跑自定义插件打印输出对比。python程序的输出结果为c程序的推理结果为在我的实现里Python 与 C 的推理结果完全一致证明双输入导出 → 双输入解析 → 双输入 enqueue → 双输入 kernel这条链路没有断裂。九、参数的传递过程参数的传递过程与上篇是一致的可以再重新梳理一遍可以看到只有参数的名称发生了变化整个传递过程与以前是一直的。十、实践经验supportsFormatCombination 开了 FP16但 kernel 没有 FP16 版本 —— 这是个隐患。我的 supportsFormatCombination 三个 case 都允许 kFLOAT || kHALF但 enqueue 无条件调用 FP32 的 customGatedTanhImplcu 里也没有 __half kernel。这意味着如果 TRT 真按 FP16 建了引擎enqueue 会把半精度数据当 float 读结果就错了。第一篇我为 FP16 单独写了 __half kernel 和 enqueue 的类型分支本篇并未补充——正确的做法要么删掉 kHALF声明只支持 FP32要么像第一篇那样补上 FP16 分支。这也是声明的格式支持和kernel 实际实现必须保持一致的真实教训。多输入的下标仍可逐元素对齐的前提两个输入必须形状一致、且和输出一致本文 x1、x2、output 都是 [1,2,5,5]才能用一个 index 同时索引三者。如果两个输入形状不同getOutputDimensions 和 kernel 的索引逻辑就要重写——那是更高阶的情况。多输入不是重写某个输入个数接口IPluginV2DynamicExt 没有 getNbInputs()多输入只体现在 enqueue 取指针、supportsFormatCombination 加 case 两处。十一、总结与展望本篇在上一篇的插件骨架上把算子从单输入、逐元素升级到双输入、融合并重点对比出真正的差异Python 端symbolic 多收一个张量 yg.op 传两个输入模型里是两个独立卷积并联 → 双输入节点的融合拓扑。CUDA 端kernel 读两个指针一个线程同时读两个输入的同一位置算一个输出sigmoid 用 1/(1expf(-z)) 手写。C 端enqueue 多取 inputs[1]supportsFormatCombination 多一个 case输入 1其余外壳逐字不变。核心认知多输入对插件外壳几乎无影响差异全在取指针和声明格式两处而融合的价值是把多个运算合并成一个 kernel。下一篇我实现一个更难的形态customMaxPool——输出尺寸变化的窗口算子2×2 最大池化、无参数。它的难点不再多输入而是 getOutputDimensions 要真正计算输出尺寸、kernel 要做线程输出元素、反查输入窗口的索引映射是空间邻域算子的骨架。敬请期待

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑