oskarkuuse-pruna's picture
commit repo mirror
6979012
Raw History Blame Contribute Delete
3.79 kB
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))}"
)