Migrate tkDNN max pooling plugin creation to pluginRegistry from the default method
This commit is contained in:
@@ -500,6 +500,7 @@ public:
|
|||||||
int winH, winW;
|
int winH, winW;
|
||||||
int strideH, strideW;
|
int strideH, strideW;
|
||||||
int paddingH, paddingW;
|
int paddingH, paddingW;
|
||||||
|
int padding;
|
||||||
bool size;
|
bool size;
|
||||||
tkdnnPoolingMode_t pool_mode;
|
tkdnnPoolingMode_t pool_mode;
|
||||||
|
|
||||||
|
|||||||
+15
-2
@@ -369,8 +369,21 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Pooling *l) {
|
|||||||
|
|
||||||
if(l->pool_mode == tkdnnPoolingMode_t::POOLING_MAX_FIXEDSIZE)
|
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);
|
auto creator = getPluginRegistry()->getPluginCreator("MaxPoolingFixedSizeRT_tkDNN","1");
|
||||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
std::vector<PluginField> 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);
|
checkNULL(lRT);
|
||||||
return lRT;
|
return lRT;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ Pooling::Pooling( Network *net, int winH, int winW, int strideH, int strideW,
|
|||||||
this->pool_mode = pool_mode;
|
this->pool_mode = pool_mode;
|
||||||
this->paddingH = paddingH;
|
this->paddingH = paddingH;
|
||||||
this->paddingW = paddingW;
|
this->paddingW = paddingW;
|
||||||
|
this->padding = winH -1;
|
||||||
|
|
||||||
checkCUDNN( cudnnCreatePoolingDescriptor(&poolingDesc) );
|
checkCUDNN( cudnnCreatePoolingDescriptor(&poolingDesc) );
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,9 @@ using namespace nvinfer1;
|
|||||||
std::vector<PluginField> MaxPoolFixedSizeRTPluginCreator::mPluginAttributes;
|
std::vector<PluginField> MaxPoolFixedSizeRTPluginCreator::mPluginAttributes;
|
||||||
PluginFieldCollection MaxPoolFixedSizeRTPluginCreator::mFC{};
|
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){
|
MaxPoolFixedSizeRT::MaxPoolFixedSizeRT(int c, int h, int w, int n, int strideH, int strideW, int winSize, int padding){
|
||||||
this->c = c;
|
this->c = c;
|
||||||
this->h = h;
|
this->h = h;
|
||||||
@@ -105,11 +108,11 @@ void MaxPoolFixedSizeRT::setPluginNamespace(const char *pluginNamespace) NOEXCEP
|
|||||||
}
|
}
|
||||||
|
|
||||||
const char *MaxPoolFixedSizeRT::getPluginType() const NOEXCEPT {
|
const char *MaxPoolFixedSizeRT::getPluginType() const NOEXCEPT {
|
||||||
return "MaxPoolingFixedSizeRT_tkDNN";
|
return MAXPOOLFIXEDSIZERT_PLUGIN_NAME;
|
||||||
}
|
}
|
||||||
|
|
||||||
const char *MaxPoolFixedSizeRT::getPluginVersion() const NOEXCEPT {
|
const char *MaxPoolFixedSizeRT::getPluginVersion() const NOEXCEPT {
|
||||||
return "1";
|
return MAXPOOLFIXEDSIZERT_PLUGIN_VERSION;
|
||||||
}
|
}
|
||||||
|
|
||||||
IPluginV2Ext *MaxPoolFixedSizeRT::clone() const NOEXCEPT {
|
IPluginV2Ext *MaxPoolFixedSizeRT::clone() const NOEXCEPT {
|
||||||
@@ -185,11 +188,11 @@ IPluginV2Ext *MaxPoolFixedSizeRTPluginCreator::createPlugin(const char *name, co
|
|||||||
}
|
}
|
||||||
|
|
||||||
const char *MaxPoolFixedSizeRTPluginCreator::getPluginName() const NOEXCEPT {
|
const char *MaxPoolFixedSizeRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||||
return "MaxPoolingFixedSizeRT_tkDNN";
|
return MAXPOOLFIXEDSIZERT_PLUGIN_NAME;
|
||||||
}
|
}
|
||||||
|
|
||||||
const char *MaxPoolFixedSizeRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
const char *MaxPoolFixedSizeRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||||
return "1";
|
return MAXPOOLFIXEDSIZERT_PLUGIN_VERSION;
|
||||||
}
|
}
|
||||||
|
|
||||||
const PluginFieldCollection *MaxPoolFixedSizeRTPluginCreator::getFieldNames() NOEXCEPT {
|
const PluginFieldCollection *MaxPoolFixedSizeRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||||
|
|||||||
Reference in New Issue
Block a user