diff --git a/.gitignore b/.gitignore index 85a8b85..d62c96d 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,5 @@ build/ *.hdf5 *.pk *.table +demo/COCO_val2017 +demo/BDD100k_val \ No newline at end of file diff --git a/include/tkDNN/DarknetParser.h b/include/tkDNN/DarknetParser.h index 61040f9..0f0e12a 100644 --- a/include/tkDNN/DarknetParser.h +++ b/include/tkDNN/DarknetParser.h @@ -127,7 +127,7 @@ namespace tk { namespace dnn { } - void darknetAddLayer(tk::dnn::Network *net, darknetFields_t &f, std::string wgs_path, std::vector &netLayers) { + void darknetAddLayer(tk::dnn::Network *net, darknetFields_t &f, std::string wgs_path, std::vector &netLayers, const std::vector& names) { if(net == nullptr) FatalError("Cant add a layer without a Net\n"); @@ -180,14 +180,30 @@ namespace tk { namespace dnn { } else if(f.type == "yolo") { std::string wgs = wgs_path + "/g" + std::to_string(netLayers.size()) + ".bin"; printf("%d %d %s %d %f\n", f.classes, f.num/f.n_mask, wgs.c_str(), f.n_mask, f.scale_xy); - netLayers.push_back(new tk::dnn::Yolo(net, f.classes, f.num/f.n_mask, wgs, f.n_mask, f.scale_xy)); + tk::dnn::Yolo *l = new tk::dnn::Yolo(net, f.classes, f.num/f.n_mask, wgs, f.n_mask, f.scale_xy); + l->classesNames = names; + netLayers.push_back(l); } else{ FatalError("layer not supported: " + f.type); } } - tk::dnn::Network* darknetParser(std::string cfg_file, std::string wgs_path) { + std::vector darknetReadNames(const std::string& names_file){ + std::ifstream if_names(names_file); + if(!if_names.is_open()) + FatalError("cloud not open names file: " + names_file); + + std::vector names; + std::string line; + while(std::getline(if_names, line)) + names.push_back(line); + + if_names.close(); + return names; + } + + tk::dnn::Network* darknetParser(const std::string& cfg_file, const std::string& wgs_path, const std::string& names_file) { tk::dnn::Network *net = nullptr; @@ -198,6 +214,8 @@ namespace tk { namespace dnn { if(!if_cfg.is_open()) FatalError("cloud not open cfg file: " + cfg_file); + std::vector names = darknetReadNames(names_file); + darknetFields_t fields; // will be filled with layers fields std::string line; while(std::getline(if_cfg, line)) { @@ -218,7 +236,7 @@ namespace tk { namespace dnn { if(fields.type == "net") net = darknetAddNet(fields); else - darknetAddLayer(net, fields, wgs_path, netLayers); + darknetAddLayer(net, fields, wgs_path, netLayers, names); } // new type @@ -237,7 +255,7 @@ namespace tk { namespace dnn { // end of filled type if(fields.type != "") { - darknetAddLayer(net, fields, wgs_path, netLayers); + darknetAddLayer(net, fields, wgs_path, netLayers, names); } if(net == nullptr) { diff --git a/tests/yolo3/coco.names b/tests/yolo3/coco.names new file mode 100644 index 0000000..ca76c80 --- /dev/null +++ b/tests/yolo3/coco.names @@ -0,0 +1,80 @@ +person +bicycle +car +motorbike +aeroplane +bus +train +truck +boat +traffic light +fire hydrant +stop sign +parking meter +bench +bird +cat +dog +horse +sheep +cow +elephant +bear +zebra +giraffe +backpack +umbrella +handbag +tie +suitcase +frisbee +skis +snowboard +sports ball +kite +baseball bat +baseball glove +skateboard +surfboard +tennis racket +bottle +wine glass +cup +fork +knife +spoon +bowl +banana +apple +sandwich +orange +broccoli +carrot +hot dog +pizza +donut +cake +chair +sofa +pottedplant +bed +diningtable +toilet +tvmonitor +laptop +mouse +remote +keyboard +cell phone +microwave +oven +toaster +sink +refrigerator +book +clock +vase +scissors +teddy bear +hair drier +toothbrush diff --git a/tests/yolo3/yolo3.cpp b/tests/yolo3/yolo3.cpp index 1036261..89b8283 100644 --- a/tests/yolo3/yolo3.cpp +++ b/tests/yolo3/yolo3.cpp @@ -9,7 +9,7 @@ int main() { std::string bin_path = "yolo3"; downloadWeightsifDoNotExist("yolo3/layers/input.bin", bin_path, "https://cloud.hipert.unimore.it/s/jPXmHyptpLoNdNR/download"); - tk::dnn::Network *net = tk::dnn::darknetParser("../tests/yolo3/yolov3.cfg", "yolo3/layers"); + tk::dnn::Network *net = tk::dnn::darknetParser("../tests/yolo3/yolov3.cfg", "yolo3/layers", "../tests/yolo3/coco.names"); net->print(); std::vector yolo; @@ -18,11 +18,6 @@ int main() { yolo.push_back((tk::dnn::Yolo*)net->layers[i]); } - // fill classes names - for(int i=0; i<3; i++) { - yolo[i]->classesNames = {"person" , "bicycle" , "car" , "motorbike" , "aeroplane" , "bus" , "train" , "truck" , "boat" , "traffic light" , "fire hydrant" , "stop sign" , "parking meter" , "bench" , "bird" , "cat" , "dog" , "horse" , "sheep" , "cow" , "elephant" , "bear" , "zebra" , "giraffe" , "backpack" , "umbrella" , "handbag" , "tie" , "suitcase" , "frisbee" , "skis" , "snowboard" , "sports ball" , "kite" , "baseball bat" , "baseball glove" , "skateboard" , "surfboard" , "tennis racket" , "bottle" , "wine glass" , "cup" , "fork" , "knife" , "spoon" , "bowl" , "banana" , "apple" , "sandwich" , "orange" , "broccoli" , "carrot" , "hot dog" , "pizza" , "donut" , "cake" , "chair" , "sofa" , "pottedplant" , "bed" , "diningtable" , "toilet" , "tvmonitor" , "laptop" , "mouse" , "remote" , "keyboard" , "cell phone" , "microwave" , "oven" , "toaster" , "sink" , "refrigerator" , "book" , "clock" , "vase" , "scissors" , "teddy bear" , "hair drier" , "toothbrush"}; - } - //convert network to tensorRT tk::dnn::NetworkRT netRT(net, net->getNetworkRTName("yolo3"));