Add grouped convolutions in CUDNN and tensorRT.

Signed-off-by: Davide Sapienza <sapienza.dav@gmail.com>
This commit is contained in:
Davide Sapienza
2020-02-05 14:32:58 +01:00
parent fe85c26888
commit 4616be0738
5 changed files with 25 additions and 15 deletions
+5 -4
View File
@@ -85,7 +85,7 @@ class LayerWgs : public Layer {
public: public:
LayerWgs(Network *net, int inputs, int outputs, int kh, int kw, int kt, LayerWgs(Network *net, int inputs, int outputs, int kh, int kw, int kt,
std::string fname_weights, bool batchnorm = false, bool additional_bias = false, bool final = false); std::string fname_weights, bool batchnorm = false, bool additional_bias = false, bool final = false, bool deConv = false, int groups = 1);
virtual ~LayerWgs(); virtual ~LayerWgs();
int inputs, outputs; int inputs, outputs;
@@ -165,7 +165,7 @@ class Conv2d : public LayerWgs {
public: public:
Conv2d( Network *net, int out_ch, int kernelH, int kernelW, Conv2d( Network *net, int out_ch, int kernelH, int kernelW,
int strideH, int strideW, int paddingH, int paddingW, int strideH, int strideW, int paddingH, int paddingW,
std::string fname_weights, bool batchnorm = false, bool deConv = false, bool final = false); std::string fname_weights, bool batchnorm = false, bool deConv = false, bool final = false, int groups = 1);
virtual ~Conv2d(); virtual ~Conv2d();
virtual layerType_t getLayerType() { return LAYER_CONV2D; }; virtual layerType_t getLayerType() { return LAYER_CONV2D; };
@@ -173,6 +173,7 @@ public:
int kernelH, kernelW, strideH, strideW, paddingH, paddingW; int kernelH, kernelW, strideH, strideW, paddingH, paddingW;
bool deConv; bool deConv;
int groups;
protected: protected:
cudnnFilterDescriptor_t filterDesc; cudnnFilterDescriptor_t filterDesc;
@@ -196,8 +197,8 @@ class DeConv2d : public Conv2d {
public: public:
DeConv2d( Network *net, int out_ch, int kernelH, int kernelW, DeConv2d( Network *net, int out_ch, int kernelH, int kernelW,
int strideH, int strideW, int paddingH, int paddingW, int strideH, int strideW, int paddingH, int paddingW,
std::string fname_weights, bool batchnorm = false) : std::string fname_weights, bool batchnorm = false, int groups = 1) :
Conv2d(net, out_ch, kernelH, kernelW, strideH, strideW, paddingH, paddingW, fname_weights, batchnorm, true) {} Conv2d(net, out_ch, kernelH, kernelW, strideH, strideW, paddingH, paddingW, fname_weights, batchnorm, true, false, groups) {}
virtual ~DeConv2d() {} virtual ~DeConv2d() {}
virtual layerType_t getLayerType() { return LAYER_DECONV2D; }; virtual layerType_t getLayerType() { return LAYER_DECONV2D; };
+9 -7
View File
@@ -17,8 +17,6 @@ void Conv2d::initCUDNN(bool back) {
idim = output_dim; idim = output_dim;
odim = input_dim; odim = input_dim;
} }
//idim.print();
//odim.print();
checkCUDNN( cudnnCreateFilterDescriptor(&filterDesc) ); checkCUDNN( cudnnCreateFilterDescriptor(&filterDesc) );
checkCUDNN( cudnnCreateConvolutionDescriptor(&convDesc) ); checkCUDNN( cudnnCreateConvolutionDescriptor(&convDesc) );
@@ -29,7 +27,7 @@ void Conv2d::initCUDNN(bool back) {
net->tensorFormat, net->dataType, idim.n, idim.c, idim.h, idim.w) ); net->tensorFormat, net->dataType, idim.n, idim.c, idim.h, idim.w) );
checkCUDNN( cudnnSetFilter4dDescriptor(filterDesc, checkCUDNN( cudnnSetFilter4dDescriptor(filterDesc,
net->dataType, net->tensorFormat, odim.c, idim.c, net->dataType, net->tensorFormat, odim.c, idim.c/groups,
kernelH, kernelW) ); kernelH, kernelW) );
checkCUDNN( cudnnSetConvolution2dDescriptor(convDesc, checkCUDNN( cudnnSetConvolution2dDescriptor(convDesc,
@@ -38,16 +36,20 @@ void Conv2d::initCUDNN(bool back) {
1,1, // upscale 1,1, // upscale
CUDNN_CROSS_CORRELATION, CUDNN_DATA_FLOAT) ); CUDNN_CROSS_CORRELATION, CUDNN_DATA_FLOAT) );
checkCUDNN( cudnnSetConvolutionGroupCount(convDesc,
groups) );
// check dimension of convolution output // check dimension of convolution output
dataDim_t tmpdim; dataDim_t tmpdim;
checkCUDNN( cudnnGetConvolution2dForwardOutputDim( checkCUDNN( cudnnGetConvolution2dForwardOutputDim(
convDesc, srcTensor, filterDesc, convDesc, srcTensor, filterDesc,
&tmpdim.n, &tmpdim.c, &tmpdim.h, &tmpdim.w) ); &tmpdim.n, &tmpdim.c, &tmpdim.h, &tmpdim.w) );
if(odim.n != tmpdim.n || odim.c != tmpdim.c || odim.h != tmpdim.h || odim.w != tmpdim.w) { if(odim.n != tmpdim.n || odim.c != tmpdim.c || odim.h != tmpdim.h || odim.w != tmpdim.w) {
std::cout<<"tkdim input: "; idim.print(); std::cout<<"tkdim input: "; idim.print();
std::cout<<"tkdim output: "; odim.print(); std::cout<<"tkdim output: "; odim.print();
std::cout<<"cudnndim: "; tmpdim.print(); std::cout<<"cudnndim: "; tmpdim.print();
FatalError("Eror conv dimension mismatch"); FatalError("Error conv dimension mismatch");
} }
checkCUDNN( cudnnSetTensor4dDescriptor(dstTensor, checkCUDNN( cudnnSetTensor4dDescriptor(dstTensor,
@@ -119,11 +121,10 @@ void Conv2d::inferCUDNN(dnnType* srcData, bool back) {
Conv2d::Conv2d( Network *net, int out_ch, int kernelH, int kernelW, Conv2d::Conv2d( Network *net, int out_ch, int kernelH, int kernelW,
int strideH, int strideW, int paddingH, int paddingW, int strideH, int strideW, int paddingH, int paddingW,
std::string fname_weights, bool batchnorm, bool deConv, bool final) : std::string fname_weights, bool batchnorm, bool deConv, bool final, int groups) :
LayerWgs(net, net->getOutputDim().c, out_ch, kernelH, kernelW, 1, LayerWgs(net, net->getOutputDim().c, out_ch, kernelH, kernelW, 1,
fname_weights, batchnorm, false, final) { fname_weights, batchnorm, false, final, deConv, groups) {
this->kernelH = kernelH; this->kernelH = kernelH;
this->kernelW = kernelW; this->kernelW = kernelW;
this->strideH = strideH; this->strideH = strideH;
@@ -131,6 +132,7 @@ Conv2d::Conv2d( Network *net, int out_ch, int kernelH, int kernelW,
this->paddingH = paddingH; this->paddingH = paddingH;
this->paddingW = paddingW; this->paddingW = paddingW;
this->deConv = deConv; this->deConv = deConv;
this->groups = groups;
if(!deConv) { if(!deConv) {
output_dim.n = input_dim.n; output_dim.n = input_dim.n;
+2 -2
View File
@@ -76,8 +76,8 @@ DeformConv2d::~DeformConv2d() {
checkCUDNN( cudnnDestroyTensorDescriptor(biasTensorDesc) ); checkCUDNN( cudnnDestroyTensorDescriptor(biasTensorDesc) );
checkCuda( cudaFree(dstData) ); checkCuda( cudaFree(dstData) );
checkCuda( cudaFreeHost(ones_d1) ); checkCuda( cudaFree(ones_d1) );
checkCuda( cudaFreeHost(ones_d2) ); checkCuda( cudaFree(ones_d2) );
checkCuda( cudaFree(offset) ); checkCuda( cudaFree(offset) );
checkCuda( cudaFree(mask) ); checkCuda( cudaFree(mask) );
checkCuda( cudaFree(output_conv) ); checkCuda( cudaFree(output_conv) );
+6 -1
View File
@@ -8,7 +8,12 @@ namespace tk { namespace dnn {
LayerWgs::LayerWgs(Network *net, int inputs, int outputs, LayerWgs::LayerWgs(Network *net, int inputs, int outputs,
int kh, int kw, int kl, int kh, int kw, int kl,
std::string fname_weights, bool batchnorm, bool additional_bias, bool final) : Layer(net, final) { std::string fname_weights, bool batchnorm, bool additional_bias, bool final, bool deConv, int groups) : Layer(net, final) {
if(deConv)
inputs = inputs/groups;
else
outputs = outputs/groups;
this->inputs = inputs; this->inputs = inputs;
this->outputs = outputs; this->outputs = outputs;
+2
View File
@@ -245,6 +245,7 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Conv2d *l) {
checkNULL(lRTconv); checkNULL(lRTconv);
lRTconv->setStride(DimsHW{l->strideH, l->strideW}); lRTconv->setStride(DimsHW{l->strideH, l->strideW});
lRTconv->setPadding(DimsHW{l->paddingH, l->paddingW}); lRTconv->setPadding(DimsHW{l->paddingH, l->paddingW});
lRTconv->setNbGroups(l->groups);
lRT = (ILayer*) lRTconv; lRT = (ILayer*) lRTconv;
} else { } else {
IDeconvolutionLayer *lRTconv = networkRT->addDeconvolution(*input, IDeconvolutionLayer *lRTconv = networkRT->addDeconvolution(*input,
@@ -252,6 +253,7 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Conv2d *l) {
checkNULL(lRTconv); checkNULL(lRTconv);
lRTconv->setStride(DimsHW{l->strideH, l->strideW}); lRTconv->setStride(DimsHW{l->strideH, l->strideW});
lRTconv->setPadding(DimsHW{l->paddingH, l->paddingW}); lRTconv->setPadding(DimsHW{l->paddingH, l->paddingW});
lRTconv->setNbGroups(l->groups);
lRT = (ILayer*) lRTconv; lRT = (ILayer*) lRTconv;
Dims d = lRTconv->getOutput(0)->getDimensions(); Dims d = lRTconv->getOutput(0)->getDimensions();