Download src/superpoint_pruning/distillation/setup.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 5.61 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/distillation/setup.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint/src/superpoint_pruning/distillation/setup.py
-
curl -L -o setup.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/distillation/setup.py
5.61 kB
| 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}" | |
| ) | |
| 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()) | |