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