diff --git a/include/tkDNN/pluginsRT/ConstantPaddingRT.h b/include/tkDNN/pluginsRT/ConstantPaddingRT.h new file mode 100644 index 0000000..15f4c0d --- /dev/null +++ b/include/tkDNN/pluginsRT/ConstantPaddingRT.h @@ -0,0 +1,109 @@ +// +// Created by perseusdg on 1/7/22. +// + +#ifndef _CONSTANTPADDINGRT_PLUGIN_H +#define _CONSTANTPADDINGRT_PLUGIN_H + +#include +#include +#include +#include +#include + +namespace nvinfer1{ + class ConstantPaddingRT : public IPluginV2Ext { + public: + ConstantPaddingRT(int32_t padH,int32_t padW,int32_t n,int32_t c,int32_t i_h,int32_t i_w,int32_t o_h,int32_t o_w,float constant); + + ConstantPaddingRT(const void *data,size_t length); + + ~ConstantPaddingRT(); + + int getNbOutputs() const NOEXCEPT override; + + Dims getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT override ; + + int initialize() NOEXCEPT override ; + + void terminate() NOEXCEPT override ; + + size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override ; + + +#if NV_TENSORRT_MAJOR > 7 + int enqueue(int batchSize, const void *const *inputs, void *const *outputs, void *workspace, cudaStream_t stream) NOEXCEPT override ; +#elif NV_TENSORRT_MAJOR <= 7 + int32_t enqueue (int32_t batchSize, const void *const *inputs, void **outputs, void *workspace, cudaStream_t stream) override; +#endif + + size_t getSerializationSize() const NOEXCEPT override ; + + void serialize(void *buffer) const NOEXCEPT override ; + + void destroy() NOEXCEPT override ; + + const char *getPluginType() const NOEXCEPT override ; + + const char *getPluginVersion() const NOEXCEPT override; + + const char *getPluginNamespace() const NOEXCEPT override ; + + void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override ; + + IPluginV2Ext *clone() const NOEXCEPT override ; + + DataType getOutputDataType(int index, const nvinfer1::DataType* inputTypes, int nbInputs) const NOEXCEPT override; + + void attachToContext(cudnnContext* cudnnContext, cublasContext* cublasContext, IGpuAllocator* gpuAllocator) NOEXCEPT override; + + bool isOutputBroadcastAcrossBatch(int outputIndex, const bool* inputIsBroadcasted, int nbInputs) const NOEXCEPT override; + + bool canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT override; + + void configurePlugin (Dims const *inputDims, int32_t nbInputs, Dims const *outputDims, + int32_t nbOutputs, DataType const *inputTypes, DataType const *outputTypes, + bool const *inputIsBroadcast, bool const *outputIsBroadcast, PluginFormat floatFormat, + int32_t maxBatchSize) NOEXCEPT override; + + void detachFromContext() NOEXCEPT override; + + bool supportsFormat (DataType type, PluginFormat format) const NOEXCEPT override; + + int32_t i_h,i_w,o_h,o_w,n,c,padH,padW; + float constant; + private: + std::string mPluginNamespace; + + }; + + class ConstantPaddingRTPluginCreator : public IPluginCreator { + public: + ConstantPaddingRTPluginCreator(); + + void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override; + + const char *getPluginNamespace() const NOEXCEPT override; + + IPluginV2Ext *deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT override ; + + IPluginV2Ext *createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT override ; + + const char *getPluginName() const NOEXCEPT override ; + + const char *getPluginVersion() const NOEXCEPT override; + + const PluginFieldCollection *getFieldNames() NOEXCEPT override ; + + private: + static PluginFieldCollection mFC; + static std::vector mPluginAttributes; + std::string mPluginNamespace; + + }; + + REGISTER_TENSORRT_PLUGIN(ConstantPaddingRTPluginCreator); +}; + + +#endif //TKDNN_CONSTANTPADDINGRT_H diff --git a/src/pluginsRT/ConstantPaddingRT.cpp b/src/pluginsRT/ConstantPaddingRT.cpp new file mode 100644 index 0000000..ac37e9d --- /dev/null +++ b/src/pluginsRT/ConstantPaddingRT.cpp @@ -0,0 +1,201 @@ +#include + +using namespace nvinfer1; + +std::vector ConstantPaddingRTPluginCreator::mPluginAttributes; +PluginFieldCollection ConstantPaddingRTPluginCreator::mFC{}; + +static const char* CONSTANTPADDINGRT_PLUGIN_VERSION{"1"}; +static const char* CONSTANTPADDINGRT_PLUGIN_NAME{"ConstantPaddingRT_tkDNN"}; + +ConstantPaddingRT::ConstantPaddingRT(int32_t padH, int32_t padW, int32_t n, int32_t c, int32_t i_h, int32_t i_w, + int32_t o_h, int32_t o_w, float constant) { + this->padH = padH; + this->padW = padW; + this->n = n; + this->c = c; + this->i_h = i_h; + this->i_w = i_w; + this->o_h = o_h; + this->o_w = o_w; + this->constant = constant; + +} + +ConstantPaddingRT::ConstantPaddingRT(const void *data, size_t length) { + const char* buf = reinterpret_cast(data),*bufcheck=buf; + padH = readBUF(buf); + padW = readBUF(buf); + i_h = readBUF(buf); + i_w = readBUF(buf); + o_h = readBUF(buf); + o_w = readBUF(buf); + n = readBUF(buf); + c = readBUF(buf); + constant = readBUF(buf); + assert(buf = bufcheck + length); +} + +ConstantPaddingRT::~ConstantPaddingRT() {} + +int ConstantPaddingRT::getNbOutputs() const NOEXCEPT{ + return 1; +} + +Dims ConstantPaddingRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT { + return Dims3{c,o_h,o_w}; +} + +int ConstantPaddingRT::initialize() NOEXCEPT { + return 0; +} + +void ConstantPaddingRT::terminate() NOEXCEPT { + +} + +size_t ConstantPaddingRT::getWorkspaceSize(int maxBatchSize) const NOEXCEPT { + return 0; +} + +#if NV_TENSORRT_MAJOR > 7 +int ConstantPaddingRT::enqueue(int batchSize, const void *const *inputs, void *const *outputs, void *workspace, cudaStream_t stream) NOEXCEPT { + dnnType* srcData = (dnnType*)reinterpret_cast(inputs[0]); + dnnType* dstData = reinterpret_cast(outputs[0]); + constant_pad2d_forward(srcData,dstData,i_h,i_w,o_h,o_w,c,n,padH,padW,constant,stream); + return 0; +} +#elif NV_TENSORRT_MAJOR <= 7 + int32_t enqueue (int32_t batchSize, const void *const *inputs, void **outputs, void *workspace, cudaStream_t stream) { + dnnType* srcData = (dnnType*)reinterpret_cast(inputs[0]); + dnnType* dstData = reinterpret_cast(outputs[0]); + constant_pad2d_forward(srcData,dstData,i_h,i_w,o_h,o_w,c,n,padH,padW,constant,stream); + return 0; +} +#endif + +size_t ConstantPaddingRT::getSerializationSize() const NOEXCEPT { + return (8*sizeof(int32_t) + 1*sizeof(float)); +} + +void ConstantPaddingRT::serialize(void *buffer) const NOEXCEPT { + char *buf = reinterpret_cast(buffer),*a=buf; + writeBUF(buf,padH); + writeBUF(buf,padW); + writeBUF(buf,i_h); + writeBUF(buf,i_w); + writeBUF(buf,o_h); + writeBUF(buf,o_w); + writeBUF(buf,n); + writeBUF(buf,c); + writeBUF(buf,constant); +} + +void ConstantPaddingRT::destroy() NOEXCEPT { + delete this; +} + +const char* ConstantPaddingRT::getPluginType() const NOEXCEPT { + return CONSTANTPADDINGRT_PLUGIN_NAME; +} + +const char* ConstantPaddingRT::getPluginVersion() const NOEXCEPT { + return CONSTANTPADDINGRT_PLUGIN_VERSION; +} + +const char* ConstantPaddingRT::getPluginNamespace() const NOEXCEPT { + return mPluginNamespace.c_str(); +} + +void ConstantPaddingRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT { + mPluginNamespace = pluginNamespace; +} + +IPluginV2Ext *ConstantPaddingRT::clone() const NOEXCEPT { + auto *p = new ConstantPaddingRT(padH,padW,n,c,i_h,i_w,o_h,o_w,constant); + p->setPluginNamespace(mPluginNamespace.c_str()); + return p; +} + +DataType ConstantPaddingRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, + int nbInputs) const NOEXCEPT { + return DataType::kFLOAT; +} + +void ConstantPaddingRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext, + IGpuAllocator *gpuAllocator) NOEXCEPT { + +} + +bool ConstantPaddingRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted, + int nbInputs) const NOEXCEPT { + return false; +} + +bool ConstantPaddingRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT { + return false; +} + +void ConstantPaddingRT::configurePlugin(const Dims *inputDims, int32_t nbInputs, const Dims *outputDims, + int32_t nbOutputs, const DataType *inputTypes, const DataType *outputTypes, + const bool *inputIsBroadcast, const bool *outputIsBroadcast, + PluginFormat floatFormat, int32_t maxBatchSize) NOEXCEPT { + +} + +void ConstantPaddingRT::detachFromContext() NOEXCEPT { + +} + +bool ConstantPaddingRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT { + return (type == DataType::kFLOAT && format == PluginFormat::kLINEAR); +} + +ConstantPaddingRTPluginCreator::ConstantPaddingRTPluginCreator() { + mPluginAttributes.clear(); + mFC.nbFields = mPluginAttributes.size(); + mFC.fields = mPluginAttributes.data(); +} + +void ConstantPaddingRTPluginCreator::setPluginNamespace(const char *pluginNamespace) NOEXCEPT { + mPluginNamespace = pluginNamespace; +} + +const char *ConstantPaddingRTPluginCreator::getPluginNamespace() const NOEXCEPT { + return mPluginNamespace.c_str(); +} + +IPluginV2Ext *ConstantPaddingRTPluginCreator::deserializePlugin(const char *name, const void *serialData, + size_t serialLength) NOEXCEPT { + auto *pluginObj = new ConstantPaddingRT(serialData,serialLength); + pluginObj->setPluginNamespace(mPluginNamespace.c_str()); + return pluginObj; +} + +IPluginV2Ext *ConstantPaddingRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT { + const PluginField *fields = fc->fields; + int padH = *(static_cast(fields[0].data)); + int padW = *(static_cast(fields[1].data)); + int inputH = *(static_cast(fields[2].data)); + int inputW = *(static_cast(fields[3].data)); + int outputH = *(static_cast(fields[4].data)); + int outputW = *(static_cast(fields[5].data)); + int n = *(static_cast(fields[6].data)); + int c = *(static_cast(fields[7].data)); + float constant = *(static_cast(fields[8].data)); + auto *pluginObj = new ConstantPaddingRT(padH,padW,n,c,inputH,inputW,outputH,outputW,constant); + pluginObj->setPluginNamespace(mPluginNamespace.c_str()); + return pluginObj; +} + +const char *ConstantPaddingRTPluginCreator::getPluginName() const NOEXCEPT { + return CONSTANTPADDINGRT_PLUGIN_NAME; +} + +const char *ConstantPaddingRTPluginCreator::getPluginVersion() const NOEXCEPT { + return CONSTANTPADDINGRT_PLUGIN_VERSION; +} + +const PluginFieldCollection *ConstantPaddingRTPluginCreator::getFieldNames() NOEXCEPT { + return &mFC; +}