Download src/superpoint_pruning/export.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 4.28 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F2/src/superpoint_pruning/export.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint@refs/pr/2/src/superpoint_pruning/export.py
-
curl -L -o export.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F2/src/superpoint_pruning/export.py
4.28 kB
| """Export SuperPoint to ONNX. | |
| Examples: | |
| superpoint-pruning export --output superpoint.onnx | |
| superpoint-pruning export --backbone_0_1 32 --backbone_1_0 48 --output pruned.onnx | |
| superpoint-pruning export --pruning-config 32_48_64.ckpt --output pruned.onnx | |
| python -m superpoint_pruning.export --output superpoint.onnx | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| import torch | |
| from superpoint_pruning.models.superpoint import SuperPoint | |
| from superpoint_pruning.pruning import add_pruning_parser_args, pruning_config_from_args | |
| def validate_export_args(args: argparse.Namespace) -> None: | |
| if args.width <= 0 or args.height <= 0 or args.batch_dim <= 0: | |
| raise ValueError("Width, height, and batch_dim must be positive.") | |
| if args.width % 8 or args.height % 8: | |
| raise ValueError("Width and height must be divisible by 8.") | |
| if args.num_keypoints <= 0 or args.num_keypoints > args.width * args.height: | |
| raise ValueError("num_keypoints must be between 1 and width * height.") | |
| if args.hierarchical_topk: | |
| if args.hierarchical_tile_size <= 0: | |
| raise ValueError("hierarchical_tile_size must be positive.") | |
| if args.height % args.hierarchical_tile_size: | |
| raise ValueError( | |
| "hierarchical_tile_size must divide the input height when hierarchical_topk is enabled." | |
| ) | |
| if args.num_keypoints > args.hierarchical_tile_size * args.width: | |
| raise ValueError( | |
| "num_keypoints must not exceed hierarchical_tile_size * width when hierarchical_topk is enabled." | |
| ) | |
| def add_parser_args(parser: argparse.ArgumentParser) -> None: | |
| parser.add_argument("--num-keypoints", type=int, default=512) | |
| parser.add_argument( | |
| "--skip-refinement", | |
| action=argparse.BooleanOptionalAction, | |
| default=False, | |
| help="Skip the iterative NMS refinement.", | |
| ) | |
| parser.add_argument( | |
| "--hierarchical-topk", | |
| action=argparse.BooleanOptionalAction, | |
| default=False, | |
| ) | |
| parser.add_argument("--hierarchical-tile-size", type=int, default=32) | |
| parser.add_argument("--width", type=int, default=640) | |
| parser.add_argument("--height", type=int, default=480) | |
| parser.add_argument("--batch-dim", type=int, default=1) | |
| parser.add_argument("--output", type=Path, default=Path("superpoint.onnx")) | |
| add_pruning_parser_args(parser) | |
| parser.add_argument( | |
| "--skip-bn", action=argparse.BooleanOptionalAction, default=False | |
| ) | |
| def build_parser() -> argparse.ArgumentParser: | |
| parser = argparse.ArgumentParser(description="Export a SuperPoint model to ONNX.") | |
| add_parser_args(parser) | |
| return parser | |
| def main(args: argparse.Namespace) -> None: | |
| pruning_config, pruning_checkpoint = pruning_config_from_args(args) | |
| validate_export_args(args) | |
| print("num_keypoints: ", args.num_keypoints) | |
| print("skip_refinement: ", args.skip_refinement) | |
| print("hierarchical_topk: ", args.hierarchical_topk) | |
| print("hierarchical_tile_size: ", args.hierarchical_tile_size) | |
| print("width: ", args.width) | |
| print("height: ", args.height) | |
| print("batch_dim: ", args.batch_dim) | |
| print("output: ", args.output) | |
| print("pruning_config: ", args.pruning_config) | |
| print("skip_bn: ", args.skip_bn) | |
| model = SuperPoint( | |
| num_keypoints=args.num_keypoints, | |
| skip_refinement=args.skip_refinement, | |
| hierarchical_topk=args.hierarchical_topk, | |
| hierarchical_tile_size=args.hierarchical_tile_size, | |
| use_bn=not args.skip_bn, | |
| ) | |
| if pruning_config: | |
| model.prune_backbone(pruning_config) | |
| if pruning_checkpoint: | |
| model.load_pruned_weights(str(pruning_checkpoint)) | |
| inputs = torch.zeros(args.batch_dim, 1, args.height, args.width) | |
| model.eval() | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| with torch.no_grad(): | |
| torch.onnx.export( | |
| model.cpu(), | |
| inputs, | |
| args.output, | |
| input_names=["inputs"], | |
| output_names=["keypoints", "scores", "descriptors"], | |
| ) | |
| print(f"Exported {args.output} with input shape {tuple(inputs.shape)}.") | |
| if __name__ == "__main__": | |
| main(build_parser().parse_args()) | |