File size: 3,140 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
from __future__ import annotations

import argparse
from pathlib import Path

from superpoint_pruning.models.superpoint import SuperPoint

PRUNABLE_LAYERS = (
    "backbone_0_1",
    "backbone_1_0",
    "backbone_1_1",
    "backbone_2_0",
    "backbone_2_1",
    "backbone_3_0",
    "backbone_3_1",
)
_C = SuperPoint.default_conf["channels"]
MAX_PRUNED_CHANNELS = {
    "backbone_0_1": _C[0],
    "backbone_1_0": _C[0],
    "backbone_1_1": _C[1],
    "backbone_2_0": _C[1],
    "backbone_2_1": _C[2],
    "backbone_3_0": _C[2],
    "backbone_3_1": _C[3],
}


def parse_pruning_config(checkpoint_path: str) -> dict[str, int]:
    """Build a pruning config from a checkpoint named ``a_b_c...ckpt``."""
    checkpoint = Path(checkpoint_path)
    if checkpoint.suffix != ".ckpt":
        raise ValueError("A pruning checkpoint must have a .ckpt extension.")

    values = checkpoint.stem.split("_")
    if not 1 <= len(values) <= len(PRUNABLE_LAYERS):
        raise ValueError(
            f"Expected 1-{len(PRUNABLE_LAYERS)} underscore-separated channel counts in '{checkpoint.name}'."
        )

    try:
        config = {layer: int(value) for layer, value in zip(PRUNABLE_LAYERS, values)}
    except ValueError as error:
        raise ValueError(
            "Pruning checkpoint names must contain only integer channel counts, for example '32_48_64.ckpt'."
        ) from error

    if any(channels <= 0 for channels in config.values()):
        raise ValueError("Pruned channel counts must be positive.")
    return config


def validate_pruning_config(pruning_config: dict[str, int]) -> None:
    for layer, channels in pruning_config.items():
        max_channels = MAX_PRUNED_CHANNELS[layer]
        if channels > max_channels:
            raise ValueError(
                f"{layer} cannot be pruned to {channels} channels; it has only {max_channels} input channels."
            )


def add_pruning_parser_args(parser: argparse.ArgumentParser) -> None:
    parser.add_argument(
        "--pruning-config",
        type=Path,
        help=(
            "Checkpoint path with pruned weights. Its basename must be named like '32_48_64.ckpt'; "
            "counts map to backbone_0_1, backbone_1_0, backbone_1_1, and so on."
        ),
    )
    for layer in PRUNABLE_LAYERS:
        parser.add_argument(
            f"--{layer}",
            dest=layer,
            type=int,
            metavar="CHANNELS",
            help=f"Prune the input channels of {layer}.",
        )


def pruning_config_from_args(
    args: argparse.Namespace,
) -> tuple[dict[str, int], Path | None]:
    cli_pruning_config = {
        layer: getattr(args, layer)
        for layer in PRUNABLE_LAYERS
        if getattr(args, layer) is not None
    }
    if args.pruning_config and cli_pruning_config:
        raise ValueError(
            "Use either --pruning-config or individual layer options, not both."
        )

    pruning_config = (
        parse_pruning_config(str(args.pruning_config))
        if args.pruning_config
        else cli_pruning_config
    )
    validate_pruning_config(pruning_config)
    return pruning_config, args.pruning_config