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))}" )