Download src/superpoint_pruning/evaluation/eval.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F2/src/superpoint_pruning/evaluation/eval.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint@refs/pr/2/src/superpoint_pruning/evaluation/eval.py
-
curl -L -o eval.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F2/src/superpoint_pruning/evaluation/eval.py
10.3 kB
| import argparse | |
| from pathlib import Path | |
| import torch | |
| import numpy as np | |
| from tqdm import tqdm | |
| from superpoint_pruning.models.superpoint import SuperPoint | |
| from superpoint_pruning.distillation.utils import ( | |
| load_grayscale_image, | |
| rescale_image, | |
| plot_images_with_keypoints, | |
| ) | |
| from superpoint_pruning.evaluation.metrics import BenchmarkMetrics | |
| from superpoint_pruning.paths import DEFAULT_IMAGE_DIR, DEFAULT_TRAINING_IDS | |
| from superpoint_pruning.pruning import add_pruning_parser_args, pruning_config_from_args | |
| def _image_path(image_dir: Path | str, sample: int) -> Path: | |
| return Path(image_dir) / f"image_0_{sample}.png" | |
| def evaluate( | |
| sp_original, | |
| sp_pruned, | |
| sp_head, | |
| matcher, | |
| metrics, | |
| image_dir, | |
| start_idx=1000, | |
| end_idx=2000, | |
| device="cpu", | |
| save_preds=False, | |
| trained=False, | |
| skip=False, | |
| image_size=(1024, 768), | |
| training_ids=None, | |
| hierarchical=False, | |
| skip_refinement=False, | |
| ): | |
| try: | |
| from lightglue import LightGlue | |
| except ImportError as exc: | |
| raise ImportError( | |
| "The evaluation requires LightGlue for matching metrics" | |
| ) from exc | |
| if training_ids is not None: | |
| with open(training_ids, "r") as f: | |
| image_ids = f.read().splitlines() | |
| image_ids = [int(id.split("_")[-1].split(".")[0]) for id in image_ids] | |
| skipped = 0 | |
| sp_head_hier = sp_pruned.dense_head | |
| image_dir = Path(image_dir) | |
| for sample in tqdm(range(start_idx, end_idx)): | |
| if skip and sample in image_ids: | |
| skipped += 1 | |
| continue | |
| img = load_grayscale_image(str(_image_path(image_dir, sample))) | |
| img, scale = rescale_image(img, new_size=image_size) | |
| scale = torch.as_tensor(scale, device=device, dtype=torch.float32) | |
| img = img[None, None].astype(np.float32) | |
| img = torch.from_numpy(img).to(device) | |
| kpts_dense, desc_dense = sp_pruned(img) | |
| original_kpts_dense, original_desc_dense = sp_original(img) | |
| kpts, _, desc = sp_head_hier(original_kpts_dense, desc_dense) | |
| kpts = (kpts.to(torch.float32) + 0.5) / scale - 0.5 | |
| img2 = load_grayscale_image(str(_image_path(image_dir, sample + 10))) | |
| img2, scale = rescale_image(img2, new_size=image_size) | |
| scale = torch.as_tensor(scale, device=device, dtype=torch.float32) | |
| img2 = img2[None, None].astype(np.float32) | |
| img2 = torch.from_numpy(img2).to(device) | |
| kpts_dense2, desc_dense2 = sp_pruned(img2) | |
| original_kpts_dense2, original_desc_dense2 = sp_original(img2) | |
| kpts2, _, desc2 = sp_head_hier(original_kpts_dense2, desc_dense2) | |
| kpts2 = (kpts2.to(torch.float32) + 0.5) / scale - 0.5 | |
| original_kpts, _, original_desc = sp_head( | |
| original_kpts_dense, original_desc_dense | |
| ) | |
| original_kpts = (original_kpts.to(torch.float32) + 0.5) / scale - 0.5 | |
| original_kpts2, _, original_desc2 = sp_head( | |
| original_kpts_dense2, original_desc_dense2 | |
| ) | |
| original_kpts2 = (original_kpts2.to(torch.float32) + 0.5) / scale - 0.5 | |
| pruned_kpts, _, pruned_desc = sp_head_hier(kpts_dense, desc_dense) | |
| pruned_kpts = (pruned_kpts.to(torch.float32) + 0.5) / scale - 0.5 | |
| pruned_kpts2, _, pruned_desc2 = sp_head_hier(kpts_dense2, desc_dense2) | |
| pruned_kpts2 = (pruned_kpts2.to(torch.float32) + 0.5) / scale - 0.5 | |
| metrics.update_keypoints( | |
| pruned_kpts.detach().cpu(), original_kpts.detach().cpu() | |
| ) | |
| original_matches = matcher( | |
| { | |
| "image0": { | |
| "keypoints": original_kpts, | |
| "descriptors": original_desc, | |
| "image_size": torch.tensor([[640.0, 480.0]], device=device), | |
| }, | |
| "image1": { | |
| "keypoints": original_kpts2, | |
| "descriptors": original_desc2, | |
| "image_size": torch.tensor([[640.0, 480.0]], device=device), | |
| }, | |
| } | |
| )["matches"][0] | |
| matches = matcher( | |
| { | |
| "image0": { | |
| "keypoints": kpts, | |
| "descriptors": desc, | |
| "image_size": torch.tensor([[640.0, 480.0]], device=device), | |
| }, | |
| "image1": { | |
| "keypoints": kpts2, | |
| "descriptors": desc2, | |
| "image_size": torch.tensor([[640.0, 480.0]], device=device), | |
| }, | |
| } | |
| )["matches"][0] | |
| coord_matches = torch.cat( | |
| (kpts[0][matches[:, 0]], kpts2[0][matches[:, 1]]), dim=1 | |
| ) | |
| coord_original_matches = torch.cat( | |
| ( | |
| original_kpts[0][original_matches[:, 0]], | |
| original_kpts2[0][original_matches[:, 1]], | |
| ), | |
| dim=1, | |
| ) | |
| matches_pruned = matcher( | |
| { | |
| "image0": { | |
| "keypoints": pruned_kpts, | |
| "descriptors": pruned_desc, | |
| "image_size": torch.tensor([[640.0, 480.0]], device=device), | |
| }, | |
| "image1": { | |
| "keypoints": pruned_kpts2, | |
| "descriptors": pruned_desc2, | |
| "image_size": torch.tensor([[640.0, 480.0]], device=device), | |
| }, | |
| } | |
| )["matches"][0] | |
| metrics.update_matches( | |
| coord_matches.detach().cpu(), coord_original_matches.detach().cpu() | |
| ) | |
| if hasattr(metrics, "pruned_matches"): | |
| metrics.pruned_matches.append(len(matches_pruned.detach().cpu())) | |
| else: | |
| metrics.pruned_matches = [len(matches_pruned.detach().cpu())] | |
| metrics.print_metrics() | |
| print( | |
| f"Average number of pruned matches (pruned keypoints + pruned descriptors): {np.mean(np.array(metrics.pruned_matches))}" | |
| ) | |
| print("\n") | |
| print(f"Skipped {skipped} images") | |
| def evaluate_simple( | |
| sp_original, | |
| sp_pruned, | |
| sp_head, | |
| metrics, | |
| image_dir, | |
| start_idx=1000, | |
| end_idx=2000, | |
| device="cpu", | |
| save_preds=False, | |
| trained=False, | |
| skip=False, | |
| image_size=(1024, 768), | |
| training_ids=None, | |
| hierarchical=False, | |
| skip_refinement=False, | |
| ): | |
| if training_ids is not None: | |
| with open(training_ids, "r") as f: | |
| image_ids = f.read().splitlines() | |
| image_ids = [int(id.split("_")[-1].split(".")[0]) for id in image_ids] | |
| skipped = 0 | |
| desc_l2s = [] | |
| sp_head_hier = sp_pruned.dense_head | |
| image_dir = Path(image_dir) | |
| for sample in tqdm(range(start_idx, end_idx)): | |
| if skip and sample in image_ids: | |
| skipped += 1 | |
| continue | |
| img = load_grayscale_image(str(_image_path(image_dir, sample))) | |
| img, scale = rescale_image(img, new_size=image_size) | |
| scale = torch.as_tensor(scale, device=device, dtype=torch.float32) | |
| img = img[None, None].astype(np.float32) | |
| img = torch.from_numpy(img).to(device) | |
| kpts_dense, desc_dense = sp_pruned(img) | |
| original_kpts_dense, original_desc_dense = sp_original(img) | |
| kpts, _, desc = sp_head_hier(original_kpts_dense, desc_dense) | |
| kpts = (kpts.to(torch.float32) + 0.5) / scale - 0.5 | |
| original_kpts, _, original_desc = sp_head( | |
| original_kpts_dense, original_desc_dense | |
| ) | |
| original_kpts = (original_kpts.to(torch.float32) + 0.5) / scale - 0.5 | |
| pruned_kpts, _, pruned_desc = sp_head_hier(kpts_dense, desc_dense) | |
| pruned_kpts = (pruned_kpts.to(torch.float32) + 0.5) / scale - 0.5 | |
| desc_l2 = torch.norm(desc - original_desc, p=2, dim=-1).mean() | |
| desc_l2s.append(desc_l2.item()) | |
| metrics.update_keypoints( | |
| pruned_kpts.detach().cpu(), original_kpts.detach().cpu() | |
| ) | |
| metrics.print_metrics() | |
| print(f"Average descriptor L2: {np.mean(desc_l2s):.6f}") | |
| print(f"Skipped {skipped} images") | |
| def add_parser_args(parser: argparse.ArgumentParser) -> None: | |
| parser.add_argument("--image-dir", type=Path, default=DEFAULT_IMAGE_DIR) | |
| parser.add_argument("--training-ids", type=Path, default=DEFAULT_TRAINING_IDS) | |
| parser.add_argument("--start-idx", type=int, default=1000) | |
| parser.add_argument("--end-idx", type=int, default=2000) | |
| parser.add_argument("--num-keypoints", type=int, default=1024) | |
| parser.add_argument("--width", type=int, default=640) | |
| parser.add_argument("--height", type=int, default=480) | |
| parser.add_argument("--skip", action=argparse.BooleanOptionalAction, default=True) | |
| parser.add_argument( | |
| "--hierarchical", action=argparse.BooleanOptionalAction, default=False | |
| ) | |
| parser.add_argument( | |
| "--skip-refinement", action=argparse.BooleanOptionalAction, default=False | |
| ) | |
| add_pruning_parser_args(parser) | |
| def main(args: argparse.Namespace) -> None: | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| pruning_config, pruning_checkpoint = pruning_config_from_args(args) | |
| sp_original = ( | |
| SuperPoint(num_keypoints=args.num_keypoints, return_dense=True) | |
| .eval() | |
| .to(device) | |
| ) | |
| sp_pruned = SuperPoint( | |
| num_keypoints=args.num_keypoints, | |
| return_dense=True, | |
| hierarchical_topk=args.hierarchical, | |
| skip_refinement=args.skip_refinement, | |
| ).eval() | |
| if pruning_config: | |
| sp_pruned.prune_backbone(pruning_config) | |
| if pruning_checkpoint is not None: | |
| sp_pruned.load_pruned_weights(str(pruning_checkpoint)) | |
| sp_pruned = sp_pruned.to(device).eval() | |
| evaluate_simple( | |
| sp_original, | |
| sp_pruned, | |
| sp_original.dense_head, | |
| BenchmarkMetrics(), | |
| args.image_dir, | |
| start_idx=args.start_idx, | |
| end_idx=args.end_idx, | |
| device=device, | |
| skip=args.skip, | |
| training_ids=args.training_ids, | |
| image_size=(args.width, args.height), | |
| hierarchical=args.hierarchical, | |
| skip_refinement=args.skip_refinement, | |
| ) | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser(description="Evaluate a SuperPoint model.") | |
| add_parser_args(parser) | |
| main(parser.parse_args()) | |