From 9e328c0daa159de69d744ea71e0b61d053c363c5 Mon Sep 17 00:00:00 2001 From: perseusdg Date: Wed, 10 Nov 2021 20:09:57 +0530 Subject: [PATCH] Migrate tkDNN max pooling plugin creation to pluginRegistry from the default method --- include/tkDNN/Layer.h | 1 + src/NetworkRT.cpp | 17 +++++++++++++++-- src/Pooling.cpp | 1 + src/pluginsRT/MaxPoolingSizeRT.cpp | 11 +++++++---- 4 files changed, 24 insertions(+), 6 deletions(-) diff --git a/include/tkDNN/Layer.h b/include/tkDNN/Layer.h index e09e5fa..5273c83 100644 --- a/include/tkDNN/Layer.h +++ b/include/tkDNN/Layer.h @@ -500,6 +500,7 @@ public: int winH, winW; int strideH, strideW; int paddingH, paddingW; + int padding; bool size; tkdnnPoolingMode_t pool_mode; diff --git a/src/NetworkRT.cpp b/src/NetworkRT.cpp index dc2e337..472b5dd 100644 --- a/src/NetworkRT.cpp +++ b/src/NetworkRT.cpp @@ -369,8 +369,21 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Pooling *l) { if(l->pool_mode == tkdnnPoolingMode_t::POOLING_MAX_FIXEDSIZE) { - IPluginV2 *plugin = new MaxPoolFixedSizeRT(l->output_dim.c, l->output_dim.h, l->output_dim.w, l->output_dim.n, l->strideH, l->strideW, l->winH, l->winH-1); - IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin); + auto creator = getPluginRegistry()->getPluginCreator("MaxPoolingFixedSizeRT_tkDNN","1"); + std::vector mPluginAttributes; + PluginFieldCollection mFC{}; + mPluginAttributes.emplace_back(PluginField("c",&l->output_dim.c,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("h",&l->output_dim.h,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("w",&l->output_dim.w,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("n",&l->output_dim.n,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("strideH",&l->strideH,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("strideW",&l->strideW,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("winSize",&l->winH,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("padding",&l->padding,PluginFieldType::kINT32,1)); + mFC.nbFields = mPluginAttributes.size(); + mFC.fields = mPluginAttributes.data(); + auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC); + auto *lRT = networkRT->addPluginV2(&input, 1, *plugin); checkNULL(lRT); return lRT; } diff --git a/src/Pooling.cpp b/src/Pooling.cpp index 0838806..2cea2f6 100644 --- a/src/Pooling.cpp +++ b/src/Pooling.cpp @@ -17,6 +17,7 @@ Pooling::Pooling( Network *net, int winH, int winW, int strideH, int strideW, this->pool_mode = pool_mode; this->paddingH = paddingH; this->paddingW = paddingW; + this->padding = winH -1; checkCUDNN( cudnnCreatePoolingDescriptor(&poolingDesc) ); diff --git a/src/pluginsRT/MaxPoolingSizeRT.cpp b/src/pluginsRT/MaxPoolingSizeRT.cpp index c45eae3..4254ac5 100644 --- a/src/pluginsRT/MaxPoolingSizeRT.cpp +++ b/src/pluginsRT/MaxPoolingSizeRT.cpp @@ -4,6 +4,9 @@ using namespace nvinfer1; std::vector MaxPoolFixedSizeRTPluginCreator::mPluginAttributes; PluginFieldCollection MaxPoolFixedSizeRTPluginCreator::mFC{}; +static const char* MAXPOOLFIXEDSIZERT_PLUGIN_VERSION{"1"}; +static const char* MAXPOOLFIXEDSIZERT_PLUGIN_NAME{"MaxPoolingFixedSizeRT_tkDNN"}; + MaxPoolFixedSizeRT::MaxPoolFixedSizeRT(int c, int h, int w, int n, int strideH, int strideW, int winSize, int padding){ this->c = c; this->h = h; @@ -105,11 +108,11 @@ void MaxPoolFixedSizeRT::setPluginNamespace(const char *pluginNamespace) NOEXCEP } const char *MaxPoolFixedSizeRT::getPluginType() const NOEXCEPT { - return "MaxPoolingFixedSizeRT_tkDNN"; + return MAXPOOLFIXEDSIZERT_PLUGIN_NAME; } const char *MaxPoolFixedSizeRT::getPluginVersion() const NOEXCEPT { - return "1"; + return MAXPOOLFIXEDSIZERT_PLUGIN_VERSION; } IPluginV2Ext *MaxPoolFixedSizeRT::clone() const NOEXCEPT { @@ -185,11 +188,11 @@ IPluginV2Ext *MaxPoolFixedSizeRTPluginCreator::createPlugin(const char *name, co } const char *MaxPoolFixedSizeRTPluginCreator::getPluginName() const NOEXCEPT { - return "MaxPoolingFixedSizeRT_tkDNN"; + return MAXPOOLFIXEDSIZERT_PLUGIN_NAME; } const char *MaxPoolFixedSizeRTPluginCreator::getPluginVersion() const NOEXCEPT { - return "1"; + return MAXPOOLFIXEDSIZERT_PLUGIN_VERSION; } const PluginFieldCollection *MaxPoolFixedSizeRTPluginCreator::getFieldNames() NOEXCEPT {