Download src/superpoint_pruning/pruning.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 3.14 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F1/src/superpoint_pruning/pruning.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint@refs/pr/1/src/superpoint_pruning/pruning.py
-
curl -L -o pruning.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/refs%2Fpr%2F1/src/superpoint_pruning/pruning.py
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 | |