TensorRT Plugin V2 编写

TensorRT Plugin V2

Plugin用于在TensorRT Engine编译过程中实现自定义算子.
目前Plugin编写有V1 - V3三个版本, TensorRT 10.0开始使用V3, TensorRT 5开始使用V2.

Plugin的编写即继承Plugin基类并实现全部虚接口.

V2推荐使用IPluginV2DynamicExt类, 该类派生于IPluginV2Ext, IPluginV2Ext派生于IPluginV2.

参考Document

IPluginV2: https://archive.docs.nvidia.com/tensorrt/tensorrt-861/api/c_api/index.html

IPluginV3: https://docs.nvidia.com/deeplearning/tensorrt/10.x.x/_static/c-api/index.html

Plugin 开发流程

  • 头文件编写
  • CUDA Kernel编写
  • 实现Plugin类IPluginV2DynamicExt
  • 实现Creator类IPluginCreator
  • 实现方法getPluginCreators
  • 注册Plugin
  • 将Plugin打包为库*.so, *.a
  • (optional) 编译TensorRT Engine
  • (optional) 静态链接库打包进执行程序

每个Plugin独立创建文件目录

CMakeLists.txt
Plugin.h
Plugin.cu
Plugin.cpp

头文件编写

#ifndef MY_PLUGIN_H_
#define MY_PLUGIN_H_

#include <string>

#include <cuda.h>
#include <NvInfer.h>
#include <NvInferPlugin.h>
#include <NvInferRuntimeCommon.h>


class MyPlugin : public nvinfer1::IPluginV2DynamicExt
{
public:
    MyPlugin();
    MyPlugin() = delete;

    // IPluginV2DynamicExt methods
    nvinfer1::IPluginV2DynamicExt* clone() const noexcept override;
    nvinfer1::DimsExprs getOutputDimensions(
        int32_t outputIndex,
        nvinfer1::DimsExprs const* inputs,
        int32_t nbInputs,
        nvinfer1::IExprBuilder& exprBuilder
    ) noexcept override;
    bool supportsFormatCombination(
        int32_t pos, 
        nvinfer1::PluginTensorDesc const* inOut,
        int32_t nbInputs, 
        int32_t nbOutputs
    ) noexcept override;
    void configurePlugin(
        nvinfer1::DynamicPluginTensorDesc const* in,
        int32_t nbInputs,
        nvinfer1::DynamicPluginTensorDesc const* out,
        int32_t nbOutputs
    ) noexcept override;
    size_t getWorkspaceSize(
        nvinfer1::PluginTensorDesc const* inputs,
        int32_t nbInputs,
        nvinfer1::PluginTensorDesc const* outputs,
        int32_t nbOutputs
    ) const noexcept override;
    int32_t enqueue(
        nvinfer1::PluginTensorDesc const* inputDesc,
        nvinfer1::PluginTensorDesc const* outputDesc,
        void const* const* inputs,
        void* const* outputs,
        void* workspace,
        cudaStream_t stream
    ) noexcept override;

    // IPluginV2Ext methods
    nvinfer1::DataType getOutputDataType(
        int32_t index, 
        nvinfer1::DataType const* inputTypes,
        int32_t nbInputs
    ) const noexcept override;

    // IPluginV2 methods
    char const* getPluginType() const noexcept override;
    char const* getPluginVersion() const noexcept override;
    char const* getPluginNamespace() const noexcept override;
    void setPluginNamespace(char const* pluginNamespace) noexcept override;
    int32_t getNbOutputs() const noexcept override;
    int32_t initialize() noexcept override;
    void terminate() noexcept override;
    size_t getSerializationSize() const noexcept override;
    void serialize(void* buffer) const noexcept override;
    void destroy() noexcept override;

private:
    std::string mNameSpace {""};

};


class MyPluginCreator : public nvinfer1::IPluginCreator
{
public:
    MyPluginCreator();
    char const* getPluginName() const noexcept override;
    char const* getPluginVersion() const noexcept override;
    char const* getPluginNamespace() const noexcept override;
    void setPluginNamespace(char const* pluginNamespace) noexcept override;
    
    nvinfer1::PluginFieldCollection const* getFieldNames() noexcept override;
    
    nvinfer1::IPluginV2* createPlugin(
        char const* name,
        nvinfer1::PluginFieldCollection const* fc
    ) noexcept override;
    
    nvinfer1::IPluginV2* deserializePlugin(
        char const* name,
        void const* serialData,
        size_t serialLength
    ) noexcept override;

private:
    std::string mNameSpace {""};
    nvinfer1::PluginFieldCollection mFC;
};

extern "C" int32_t ComputeHSwish(
    const float* __restrict__ inputs,
    float* __restrict__ outputs,
    int n_elements,
    cudaStream_t stream 
);

nvinfer1::ILogger* getPluginLogger();

extern "C" void setLoggerFinder(nvinfer1::ILoggerFinder* finder);
extern "C" nvinfer1::IPluginCreator* const* getPluginCreators(int32_t& nbCreators);

#endif

CUDA Kernel编写

in Plugin.cu

为了尽可能提高效率, 一般需要针对FP32, FP16格式编写两个Kernel. 根据传入数据类型分别调用不同的Kernel.

#include <cuda.h>
#include <NvInfer.h>

#include "MyPlugin.h"

// Kernel
// FP32
__global__ void kernel_fp32(
    const float* __restrict__ input,
    float* __restrict__ output,
    int n
)
{ ... }

// FP16
__global__ void kernel_fp16(
    const float* __restrict__ input,
    float* __restrict__ output,
    int n
)
{ ... }

// 供外部调用Kernel的函数
int32_t ComputeKernel(
    const float* __restrict__ inputs
    float* __restrict__ outputs,
    int n_elements,
    cudaStream_t stream
)
{
    nvinfer1::DataType mDataType = ...; // 获取输入数据类型

    constexpr int32_t kBlock = 256; // 32为单位

    if (mDataType == nvinfer1::DataType::kFLOAT)
    {
        // inputs[0], outputs[0] 指向当前batch, 不一定是0
        const float* input = static_cast<const float*>(inputs[0]);
        float*       output = static_cast<float*>(outputs[0]);
        int32_t kGrid = (n_elements + kBlock - 1) / kBlock; // n_elements / kBlock 向上取整
        Kernel_fp32<<<kGrid, kBlock, 0, stream>>>(input, output, n_elements);
    } 
    else if (mDataType == nvinfer1::DataType::kHALF)
    {
        const __half* input = static_cast<const __half*>(inputs[0]);
        __half*       output = static_cast<__half*>(outputs[0]);
        int32_t n_half = (n_elements + 1) / 2;
        int32_t kGrid = (n_half + kBlock - 1) / kBlock; 
        Kernel_fp16<<<kGrid, kBlock, 0, stream>>>(input, output, n_elements);
    }
    else
    {
        return -1; // invalid type
    }
}

实现Plugin类IPluginV2DynamicExt

In Plugin.cpp

import

#include <vector>

#include <cuda.h>
#include "NvInfer.h"
#include "NvInferPlugin.h"

#include "MyPlugin.h"
/**
 * class MyPlugin 
 */

MyPlugin::MyPlugin() {}

// IPluginV2DynamicExt methods


nvinfer1::IPluginV2DynamicExt* MyPlugin::clone() const noexcept
{
    // 深拷贝
    // e.g. 
    auto* p = new MyPlugin();
    p->mNameSpace = mNameSpace;
    return p;
}


nvinfer1::DimsExprs MyPlugin::getOutputDimensions(
    int32_t outputIndex,
    nvinfer1::DimsExprs const* inputs,
    int32_t nbInputs,
    nvinfer1::IExprBuilder& exprBuilder
) noexcept 
{
    // 计算op输出Tensor shape
    return inputs[0];
}


bool MyPlugin::supportsFormatCombination(
    int32_t pos, 
    nvinfer1::PluginTensorDesc const* inOut,
    int32_t nbInputs, 
    int32_t nbOutputs
) noexcept 
{
    /**
     * 检查inOut中pos索引的输入输出类型和Tensor格式是否支持
     * 类型type, e.g. FP32, FP16
     * 格式format, e.g. kLINEAR{N, C, H, W}, kCHW2...
     * For this method inputs are numbered 0..(nbInputs-1) and 
     * outputs are numbered nbInputs..(nbInputs+nbOutputs-1)
     */
    return
           (inOut[pos].type == nvinfer1::DataType::kFLOAT ||
            inOut[pos].type == nvinfer1::DataType::kHALF) &&
            // NCHW format
            inOut[pos].format == nvinfer1::TensorFormat::kLINEAR &&
            // input & output should be the same
            inOut[pos].type == inOut[0].type;
}


void MyPlugin::configurePlugin(
    nvinfer1::DynamicPluginTensorDesc const* in,
    int32_t nbInputs,
    nvinfer1::DynamicPluginTensorDesc const* out,
    int32_t nbOutputs
) noexcept {}


size_t MyPlugin::getWorkspaceSize(
    nvinfer1::PluginTensorDesc const* inputs,
    int32_t nbInputs,
    nvinfer1::PluginTensorDesc const* outputs,
    int32_t nbOutputs
) const noexcept
{
    return 0; // activation need no extra space
}


int32_t MyPlugin::enqueue(
    nvinfer1::PluginTensorDesc const* inputDesc,
    nvinfer1::PluginTensorDesc const* outputDesc,
    void const* const* inputs,
    void* const* outputs,
    void* workspace,
    cudaStream_t stream
) noexcept
{
    // 执行Plugin

    // count elements of input frature map
    int num_elements = 1;
    for (int32_t i = 0; i < inputDesc[0].dims.nbDims; ++i)
        num_elements *= inputDesc[0].dims.d[i];
    
    const float* input = static_cast<const float*>(inputs[0]);
    float* output = static_cast<float*>(outputs[0]);

    ComputeHSwish(input, output, num_elements, stream);
    
    return 0;
}



// IPluginV2Ext methods
nvinfer1::DataType MyPlugin::getOutputDataType(
    int32_t index, 
    nvinfer1::DataType const* inputTypes,
    int32_t nbInputs
) const noexcept
{
    return inputTypes[0];
}


// IPluginV2 methods
char const* MyPlugin::getPluginType() const noexcept 
{
    // 返回Plugin名称
    return "";
}


char const* MyPlugin::getPluginVersion() const noexcept 
{
    return "1";
}


char const* MyPlugin::getPluginNamespace() const noexcept 
{
    return mNameSpace.c_str();
}


void MyPlugin::setPluginNamespace(char const* pluginNamespace) noexcept 
{
    mNameSpace = pluginNamespace;
}


int32_t MyPlugin::getNbOutputs() const noexcept 
{
    return 1; // return number of outputs from this op
}


int32_t MyPlugin::initialize() noexcept 
{
    return 0;
}


void MyPlugin::terminate() noexcept {}


size_t MyPlugin::getSerializationSize() const noexcept 
{
    // 返回序列化时需要多少byte的空间
    return 0;
}


// 将参数按照顺序写入buffer中
void MyPlugin::serialize(void* buffer) const noexcept {}


void MyPlugin::destroy() noexcept 
{
    delete this;
}

/**
 * class MyPlugin end
 */

实现Creator类IPluginCreator

In Plugin.cpp

/**
 * class MyPluginCreator
 */
MyPluginCreator::MyPluginCreator() {}


char const* MyPluginCreator::getPluginName() const noexcept 
{
    return "MyPlugin"; // 严格匹配ONNX模型中自定义算子的名称
}


char const* MyPluginCreator::getPluginVersion() const noexcept 
{
    return "1";
}


char const* MyPluginCreator::getPluginNamespace() const noexcept 
{
    return mNameSpace.c_str();
}


void MyPluginCreator::setPluginNamespace(char const* pluginNamespace) noexcept 
{
    mNameSpace = pluginNamespace;
}


nvinfer1::PluginFieldCollection const* MyPluginCreator::getFieldNames() noexcept 
{
    return &mFC;
}


// 从fc->fields中取出参数创建Plugin实例
nvinfer1::IPluginV2* MyPluginCreator::createPlugin(
    char const* name,
    nvinfer1::PluginFieldCollection const* fc
) noexcept
{
    auto* p = new MyPlugin();
    p->setPluginNamespace(mNameSpace.c_str());
    return p;
}


nvinfer1::IPluginV2* MyPluginCreator::deserializePlugin(
    char const* name,
    void const* serialData,
    size_t serialLength
) noexcept 
{
    // 从serialData中创建Plugin实例
    // e.g. 目前Plugin为激活函数, 无需权重初始化
    auto* p = new MyPlugin();
    p->setPluginNamespace(mNameSpace.c_str());
    return p;
}

/**
 * class MyPluginCreator end
 */

实现方法getPluginCreators

In Plugin.cpp

class ThreadSafeLoggerFinder
{
private:
    nvinfer1::ILoggerFinder* mLoggerFinder {nullptr};

public:
    void setLoggerFinder(nvinfer1::ILoggerFinder* finder)
    {
        if (mLoggerFinder == nullptr && finder != nullptr)
            { mLoggerFinder = finder; }
    }

    nvinfer1::ILogger* getLogger() noexcept
    {
        if (mLoggerFinder != nullptr)
            { return mLoggerFinder->findLogger(); }\
        return nullptr;
    }
};

ThreadSafeLoggerFinder gLoggerFinder;


nvinfer1::ILogger* getPluginLogger()
{
    return gLoggerFinder.getLogger();
}

extern "C" void setLoggerFinder(nvinfer1::ILoggerFinder* finder)
{   gLoggerFinder.setLoggerFinder(finder);
}


extern "C" nvinfer1::IPluginCreator* const* getPluginCreators(int32_t& nbCreators)
{
    nbCreators = 1;
    static MyPluginCreator HSwishCreator;
    static nvinfer1::IPluginCreator* const pluginCreatorList[] = {&HSwishCreator};
    return pluginCreatorList;
}

注册Plugin

In Plugin.cpp 末尾

REGISTER_TENSORRT_PLUGIN(MyPluginCreator);

将Plugin打包为库*.so, *.a

静态库的打包过程限制在Plugin实现的必要实例和必须函数. 无关内容不需打包, 避免最终可执行文件体积过大.

对于IPluginV2DynamicExt插件, 只需包括CUDA Kerne, IPluginV2DynamicExt, IPluginCreator, setLoggerFinder, getPluginCreators, REGISTER_TENSORRT_PLUGIN

NVIDIA compute capability: https://developer.nvidia.com/cuda/gpus

e.g. compute capability = 8.9, CMAKE_CUDA_ARCHITECTURES=89

cmake_minimum_required(VERSION 3.10)

project(MyPlugin LANGUAGES CXX CUDA)
enable_language(CUDA)
set(CMAKE_CUDA_ARCHITECTURES 89) # 查对应架构, NVIDIA compute capability


set(CMAKE_CXX_STANDARD 14)
set(CMAKE_CXX_STANDARD_REQUIRED ON)

# 路径按实际配置
# TensorRT
set(TENSORRT_ROOT /usr)

include_directories(
    ${TENSORRT_ROOT}/include
    /usr/local/cuda/include 
    ${PROJECT_SOURCE_DIR}
)

link_directories(
    ${TENSORRT_ROOT}/lib
    ${TENSORRT_ROOT}/lib64
    /usr/local/cuda/lib64
)

# source files
set(SOURCES
    Plugin.cpp
    Plugin.cu
)

# build .so
add_library(MyPlugin SHARED 
    ${SOURCES}
)

# build .a
add_library(MyPlugin_static STATIC
    ${SOURCES}
)
set_target_properties(MyPlugin PROPERTIES POSITION_INDEPENDENT_CODE ON)

# CUDA
set_target_properties(MyPlugin PROPERTIES 
    CUDA_SEPARABLE_COMPILATION ON
    POSITION_INDEPENDENT_CODE ON 
)

target_link_libraries(MyPlugin
    nvinfer
    cudart
    nvinfer_plugin
)

set(CMAKE_LIBRARY_OUTPUT_DIRECTORY ${PROJECT_SOURCE_DIR})

(optional) 编译TensorRT Engine

#!/bin/bash

trtexec --onnx=... \# ONNX 模型文件
        --saveEngine=... \# TensorRT Engine文件名
        --staticPlugins=... \# Plugin共享库*.so
        --fp16

(optional) 静态链接库打包进执行程序

Plugin在使用前需要先被注册, 为了避免注册逻辑被链接器优化去除, 静态链接库需要完整地被链接进可执行文件中, 以保证Plugin被注册.

CMakeLists.txt

add_library(Plugin STATIC IMPORTED)
set_target_properties(Plugin PROPERTIES
    IMPORTED_LOCATION /to/path/Plugin.a
)

...

target_link_libraries(test PRIVATE
    # 将静态链接库完整强制链接进执行文件, 避免Plugin注册机制被删除
    -Wl,--whole-archive Plugin -Wl,--no-whole-archive
    ...
)
posted @ 2026-06-04 01:42  CINKUK  阅读(25)  评论(0)    收藏  举报