import argparse from pathlib import Path import matplotlib.pyplot as plt import numpy as np import torch from superpoint_pruning.distillation.utils import load_grayscale_image, rescale_image from superpoint_pruning.models.superpoint import SuperPoint from superpoint_pruning.paths import DEFAULT_IMAGE_DIR from superpoint_pruning.pruning import add_pruning_parser_args, pruning_config_from_args SHARED_COLOR = "lime" ORIGINAL_ONLY_COLOR = "deepskyblue" PRUNED_ONLY_COLOR = "orangered" def resolve_image_path(image_dir: Path, image_name: Path) -> Path: if image_name.is_absolute(): return image_name return image_dir / image_name def classify_keypoints( original: torch.Tensor, pruned: torch.Tensor, image_width: int, ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: original = original.detach().cpu().round().to(torch.int64) pruned = pruned.detach().cpu().round().to(torch.int64) original_ids = original[:, 1] * image_width + original[:, 0] pruned_ids = pruned[:, 1] * image_width + pruned[:, 0] shared_mask = torch.isin(original_ids, pruned_ids) original_only_mask = ~shared_mask pruned_only_mask = ~torch.isin(pruned_ids, original_ids) return ( original[shared_mask].numpy(), original[original_only_mask].numpy(), pruned[pruned_only_mask].numpy(), ) @torch.inference_mode() def extract_keypoints( model: SuperPoint, image: torch.Tensor, scale: torch.Tensor, ) -> torch.Tensor: keypoints, _, _ = model(image) keypoints = (keypoints.to(torch.float32) + 0.5) / scale - 0.5 return keypoints[0].detach().cpu() def _scatter(ax, points: np.ndarray, color: str, label: str, size: int) -> None: if len(points) == 0: return ax.scatter(points[:, 0], points[:, 1], s=size, marker=".", color=color, label=label) def plot_side_by_side( image: np.ndarray, original: np.ndarray, pruned: np.ndarray, output: Path, keypoint_size: int, ) -> None: fig, axes = plt.subplots(1, 2, figsize=(12, 5)) for ax, points, title in ( (axes[0], original, f"Original ({len(original)})"), (axes[1], pruned, f"Pruned ({len(pruned)})"), ): ax.imshow(image, cmap="gray") _scatter(ax, points, SHARED_COLOR, title, keypoint_size) ax.set_title(title) ax.axis("off") fig.tight_layout() output.parent.mkdir(parents=True, exist_ok=True) fig.savefig(output, dpi=150, bbox_inches="tight") plt.close(fig) def plot_overlay( image: np.ndarray, original: torch.Tensor, pruned: torch.Tensor, output: Path, keypoint_size: int, ) -> None: height, width = image.shape[:2] shared, original_only, pruned_only = classify_keypoints(original, pruned, width) fig, ax = plt.subplots(figsize=(8, 6)) ax.imshow(image, cmap="gray") _scatter(ax, shared, SHARED_COLOR, f"Shared ({len(shared)})", keypoint_size) _scatter( ax, original_only, ORIGINAL_ONLY_COLOR, f"Original only ({len(original_only)})", keypoint_size, ) _scatter( ax, pruned_only, PRUNED_ONLY_COLOR, f"Pruned only ({len(pruned_only)})", keypoint_size, ) ax.set_title("Shared vs unique keypoints") ax.axis("off") ax.legend(loc="upper right", framealpha=0.8) fig.tight_layout() output.parent.mkdir(parents=True, exist_ok=True) fig.savefig(output, dpi=150, bbox_inches="tight") plt.close(fig) def add_parser_args(parser: argparse.ArgumentParser) -> None: parser.add_argument("--image-dir", type=Path, default=DEFAULT_IMAGE_DIR) parser.add_argument( "--image-name", type=Path, required=True, help="Image filename or absolute path.", ) parser.add_argument("--output", type=Path, default=None, help="Output figure path.") parser.add_argument( "--overlay", action=argparse.BooleanOptionalAction, default=False, help="Plot both models on one image: shared, original-only, and pruned-only keypoints.", ) parser.add_argument("--num-keypoints", type=int, default=512) parser.add_argument("--width", type=int, default=640) parser.add_argument("--height", type=int, default=480) parser.add_argument( "--hierarchical", action=argparse.BooleanOptionalAction, default=False ) parser.add_argument( "--skip-refinement", action=argparse.BooleanOptionalAction, default=False ) parser.add_argument("--keypoint-size", type=int, default=5) add_pruning_parser_args(parser) def main(args: argparse.Namespace) -> None: image_path = resolve_image_path(args.image_dir, args.image_name) output = args.output if output is None: suffix = "overlay" if args.overlay else "keypoints" output = Path(f"{image_path.stem}_{suffix}.png") device = "cuda" if torch.cuda.is_available() else "cpu" pruning_config, pruning_checkpoint = pruning_config_from_args(args) original_image = load_grayscale_image(str(image_path)) image, scale = rescale_image(original_image, new_size=(args.width, args.height)) scale = torch.as_tensor(scale, device=device, dtype=torch.float32) batch = torch.from_numpy(image[None, None].astype(np.float32)).to(device) sp_original = SuperPoint(num_keypoints=args.num_keypoints).eval().to(device) sp_pruned = SuperPoint( num_keypoints=args.num_keypoints, hierarchical_topk=args.hierarchical, skip_refinement=args.skip_refinement, ) 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() original_kpts = extract_keypoints(sp_original, batch, scale) pruned_kpts = extract_keypoints(sp_pruned, batch, scale) if args.overlay: plot_overlay( original_image, original_kpts, pruned_kpts, output, args.keypoint_size ) else: plot_side_by_side( original_image, original_kpts.numpy(), pruned_kpts.numpy(), output, args.keypoint_size, ) print(f"Saved keypoint plot to {output.resolve()}") if __name__ == "__main__": parser = argparse.ArgumentParser( description="Plot original and pruned SuperPoint keypoints." ) add_parser_args(parser) main(parser.parse_args())