YoloRT save bias, mask and clasesName into RT file

This commit is contained in:
Francesco Gatti
2022-03-30 20:46:51 +02:00
parent fa9db167b8
commit 5e71b99265
15 changed files with 104 additions and 79 deletions
+16
View File
@@ -15,6 +15,9 @@
using namespace nvinfer1;
extern std::mutex gYoloPlugins_mutex;
extern std::vector<YoloRT*> gYoloPlugins;
// Logger for info/warning/errors
class Logger : public ILogger {
void log(Severity severity, const char* msg) NOEXCEPT override {
@@ -826,6 +829,12 @@ IPluginV2Layer* NetworkRT::convert_layer(ITensor *input, Yolo *l) {
mPluginAttributes.emplace_back(PluginField("nms_thresh",&l->nms_thresh,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("nms_kins",&l->nsm_kind,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("new_coords",&l->new_coords,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("mask",l->mask_h,PluginFieldType::kFLOAT32,l->n_masks));
mPluginAttributes.emplace_back(PluginField("bias",l->bias_h,PluginFieldType::kFLOAT32,l->n_masks*2*l->num));
for(int i=0; i<l->classes; i++) {
mPluginAttributes.emplace_back(PluginField("class_name",l->classesNames[i].data(),PluginFieldType::kCHAR,l->classesNames[i].size()));
}
mFC.nbFields = mPluginAttributes.size();
mFC.fields = mPluginAttributes.data();
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
@@ -1001,7 +1010,14 @@ bool NetworkRT::deserialize(const char *filename) {
}
runtimeRT = createInferRuntime(loggerRT);
gYoloPlugins_mutex.lock();
gYoloPlugins.clear();
engineRT = runtimeRT->deserializeCudaEngine(gieModelStream, size);
yolo_plugins = gYoloPlugins;
gYoloPlugins.clear();
gYoloPlugins_mutex.unlock();
std::cout<<size<<std::endl;
//if (gieModelStream) delete [] gieModelStream;