"""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())