YoloRT save bias, mask and clasesName into RT file
This commit is contained in:
@@ -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;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user