tensorRT serialization OK

This commit is contained in:
Francesco Gatti
2017-08-11 17:17:05 +02:00
parent 57c9a6ec99
commit 3b2f062dd9
7 changed files with 54 additions and 14 deletions
+1 -1
View File
@@ -25,7 +25,7 @@ public:
dnnType *output; dnnType *output;
cudaStream_t stream; cudaStream_t stream;
NetworkRT(Network *net); NetworkRT(Network *net, const char *name);
virtual ~NetworkRT(); virtual ~NetworkRT();
/** /**
+35 -8
View File
@@ -1,5 +1,7 @@
#include <iostream> #include <iostream>
#include <map> #include <map>
#include <errno.h>
#include "NvInfer.h" #include "NvInfer.h"
#include "NetworkRT.h" #include "NetworkRT.h"
@@ -22,7 +24,7 @@ namespace tkDNN {
std::map<Layer*, nvinfer1::ITensor*>tensors; std::map<Layer*, nvinfer1::ITensor*>tensors;
NetworkRT::NetworkRT(Network *net) { NetworkRT::NetworkRT(Network *net, const char *name) {
float rt_ver = float(NV_TENSORRT_MAJOR) + float rt_ver = float(NV_TENSORRT_MAJOR) +
float(NV_TENSORRT_MINOR)/10 + float(NV_TENSORRT_MINOR)/10 +
@@ -35,7 +37,7 @@ NetworkRT::NetworkRT(Network *net) {
//add input layer //add input layer
dataDim_t dim = net->layers[0]->input_dim; dataDim_t dim = net->layers[0]->input_dim;
if(!fileExist("net.rt")) { if(!fileExist(name)) {
ITensor *input = networkRT->addInput("data", dtRT, ITensor *input = networkRT->addInput("data", dtRT,
DimsCHW{ dim.c, dim.h, dim.w}); DimsCHW{ dim.c, dim.h, dim.w});
checkNULL(input); checkNULL(input);
@@ -64,9 +66,9 @@ NetworkRT::NetworkRT(Network *net) {
engineRT = builderRT->buildCudaEngine(*networkRT); engineRT = builderRT->buildCudaEngine(*networkRT);
// we don't need the network any more // we don't need the network any more
//networkRT->destroy(); //networkRT->destroy();
serialize("net.rt"); serialize(name);
} else { } else {
deserialize("net.rt"); deserialize(name);
} }
std::cout<<"create execution context\n"; std::cout<<"create execution context\n";
@@ -220,11 +222,16 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Activation *l) {
IPluginLayer *lRT = networkRT->addPlugin(&input, 1, *plugin); IPluginLayer *lRT = networkRT->addPlugin(&input, 1, *plugin);
checkNULL(lRT); checkNULL(lRT);
return lRT; return lRT;
}
IActivationLayer *lRT = networkRT->addActivation(*input, ActivationType::kRELU); } else if(l->act_mode == CUDNN_ACTIVATION_RELU) {
checkNULL(lRT); IActivationLayer *lRT = networkRT->addActivation(*input, ActivationType::kRELU);
return lRT; checkNULL(lRT);
return lRT;
} else {
FatalError("this Activation mode is not yet implemented");
return NULL;
}
} }
ILayer* NetworkRT::convert_layer(ITensor *input, Softmax *l) { ILayer* NetworkRT::convert_layer(ITensor *input, Softmax *l) {
@@ -298,6 +305,26 @@ public:
a->size = readBUF<int>(buf); a->size = readBUF<int>(buf);
return a; return a;
} }
if(name.find("Region") == 0) {
RegionRT *r = new RegionRT(readBUF<int>(buf), //classes
readBUF<int>(buf), //coords
readBUF<int>(buf), //num
readBUF<float>(buf)); //thesh
r->c = readBUF<int>(buf);
r->h = readBUF<int>(buf);
r->w = readBUF<int>(buf);
return r;
}
if(name.find("Reorg") == 0) {
ReorgRT *r = new ReorgRT(readBUF<int>(buf)); //stride
r->c = readBUF<int>(buf);
r->h = readBUF<int>(buf);
r->w = readBUF<int>(buf);
return r;
}
FatalError("Cant deserialize Plugin"); FatalError("Cant deserialize Plugin");
return NULL; return NULL;
} }
+9 -1
View File
@@ -70,10 +70,18 @@ public:
virtual size_t getSerializationSize() override { virtual size_t getSerializationSize() override {
return 0; return 6*sizeof(int) + 1*sizeof(float);
} }
virtual void serialize(void* buffer) override { virtual void serialize(void* buffer) override {
char *buf = reinterpret_cast<char*>(buffer);
tkDNN::writeBUF(buf, classes);
tkDNN::writeBUF(buf, coords);
tkDNN::writeBUF(buf, num);
tkDNN::writeBUF(buf, thresh);
tkDNN::writeBUF(buf, c);
tkDNN::writeBUF(buf, h);
tkDNN::writeBUF(buf, w);
} }
int c, h, w; int c, h, w;
+6 -1
View File
@@ -48,10 +48,15 @@ public:
virtual size_t getSerializationSize() override { virtual size_t getSerializationSize() override {
return 0; return 4*sizeof(int);
} }
virtual void serialize(void* buffer) override { virtual void serialize(void* buffer) override {
char *buf = reinterpret_cast<char*>(buffer);
tkDNN::writeBUF(buf, stride);
tkDNN::writeBUF(buf, c);
tkDNN::writeBUF(buf, h);
tkDNN::writeBUF(buf, w);
} }
int c, h, w, stride; int c, h, w, stride;
+1 -1
View File
@@ -22,7 +22,7 @@ int main() {
tkDNN::Dense l6(&net, 10, d3_bin); tkDNN::Dense l6(&net, 10, d3_bin);
tkDNN::Softmax l7(&net); tkDNN::Softmax l7(&net);
tkDNN::NetworkRT netRT(&net); tkDNN::NetworkRT netRT(&net, "mnist.rt");
// Load input // Load input
dnnType *data; dnnType *data;
+1 -1
View File
@@ -60,7 +60,7 @@ int main() {
net.print(); net.print();
//convert network to tensorRT //convert network to tensorRT
tkDNN::NetworkRT netRT(&net); tkDNN::NetworkRT netRT(&net, "yolo-tiny.rt");
dnnType *out_data, *out_data2; // cudnn output, tensorRT output dnnType *out_data, *out_data2; // cudnn output, tensorRT output
+1 -1
View File
@@ -108,7 +108,7 @@ int main() {
net.print(); net.print();
//convert network to tensorRT //convert network to tensorRT
tkDNN::NetworkRT netRT(&net); tkDNN::NetworkRT netRT(&net, "yolo.rt");
dnnType *out_data, *out_data2; // cudnn output, tensorRT output dnnType *out_data, *out_data2; // cudnn output, tensorRT output