update tensorrt8 branch

This commit is contained in:
perseusdg
2021-08-30 19:04:26 +05:30
parent 2ffe07057e
commit de83ae5d25
20 changed files with 96 additions and 97 deletions
+1
View File
@@ -119,6 +119,7 @@ public:
bool serialize(const char *filename);
bool deserialize(const char *filename);
void destroy();
+3 -2
View File
@@ -8,6 +8,7 @@ public:
ActivationLeakyRT(const void *data, size_t length)
{
std::cout<<"DESERIALIZE LEAKYRT"<<std::endl;
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
slope = readBUF<float>(buf);
size = readBUF<int>(buf);
@@ -139,8 +140,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
@@ -139,8 +139,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+3 -2
View File
@@ -9,6 +9,7 @@ public:
~ActivationMishRT() {}
ActivationMishRT(const void *data, size_t length) {
std::cout<<"DESERIALIZE MISH"<<std::endl;
const char *buf = reinterpret_cast<const char *>(data), *bufCheck = buf;
size = readBUF<int>(buf);
assert(buf == bufCheck + length);
@@ -126,8 +127,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
@@ -14,6 +14,7 @@ public:
}
ActivationReLUCeiling(const void *data, size_t length) {
std::cout<<"RELU CEILING DESERIALIZE"<<std::endl;
const char *buf = reinterpret_cast<const char *>(data), *bufCheck = buf;
ceiling = readBUF<float>(buf);
size = readBUF<int>(buf);
@@ -140,9 +141,9 @@ public:
return &mFC;
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
public:
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+2 -3
View File
@@ -352,10 +352,9 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
REGISTER_TENSORRT_PLUGIN(DeformableConvRTPluginCreator);
+2 -2
View File
@@ -154,8 +154,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
@@ -173,8 +173,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+2 -2
View File
@@ -179,8 +179,8 @@ public:
return &mFC;
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+2 -2
View File
@@ -134,8 +134,8 @@ public:
return &mFC;
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+2 -2
View File
@@ -140,8 +140,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+2 -2
View File
@@ -155,8 +155,8 @@ public:
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+2 -2
View File
@@ -173,8 +173,8 @@ public:
return &mFC;
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+3 -3
View File
@@ -157,9 +157,9 @@ public:
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
return &mFC;
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
public:
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
+7 -4
View File
@@ -73,7 +73,8 @@ public:
}
const char *getPluginVersion() const NOEXCEPT override{
return "UpsampleRT_tkDNN";
static const char* UPSAMPLE_RT_PLUGIN = "UpsampleRT_TRT";
return UPSAMPLE_RT_PLUGIN;
}
void destroy() NOEXCEPT override{delete this;}
@@ -128,7 +129,8 @@ public:
}
const char *getPluginName() const NOEXCEPT override{
return "UpsampleRT_tkDNN";
static const char* UPSAMPLE_RT_PLUGIN = "UpsampleRT_TRT";
return UPSAMPLE_RT_PLUGIN;
}
const char *getPluginVersion() const NOEXCEPT override{
@@ -139,9 +141,10 @@ public:
return &mFC;
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};
REGISTER_TENSORRT_PLUGIN(UpsampleRTPluginCreator);
+16 -7
View File
@@ -1,4 +1,5 @@
#include<cassert>
#include <vector>
#include "../kernels.h"
#define YOLORT_CLASSNAME_W 256
@@ -27,11 +28,12 @@ public:
}
YoloRT(const void *data,size_t length){
std::vector<float> maskTemp,biasTemp;
std::cout<<"LENGTH : "<<length<<std::endl;
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
classes = readBUF<int>(buf);
num = readBUF<int>(buf);
n_masks = readBUF<int>(buf);
std::cout<<n_masks<<std::endl;
scaleXY = readBUF<float>(buf);
nms_thresh = readBUF<float>(buf);
nms_kind = readBUF<int>(buf);
@@ -39,10 +41,16 @@ public:
c = readBUF<int>(buf);
h = readBUF<int>(buf);
w = readBUF<int>(buf);
for(int i=0;i<n_masks;i++)
mask[i] = readBUF<dnnType>(buf);
for(int i=0;i<n_masks*2*num;i++)
bias[i] = readBUF<dnnType>(buf);
for(int i=0;i<n_masks;i++){
maskTemp.push_back(readBUF<dnnType>(buf));
std::cout<<maskTemp[i]<<std::endl;
}
for(int i=0;i<n_masks*2*num;i++){
biasTemp.push_back(readBUF<dnnType>(buf));
std::cout<<biasTemp[i]<<std::endl;
}
mask = maskTemp.data();
bias = biasTemp.data();
classesNames.resize(classes);
for(int i=0;i<classes;i++){
char tmp[YOLORT_CLASSNAME_W];
@@ -133,6 +141,7 @@ public:
char *buf = reinterpret_cast<char *>(buffer), *a = buf;
tk::dnn::writeBUF(buf, classes); //std::cout << "Classes :" << classes << std::endl;
tk::dnn::writeBUF(buf, num); //std::cout << "Num : " << num << std::endl;
std::cout<<num<<std::endl;
tk::dnn::writeBUF(buf, n_masks); //std::cout << "N_Masks" << n_masks << std::endl;
tk::dnn::writeBUF(buf, scaleXY); //std::cout << "ScaleXY :" << scaleXY << std::endl;
tk::dnn::writeBUF(buf, nms_thresh); //std::cout << "nms_thresh :" << nms_thresh << std::endl;
@@ -265,8 +274,8 @@ public:
}
private:
static PluginFieldCollection mFC;
static std::vector<PluginField> mPluginAttributes;
PluginFieldCollection mFC;
std::vector<PluginField> mPluginAttributes;
std::string mPluginNamespace;
};