| import os
|
| import sys
|
| import json
|
| import argparse
|
| import pandas as pd
|
| from tqdm import tqdm
|
| from pycocotools_ovod.semantic_matching import is_semantic_match, gt_cat_match_path
|
|
|
| gt_dataset = None
|
| preds = None
|
| cat_id_to_name = None
|
| WHOLE_DATASET_PATH = './gts/det/semantics/union3_test.json'
|
|
|
|
|
| def get_obj_name(cat):
|
| return cat.split('-')[0]
|
|
|
|
|
| def is_category_interactable(cat):
|
| if isinstance(cat, str):
|
| return not cat.endswith('-n')
|
| elif isinstance(cat, dict):
|
| return not cat['name'].endswith('-n')
|
| else:
|
| raise ValueError("Invalid input type")
|
|
|
|
|
| def iou(bbox1, bbox2):
|
| x1, y1, w1, h1 = bbox1
|
| x2, y2, w2, h2 = bbox2
|
| union = w1 * h1 + w2 * h2
|
| inter = max(0, min(x1 + w1, x2 + w2) - max(x1, x2)) * \
|
| max(0, min(y1 + h1, y2 + h2) - max(y1, y2))
|
| return inter / (union - inter)
|
|
|
|
|
| def best_match_gt(pred, anns):
|
| if len(anns) == 0:
|
| return 0, None
|
| best_iou = -1
|
| best_match = None
|
| for ann in anns:
|
| iou_score = iou(pred['bbox'], ann['bbox'])
|
| if iou_score > best_iou:
|
| best_iou = iou_score
|
| best_match = ann
|
| return best_iou, best_match
|
|
|
|
|
| def match_cats(gt_cats, preds, eval_dimension):
|
| if os.path.exists(gt_cat_match_path):
|
| os.remove(gt_cat_match_path)
|
| dt_cats = set()
|
| for pred in preds:
|
| dt_cats.add(pred['category_id'])
|
| dt_cats = list(dt_cats)
|
| dt_cats.sort()
|
| gt_cats.sort()
|
|
|
| gt_cat_match = {gt_cat: [] for gt_cat in gt_cats}
|
| print('matching dt cats to gt cats...', file=sys.stderr)
|
| for gt_cat in tqdm(gt_cats):
|
| for dt_cat in dt_cats:
|
| if is_semantic_match(gt_cat, dt_cat, eval_dimension=eval_dimension):
|
| gt_cat_match[gt_cat].append(dt_cat)
|
|
|
| with open(gt_cat_match_path, 'w') as f:
|
| json.dump(gt_cat_match, f)
|
|
|
|
|
| def eval_category(imgs_anns, imgs_preds, iou_threshold):
|
| global cat_id_to_name
|
| tp = 0
|
| fp = 0
|
| tn = 0
|
| fn = 0
|
| bg = 0
|
| for img_id in imgs_anns:
|
| anns_match_flags = {ann['id']: False for ann in imgs_anns[img_id]}
|
| if img_id in imgs_preds:
|
| for pred in imgs_preds[img_id]:
|
| iou_score, best_match_ann = best_match_gt(pred, imgs_anns[img_id])
|
| if iou_score < iou_threshold:
|
| bg += 1
|
| continue
|
| if is_category_interactable(cat_id_to_name[best_match_ann['category_id']]):
|
| tp += 1
|
| else:
|
| fp += 1
|
| anns_match_flags[best_match_ann['id']] = True
|
| fn += sum([not anns_match_flags[ann['id']] for ann in imgs_anns[img_id]
|
| if is_category_interactable(cat_id_to_name[ann['category_id']])])
|
| tn += sum([not anns_match_flags[ann['id']] for ann in imgs_anns[img_id]
|
| if not is_category_interactable(cat_id_to_name[ann['category_id']])])
|
| return tp, fp, tn, fn, bg
|
|
|
|
|
| def main(args):
|
| global gt_dataset, preds, cat_id_to_name
|
|
|
| with open(args.gt, 'r') as f:
|
| gt_dataset = json.load(f)
|
| with open(args.pred, 'r') as f:
|
| preds = json.load(f)
|
|
|
| if args.num_cat:
|
| with open(WHOLE_DATASET_PATH, 'r') as f:
|
| whole_dataset = json.load(f)
|
| cat_id_to_name = {cat['id']: cat['name'] for cat in whole_dataset['categories']}
|
| for pred in preds:
|
| pred['category_id'] = cat_id_to_name[pred['category_id']]
|
|
|
|
|
| cat_id_to_name = {cat['id']: cat['name']
|
| for cat in gt_dataset['categories']}
|
|
|
| obj_names = [cat['name'] for cat in gt_dataset['categories']
|
| if is_category_interactable(cat)]
|
|
|
| match_cats(obj_names, preds, args.dimension)
|
|
|
|
|
| objs_imgs_anns = {obj_name: {} for obj_name in obj_names}
|
| for ann in gt_dataset['annotations']:
|
| obj_name = get_obj_name(cat_id_to_name[ann['category_id']])
|
| if ann['image_id'] not in objs_imgs_anns[obj_name]:
|
| objs_imgs_anns[obj_name][ann['image_id']] = []
|
| objs_imgs_anns[obj_name][ann['image_id']].append(ann)
|
|
|
| for obj_name in obj_names:
|
| if objs_imgs_anns[obj_name] == {}:
|
| print(f'No annotation for {obj_name}')
|
| obj_names.remove(obj_name)
|
|
|
|
|
| objs_imgs_preds = {obj_name: {} for obj_name in obj_names}
|
| for obj_name in tqdm(obj_names, total=len(obj_names)):
|
| for pred in preds:
|
| if is_semantic_match(obj_name, pred['category_id'], eval_dimension=args.dimension):
|
| if pred['image_id'] not in objs_imgs_preds[obj_name]:
|
| objs_imgs_preds[obj_name][pred['image_id']] = []
|
| objs_imgs_preds[obj_name][pred['image_id']].append(pred)
|
|
|
| result_list = []
|
| P_avg, R_avg, f1_avg, bgr_avg = 0, 0, 0, 0
|
| tp_avg, fp_avg, tn_avg, fn_avg, bg_avg = 0, 0, 0, 0, 0
|
|
|
| for obj_name in obj_names:
|
| tp, fp, tn, fn, bg = eval_category(objs_imgs_anns[obj_name], objs_imgs_preds[obj_name], args.iou)
|
| tp_avg += tp
|
| fp_avg += fp
|
| tn_avg += tn
|
| fn_avg += fn
|
| bg_avg += bg
|
| precision = tp / (tp + fp) if tp + fp > 0 else 0
|
| P_avg += precision
|
| recall = tp / (tp + fn) if tp + fn > 0 else 0
|
| R_avg += recall
|
| f1 = 2 * precision * recall / (precision + recall) if precision + recall > 0 else 0
|
| f1_avg += f1
|
| bg_rate = bg / (tp + fp + bg) if tp + fp + bg > 0 else 0
|
| bgr_avg += bg_rate
|
| result_list.append([obj_name, precision, recall, f1, bg_rate, tp, fp, tn, fn, bg])
|
|
|
| P_avg /= len(obj_names)
|
| R_avg /= len(obj_names)
|
| f1_avg /= len(obj_names)
|
| bgr_avg /= len(obj_names)
|
| tp_avg /= len(obj_names)
|
| fp_avg /= len(obj_names)
|
| tn_avg /= len(obj_names)
|
| fn_avg /= len(obj_names)
|
| bg_avg /= len(obj_names)
|
|
|
| result_list.append(['average', P_avg, R_avg, f1_avg, bgr_avg, tp_avg, fp_avg, tn_avg, fn_avg, bg_avg])
|
|
|
| df = pd.DataFrame(result_list, columns=['object', 'precision', 'recall', 'f1', 'bg_rate', 'tp', 'fp', 'tn', 'fn', 'bg'])
|
| df.to_csv(args.output, index=False)
|
|
|
| if os.path.exists(gt_cat_match_path):
|
| os.remove(gt_cat_match_path)
|
|
|
|
|
| if __name__ == '__main__':
|
| parser = argparse.ArgumentParser()
|
| parser.add_argument('-g', '--gt', type=str, required=True)
|
| parser.add_argument('-p', '--pred', type=str, required=True)
|
| parser.add_argument('-o', '--output', type=str)
|
| parser.add_argument('-d', '--dimension', type=str, default='s')
|
| parser.add_argument('-n', '--num_cat', action='store_true')
|
| parser.add_argument('-i', '--iou', type=float)
|
| args = parser.parse_args()
|
| main(args)
|
|
|