Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# -*- coding: utf-8 -*-2# Copyright (c) Facebook, Inc. and its affiliates.3 4import logging5import numpy as np6import os7import tempfile8import xml.etree.ElementTree as ET9from collections import OrderedDict, defaultdict10from functools import lru_cache11import torch12 13from detectron2.data import MetadataCatalog14from detectron2.utils import comm15from detectron2.utils.file_io import PathManager16 17from .evaluator import DatasetEvaluator18 19 20class PascalVOCDetectionEvaluator(DatasetEvaluator):21 """22 Evaluate Pascal VOC style AP for Pascal VOC dataset.23 It contains a synchronization, therefore has to be called from all ranks.24 25 Note that the concept of AP can be implemented in different ways and may not26 produce identical results. This class mimics the implementation of the official27 Pascal VOC Matlab API, and should produce similar but not identical results to the28 official API.29 """30 31 def __init__(self, dataset_name):32 """33 Args:34 dataset_name (str): name of the dataset, e.g., "voc_2007_test"35 """36 self._dataset_name = dataset_name37 meta = MetadataCatalog.get(dataset_name)38 39 # Too many tiny files, download all to local for speed.40 annotation_dir_local = PathManager.get_local_path(41 os.path.join(meta.dirname, "Annotations/")42 )43 self._anno_file_template = os.path.join(annotation_dir_local, "{}.xml")44 self._image_set_path = os.path.join(meta.dirname, "ImageSets", "Main", meta.split + ".txt")45 self._class_names = meta.thing_classes46 assert meta.year in [2007, 2012], meta.year47 self._is_2007 = meta.year == 200748 self._cpu_device = torch.device("cpu")49 self._logger = logging.getLogger(__name__)50 51 def reset(self):52 self._predictions = defaultdict(list) # class name -> list of prediction strings53 54 def process(self, inputs, outputs):55 for input, output in zip(inputs, outputs):56 image_id = input["image_id"]57 instances = output["instances"].to(self._cpu_device)58 boxes = instances.pred_boxes.tensor.numpy()59 scores = instances.scores.tolist()60 classes = instances.pred_classes.tolist()61 for box, score, cls in zip(boxes, scores, classes):62 xmin, ymin, xmax, ymax = box63 # The inverse of data loading logic in `datasets/pascal_voc.py`64 xmin += 165 ymin += 166 self._predictions[cls].append(67 f"{image_id} {score:.3f} {xmin:.1f} {ymin:.1f} {xmax:.1f} {ymax:.1f}"68 )69 70 def evaluate(self):71 """72 Returns:73 dict: has a key "segm", whose value is a dict of "AP", "AP50", and "AP75".74 """75 all_predictions = comm.gather(self._predictions, dst=0)76 if not comm.is_main_process():77 return78 predictions = defaultdict(list)79 for predictions_per_rank in all_predictions:80 for clsid, lines in predictions_per_rank.items():81 predictions[clsid].extend(lines)82 del all_predictions83 84 self._logger.info(85 "Evaluating {} using {} metric. "86 "Note that results do not use the official Matlab API.".format(87 self._dataset_name, 2007 if self._is_2007 else 201288 )89 )90 91 with tempfile.TemporaryDirectory(prefix="pascal_voc_eval_") as dirname:92 res_file_template = os.path.join(dirname, "{}.txt")93 94 aps = defaultdict(list) # iou -> ap per class95 for cls_id, cls_name in enumerate(self._class_names):96 lines = predictions.get(cls_id, [""])97 98 with open(res_file_template.format(cls_name), "w") as f:99 f.write("\n".join(lines))100 101 for thresh in range(50, 100, 5):102 rec, prec, ap = voc_eval(103 res_file_template,104 self._anno_file_template,105 self._image_set_path,106 cls_name,107 ovthresh=thresh / 100.0,108 use_07_metric=self._is_2007,109 )110 aps[thresh].append(ap * 100)111 112 ret = OrderedDict()113 mAP = {iou: np.mean(x) for iou, x in aps.items()}114 ret["bbox"] = {"AP": np.mean(list(mAP.values())), "AP50": mAP[50], "AP75": mAP[75]}115 return ret116 117 118##############################################################################119#120# Below code is modified from121# https://github.com/rbgirshick/py-faster-rcnn/blob/master/lib/datasets/voc_eval.py122# --------------------------------------------------------123# Fast/er R-CNN124# Licensed under The MIT License [see LICENSE for details]125# Written by Bharath Hariharan126# --------------------------------------------------------127 128"""Python implementation of the PASCAL VOC devkit's AP evaluation code."""129 130 131@lru_cache(maxsize=None)132def parse_rec(filename):133 """Parse a PASCAL VOC xml file."""134 with PathManager.open(filename) as f:135 tree = ET.parse(f)136 objects = []137 for obj in tree.findall("object"):138 obj_struct = {}139 obj_struct["name"] = obj.find("name").text140 obj_struct["pose"] = obj.find("pose").text141 obj_struct["truncated"] = int(obj.find("truncated").text)142 obj_struct["difficult"] = int(obj.find("difficult").text)143 bbox = obj.find("bndbox")144 obj_struct["bbox"] = [145 int(bbox.find("xmin").text),146 int(bbox.find("ymin").text),147 int(bbox.find("xmax").text),148 int(bbox.find("ymax").text),149 ]150 objects.append(obj_struct)151 152 return objects153 154 155def voc_ap(rec, prec, use_07_metric=False):156 """Compute VOC AP given precision and recall. If use_07_metric is true, uses157 the VOC 07 11-point method (default:False).158 """159 if use_07_metric:160 # 11 point metric161 ap = 0.0162 for t in np.arange(0.0, 1.1, 0.1):163 if np.sum(rec >= t) == 0:164 p = 0165 else:166 p = np.max(prec[rec >= t])167 ap = ap + p / 11.0168 else:169 # correct AP calculation170 # first append sentinel values at the end171 mrec = np.concatenate(([0.0], rec, [1.0]))172 mpre = np.concatenate(([0.0], prec, [0.0]))173 174 # compute the precision envelope175 for i in range(mpre.size - 1, 0, -1):176 mpre[i - 1] = np.maximum(mpre[i - 1], mpre[i])177 178 # to calculate area under PR curve, look for points179 # where X axis (recall) changes value180 i = np.where(mrec[1:] != mrec[:-1])[0]181 182 # and sum (\Delta recall) * prec183 ap = np.sum((mrec[i + 1] - mrec[i]) * mpre[i + 1])184 return ap185 186 187def voc_eval(detpath, annopath, imagesetfile, classname, ovthresh=0.5, use_07_metric=False):188 """rec, prec, ap = voc_eval(detpath,189 annopath,190 imagesetfile,191 classname,192 [ovthresh],193 [use_07_metric])194 195 Top level function that does the PASCAL VOC evaluation.196 197 detpath: Path to detections198 detpath.format(classname) should produce the detection results file.199 annopath: Path to annotations200 annopath.format(imagename) should be the xml annotations file.201 imagesetfile: Text file containing the list of images, one image per line.202 classname: Category name (duh)203 [ovthresh]: Overlap threshold (default = 0.5)204 [use_07_metric]: Whether to use VOC07's 11 point AP computation205 (default False)206 """207 # assumes detections are in detpath.format(classname)208 # assumes annotations are in annopath.format(imagename)209 # assumes imagesetfile is a text file with each line an image name210 211 # first load gt212 # read list of images213 with PathManager.open(imagesetfile, "r") as f:214 lines = f.readlines()215 imagenames = [x.strip() for x in lines]216 217 # load annots218 recs = {}219 for imagename in imagenames:220 recs[imagename] = parse_rec(annopath.format(imagename))221 222 # extract gt objects for this class223 class_recs = {}224 npos = 0225 for imagename in imagenames:226 R = [obj for obj in recs[imagename] if obj["name"] == classname]227 bbox = np.array([x["bbox"] for x in R])228 difficult = np.array([x["difficult"] for x in R]).astype(bool)229 # difficult = np.array([False for x in R]).astype(bool) # treat all "difficult" as GT230 det = [False] * len(R)231 npos = npos + sum(~difficult)232 class_recs[imagename] = {"bbox": bbox, "difficult": difficult, "det": det}233 234 # read dets235 detfile = detpath.format(classname)236 with open(detfile, "r") as f:237 lines = f.readlines()238 239 splitlines = [x.strip().split(" ") for x in lines]240 image_ids = [x[0] for x in splitlines]241 confidence = np.array([float(x[1]) for x in splitlines])242 BB = np.array([[float(z) for z in x[2:]] for x in splitlines]).reshape(-1, 4)243 244 # sort by confidence245 sorted_ind = np.argsort(-confidence)246 BB = BB[sorted_ind, :]247 image_ids = [image_ids[x] for x in sorted_ind]248 249 # go down dets and mark TPs and FPs250 nd = len(image_ids)251 tp = np.zeros(nd)252 fp = np.zeros(nd)253 for d in range(nd):254 R = class_recs[image_ids[d]]255 bb = BB[d, :].astype(float)256 ovmax = -np.inf257 BBGT = R["bbox"].astype(float)258 259 if BBGT.size > 0:260 # compute overlaps261 # intersection262 ixmin = np.maximum(BBGT[:, 0], bb[0])263 iymin = np.maximum(BBGT[:, 1], bb[1])264 ixmax = np.minimum(BBGT[:, 2], bb[2])265 iymax = np.minimum(BBGT[:, 3], bb[3])266 iw = np.maximum(ixmax - ixmin + 1.0, 0.0)267 ih = np.maximum(iymax - iymin + 1.0, 0.0)268 inters = iw * ih269 270 # union271 uni = (272 (bb[2] - bb[0] + 1.0) * (bb[3] - bb[1] + 1.0)273 + (BBGT[:, 2] - BBGT[:, 0] + 1.0) * (BBGT[:, 3] - BBGT[:, 1] + 1.0)274 - inters275 )276 277 overlaps = inters / uni278 ovmax = np.max(overlaps)279 jmax = np.argmax(overlaps)280 281 if ovmax > ovthresh:282 if not R["difficult"][jmax]:283 if not R["det"][jmax]:284 tp[d] = 1.0285 R["det"][jmax] = 1286 else:287 fp[d] = 1.0288 else:289 fp[d] = 1.0290 291 # compute precision recall292 fp = np.cumsum(fp)293 tp = np.cumsum(tp)294 rec = tp / float(npos)295 # avoid divide by zero in case the first detection matches a difficult296 # ground truth297 prec = tp / np.maximum(tp + fp, np.finfo(np.float64).eps)298 ap = voc_ap(rec, prec, use_07_metric)299 300 return rec, prec, ap301 