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