File size: 4,278 Bytes
6979012
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
"""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())