TRT8 works with almost every nerual network now!!!!(including demo3d)

This commit is contained in:
perseusdg
2021-10-28 23:35:37 +05:30
parent 8c36dd0431
commit c5e66c6bf6
37 changed files with 1042 additions and 700 deletions
+58 -21
View File
@@ -1,12 +1,19 @@
#include <tkDNN/pluginsRT/RegionRT.h>
using namespace nvinfer1;
std::vector<PluginField> RegionRTPluginCreator::mPluginAttributes;
PluginFieldCollection RegionRTPluginCreator::mFC{};
RegionRT::RegionRT(int classes, int coords, int num) {
static const char* REGIONRT_PLUGIN_VERSION{"1"};
static const char* REGIONRT_PLUGIN_NAME{"RegionRT_tkDNN"};
RegionRT::RegionRT(int classes, int coords, int num,int c,int h,int w) {
this->classes = classes;
this->coords = coords;
this->num = num;
this->c = c;
this->h = h;
this->w = w;
}
RegionRT::~RegionRT() {}
@@ -30,12 +37,6 @@ Dims RegionRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDim
return inputs[0];
}
void RegionRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type,
PluginFormat format, int maxBatchSize) NOEXCEPT {
c = inputDims[0].d[0];
h = inputDims[0].d[1];
w = inputDims[0].d[2];
}
int RegionRT::initialize() NOEXCEPT {return 0;}
@@ -112,11 +113,11 @@ void RegionRT::serialize(void *buffer) const NOEXCEPT {
}
const char *RegionRT::getPluginType() const NOEXCEPT {
return "RegionRT_tkDNN";
return REGIONRT_PLUGIN_NAME;
}
const char *RegionRT::getPluginVersion() const NOEXCEPT {
return "1";
return REGIONRT_PLUGIN_VERSION;
}
void RegionRT::destroy() NOEXCEPT { delete this; }
@@ -133,14 +134,48 @@ bool RegionRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT
return true;
}
IPluginV2 *RegionRT::clone() const NOEXCEPT {
auto *p = new RegionRT(classes,coords,num);
IPluginV2Ext *RegionRT::clone() const NOEXCEPT {
auto *p = new RegionRT(classes,coords,num,c,h,w);
p->setPluginNamespace(mPluginNamespace.c_str());
return p;
}
DataType RegionRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
return DataType::kFLOAT;
}
void RegionRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
IGpuAllocator *gpuAllocator) NOEXCEPT {
}
bool RegionRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted, int nbInputs) const NOEXCEPT {
return false;
}
bool RegionRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
return false;
}
void RegionRT::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 RegionRT::detachFromContext() NOEXCEPT {
}
RegionRTPluginCreator::RegionRTPluginCreator() {
mPluginAttributes.clear();
mPluginAttributes.emplace_back(PluginField("classes", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("coords", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("num", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
mFC.nbFields = mPluginAttributes.size();
mFC.fields = mPluginAttributes.data();
}
@@ -153,32 +188,35 @@ const char *RegionRTPluginCreator::getPluginNamespace() const NOEXCEPT {
return mPluginNamespace.c_str();
}
IPluginV2 *RegionRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
IPluginV2Ext *RegionRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
auto *pluginObj = new RegionRT(serialData,serialLength);
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
return pluginObj;
}
IPluginV2 *RegionRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
IPluginV2Ext *RegionRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
const PluginField *fields = fc->fields;
assert(fc->nbFields == 3);
assert(fields[0].type == PluginFieldType::kINT32);
assert(fields[1].type == PluginFieldType::kINT32);
assert(fields[2].type == PluginFieldType::kINT32);
assert(fc->nbFields == 6);
for(int i=0;i<6;i++){
assert(fields[i].type == PluginFieldType::kINT32);
}
int classes = *(static_cast<const int*>(fields[0].data));
int coords = *(static_cast<const int*>(fields[1].data));
int num = *(static_cast<const int*>(fields[2].data));
RegionRT *pluginObj = new RegionRT(classes,coords,num);
int c = *(static_cast<const int*>(fields[3].data));
int h = *(static_cast<const int*>(fields[4].data));
int w = *(static_cast<const int*>(fields[5].data));
auto *pluginObj = new RegionRT(classes,coords,num,c,h,w);
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
return pluginObj;
}
const char *RegionRTPluginCreator::getPluginName() const NOEXCEPT {
return "RegionRT_tkDNN";
return REGIONRT_PLUGIN_NAME;
}
const char *RegionRTPluginCreator::getPluginVersion() const NOEXCEPT {
return "1";
return REGIONRT_PLUGIN_VERSION;
}
const PluginFieldCollection *RegionRTPluginCreator::getFieldNames() NOEXCEPT {
@@ -195,4 +233,3 @@ const PluginFieldCollection *RegionRTPluginCreator::getFieldNames() NOEXCEPT {