Download src/superpoint_pruning/distillation/lightning_trainer.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 9.76 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/distillation/lightning_trainer.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint/src/superpoint_pruning/distillation/lightning_trainer.py
-
curl -L -o lightning_trainer.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/distillation/lightning_trainer.py
9.76 kB
| import argparse | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import torch | |
| from torch.utils.data import DataLoader, Dataset | |
| import yaml | |
| import lightning as pl | |
| import os | |
| from superpoint_pruning.models.superpoint import SuperPoint | |
| from superpoint_pruning.distillation.utils import rescale_image, load_grayscale_image | |
| from superpoint_pruning.distillation.losses import ( | |
| detector_loss_simple, | |
| descriptor_loss_simple, | |
| detector_kd_kl, | |
| ) | |
| from superpoint_pruning.paths import DEFAULT_CONFIG_PATH | |
| DATA_PATH_KEYS = ( | |
| "train_image_dir", | |
| "train_image_ids_file", | |
| "val_image_dir", | |
| "val_image_ids_file", | |
| "ground_truth_dir", | |
| ) | |
| def resolve_path(base: Path, value: str | Path) -> Path: | |
| path = Path(value) | |
| if not path.is_absolute(): | |
| path = base / path | |
| return path.resolve() | |
| def load_config(path: str | Path, data_root: Path | None = None) -> dict[str, Any]: | |
| config_path = Path(path).resolve() | |
| with open(config_path, encoding="utf-8") as f: | |
| cfg = yaml.safe_load(f) | |
| base = data_root.resolve() if data_root is not None else config_path.parent | |
| for key in DATA_PATH_KEYS: | |
| cfg["data"][key] = str(resolve_path(base, cfg["data"][key])) | |
| cfg["trainer"]["default_root_dir"] = str( | |
| resolve_path(config_path.parent, cfg["trainer"]["default_root_dir"]) | |
| ) | |
| return cfg | |
| class ImageFolderDataset(Dataset): | |
| """Simple grayscale dataset from an image directory.""" | |
| def __init__( | |
| self, | |
| image_dir: str, | |
| image_ids_file: str, | |
| ground_truth_dir: str, | |
| keypoints_file: str, | |
| descriptors_file: str, | |
| array_ids_file: str, | |
| image_size: tuple[int, int], | |
| ) -> None: | |
| self.image_ids = Path(image_ids_file).read_text().splitlines() | |
| self.image_dir = image_dir | |
| self.ground_truth_dir = Path(ground_truth_dir) | |
| self.keypoint_logits = np.load( | |
| self.ground_truth_dir / keypoints_file, mmap_mode="r" | |
| ) | |
| self.descriptor_logits = np.load( | |
| self.ground_truth_dir / descriptors_file, mmap_mode="r" | |
| ) | |
| self.image_size = image_size | |
| array_ids_file = self.ground_truth_dir / array_ids_file | |
| if array_ids_file.exists(): | |
| array_ids = array_ids_file.read_text().splitlines() | |
| self.array_indices = { | |
| image_id: idx for idx, image_id in enumerate(array_ids) | |
| } | |
| else: | |
| if len(self.image_ids) > len(self.keypoint_logits): | |
| raise ValueError( | |
| "The image id list is longer than the ground-truth arrays." | |
| ) | |
| self.array_indices = { | |
| image_id: idx for idx, image_id in enumerate(self.image_ids) | |
| } | |
| def __len__(self) -> int: | |
| return len(self.image_ids) | |
| def __getitem__( | |
| self, idx: int | |
| ) -> tuple[torch.Tensor, Any, torch.Tensor, torch.Tensor]: | |
| image_id = self.image_ids[idx] | |
| image_path = os.path.join(self.image_dir, image_id) | |
| image = load_grayscale_image(image_path) | |
| img, scale = rescale_image(image, self.image_size) | |
| img = img[None, None].astype(np.float32) | |
| img = torch.from_numpy(img) | |
| array_idx = self.array_indices[image_id] | |
| keypoint_logits = torch.from_numpy( | |
| np.array(self.keypoint_logits[array_idx], copy=True) | |
| ) | |
| descriptor_logits = torch.from_numpy( | |
| np.array(self.descriptor_logits[array_idx], copy=True) | |
| ) | |
| return img[0], scale, keypoint_logits, descriptor_logits | |
| class SuperPointDataModule(pl.LightningDataModule): | |
| def __init__(self, cfg: dict[str, Any]) -> None: | |
| super().__init__() | |
| self.cfg = cfg | |
| self.train_ds: ImageFolderDataset | None = None | |
| self.val_ds: ImageFolderDataset | None = None | |
| def setup(self, stage: str | None = None) -> None: | |
| self.train_ds = ImageFolderDataset( | |
| image_dir=self.cfg["train_image_dir"], | |
| image_ids_file=self.cfg["train_image_ids_file"], | |
| ground_truth_dir=self.cfg["ground_truth_dir"], | |
| keypoints_file=self.cfg["keypoints_file"], | |
| descriptors_file=self.cfg["descriptors_file"], | |
| array_ids_file=self.cfg["array_ids_file"], | |
| image_size=self.cfg["image_size"], | |
| ) | |
| self.val_ds = ImageFolderDataset( | |
| image_dir=self.cfg["val_image_dir"], | |
| image_ids_file=self.cfg["val_image_ids_file"], | |
| ground_truth_dir=self.cfg["ground_truth_dir"], | |
| keypoints_file=self.cfg["keypoints_file"], | |
| descriptors_file=self.cfg["descriptors_file"], | |
| array_ids_file=self.cfg["array_ids_file"], | |
| image_size=self.cfg["image_size"], | |
| ) | |
| def train_dataloader(self) -> DataLoader: | |
| if self.train_ds is None: | |
| raise RuntimeError("DataModule is not set up.") | |
| return DataLoader( | |
| self.train_ds, | |
| batch_size=self.cfg["batch_size"], | |
| shuffle=True, | |
| num_workers=self.cfg["num_workers"], | |
| pin_memory=self.cfg["pin_memory"], | |
| drop_last=True, | |
| ) | |
| def val_dataloader(self) -> DataLoader: | |
| if self.val_ds is None: | |
| raise RuntimeError("DataModule is not set up.") | |
| return DataLoader( | |
| self.val_ds, | |
| batch_size=self.cfg["batch_size"], | |
| shuffle=False, | |
| num_workers=self.cfg["num_workers"], | |
| pin_memory=self.cfg["pin_memory"], | |
| drop_last=False, | |
| ) | |
| class SuperPointLightningModule(pl.LightningModule): | |
| def __init__(self, cfg: dict[str, Any]) -> None: | |
| super().__init__() | |
| self.save_hyperparameters(cfg) | |
| self.cfg = cfg | |
| self.model = SuperPoint( | |
| num_keypoints=cfg["model"]["num_keypoints"], return_dense=True | |
| ) | |
| pruning_config = cfg["model"]["prune"] | |
| self.model.prune_backbone(pruning_config) | |
| print(self.model) | |
| def forward(self, image: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: | |
| keypoints, descriptors = self.model(image) | |
| return keypoints, descriptors | |
| def _compute_loss( | |
| self, | |
| keypoints: torch.Tensor, | |
| descriptors: torch.Tensor, | |
| keypoints_gt: torch.Tensor, | |
| descriptors_gt: torch.Tensor, | |
| ) -> torch.Tensor: | |
| score_term_hard = detector_loss_simple(keypoints_gt, keypoints) | |
| score_term = detector_kd_kl(keypoints_gt, keypoints) | |
| desc_term = descriptor_loss_simple(descriptors, descriptors_gt) | |
| self.log( | |
| f"train/loss_ce", | |
| score_term_hard, | |
| prog_bar=True, | |
| on_step=True, | |
| on_epoch=True, | |
| ) | |
| self.log( | |
| f"train/loss_kl", score_term, prog_bar=True, on_step=True, on_epoch=True | |
| ) | |
| self.log( | |
| f"train/loss_desc", desc_term, prog_bar=True, on_step=True, on_epoch=True | |
| ) | |
| return ( | |
| self.cfg["loss"]["cross_entropy_coef"] * score_term_hard | |
| + self.cfg["loss"]["kl_coef"] * score_term | |
| + self.cfg["loss"]["descriptor_coef"] * desc_term | |
| ) | |
| def _shared_step(self, batch: dict[str, Any], stage: str) -> torch.Tensor: | |
| images, scales, keypoints_gt, descriptors_gt = batch | |
| keypoint_logits, descriptor_logits = self.forward(images) | |
| loss = self._compute_loss( | |
| keypoint_logits, descriptor_logits, keypoints_gt, descriptors_gt | |
| ) | |
| self.log( | |
| f"{stage}/loss", | |
| loss, | |
| prog_bar=True, | |
| on_step=(stage == "train"), | |
| on_epoch=True, | |
| ) | |
| self.log( | |
| f"{stage}/avg_score", | |
| loss.mean(), | |
| prog_bar=False, | |
| on_step=False, | |
| on_epoch=True, | |
| ) | |
| return loss | |
| def training_step(self, batch: dict[str, Any], batch_idx: int) -> torch.Tensor: | |
| return self._shared_step(batch, stage="train") | |
| def validation_step(self, batch: dict[str, Any], batch_idx: int) -> None: | |
| self._shared_step(batch, stage="val") | |
| def configure_optimizers(self) -> torch.optim.Optimizer: | |
| trainable_params = [p for p in self.parameters() if p.requires_grad] | |
| return torch.optim.AdamW( | |
| trainable_params, | |
| lr=self.cfg["optimizer"]["lr"], | |
| weight_decay=self.cfg["optimizer"]["weight_decay"], | |
| ) | |
| def add_parser_args(parser: argparse.ArgumentParser) -> None: | |
| parser.add_argument( | |
| "--config", type=Path, default=DEFAULT_CONFIG_PATH, help="Path to YAML config" | |
| ) | |
| parser.add_argument( | |
| "--data-root", | |
| type=Path, | |
| default=None, | |
| help="Base for relative data paths in the config. Default: directory of --config.", | |
| ) | |
| def main(args: argparse.Namespace) -> None: | |
| cfg = load_config(args.config, data_root=args.data_root) | |
| pl.seed_everything(cfg["seed"], workers=True) | |
| data_module = SuperPointDataModule(cfg["data"]) | |
| lightning_module = SuperPointLightningModule(cfg) | |
| trainer = pl.Trainer( | |
| max_epochs=cfg["trainer"]["max_epochs"], | |
| accelerator=cfg["trainer"]["accelerator"], | |
| devices=cfg["trainer"]["devices"], | |
| precision=cfg["trainer"]["precision"], | |
| log_every_n_steps=cfg["trainer"]["log_every_n_steps"], | |
| default_root_dir=cfg["trainer"]["default_root_dir"], | |
| # enable_checkpointing=False, | |
| logger=True, | |
| limit_val_batches=0, | |
| ) | |
| trainer.fit(model=lightning_module, datamodule=data_module) | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser() | |
| add_parser_args(parser) | |
| main(parser.parse_args()) | |