import random import os from pathlib import Path from superpoint_pruning.distillation.utils import rescale_image, load_grayscale_image from superpoint_pruning.models.superpoint import SuperPoint from superpoint_pruning.paths import DATA_ROOT, DEFAULT_DATASET_NAME import numpy as np import torch from tqdm import tqdm import urllib.request import zipfile import argparse DEFAULT_DATASET_URL = "http://rpg.ifi.uzh.ch/datasets/uzh-fpv-newer-versions/v3/indoor_forward_3_snapdragon_with_gt.zip" DEFAULT_GT_INDEX_FILE = "train_base.txt" def download_dataset_from_url(url: str, output_dir: Path) -> None: zip_path = DATA_ROOT / "download.zip" if output_dir.exists(): print( f"Dataset directory {output_dir.resolve()} already exists. Skipping download." ) return output_dir.mkdir(parents=True, exist_ok=True) print(f"Downloading dataset from {url} to {zip_path.resolve()}") urllib.request.urlretrieve(url, zip_path) print(f"Extracting files from {zip_path.resolve()} to {output_dir.resolve()}") with zipfile.ZipFile(zip_path, "r") as archive: archive.extractall(output_dir) zip_path.unlink() print(f"Files extracted to {output_dir.resolve()}") def generate_data_split( image_dir: Path, split_ratio: float = 0.8 ) -> tuple[list[str], list[str]]: random.seed(42) images = [path.name for path in image_dir.glob("*.png")] random.shuffle(images) split_index = int(len(images) * split_ratio) train_images = images[:split_index] val_images = images[split_index:] with open(DATA_ROOT / "train.txt", "w") as f: for image in train_images: f.write(image + "\n") with open(DATA_ROOT / "val.txt", "w") as f: for image in val_images: f.write(image + "\n") print( f"Train images: {len(train_images)}, Val images: {len(val_images)} - Saved to {DATA_ROOT}" ) @torch.inference_mode() def generate_ground_truth( image_dir: Path, id_file: Path, num_keypoints: int = 512, image_size: tuple[int, int] = (640, 480), num_images: int = 250, ) -> None: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") images = [path.name for path in image_dir.glob("*.png")] model = SuperPoint(num_keypoints=num_keypoints, return_dense=True) model.eval().to(device) ids = id_file.read_text().splitlines() ids = ids[:num_images] id_set = set(ids) selected_images = [image_id for image_id in images if image_id in id_set] Path(DATA_ROOT / f"gt_id_file_{Path(id_file).stem}_{num_images}.txt").write_text( "\n".join(selected_images) + "\n" ) desc = None kpts = None for index, image_id in enumerate(tqdm(selected_images)): image_path = os.path.join(image_dir, image_id) image = load_grayscale_image(image_path) img, _ = rescale_image(image, image_size) img = img[None, None].astype(np.float32) img = torch.from_numpy(img).to(device) keypoint_logits, descriptor_logits = model(img) desc_np = descriptor_logits.detach().cpu().numpy()[0] kpts_np = keypoint_logits.detach().cpu().numpy()[0] if desc is None or kpts is None: desc = np.lib.format.open_memmap( DATA_ROOT / f"desc_{id_file.stem}_{num_images}.npy", mode="w+", dtype=desc_np.dtype, shape=(len(selected_images), *desc_np.shape), ) kpts = np.lib.format.open_memmap( DATA_ROOT / f"kpts_{id_file.stem}_{num_images}.npy", mode="w+", dtype=kpts_np.dtype, shape=(len(selected_images), *kpts_np.shape), ) desc[index] = desc_np kpts[index] = kpts_np if desc is not None: desc.flush() if kpts is not None: kpts.flush() print(f"Ground truth files generated to {DATA_ROOT}") def add_parser_args(parser: argparse.ArgumentParser) -> None: parser.add_argument("--download-url", type=str, default=DEFAULT_DATASET_URL) parser.add_argument("--output-dir-name", type=str, default=DEFAULT_DATASET_NAME) parser.add_argument("--gt-index-file", type=str, default=DEFAULT_GT_INDEX_FILE) parser.add_argument("--generate-data-split", action="store_true", default=False) parser.add_argument("--generate-gt", action="store_true", default=False) 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("--num-images", type=int, default=250) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description="Setup the dataset for distillation and evaluation." ) add_parser_args(parser) return parser def main(args: argparse.Namespace) -> None: if args.download_url and args.output_dir_name: download_dataset_from_url( args.download_url, DATA_ROOT / "datasets" / args.output_dir_name ) if args.generate_data_split: generate_data_split(DATA_ROOT / "datasets" / args.output_dir_name / "img", 0.8) if args.generate_gt: generate_ground_truth( DATA_ROOT / "datasets" / args.output_dir_name / "img", DATA_ROOT / args.gt_index_file, num_images=args.num_images, image_size=(args.width, args.height), num_keypoints=args.num_keypoints, ) if __name__ == "__main__": main(build_parser().parse_args())