depth->tensorrt8 patches
This commit is contained in:
+3
-2
@@ -85,12 +85,14 @@ endif()
|
|||||||
find_package(CUDNN REQUIRED)
|
find_package(CUDNN REQUIRED)
|
||||||
include_directories(${CUDNN_INCLUDE_DIR})
|
include_directories(${CUDNN_INCLUDE_DIR})
|
||||||
|
|
||||||
|
find_package(yaml-cpp REQUIRED)
|
||||||
|
|
||||||
|
|
||||||
# compile
|
# compile
|
||||||
file(GLOB tkdnn_CUSRC "src/kernels/*.cu" "src/sorting.cu" "src/pluginsRT/*.cpp")
|
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_include_directories(${CMAKE_CURRENT_SOURCE_DIR}/include ${CUDA_INCLUDE_DIRS} ${CUDNN_INCLUDE_DIRS})
|
||||||
cuda_add_library(kernels SHARED ${tkdnn_CUSRC})
|
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()
|
# endif()
|
||||||
|
|
||||||
# gives problems in cross-compiling, probably malformed cmake config
|
# gives problems in cross-compiling, probably malformed cmake config
|
||||||
find_package(yaml-cpp REQUIRED)
|
|
||||||
|
|
||||||
#-------------------------------------------------------------------------------
|
#-------------------------------------------------------------------------------
|
||||||
# Build Libraries
|
# Build Libraries
|
||||||
|
|||||||
@@ -23,6 +23,8 @@
|
|||||||
#include <pluginsRT/ShortcutRT.h>
|
#include <pluginsRT/ShortcutRT.h>
|
||||||
#include <pluginsRT/UpsampleRT.h>
|
#include <pluginsRT/UpsampleRT.h>
|
||||||
#include <pluginsRT/YoloRT.h>
|
#include <pluginsRT/YoloRT.h>
|
||||||
|
#include <pluginsRT/ConstantPaddingRT.h>
|
||||||
|
#include <pluginsRT/ReflectionPadding.h>
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -94,6 +94,8 @@ namespace nvinfer1{
|
|||||||
static std::vector<PluginField> mPluginAttributes;
|
static std::vector<PluginField> mPluginAttributes;
|
||||||
std::string mPluginNamespace;
|
std::string mPluginNamespace;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
REGISTER_TENSORRT_PLUGIN(ReflectionPaddingRTPluginCreator);
|
||||||
};
|
};
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
|
|||||||
+39
-16
@@ -474,23 +474,46 @@ ILayer* NetworkRT::convert_layer(ITensor *input,Padding *l){
|
|||||||
#else
|
#else
|
||||||
//todo add PADDING_MODE_CONSTANT AND PADDING_MODE_ZERO for tensorrt versions < 8.2
|
//todo add PADDING_MODE_CONSTANT AND PADDING_MODE_ZERO for tensorrt versions < 8.2
|
||||||
if(l->padding_mode == PADDING_MODE_REFLECTION){
|
if(l->padding_mode == PADDING_MODE_REFLECTION){
|
||||||
auto creator = getPluginRegistry()->getPluginCreator("ReflectionPaddingRT_tkDNN","1");
|
auto creator = getPluginRegistry()->getPluginCreator("ReflectionPaddingRT_tkDNN","1");
|
||||||
std::vector<PluginField> mPluginAttributes;
|
std::vector<PluginField> mPluginAttributes;
|
||||||
PluginFieldCollection mFC{};
|
PluginFieldCollection mFC{};
|
||||||
mPluginAttributes.emplace_back(PluginField("padH",&l->paddingH,PluginFieldType::kINT32,1));
|
mPluginAttributes.emplace_back(PluginField("padH",&l->paddingH,PluginFieldType::kINT32,1));
|
||||||
mPluginAttributes.emplace_back(PluginField("padW",&l->paddingW,PluginFieldType::kINT32,1));
|
mPluginAttributes.emplace_back(PluginField("padW",&l->paddingW,PluginFieldType::kINT32,1));
|
||||||
mPluginAttributes.emplace_back(PluginField("inputH",&l->input_dim.h,PluginFieldType::kINT32,1));
|
mPluginAttributes.emplace_back(PluginField("inputH",&l->input_dim.h,PluginFieldType::kINT32,1));
|
||||||
mPluginAttributes.emplace_back(PluginField("inputW",&l->input_dim.w,PluginFieldType::kINT32,1));
|
mPluginAttributes.emplace_back(PluginField("inputW",&l->input_dim.w,PluginFieldType::kINT32,1));
|
||||||
mPluginAttributes.emplace_back(PluginField("outputH",&l->output_dim.h,PluginFieldType::kINT32,1));
|
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("outputW",&l->output_dim.w,PluginFieldType::kINT32,1));
|
||||||
mPluginAttributes.emplace_back(PluginField("n",&l->input_dim.n,PluginFieldType::kINT32,1));
|
mPluginAttributes.emplace_back(PluginField("n",&l->input_dim.n,PluginFieldType::kINT32,1));
|
||||||
mFC.nbFields = mPluginAttributes.size();
|
mPluginAttributes.emplace_back(PluginField("c",&l->input_dim.c,PluginFieldType::kINT32,1));
|
||||||
mFC.fields = mPluginAttributes.data();
|
mFC.nbFields = mPluginAttributes.size();
|
||||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
mFC.fields = mPluginAttributes.data();
|
||||||
|
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||||
|
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||||
|
checkNULL(lRT);
|
||||||
|
return lRT;
|
||||||
|
}else if(l->padding_mode == PADDING_MODE_CONSTANT || l->padding_mode == PADDING_MODE_ZERO){
|
||||||
|
auto creator = getPluginRegistry()->getPluginCreator("ConstantPaddingRT_tkDNN","1");
|
||||||
|
std::vector<PluginField> mPluginAttributes;
|
||||||
|
PluginFieldCollection mFC{};
|
||||||
|
mPluginAttributes.emplace_back(PluginField("padH",&l->paddingH,PluginFieldType::kINT32,1));
|
||||||
|
mPluginAttributes.emplace_back(PluginField("padW",&l->paddingW,PluginFieldType::kINT32,1));
|
||||||
|
mPluginAttributes.emplace_back(PluginField("inputH",&l->input_dim.h,PluginFieldType::kINT32,1));
|
||||||
|
mPluginAttributes.emplace_back(PluginField("inputW",&l->input_dim.w,PluginFieldType::kINT32,1));
|
||||||
|
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));
|
||||||
|
mPluginAttributes.emplace_back(PluginField("constant",&l->constant,PluginFieldType::kFLOAT32,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;
|
||||||
}
|
}
|
||||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
|
||||||
checkNULL(lRT);
|
return nullptr;
|
||||||
return lRT;
|
|
||||||
#endif
|
#endif
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user