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