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" @torch.inference_mode() 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") @torch.inference_mode() 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())