From 04a5823873516def7f7ec8bf88fcd56387874c95 Mon Sep 17 00:00:00 2001 From: oren Date: Tue, 12 Jan 2021 09:43:54 -0600 Subject: [PATCH] Added options to specify dimension ordering --- include/tkDNN/Network.h | 34 +++++++++++++++++++++++++++++++--- include/tkDNN/NetworkRT.h | 2 +- src/NetworkRT.cpp | 8 ++++---- 3 files changed, 36 insertions(+), 8 deletions(-) diff --git a/include/tkDNN/Network.h b/include/tkDNN/Network.h index 571e728..05f1bcf 100644 --- a/include/tkDNN/Network.h +++ b/include/tkDNN/Network.h @@ -7,6 +7,12 @@ namespace tk { namespace dnn { +enum dimFormat_t { + CHW, + NCHW, + //NHWC +}; + /** Data representation between layers n = batch size @@ -21,9 +27,31 @@ struct dataDim_t { dataDim_t() : n(1), c(1), h(1), w(1), l(1) {}; - dataDim_t(nvinfer1::Dims &d) : - n(1), c(d.d[0] ? d.d[0] : 1), h(d.d[1] ? d.d[1] : 1), - w(d.d[2] ? d.d[2] : 1), l(d.d[3] ? d.d[3] : 1) {}; + dataDim_t(nvinfer1::Dims &d, dimFormat_t df) { + switch(df) { + case CHW: + n=1; + c = d.d[0] ? d.d[0] : 1; + h = d.d[1] ? d.d[1] : 1; + w = d.d[2] ? d.d[2] : 1; + l = d.d[3] ? d.d[3] : 1; + break; + case NCHW: + n = d.d[0] ? d.d[0] : 1; + c = d.d[1] ? d.d[1] : 1; + h = d.d[2] ? d.d[2] : 1; + w = d.d[3] ? d.d[3] : 1; + l = d.d[4] ? d.d[4] : 1; + break; + // case NHWC: + // n = d.d[0] ? d.d[0] : 1; + // h = d.d[1] ? d.d[1] : 1; + // w = d.d[2] ? d.d[2] : 1; + // c = d.d[3] ? d.d[3] : 1; + // l = d.d[4] ? d.d[4] : 1; + // break; + } + }; dataDim_t(int _n, int _c, int _h, int _w, int _l = 1) : n(_n), c(_c), h(_h), w(_w), l(_l) {}; diff --git a/include/tkDNN/NetworkRT.h b/include/tkDNN/NetworkRT.h index d28b96b..39f5db3 100644 --- a/include/tkDNN/NetworkRT.h +++ b/include/tkDNN/NetworkRT.h @@ -73,7 +73,7 @@ public: PluginFactory *pluginFactory; - NetworkRT(Network *net, const char *name, const char *input_name="data", const char *output_name="out"); + NetworkRT(Network *net, const char *name, dimFormat_t dim_format=CHW, const char *input_name="data", const char *output_name="out"); virtual ~NetworkRT(); int getMaxBatchSize() { diff --git a/src/NetworkRT.cpp b/src/NetworkRT.cpp index bc657ae..45e7d49 100644 --- a/src/NetworkRT.cpp +++ b/src/NetworkRT.cpp @@ -26,7 +26,7 @@ namespace tk { namespace dnn { std::maptensors; -NetworkRT::NetworkRT(Network *net, const char *name, const char *input_name, const char *output_name) { +NetworkRT::NetworkRT(Network *net, const char *name, dimFormat_t dim_format, const char *input_name, const char *output_name) { float rt_ver = float(NV_TENSORRT_MAJOR) + float(NV_TENSORRT_MINOR)/10 + @@ -167,17 +167,17 @@ NetworkRT::NetworkRT(Network *net, const char *name, const char *input_name, con Dims iDim = engineRT->getBindingDimensions(buf_input_idx); - input_dim = dataDim_t(iDim); + input_dim = dataDim_t(iDim, dim_format); input_dim.print(); Dims oDim = engineRT->getBindingDimensions(buf_output_idx); - output_dim = dataDim_t(oDim); + output_dim = dataDim_t(oDim, dim_format); output_dim.print(); // create GPU buffers and a stream for(int i=0; igetNbBindings(); i++) { Dims dim = engineRT->getBindingDimensions(i); - buffersDIM[i] = dataDim_t(dim); + buffersDIM[i] = dataDim_t(dim, dim_format); std::cout<<"RtBuffer "<getMaxBatchSize()*buffersDIM[i].tot()*sizeof(dnnType))); }