oskarkuuse-pruna's picture
commit repo mirror
6979012
Raw History Blame
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(),
)
@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())