mobilenet works with trt7(IPluginV2IOExt) ,need to test it with trt8

This commit is contained in:
perseusdg
2021-10-27 18:00:33 +05:30
parent 54e7af11ed
commit 8c36dd0431
7 changed files with 88 additions and 40 deletions
+5
View File
@@ -15,6 +15,11 @@ Flatten::Flatten(Network *net) : Layer(net) {
output_dim.w = 1;
output_dim.l = 1;
this->h = 1;
this->w = 1;
this->rows = input_dim.w;
this->cols = input_dim.h * input_dim.c;
this->c = input_dim.w * input_dim.h * input_dim.c;
}
Flatten::~Flatten() {
+3 -3
View File
@@ -18,9 +18,9 @@ using namespace nvinfer1;
// Logger for info/warning/errors
class Logger : public ILogger {
void log(Severity severity, const char* msg) NOEXCEPT override {
#ifdef DEBUG
//#ifdef DEBUG
std::cout <<"TENSORRT LOG: "<< msg << std::endl;
#endif
//#endif
}
} loggerRT;
@@ -473,7 +473,7 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Route *l) {
ILayer* NetworkRT::convert_layer(ITensor *input, Flatten *l) {
IPluginV2 *plugin = new FlattenConcatRT();
IPluginV2IOExt *plugin = new FlattenConcatRT(l->c,l->h,l->w,l->rows,l->cols);
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
checkNULL(lRT);
return lRT;
+50 -19
View File
@@ -4,12 +4,17 @@ using namespace nvinfer1;
std::vector<PluginField> FlattenConcatRTPluginCreator::mPluginAttributes;
PluginFieldCollection FlattenConcatRTPluginCreator::mFC{};
FlattenConcatRT::FlattenConcatRT() {
FlattenConcatRT::FlattenConcatRT(int c, int h, int w, int rows, int cols) {
stat = cublasCreate(&handle);
if (stat != CUBLAS_STATUS_SUCCESS) {
printf ("CUBLAS initialization failed\n");
return;
}
this->c = c;
this->h = h;
this->w = w;
this->rows = rows;
this->cols = cols;
}
FlattenConcatRT::FlattenConcatRT(const void *data, size_t length) {
@@ -32,16 +37,6 @@ Dims FlattenConcatRT::getOutputDimensions(int index, const Dims *inputs, int nbI
return Dims3{ inputs[0].d[0] * inputs[0].d[1] * inputs[0].d[2], 1, 1};
}
void FlattenConcatRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,
DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {
assert(nbOutputs == 1 && nbInputs ==1);
rows = inputDims[0].d[0];
cols = inputDims[0].d[1] * inputDims[0].d[2];
c = inputDims[0].d[0] * inputDims[0].d[1] * inputDims[0].d[2];
h = 1;
w = 1;
}
int FlattenConcatRT::initialize() NOEXCEPT {
return 0;
}
@@ -107,9 +102,7 @@ void FlattenConcatRT::destroy() NOEXCEPT {
delete this;
}
bool FlattenConcatRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT {
return true;
}
const char *FlattenConcatRT::getPluginType() const NOEXCEPT {
return "FlattenConcatRT_tkDNN";
@@ -127,12 +120,44 @@ void FlattenConcatRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
mPluginNamespace = pluginNamespace;
}
IPluginV2 *FlattenConcatRT::clone() const NOEXCEPT {
auto *p = new FlattenConcatRT();
IPluginV2IOExt *FlattenConcatRT::clone() const NOEXCEPT {
auto* p = new FlattenConcatRT(c, h, w, rows, cols);
p->setPluginNamespace(mPluginNamespace.c_str());
return p;
}
DataType FlattenConcatRT::getOutputDataType(int index, const nvinfer1::DataType* inputTypes, int nbInputs) const NOEXCEPT
{
return DataType::kFLOAT;
}
void FlattenConcatRT::configurePlugin(const PluginTensorDesc* in, int nbInput, const PluginTensorDesc* out, int nbOutput) NOEXCEPT
{
}
void FlattenConcatRT::attachToContext(cudnnContext* cudnnContext, cublasContext* cublasContext, IGpuAllocator* gpuAllocator) NOEXCEPT
{
}
bool FlattenConcatRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool* inputIsBroadcasted, int nbInputs) const NOEXCEPT
{
return false;
}
bool FlattenConcatRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT
{
return false;
}
bool FlattenConcatRT::supportsFormatCombination(int pos, const PluginTensorDesc* inOut, int nbInputs, int nbOutputs) const NOEXCEPT
{
return true;
}
void FlattenConcatRT::detachFromContext() NOEXCEPT
{
}
FlattenConcatRTPluginCreator::FlattenConcatRTPluginCreator() {
mPluginAttributes.clear();
mFC.nbFields = mPluginAttributes.size();
@@ -147,15 +172,21 @@ const char *FlattenConcatRTPluginCreator::getPluginNamespace() const NOEXCEPT {
return mPluginNamespace.c_str();
}
IPluginV2 *FlattenConcatRTPluginCreator::deserializePlugin(const char *name, const void *serialData,
IPluginV2IOExt *FlattenConcatRTPluginCreator::deserializePlugin(const char *name, const void *serialData,
size_t serialLength) NOEXCEPT {
auto *pluginObj = new FlattenConcatRT(serialData,serialLength);
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
return pluginObj;
}
IPluginV2 *FlattenConcatRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
auto *pluginObj = new FlattenConcatRT();
IPluginV2IOExt *FlattenConcatRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
const PluginField* fields = fc->fields;
int c = *(static_cast<const int*>(fields[0].data));
int h = *(static_cast<const int*>(fields[1].data));
int w = *(static_cast<const int*>(fields[2].data));
int rows = *(static_cast<const int*>(fields[3].data));
int cols = *(static_cast<const int*>(fields[4].data));
auto* pluginObj = new FlattenConcatRT(c, h, w, rows, cols);
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
return pluginObj;
}
+3 -2
View File
@@ -54,10 +54,11 @@ int ReshapeRT::enqueue(int batchSize, const void *const *inputs, void *const *ou
}
#elif NV_TENSORRT_MAJOR == 7
int32_t ReshapeRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace, cudaStream_t stream) {
std::cout << new_dim.c << ":" << new_dim.h << std::endl;
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*c*h*w*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
std::cout << "C : " << c << "H : " << h << "w :" << w << std::endl;
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*new_dim.c*new_dim.h*new_dim.w*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
return 0;
}
#endif