MLVLab/Human_Object_Interaction
1
1# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved2"""3COCO evaluator that works in distributed mode.4Mostly copy-paste from https://github.com/pytorch/vision/blob/edfd5a7/references/detection/coco_eval.py5The difference is that there is less copy-pasting from pycocotools6in the end of the file, as python3 can suppress prints with contextlib7"""8import os9import contextlib10import copy11import numpy as np12import torch13 14from pycocotools.cocoeval import COCOeval15from pycocotools.coco import COCO16import pycocotools.mask as mask_util17 18from hotr.util.misc import all_gather19 20 21class CocoEvaluator(object):22 def __init__(self, coco_gt, iou_types):23 assert isinstance(iou_types, (list, tuple))24 coco_gt = copy.deepcopy(coco_gt)25 self.coco_gt = coco_gt26 27 self.iou_types = iou_types28 self.coco_eval = {}29 for iou_type in iou_types:30 self.coco_eval[iou_type] = COCOeval(coco_gt, iouType=iou_type)31 32 self.img_ids = []33 self.eval_imgs = {k: [] for k in iou_types}34 35 def update(self, predictions):36 img_ids = list(np.unique(list(predictions.keys())))37 self.img_ids.extend(img_ids)38 39 for iou_type in self.iou_types:40 results = self.prepare(predictions, iou_type)41 42 # suppress pycocotools prints43 with open(os.devnull, 'w') as devnull:44 with contextlib.redirect_stdout(devnull):45 coco_dt = COCO.loadRes(self.coco_gt, results) if results else COCO()46 coco_eval = self.coco_eval[iou_type]47 48 coco_eval.cocoDt = coco_dt49 coco_eval.params.imgIds = list(img_ids)50 img_ids, eval_imgs = evaluate(coco_eval)51 52 self.eval_imgs[iou_type].append(eval_imgs)53 54 def synchronize_between_processes(self):55 for iou_type in self.iou_types:56 self.eval_imgs[iou_type] = np.concatenate(self.eval_imgs[iou_type], 2)57 create_common_coco_eval(self.coco_eval[iou_type], self.img_ids, self.eval_imgs[iou_type])58 59 def accumulate(self):60 for coco_eval in self.coco_eval.values():61 coco_eval.accumulate()62 63 def summarize(self):64 for iou_type, coco_eval in self.coco_eval.items():65 print("IoU metric: {}".format(iou_type))66 coco_eval.summarize()67 68 def prepare(self, predictions, iou_type):69 if iou_type == "bbox":70 return self.prepare_for_coco_detection(predictions)71 elif iou_type == "segm":72 return self.prepare_for_coco_segmentation(predictions)73 elif iou_type == "keypoints":74 return self.prepare_for_coco_keypoint(predictions)75 else:76 raise ValueError("Unknown iou type {}".format(iou_type))77 78 def prepare_for_coco_detection(self, predictions):79 coco_results = []80 for original_id, prediction in predictions.items():81 if len(prediction) == 0:82 continue83 84 boxes = prediction["boxes"]85 boxes = convert_to_xywh(boxes).tolist()86 scores = prediction["scores"].tolist()87 labels = prediction["labels"].tolist()88 89 coco_results.extend(90 [91 {92 "image_id": original_id,93 "category_id": labels[k],94 "bbox": box,95 "score": scores[k],96 }97 for k, box in enumerate(boxes)98 ]99 )100 return coco_results101 102 def prepare_for_coco_segmentation(self, predictions):103 coco_results = []104 for original_id, prediction in predictions.items():105 if len(prediction) == 0:106 continue107 108 scores = prediction["scores"]109 labels = prediction["labels"]110 masks = prediction["masks"]111 112 masks = masks > 0.5113 114 scores = prediction["scores"].tolist()115 labels = prediction["labels"].tolist()116 117 rles = [118 mask_util.encode(np.array(mask[0, :, :, np.newaxis], dtype=np.uint8, order="F"))[0]119 for mask in masks120 ]121 for rle in rles:122 rle["counts"] = rle["counts"].decode("utf-8")123 124 coco_results.extend(125 [126 {127 "image_id": original_id,128 "category_id": labels[k],129 "segmentation": rle,130 "score": scores[k],131 }132 for k, rle in enumerate(rles)133 ]134 )135 return coco_results136 137 def prepare_for_coco_keypoint(self, predictions):138 coco_results = []139 for original_id, prediction in predictions.items():140 if len(prediction) == 0:141 continue142 143 boxes = prediction["boxes"]144 boxes = convert_to_xywh(boxes).tolist()145 scores = prediction["scores"].tolist()146 labels = prediction["labels"].tolist()147 keypoints = prediction["keypoints"]148 keypoints = keypoints.flatten(start_dim=1).tolist()149 150 coco_results.extend(151 [152 {153 "image_id": original_id,154 "category_id": labels[k],155 'keypoints': keypoint,156 "score": scores[k],157 }158 for k, keypoint in enumerate(keypoints)159 ]160 )161 return coco_results162 163 164def convert_to_xywh(boxes):165 xmin, ymin, xmax, ymax = boxes.unbind(1)166 return torch.stack((xmin, ymin, xmax - xmin, ymax - ymin), dim=1)167 168 169def merge(img_ids, eval_imgs):170 all_img_ids = all_gather(img_ids)171 all_eval_imgs = all_gather(eval_imgs)172 173 merged_img_ids = []174 for p in all_img_ids:175 merged_img_ids.extend(p)176 177 merged_eval_imgs = []178 for p in all_eval_imgs:179 merged_eval_imgs.append(p)180 181 merged_img_ids = np.array(merged_img_ids)182 merged_eval_imgs = np.concatenate(merged_eval_imgs, 2)183 184 # keep only unique (and in sorted order) images185 merged_img_ids, idx = np.unique(merged_img_ids, return_index=True)186 merged_eval_imgs = merged_eval_imgs[..., idx]187 188 return merged_img_ids, merged_eval_imgs189 190 191def create_common_coco_eval(coco_eval, img_ids, eval_imgs):192 img_ids, eval_imgs = merge(img_ids, eval_imgs)193 img_ids = list(img_ids)194 eval_imgs = list(eval_imgs.flatten())195 196 coco_eval.evalImgs = eval_imgs197 coco_eval.params.imgIds = img_ids198 coco_eval._paramsEval = copy.deepcopy(coco_eval.params)199 200 201#################################################################202# From pycocotools, just removed the prints and fixed203# a Python3 bug about unicode not defined204#################################################################205 206 207def evaluate(self):208 '''209 Run per image evaluation on given images and store results (a list of dict) in self.evalImgs210 :return: None211 '''212 # tic = time.time()213 # print('Running per image evaluation...')214 p = self.params215 # add backward compatibility if useSegm is specified in params216 if p.useSegm is not None:217 p.iouType = 'segm' if p.useSegm == 1 else 'bbox'218 print('useSegm (deprecated) is not None. Running {} evaluation'.format(p.iouType))219 # print('Evaluate annotation type *{}*'.format(p.iouType))220 p.imgIds = list(np.unique(p.imgIds))221 if p.useCats:222 p.catIds = list(np.unique(p.catIds))223 p.maxDets = sorted(p.maxDets)224 self.params = p225 226 self._prepare()227 # loop through images, area range, max detection number228 catIds = p.catIds if p.useCats else [-1]229 230 if p.iouType == 'segm' or p.iouType == 'bbox':231 computeIoU = self.computeIoU232 elif p.iouType == 'keypoints':233 computeIoU = self.computeOks234 self.ious = {235 (imgId, catId): computeIoU(imgId, catId)236 for imgId in p.imgIds237 for catId in catIds}238 239 evaluateImg = self.evaluateImg240 maxDet = p.maxDets[-1]241 evalImgs = [242 evaluateImg(imgId, catId, areaRng, maxDet)243 for catId in catIds244 for areaRng in p.areaRng245 for imgId in p.imgIds246 ]247 # this is NOT in the pycocotools code, but could be done outside248 evalImgs = np.asarray(evalImgs).reshape(len(catIds), len(p.areaRng), len(p.imgIds))249 self._paramsEval = copy.deepcopy(self.params)250 # toc = time.time()251 # print('DONE (t={:0.2f}s).'.format(toc-tic))252 return p.imgIds, evalImgs253 254#################################################################255# end of straight copy from pycocotools, just removing the prints256#################################################################257 