Team Ai
Apppublic

MLVLab/Human_Object_Interaction

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
1likes
hico_eval.py242 linesDownload Raw Back to evaluators
1# ------------------------------------------------------------------------2# HOTR official code : hotr/data/evaluators/hico_eval.py3# Copyright (c) Kakao Brain, Inc. and its affiliates. All Rights Reserved4# ------------------------------------------------------------------------5# Modified from QPIC (https://github.com/hitachi-rd-cv/qpic)6# Copyright (c) Hitachi, Ltd. All Rights Reserved.7# Licensed under the Apache License, Version 2.0 [see LICENSE for details]8# ------------------------------------------------------------------------9import numpy as np10from collections import defaultdict11 12class HICOEvaluator():13    def __init__(self, preds, gts, rare_triplets, non_rare_triplets, correct_mat):14        self.overlap_iou = 0.515        self.max_hois = 10016 17        self.rare_triplets = rare_triplets18        self.non_rare_triplets = non_rare_triplets19 20        self.fp = defaultdict(list)21        self.tp = defaultdict(list)22        self.score = defaultdict(list)23        self.sum_gts = defaultdict(lambda: 0)24        self.gt_triplets = []25 26        self.preds = []27        for img_preds in preds:28            img_preds = {k: v.to('cpu').numpy() for k, v in img_preds.items() if k != 'hoi_recognition_time'}29            bboxes = [{'bbox': bbox, 'category_id': label} for bbox, label in zip(img_preds['boxes'], img_preds['labels'])]30            hoi_scores = img_preds['verb_scores']31            verb_labels = np.tile(np.arange(hoi_scores.shape[1]), (hoi_scores.shape[0], 1))32            subject_ids = np.tile(img_preds['sub_ids'], (hoi_scores.shape[1], 1)).T33            object_ids = np.tile(img_preds['obj_ids'], (hoi_scores.shape[1], 1)).T34 35            hoi_scores = hoi_scores.ravel()36            verb_labels = verb_labels.ravel()37            subject_ids = subject_ids.ravel()38            object_ids = object_ids.ravel()39 40            if len(subject_ids) > 0:41                object_labels = np.array([bboxes[object_id]['category_id'] for object_id in object_ids])42                masks = correct_mat[verb_labels, object_labels]43                hoi_scores *= masks44 45                hois = [{'subject_id': subject_id, 'object_id': object_id, 'category_id': category_id, 'score': score} for46                        subject_id, object_id, category_id, score in zip(subject_ids, object_ids, verb_labels, hoi_scores)]47                hois.sort(key=lambda k: (k.get('score', 0)), reverse=True)48                hois = hois[:self.max_hois]49            else:50                hois = []51 52            self.preds.append({53                'predictions': bboxes,54                'hoi_prediction': hois55            })56 57        self.gts = []58        for img_gts in gts:59            img_gts = {k: v.to('cpu').numpy() for k, v in img_gts.items() if k != 'id'}60            self.gts.append({61                'annotations': [{'bbox': bbox, 'category_id': label} for bbox, label in zip(img_gts['boxes'], img_gts['labels'])],62                'hoi_annotation': [{'subject_id': hoi[0], 'object_id': hoi[1], 'category_id': hoi[2]} for hoi in img_gts['hois']]63            })64            for hoi in self.gts[-1]['hoi_annotation']:65                triplet = (self.gts[-1]['annotations'][hoi['subject_id']]['category_id'],66                           self.gts[-1]['annotations'][hoi['object_id']]['category_id'],67                           hoi['category_id'])68 69                if triplet not in self.gt_triplets:70                    self.gt_triplets.append(triplet)71 72                self.sum_gts[triplet] += 173 74    def evaluate(self):75        for img_id, (img_preds, img_gts) in enumerate(zip(self.preds, self.gts)):76            print(f"Evaluating Score Matrix... : [{(img_id+1):>4}/{len(self.gts):<4}]" ,flush=True, end="\r")77            pred_bboxes = img_preds['predictions']78            gt_bboxes = img_gts['annotations']79            pred_hois = img_preds['hoi_prediction']80            gt_hois = img_gts['hoi_annotation']81            if len(gt_bboxes) != 0:82                bbox_pairs, bbox_overlaps = self.compute_iou_mat(gt_bboxes, pred_bboxes)83                self.compute_fptp(pred_hois, gt_hois, bbox_pairs, pred_bboxes, bbox_overlaps)84            else:85                for pred_hoi in pred_hois:86                    triplet = [pred_bboxes[pred_hoi['subject_id']]['category_id'],87                               pred_bboxes[pred_hoi['object_id']]['category_id'], pred_hoi['category_id']]88                    if triplet not in self.gt_triplets:89                        continue90                    self.tp[triplet].append(0)91                    self.fp[triplet].append(1)92                    self.score[triplet].append(pred_hoi['score'])93        print(f"[stats] Score Matrix Generation completed!!          ")94        map = self.compute_map()95        return map96 97    def compute_map(self):98        ap = defaultdict(lambda: 0)99        rare_ap = defaultdict(lambda: 0)100        non_rare_ap = defaultdict(lambda: 0)101        max_recall = defaultdict(lambda: 0)102        for triplet in self.gt_triplets:103            sum_gts = self.sum_gts[triplet]104            if sum_gts == 0:105                continue106 107            tp = np.array((self.tp[triplet]))108            fp = np.array((self.fp[triplet]))109            if len(tp) == 0:110                ap[triplet] = 0111                max_recall[triplet] = 0112                if triplet in self.rare_triplets:113                    rare_ap[triplet] = 0114                elif triplet in self.non_rare_triplets:115                    non_rare_ap[triplet] = 0116                else:117                    print('Warning: triplet {} is neither in rare triplets nor in non-rare triplets'.format(triplet))118                continue119 120            score = np.array(self.score[triplet])121            sort_inds = np.argsort(-score)122            fp = fp[sort_inds]123            tp = tp[sort_inds]124            fp = np.cumsum(fp)125            tp = np.cumsum(tp)126            rec = tp / sum_gts127            prec = tp / (fp + tp)128            ap[triplet] = self.voc_ap(rec, prec)129            max_recall[triplet] = np.amax(rec)130            if triplet in self.rare_triplets:131                rare_ap[triplet] = ap[triplet]132            elif triplet in self.non_rare_triplets:133                non_rare_ap[triplet] = ap[triplet]134            else:135                print('Warning: triplet {} is neither in rare triplets nor in non-rare triplets'.format(triplet))136        m_ap = np.mean(list(ap.values())) * 100 # percentage137        m_ap_rare = np.mean(list(rare_ap.values())) * 100 # percentage138        m_ap_non_rare = np.mean(list(non_rare_ap.values())) * 100 # percentage139        m_max_recall = np.mean(list(max_recall.values()))140 141        return {'mAP': m_ap, 'mAP rare': m_ap_rare, 'mAP non-rare': m_ap_non_rare, 'mean max recall': m_max_recall}142 143    def voc_ap(self, rec, prec):144        ap = 0.145        for t in np.arange(0., 1.1, 0.1):146            if np.sum(rec >= t) == 0:147                p = 0148            else:149                p = np.max(prec[rec >= t])150            ap = ap + p / 11.151        return ap152 153    def compute_fptp(self, pred_hois, gt_hois, match_pairs, pred_bboxes, bbox_overlaps):154        pos_pred_ids = match_pairs.keys()155        vis_tag = np.zeros(len(gt_hois))156        pred_hois.sort(key=lambda k: (k.get('score', 0)), reverse=True)157        if len(pred_hois) != 0:158            for pred_hoi in pred_hois:159                is_match = 0160                if len(match_pairs) != 0 and pred_hoi['subject_id'] in pos_pred_ids and pred_hoi['object_id'] in pos_pred_ids:161                    pred_sub_ids = match_pairs[pred_hoi['subject_id']]162                    pred_obj_ids = match_pairs[pred_hoi['object_id']]163                    pred_sub_overlaps = bbox_overlaps[pred_hoi['subject_id']]164                    pred_obj_overlaps = bbox_overlaps[pred_hoi['object_id']]165                    pred_category_id = pred_hoi['category_id']166                    max_overlap = 0167                    max_gt_hoi = 0168                    for gt_hoi in gt_hois:169                        if gt_hoi['subject_id'] in pred_sub_ids and gt_hoi['object_id'] in pred_obj_ids \170                           and pred_category_id == gt_hoi['category_id']:171                            is_match = 1172                            min_overlap_gt = min(pred_sub_overlaps[pred_sub_ids.index(gt_hoi['subject_id'])],173                                                 pred_obj_overlaps[pred_obj_ids.index(gt_hoi['object_id'])])174                            if min_overlap_gt > max_overlap:175                                max_overlap = min_overlap_gt176                                max_gt_hoi = gt_hoi177                triplet = (pred_bboxes[pred_hoi['subject_id']]['category_id'], pred_bboxes[pred_hoi['object_id']]['category_id'],178                           pred_hoi['category_id'])179                if triplet not in self.gt_triplets:180                    continue181                if is_match == 1 and vis_tag[gt_hois.index(max_gt_hoi)] == 0:182                    self.fp[triplet].append(0)183                    self.tp[triplet].append(1)184                    vis_tag[gt_hois.index(max_gt_hoi)] =1185                else:186                    self.fp[triplet].append(1)187                    self.tp[triplet].append(0)188                self.score[triplet].append(pred_hoi['score'])189 190    def compute_iou_mat(self, bbox_list1, bbox_list2):191        iou_mat = np.zeros((len(bbox_list1), len(bbox_list2)))192        if len(bbox_list1) == 0 or len(bbox_list2) == 0:193            return {}194        for i, bbox1 in enumerate(bbox_list1):195            for j, bbox2 in enumerate(bbox_list2):196                iou_i = self.compute_IOU(bbox1, bbox2)197                iou_mat[i, j] = iou_i198 199        iou_mat_ov=iou_mat.copy()200        iou_mat[iou_mat>=self.overlap_iou] = 1201        iou_mat[iou_mat<self.overlap_iou] = 0202 203        match_pairs = np.nonzero(iou_mat)204        match_pairs_dict = {}205        match_pair_overlaps = {}206        if iou_mat.max() > 0:207            for i, pred_id in enumerate(match_pairs[1]):208                if pred_id not in match_pairs_dict.keys():209                    match_pairs_dict[pred_id] = []210                    match_pair_overlaps[pred_id]=[]211                match_pairs_dict[pred_id].append(match_pairs[0][i])212                match_pair_overlaps[pred_id].append(iou_mat_ov[match_pairs[0][i],pred_id])213        return match_pairs_dict, match_pair_overlaps214 215    def compute_IOU(self, bbox1, bbox2):216        if isinstance(bbox1['category_id'], str):217            bbox1['category_id'] = int(bbox1['category_id'].replace('\n', ''))218        if isinstance(bbox2['category_id'], str):219            bbox2['category_id'] = int(bbox2['category_id'].replace('\n', ''))220        if bbox1['category_id'] == bbox2['category_id']:221            rec1 = bbox1['bbox']222            rec2 = bbox2['bbox']223            # computing area of each rectangles224            S_rec1 = (rec1[2] - rec1[0]+1) * (rec1[3] - rec1[1]+1)225            S_rec2 = (rec2[2] - rec2[0]+1) * (rec2[3] - rec2[1]+1)226 227            # computing the sum_area228            sum_area = S_rec1 + S_rec2229 230            # find the each edge of intersect rectangle231            left_line = max(rec1[1], rec2[1])232            right_line = min(rec1[3], rec2[3])233            top_line = max(rec1[0], rec2[0])234            bottom_line = min(rec1[2], rec2[2])235            # judge if there is an intersect236            if left_line >= right_line or top_line >= bottom_line:237                return 0238            else:239                intersect = (right_line - left_line+1) * (bottom_line - top_line+1)240                return intersect / (sum_area - intersect)241        else:242            return 0