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