From fe2e4eae92324d4c03482e20b5746ebc6aed75c5 Mon Sep 17 00:00:00 2001 From: Micaela Verucchi Date: Tue, 30 Jun 2020 15:52:58 +0200 Subject: [PATCH] yolo4tiny works on tensorRT :dolphin: :dolphin: :dolphin: :dolphin: Signed-off-by: Micaela Verucchi --- include/tkDNN/NetworkRT.h | 2 +- include/tkDNN/pluginsRT/RouteRT.h | 17 ++++++++++++----- src/NetworkRT.cpp | 19 +++++++++++-------- tests/darknet/yolo4tiny.cpp | 6 +++--- 4 files changed, 27 insertions(+), 17 deletions(-) diff --git a/include/tkDNN/NetworkRT.h b/include/tkDNN/NetworkRT.h index ee1f728..66b4f3d 100644 --- a/include/tkDNN/NetworkRT.h +++ b/include/tkDNN/NetworkRT.h @@ -28,7 +28,7 @@ using namespace nvinfer1; #include "pluginsRT/ActivationMishRT.h" #include "pluginsRT/ReorgRT.h" #include "pluginsRT/RegionRT.h" -//#include "pluginsRT/RouteRT.h" +#include "pluginsRT/RouteRT.h" #include "pluginsRT/ShortcutRT.h" #include "pluginsRT/YoloRT.h" #include "pluginsRT/UpsampleRT.h" diff --git a/include/tkDNN/pluginsRT/RouteRT.h b/include/tkDNN/pluginsRT/RouteRT.h index 0e94a97..263893f 100644 --- a/include/tkDNN/pluginsRT/RouteRT.h +++ b/include/tkDNN/pluginsRT/RouteRT.h @@ -8,7 +8,9 @@ class RouteRT : public IPlugin { */ public: - RouteRT() { + RouteRT(int groups, int group_id) { + this->groups = groups; + this->group_id = group_id; } ~RouteRT(){ @@ -22,7 +24,7 @@ public: Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override { int out_c = 0; for(int i=0; i(inputs[i]); int in_dim = c_in[i]*h*w; - checkCuda( cudaMemcpyAsync(dstData + offset, input, in_dim*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream) ); - offset += in_dim; + int part_in_dim = in_dim / this->groups; + checkCuda( cudaMemcpyAsync(dstData + offset, input + this->group_id*part_in_dim, part_in_dim*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream) ); + offset += part_in_dim; } return 0; @@ -65,11 +69,13 @@ public: virtual size_t getSerializationSize() override { - return (4+MAX_INPUTS)*sizeof(int); + return (6+MAX_INPUTS)*sizeof(int); } virtual void serialize(void* buffer) override { char *buf = reinterpret_cast(buffer); + tk::dnn::writeBUF(buf, groups); + tk::dnn::writeBUF(buf, group_id); tk::dnn::writeBUF(buf, in); for(int i=0; iaddConcatenation(tens, l->layers_n); - //IPlugin *plugin = new RouteRT(); - //IPluginLayer *lRT = networkRT->addPlugin(tens, l->layers_n, *plugin); - checkNULL(lRT); + if(l->groups > 1){ + IPlugin *plugin = new RouteRT(l->groups, l->group_id); + IPluginLayer *lRT = networkRT->addPlugin(tens, l->layers_n, *plugin); + checkNULL(lRT); + return lRT; + } + IConcatenationLayer *lRT = networkRT->addConcatenation(tens, l->layers_n); + checkNULL(lRT); return lRT; } @@ -766,9 +769,9 @@ IPlugin* PluginFactory::createPlugin(const char* layerName, const void* serialDa r->w = readBUF(buf); return r; } -/* + if(name.find("Route") == 0) { - RouteRT *r = new RouteRT(); + RouteRT *r = new RouteRT(readBUF(buf),readBUF(buf)); r->in = readBUF(buf); for(int i=0; ic_in[i] = readBUF(buf); @@ -777,7 +780,7 @@ IPlugin* PluginFactory::createPlugin(const char* layerName, const void* serialDa r->w = readBUF(buf); return r; } -*/ + if(name.find("Deformable") == 0) { DeformableConvRT *r = new DeformableConvRT(readBUF(buf), readBUF(buf), readBUF(buf), readBUF(buf), readBUF(buf), readBUF(buf), diff --git a/tests/darknet/yolo4tiny.cpp b/tests/darknet/yolo4tiny.cpp index d9011a8..44fbac8 100644 --- a/tests/darknet/yolo4tiny.cpp +++ b/tests/darknet/yolo4tiny.cpp @@ -23,11 +23,11 @@ int main() { net->print(); //convert network to tensorRT - // tk::dnn::NetworkRT *netRT = new tk::dnn::NetworkRT(net, net->getNetworkRTName(bin_path.c_str())); + tk::dnn::NetworkRT *netRT = new tk::dnn::NetworkRT(net, net->getNetworkRTName(bin_path.c_str())); - int ret = testInference(input_bins, output_bins, net, nullptr); + int ret = testInference(input_bins, output_bins, net, netRT); net->releaseLayers(); delete net; - // delete netRT; + delete netRT; return ret; }