File size: 4,691 Bytes
1da285f | 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 | import os
import json
import pandas as pd
import argparse
from tqdm import tqdm
parser = argparse.ArgumentParser(description='Evaluate the interact points')
parser.add_argument('-a', '--annotation', type=str, help='Path to the annotation file')
parser.add_argument('-i', '--interact', type=str, help='Path to the interact points file')
parser.add_argument('-o', '--output_file', type=str, help='Path to the result directory')
parser.add_argument('-p', '--prediction', type=str, help='Path to the prediction file')
def cal_metrics(anno_data, interact_data):
# Group annotations by image id
img_annos = {}
for anno in anno_data['annotations']:
img_id = anno['image_id']
if img_id not in img_annos:
img_annos[img_id] = []
img_annos[img_id].append(anno)
current_interacts = []
effective_interacts_rate = [0]
effective_interacts_cnt = [0]
coverage_rate = [0]
# Batch k from 1 to MAX_INTERACT_PER_IMAGE
# for sk in tqdm(interact_data, desc='Calculating interact metric', total=len(interact_data)):
for sk in interact_data:
k = int(sk)
# Add current batch interact points
current_interacts.extend(interact_data[sk])
effective_interacts_rate.append(0)
effective_interacts_cnt.append(0)
coverage_rate.append(0)
# Group interact points by image id
img_interact = {}
for interact in current_interacts:
img_id = interact['img_id']
if img_id not in img_interact:
img_interact[img_id] = []
img_interact[img_id].append(interact)
# Calculate effective interact rate
for img_id in img_annos:
if img_id not in img_interact:
continue
effective_interacts_rate[k] += cnt_effective_interact(img_annos[img_id], img_interact[img_id]) / len(img_interact[img_id])
effective_interacts_cnt[k] += cnt_effective_interact(img_annos[img_id], img_interact[img_id])
effective_interacts_rate[k] /= len(img_annos)
# Calculate coverage rate
for img_id in img_annos:
if img_id not in img_interact:
continue
coverage_rate[k] += cnt_covered_anno(img_annos[img_id], img_interact[img_id]) / len(img_annos[img_id])
coverage_rate[k] /= len(img_annos)
return effective_interacts_cnt, effective_interacts_rate, coverage_rate
def cnt_effective_interact(annos, interacts):
effective_interacts = 0
for interact in interacts:
for anno in annos:
if interact['X'] - anno['bbox'][0] >= 0 and \
interact['Y'] - anno['bbox'][1] >= 0 and \
interact['X'] - anno['bbox'][0] <= anno['bbox'][2] and \
interact['Y'] - anno['bbox'][1] <= anno['bbox'][3]:
effective_interacts += 1
break
return effective_interacts
def cnt_covered_anno(annos, interacts):
coverage = 0
# each interact point can only cover one interactable object
t_interacts = []
t_interacts.extend(interacts)
for anno in annos:
for interact in t_interacts:
if interact['X'] - anno['bbox'][0] >= 0 and \
interact['Y'] - anno['bbox'][1] >= 0 and \
interact['X'] - anno['bbox'][0] <= anno['bbox'][2] and \
interact['Y'] - anno['bbox'][1] <= anno['bbox'][3]:
coverage += 1
t_interacts.remove(interact)
break
return coverage
def main(args):
if os.path.exists(args.output_file):
print(f'{args.output_file} already exists')
return
with open(args.annotation, 'r') as f:
anno_data = json.load(f)
with open(args.interact, 'r') as f:
interact_data = json.load(f)
if args.prediction:
with open(args.prediction, 'r') as f:
pred_data = json.load(f)
pred_imgs = [pred['image_id'] for pred in pred_data]
anno_data['images'] = [img for img in anno_data['images'] if img['id'] in pred_imgs]
anno_data['annotations'] = [anno for anno in anno_data['annotations'] if anno['image_id'] in pred_imgs]
# Calculate effective_interacts_rate and coverage_rate
effective_interacts_cnt, effective_interacts_rate, coverage_rate = cal_metrics(anno_data, interact_data)
interact_metric = pd.DataFrame({
'effective_interacts_cnt': effective_interacts_cnt,
'effective_interacts_rate': effective_interacts_rate,
'coverage_rate': coverage_rate
})
interact_metric.to_csv(args.output_file, index_label='K')
if __name__ == '__main__':
args = parser.parse_args()
main(args)
|