Download src/superpoint_pruning/evaluation/plot_keypoints.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 6.49 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F2/src/superpoint_pruning/evaluation/plot_keypoints.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint@refs/pr/2/src/superpoint_pruning/evaluation/plot_keypoints.py
-
curl -L -o plot_keypoints.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F2/src/superpoint_pruning/evaluation/plot_keypoints.py
6.49 kB
| 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(), | |
| ) | |
| 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()) | |