update tensorrt8 branch
This commit is contained in:
@@ -119,6 +119,7 @@ public:
|
||||
|
||||
bool serialize(const char *filename);
|
||||
bool deserialize(const char *filename);
|
||||
void destroy();
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
};
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -140,8 +140,8 @@ public:
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC;
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
|
||||
@@ -155,8 +155,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:
|
||||
return &mFC;
|
||||
}
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC;
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user