File size: 3,785 Bytes
6979012
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import numpy as np
import torch


def match_overlap(pred, gt):
    pred = pred.detach().cpu() if torch.is_tensor(pred) else torch.as_tensor(pred)
    gt = gt.detach().cpu() if torch.is_tensor(gt) else torch.as_tensor(gt)

    if len(pred.shape) > 1 and pred.shape[1] == 3:
        pred = pred[:, 1:]
    if len(gt.shape) > 1 and gt.shape[1] == 3:
        gt = gt[:, 1:]

    if len(pred) > 0:
        invalid = torch.nonzero(pred[:, 0] == -1, as_tuple=False)
        if len(invalid) > 0:
            pred = pred[: invalid[0].item()]

    num_preds = len(pred)
    if num_preds == 0 or len(gt) == 0:
        running_recall = 0
    else:
        overlap = (pred[:, None, :] == gt[None, :, :]).all(dim=2)
        running_recall = overlap.any(dim=1).sum().item()

    recall = running_recall / len(gt) if len(gt) > 0 else 0
    precision = running_recall / num_preds if num_preds > 0 else 0

    # recall: # of original matches found
    # precision: in this case, the # of correct predictions
    return recall, precision, num_preds


def keypoint_overlap(pred, gt, cell_size=8, max_x=640, max_y=480):
    if len(pred.shape) == 3:
        pred = pred[0]
    if len(gt.shape) == 3:
        gt = gt[0]

    num_cells_x = (max_x + cell_size - 1) // cell_size
    num_cells_y = (max_y + cell_size - 1) // cell_size
    num_cells = num_cells_x * num_cells_y

    def counts_per_cell(kpts):
        valid = (
            (kpts[:, 0] >= 0)
            & (kpts[:, 0] < max_x)
            & (kpts[:, 1] >= 0)
            & (kpts[:, 1] < max_y)
        )
        kpts = kpts[valid]
        cell_x = torch.div(kpts[:, 0], cell_size, rounding_mode="floor").long()
        cell_y = torch.div(kpts[:, 1], cell_size, rounding_mode="floor").long()
        linear_idx = cell_y * num_cells_x + cell_x
        return torch.bincount(linear_idx, minlength=num_cells)

    pred_counts = counts_per_cell(pred)
    gt_counts = counts_per_cell(gt)
    covered = torch.minimum(pred_counts, gt_counts).sum().item()
    return covered, pred_counts.sum().item(), gt_counts.sum().item()


class BenchmarkMetrics:
    def __init__(self):
        self.recalls = []
        self.precisions = []
        self.gt_counts = []
        self.pred_counts = []
        self.keypoints_covered = []
        self.pred_counts_kpts = []
        self.gt_counts_kpts = []

    def update_matches(self, pred, gt):
        recall, precision, num_preds = match_overlap(pred, gt)
        self.recalls.append(recall)
        self.precisions.append(precision)
        self.gt_counts.append(len(gt))
        self.pred_counts.append(num_preds)

    def update_keypoints(self, pred, gt):
        covered, pred_counts, gt_counts = keypoint_overlap(pred, gt)
        self.keypoints_covered.append(covered)
        self.pred_counts_kpts.append(pred_counts)
        self.gt_counts_kpts.append(gt_counts)

    def print_metrics(self):
        print(f"Recall: {np.mean(self.recalls)}")
        print(f"Precision: {np.mean(self.precisions)}")
        print(
            f"Average number of matches (original keypoints + pruned descriptors): {np.mean(np.array(self.pred_counts))}"
        )
        print(
            f"Average number of matches (original keypoints + original descriptors): {np.mean(np.array(self.gt_counts))}"
        )
        print(
            f"Average change in number of matches: {np.mean(np.array(self.pred_counts) - np.array(self.gt_counts))}"
        )
        print(
            f"Average number of keypoints covered: {np.mean(np.array(self.keypoints_covered))}"
        )
        print(
            f"Average number of keypoints in prediction: {np.mean(np.array(self.pred_counts_kpts))}"
        )
        print(
            f"Average number of keypoints in ground truth: {np.mean(np.array(self.gt_counts_kpts))}"
        )