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)