From f189efcbb1e3f8b3522b3e3d8c8844433c5af1e8 Mon Sep 17 00:00:00 2001 From: perseusdg Date: Fri, 7 Jan 2022 05:41:26 +0000 Subject: [PATCH] Reflection Padding native plugin fix ,forgot to added input_dim.c in the plugin creator --- CMakeLists.txt | 5 +++-- include/tkDNN/pluginsRT/ReflectionPadding.h | 2 ++ src/NetworkRT.cpp | 4 +++- 3 files changed, 8 insertions(+), 3 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 6a775ed..cadaae7 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -85,12 +85,14 @@ endif() find_package(CUDNN REQUIRED) include_directories(${CUDNN_INCLUDE_DIR}) +find_package(yaml-cpp REQUIRED) + # compile file(GLOB tkdnn_CUSRC "src/kernels/*.cu" "src/sorting.cu" "src/pluginsRT/*.cpp") cuda_include_directories(${CMAKE_CURRENT_SOURCE_DIR}/include ${CUDA_INCLUDE_DIRS} ${CUDNN_INCLUDE_DIRS}) cuda_add_library(kernels SHARED ${tkdnn_CUSRC}) -target_link_libraries(kernels ${CUDA_CUBLAS_LIBRARIES} ${CUDA_LIBRARIES} ${CUDNN_LIBRARIES}) +target_link_libraries(kernels ${CUDA_CUBLAS_LIBRARIES} ${CUDA_LIBRARIES} ${CUDNN_LIBRARIES} yaml-cpp) @@ -120,7 +122,6 @@ endif() # endif() # gives problems in cross-compiling, probably malformed cmake config -find_package(yaml-cpp REQUIRED) #------------------------------------------------------------------------------- # Build Libraries diff --git a/include/tkDNN/pluginsRT/ReflectionPadding.h b/include/tkDNN/pluginsRT/ReflectionPadding.h index 894ed98..7b13710 100644 --- a/include/tkDNN/pluginsRT/ReflectionPadding.h +++ b/include/tkDNN/pluginsRT/ReflectionPadding.h @@ -94,6 +94,8 @@ namespace nvinfer1{ static std::vector mPluginAttributes; std::string mPluginNamespace; }; + + REGISTER_TENSORRT_PLUGIN(ReflectionPaddingRTPluginCreator); }; #endif diff --git a/src/NetworkRT.cpp b/src/NetworkRT.cpp index 6192006..549b010 100644 --- a/src/NetworkRT.cpp +++ b/src/NetworkRT.cpp @@ -484,13 +484,15 @@ ILayer* NetworkRT::convert_layer(ITensor *input,Padding *l){ mPluginAttributes.emplace_back(PluginField("outputH",&l->output_dim.h,PluginFieldType::kINT32,1)); mPluginAttributes.emplace_back(PluginField("outputW",&l->output_dim.w,PluginFieldType::kINT32,1)); mPluginAttributes.emplace_back(PluginField("n",&l->input_dim.n,PluginFieldType::kINT32,1)); + mPluginAttributes.emplace_back(PluginField("c",&l->input_dim.c,PluginFieldType::kINT32,1)); mFC.nbFields = mPluginAttributes.size(); mFC.fields = mPluginAttributes.data(); auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC); - } auto *lRT = networkRT->addPluginV2(&input, 1, *plugin); checkNULL(lRT); return lRT; + } + #endif }