Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
pascal_voc_evaluation.py301 linesDownload Raw Back to evaluation
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