From f6884f07bf54d0f143012c03a30d5e8ad4a84d93 Mon Sep 17 00:00:00 2001 From: oren Date: Tue, 15 Dec 2020 16:39:08 -0600 Subject: [PATCH] Can now specify input and output names to NetworkRT --- include/tkDNN/NetworkRT.h | 2 +- src/NetworkRT.cpp | 12 ++++++------ 2 files changed, 7 insertions(+), 7 deletions(-) diff --git a/include/tkDNN/NetworkRT.h b/include/tkDNN/NetworkRT.h index 4c6c816..d28b96b 100644 --- a/include/tkDNN/NetworkRT.h +++ b/include/tkDNN/NetworkRT.h @@ -73,7 +73,7 @@ public: PluginFactory *pluginFactory; - NetworkRT(Network *net, const char *name); + NetworkRT(Network *net, const char *name, 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 889fb9f..bc657ae 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) { +NetworkRT::NetworkRT(Network *net, const char *name, const char *input_name, const char *output_name) { float rt_ver = float(NV_TENSORRT_MAJOR) + float(NV_TENSORRT_MINOR)/10 + @@ -97,13 +97,13 @@ NetworkRT::NetworkRT(Network *net, const char *name) { calibrator.reset(new Int8EntropyCalibrator(calibrationStream, 1, calib_table_name, - "data")); + input_name)); configRT->setInt8Calibrator(calibrator.get()); } #endif // add input layer - ITensor *input = networkRT->addInput("data", DataType::kFLOAT, + ITensor *input = networkRT->addInput(input_name, DataType::kFLOAT, DimsCHW{ dim.c, dim.h, dim.w}); checkNULL(input); @@ -130,7 +130,7 @@ NetworkRT::NetworkRT(Network *net, const char *name) { FatalError("conversion failed"); //build tensorRT - input->setName("out"); + input->setName(output_name); networkRT->markOutput(*input); std::cout<<"Selected maxBatchSize: "<getMaxBatchSize()<<"\n"; @@ -161,8 +161,8 @@ NetworkRT::NetworkRT(Network *net, const char *name) { // In order to bind the buffers, we need to know the names of the input and output tensors. // note that indices are guaranteed to be less than IEngine::getNbBindings() - buf_input_idx = engineRT->getBindingIndex("data"); - buf_output_idx = engineRT->getBindingIndex("out"); + buf_input_idx = engineRT->getBindingIndex(input_name); + buf_output_idx = engineRT->getBindingIndex(output_name); std::cout<<"input index = "< output index = "<