Fixed tensorrt8 branch to work jetpack 4.5 and tensorrt7

Signed-off-by: perseusdg <f20180523@goa.bits-pilani.ac.in>
This commit was merged in pull request #278.
This commit is contained in:
Harshvardhan Chandirasekar
2022-03-16 18:09:32 +05:30
committed by perseusdg
parent 480b5a9c5a
commit 40266a6c32
3 changed files with 17 additions and 7 deletions
+8 -1
View File
@@ -834,7 +834,7 @@ IPluginV2Layer* NetworkRT::convert_layer(ITensor *input, Yolo *l) {
return lRT;
}
IResizeLayer* NetworkRT::convert_layer(ITensor *input, Upsample *l) {
ILayer* NetworkRT::convert_layer(ITensor *input, Upsample *l) {
#if NV_TENSORRT_MAJOR < 8
auto creator = getPluginRegistry()->getPluginCreator("UpSample_tkDNN","1");
@@ -1008,6 +1008,7 @@ bool NetworkRT::deserialize(const char *filename) {
return true;
}
#if NV_TENSORRT_MAJOR > 7
void NetworkRT::destroy() {
delete contextRT;
if(builderActive) {
@@ -1015,5 +1016,11 @@ void NetworkRT::destroy() {
delete builderRT;
}
}
#elif NV_TENSORRT_MAJOR <=7
void NetworkRT::destroy() {
}
#endif
}}
+8 -5
View File
@@ -66,11 +66,12 @@ int ConstantPaddingRT::enqueue(int batchSize, const void *const *inputs, void *c
return 0;
}
#elif NV_TENSORRT_MAJOR <= 7
int32_t enqueue (int32_t batchSize, const void *const *inputs, void **outputs, void *workspace, cudaStream_t stream) {
dnnType* srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
dnnType* dstData = reinterpret_cast<dnnType*>(outputs[0]);
constant_pad2d_forward(srcData,dstData,i_h,i_w,o_h,o_w,c,n,padH,padW,constant,stream);
return 0;
int32_t ConstantPaddingRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
cudaStream_t stream) {
dnnType* srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
dnnType* dstData = reinterpret_cast<dnnType*>(outputs[0]);
constant_pad2d_forward(srcData,dstData,i_h,i_w,o_h,o_w,c,n,padH,padW,constant,stream);
return 0;
}
#endif
@@ -151,6 +152,8 @@ bool ConstantPaddingRT::supportsFormat(DataType type, PluginFormat format) const
return (type == DataType::kFLOAT && format == PluginFormat::kLINEAR);
}
ConstantPaddingRTPluginCreator::ConstantPaddingRTPluginCreator() {
mPluginAttributes.clear();
mFC.nbFields = mPluginAttributes.size();