Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import glob3import logging4import numpy as np5import os6import tempfile7from collections import OrderedDict8import torch9from PIL import Image10 11from detectron2.data import MetadataCatalog12from detectron2.utils import comm13from detectron2.utils.file_io import PathManager14 15from .evaluator import DatasetEvaluator16 17 18class CityscapesEvaluator(DatasetEvaluator):19 """20 Base class for evaluation using cityscapes API.21 """22 23 def __init__(self, dataset_name):24 """25 Args:26 dataset_name (str): the name of the dataset.27 It must have the following metadata associated with it:28 "thing_classes", "gt_dir".29 """30 self._metadata = MetadataCatalog.get(dataset_name)31 self._cpu_device = torch.device("cpu")32 self._logger = logging.getLogger(__name__)33 34 def reset(self):35 self._working_dir = tempfile.TemporaryDirectory(prefix="cityscapes_eval_")36 self._temp_dir = self._working_dir.name37 # All workers will write to the same results directory38 # TODO this does not work in distributed training39 assert (40 comm.get_local_size() == comm.get_world_size()41 ), "CityscapesEvaluator currently do not work with multiple machines."42 self._temp_dir = comm.all_gather(self._temp_dir)[0]43 if self._temp_dir != self._working_dir.name:44 self._working_dir.cleanup()45 self._logger.info(46 "Writing cityscapes results to temporary directory {} ...".format(self._temp_dir)47 )48 49 50class CityscapesInstanceEvaluator(CityscapesEvaluator):51 """52 Evaluate instance segmentation results on cityscapes dataset using cityscapes API.53 54 Note:55 * It does not work in multi-machine distributed training.56 * It contains a synchronization, therefore has to be used on all ranks.57 * Only the main process runs evaluation.58 """59 60 def process(self, inputs, outputs):61 from cityscapesscripts.helpers.labels import name2label62 63 for input, output in zip(inputs, outputs):64 file_name = input["file_name"]65 basename = os.path.splitext(os.path.basename(file_name))[0]66 pred_txt = os.path.join(self._temp_dir, basename + "_pred.txt")67 68 if "instances" in output:69 output = output["instances"].to(self._cpu_device)70 num_instances = len(output)71 with open(pred_txt, "w") as fout:72 for i in range(num_instances):73 pred_class = output.pred_classes[i]74 classes = self._metadata.thing_classes[pred_class]75 class_id = name2label[classes].id76 score = output.scores[i]77 mask = output.pred_masks[i].numpy().astype("uint8")78 png_filename = os.path.join(79 self._temp_dir, basename + "_{}_{}.png".format(i, classes)80 )81 82 Image.fromarray(mask * 255).save(png_filename)83 fout.write(84 "{} {} {}\n".format(os.path.basename(png_filename), class_id, score)85 )86 else:87 # Cityscapes requires a prediction file for every ground truth image.88 with open(pred_txt, "w") as fout:89 pass90 91 def evaluate(self):92 """93 Returns:94 dict: has a key "segm", whose value is a dict of "AP" and "AP50".95 """96 comm.synchronize()97 if comm.get_rank() > 0:98 return99 import cityscapesscripts.evaluation.evalInstanceLevelSemanticLabeling as cityscapes_eval100 101 self._logger.info("Evaluating results under {} ...".format(self._temp_dir))102 103 # set some global states in cityscapes evaluation API, before evaluating104 cityscapes_eval.args.predictionPath = os.path.abspath(self._temp_dir)105 cityscapes_eval.args.predictionWalk = None106 cityscapes_eval.args.JSONOutput = False107 cityscapes_eval.args.colorized = False108 cityscapes_eval.args.gtInstancesFile = os.path.join(self._temp_dir, "gtInstances.json")109 110 # These lines are adopted from111 # https://github.com/mcordts/cityscapesScripts/blob/master/cityscapesscripts/evaluation/evalInstanceLevelSemanticLabeling.py # noqa112 gt_dir = PathManager.get_local_path(self._metadata.gt_dir)113 groundTruthImgList = glob.glob(os.path.join(gt_dir, "*", "*_gtFine_instanceIds.png"))114 assert len(115 groundTruthImgList116 ), "Cannot find any ground truth images to use for evaluation. Searched for: {}".format(117 cityscapes_eval.args.groundTruthSearch118 )119 predictionImgList = []120 for gt in groundTruthImgList:121 predictionImgList.append(cityscapes_eval.getPrediction(gt, cityscapes_eval.args))122 results = cityscapes_eval.evaluateImgLists(123 predictionImgList, groundTruthImgList, cityscapes_eval.args124 )["averages"]125 126 ret = OrderedDict()127 ret["segm"] = {"AP": results["allAp"] * 100, "AP50": results["allAp50%"] * 100}128 self._working_dir.cleanup()129 return ret130 131 132class CityscapesSemSegEvaluator(CityscapesEvaluator):133 """134 Evaluate semantic segmentation results on cityscapes dataset using cityscapes API.135 136 Note:137 * It does not work in multi-machine distributed training.138 * It contains a synchronization, therefore has to be used on all ranks.139 * Only the main process runs evaluation.140 """141 142 def process(self, inputs, outputs):143 from cityscapesscripts.helpers.labels import trainId2label144 145 for input, output in zip(inputs, outputs):146 file_name = input["file_name"]147 basename = os.path.splitext(os.path.basename(file_name))[0]148 pred_filename = os.path.join(self._temp_dir, basename + "_pred.png")149 150 output = output["sem_seg"].argmax(dim=0).to(self._cpu_device).numpy()151 pred = 255 * np.ones(output.shape, dtype=np.uint8)152 for train_id, label in trainId2label.items():153 if label.ignoreInEval:154 continue155 pred[output == train_id] = label.id156 Image.fromarray(pred).save(pred_filename)157 158 def evaluate(self):159 comm.synchronize()160 if comm.get_rank() > 0:161 return162 # Load the Cityscapes eval script *after* setting the required env var,163 # since the script reads CITYSCAPES_DATASET into global variables at load time.164 import cityscapesscripts.evaluation.evalPixelLevelSemanticLabeling as cityscapes_eval165 166 self._logger.info("Evaluating results under {} ...".format(self._temp_dir))167 168 # set some global states in cityscapes evaluation API, before evaluating169 cityscapes_eval.args.predictionPath = os.path.abspath(self._temp_dir)170 cityscapes_eval.args.predictionWalk = None171 cityscapes_eval.args.JSONOutput = False172 cityscapes_eval.args.colorized = False173 174 # These lines are adopted from175 # https://github.com/mcordts/cityscapesScripts/blob/master/cityscapesscripts/evaluation/evalPixelLevelSemanticLabeling.py # noqa176 gt_dir = PathManager.get_local_path(self._metadata.gt_dir)177 groundTruthImgList = glob.glob(os.path.join(gt_dir, "*", "*_gtFine_labelIds.png"))178 assert len(179 groundTruthImgList180 ), "Cannot find any ground truth images to use for evaluation. Searched for: {}".format(181 cityscapes_eval.args.groundTruthSearch182 )183 predictionImgList = []184 for gt in groundTruthImgList:185 predictionImgList.append(cityscapes_eval.getPrediction(cityscapes_eval.args, gt))186 results = cityscapes_eval.evaluateImgLists(187 predictionImgList, groundTruthImgList, cityscapes_eval.args188 )189 ret = OrderedDict()190 ret["sem_seg"] = {191 "IoU": 100.0 * results["averageScoreClasses"],192 "iIoU": 100.0 * results["averageScoreInstClasses"],193 "IoU_sup": 100.0 * results["averageScoreCategories"],194 "iIoU_sup": 100.0 * results["averageScoreInstCategories"],195 }196 self._working_dir.cleanup()197 return ret198 