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