Orienter / evaluation /context_eval.py
stereoid's picture
Add files using upload-large-folder tool
6d35aff verified
Raw
History Blame Contribute Delete
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)