Download src/superpoint_pruning/evaluation/metrics.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 3.79 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/evaluation/metrics.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint/src/superpoint_pruning/evaluation/metrics.py
-
curl -L -o metrics.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/evaluation/metrics.py
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))}" | |
| ) | |