From 0aa9de4ce827f4477af4e1b79126fde3eacbcac1 Mon Sep 17 00:00:00 2001 From: Francesco Gatti Date: Mon, 21 Aug 2017 12:10:17 +0000 Subject: [PATCH] better rt inference --- include/NetworkRT.h | 3 ++- src/NetworkRT.cpp | 4 ++++ tests/test_rtinference/rtinference.cpp | 23 ++++++++++++++++------- 3 files changed, 22 insertions(+), 8 deletions(-) diff --git a/include/NetworkRT.h b/include/NetworkRT.h index 710c4d4..73515da 100644 --- a/include/NetworkRT.h +++ b/include/NetworkRT.h @@ -32,6 +32,7 @@ public: Do inferece */ dnnType* infer(dataDim_t &dim, dnnType* data); + void enqueue(); nvinfer1::ILayer* convert_layer(nvinfer1::ITensor *input, Layer *l); nvinfer1::ILayer* convert_layer(nvinfer1::ITensor *input, Conv2d *l); @@ -62,4 +63,4 @@ template T readBUF(const char*& buffer) } } -#endif //NETWORKRT_H \ No newline at end of file +#endif //NETWORKRT_H diff --git a/src/NetworkRT.cpp b/src/NetworkRT.cpp index 27f2901..3eb486a 100644 --- a/src/NetworkRT.cpp +++ b/src/NetworkRT.cpp @@ -139,6 +139,10 @@ dnnType* NetworkRT::infer(dataDim_t &dim, dnnType* data) { return output; } +void NetworkRT::enqueue() { + contextRT->enqueue(1, buffersRT, stream, nullptr); +} + ILayer* NetworkRT::convert_layer(ITensor *input, Layer *l) { layerType_t type = l->getLayerType(); diff --git a/tests/test_rtinference/rtinference.cpp b/tests/test_rtinference/rtinference.cpp index a731394..bd8e6ff 100644 --- a/tests/test_rtinference/rtinference.cpp +++ b/tests/test_rtinference/rtinference.cpp @@ -1,24 +1,33 @@ #include #include "tkdnn.h" +#include /* srand, rand */ int main(int argc, char *argv[]) { if(argc < 2 || !fileExist(argv[1])) FatalError("unable to read serialRT file"); + //always same test + srand (0); + //convert network to tensorRT tkDNN::NetworkRT netRT(NULL, argv[1]); - tkDNN::dataDim_t dim = netRT.input_dim; - dnnType *data; - checkCuda(cudaMalloc(&data, dim.tot()*sizeof(dnnType))); + dnnType *input = new float[netRT.input_dim.tot()]; + dnnType *output = new float[netRT.input_dim.tot()]; - printCenteredTitle(" TENSORRT inference ", '=', 30); { - dim.print(); + printCenteredTitle(" TENSORRT inference ", '=', 30); + for(int i=0; i<100; i++) { + for(int j=0; j