Download evaluation/context_eval.py from stereoid/Orienter: direct link, hf CLI and curl.
- Browser
- Download file 7.21 kB
-
https://huggingface.co/stereoid/Orienter/resolve/main/evaluation/context_eval.py
- Command line
-
hf download hf://stereoid/Orienter/evaluation/context_eval.py
-
curl -L -o context_eval.py https://huggingface.co/stereoid/Orienter/resolve/main/evaluation/context_eval.py
7.21 kB
| 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() | |
| # dt_cats = gt_cats | |
| 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) | |
| # print(gt_cat, gt_cat_match[gt_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 # predicion matching interactable annotation | |
| fp = 0 # prediction matching non-interactable annotation | |
| tn = 0 # not matched non-interactable annotation | |
| fn = 0 # not matched interactable annotation | |
| bg = 0 # background | |
| 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) | |
| # find annotations by object name and image id | |
| 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) | |
| # find predictions by object name and image id | |
| 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) | |