File size: 7,205 Bytes
6d35aff | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 | 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)
|