Download src/superpoint_pruning/models/superpoint.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 13 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/models/superpoint.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint/src/superpoint_pruning/models/superpoint.py
-
curl -L -o superpoint.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/models/superpoint.py
13 kB
| """PyTorch implementation of the SuperPoint model, | |
| derived from the TensorFlow re-implementation (2018). | |
| Authors: Rémi Pautrat, Paul-Edouard Sarlin | |
| """ | |
| from types import SimpleNamespace | |
| import torch | |
| import torch.nn as nn | |
| from superpoint_pruning.paths import DEFAULT_WEIGHTS_PATH | |
| V6_PREFIXES = [ | |
| "backbone.0.0", | |
| "backbone.0.1", | |
| "backbone.1.0", | |
| "backbone.1.1", | |
| "backbone.2.0", | |
| "backbone.2.1", | |
| "backbone.3.0", | |
| "backbone.3.1", | |
| "detector.0", | |
| "detector.1", | |
| "descriptor.0", | |
| "descriptor.1", | |
| ] | |
| def convert_v6_state_dict(state_dict: dict) -> dict: | |
| converted = {} | |
| for key, value in state_dict.items(): | |
| for prefix in V6_PREFIXES: | |
| token = prefix + "." | |
| if not key.startswith(token): | |
| continue | |
| kind, param = key[len(token) :].split(".", 1) | |
| if kind == "conv": | |
| converted[f"{prefix.replace('.', '_')}.{param}"] = value | |
| elif kind == "bn": | |
| converted[f"{prefix.replace('.', '_')}_bn.{param}"] = value | |
| else: | |
| raise KeyError(f"Unrecognized v6 submodule in '{key}'") | |
| break | |
| else: | |
| raise KeyError(f"Unrecognized v6 key: {key}") | |
| return converted | |
| def sample_descriptors(keypoints, descriptors, s: int = 8): | |
| b, c, h, w = descriptors.shape | |
| divisor = ( | |
| torch._shape_as_tensor(descriptors)[[3, 2]] | |
| .to(keypoints.dtype) | |
| .to(keypoints.device) | |
| * s | |
| ) | |
| keypoints = (keypoints + 0.5) / divisor | |
| keypoints = keypoints * 2 - 1 | |
| descriptors = torch.nn.functional.grid_sample( | |
| descriptors, keypoints.view(b, 1, -1, 2), mode="bilinear", align_corners=False | |
| ) | |
| descriptors = torch.nn.functional.normalize( | |
| descriptors.reshape(b, c, -1), p=2, dim=1 | |
| ).permute(0, 2, 1) | |
| return descriptors | |
| def batched_nms(scores, nms_radius: int, skip_refinement: bool = False): | |
| assert nms_radius >= 0 | |
| def max_pool(x): | |
| return torch.nn.functional.max_pool2d( | |
| x, kernel_size=nms_radius * 2 + 1, stride=1, padding=nms_radius | |
| ) | |
| scores = scores[:, None] | |
| zeros = torch.zeros_like(scores) | |
| max_mask = scores == max_pool(scores) | |
| if not skip_refinement: | |
| for _ in range(2): | |
| supp_mask = max_pool(max_mask.float()) > 0 | |
| supp_scores = torch.where(supp_mask, zeros, scores) | |
| new_max_mask = supp_scores == max_pool(supp_scores) | |
| max_mask = max_mask | (new_max_mask & (~supp_mask)) | |
| return torch.where(max_mask, scores, zeros)[:, 0] | |
| def hierarchical_topk(scores, tile_size: int, num_keypoints: int): | |
| B, H, W = scores.shape | |
| tile_h = tile_size | |
| assert H % tile_h == 0, "Tile size must divide the height of the scores" | |
| scores_tiled = scores.reshape(B, H // tile_h, tile_h * W) | |
| local_scores, local_idx = scores_tiled.topk(num_keypoints, dim=-1, sorted=False) | |
| # Convert local indices into global flattened indices. | |
| tile_offset = ( | |
| torch.arange(H // tile_h, device=scores.device).view(1, -1, 1) * tile_h * W | |
| ) | |
| global_idx = local_idx + tile_offset | |
| local_scores = local_scores.reshape(B, -1) | |
| global_idx = global_idx.reshape(B, -1) | |
| top_scores, sel = local_scores.topk(num_keypoints, dim=-1, sorted=True) | |
| top_indices = global_idx.gather(1, sel) | |
| return top_scores, top_indices | |
| def get_conv_layer(c_in, c_out, kernel_size, relu=True): | |
| padding = (kernel_size - 1) // 2 | |
| conv = nn.Conv2d(c_in, c_out, kernel_size=kernel_size, stride=1, padding=padding) | |
| return conv | |
| class SuperPoint(nn.Module): | |
| default_conf = { | |
| "nms_radius": 4, | |
| "num_keypoints": 1024, | |
| "remove_borders": 4, | |
| "descriptor_dim": 256, | |
| "channels": [64, 64, 128, 128, 256], | |
| "skip_refinement": False, | |
| "hierarchical_topk": False, | |
| "hierarchical_tile_size": 32, | |
| "use_bn": True, | |
| "default_weights_path": DEFAULT_WEIGHTS_PATH, | |
| "load_default_weights": True, | |
| "return_dense": False, | |
| } | |
| def __init__(self, **conf): | |
| super().__init__() | |
| conf = {**self.default_conf, **conf} | |
| self.conf = SimpleNamespace(**conf) | |
| self.stride = 2 ** (len(self.conf.channels) - 2) | |
| channels = [1, *self.conf.channels[:-1]] | |
| self.pool = nn.MaxPool2d(kernel_size=2, stride=2) | |
| self.relu = nn.ReLU(inplace=True) | |
| for i, c in enumerate(channels[1:]): | |
| self.add_module(f"backbone_{i}_0", get_conv_layer(channels[i], c, 3)) | |
| if self.conf.use_bn: | |
| self.add_module(f"backbone_{i}_0_bn", nn.BatchNorm2d(c, eps=0.001)) | |
| self.add_module(f"backbone_{i}_1", get_conv_layer(c, c, 3)) | |
| if self.conf.use_bn: | |
| self.add_module(f"backbone_{i}_1_bn", nn.BatchNorm2d(c, eps=0.001)) | |
| c = self.conf.channels[-1] | |
| self.add_module("detector_0", get_conv_layer(channels[-1], c, 3)) | |
| if self.conf.use_bn: | |
| self.add_module("detector_0_bn", nn.BatchNorm2d(c, eps=0.001)) | |
| self.add_module("detector_1", get_conv_layer(c, self.stride**2 + 1, 1)) | |
| if self.conf.use_bn: | |
| self.add_module( | |
| "detector_1_bn", nn.BatchNorm2d(self.stride**2 + 1, eps=0.001) | |
| ) | |
| self.add_module("descriptor_0", get_conv_layer(channels[-1], c, 3)) | |
| if self.conf.use_bn: | |
| self.add_module("descriptor_0_bn", nn.BatchNorm2d(c, eps=0.001)) | |
| self.add_module("descriptor_1", get_conv_layer(c, self.conf.descriptor_dim, 1)) | |
| if self.conf.use_bn: | |
| self.add_module( | |
| "descriptor_1_bn", nn.BatchNorm2d(self.conf.descriptor_dim, eps=0.001) | |
| ) | |
| if self.conf.use_bn and self.conf.load_default_weights: | |
| self.load_default_weights(self.conf.default_weights_path) | |
| def _forward_conv( | |
| self, name: str, x: torch.Tensor, relu: bool = True | |
| ) -> torch.Tensor: | |
| x = getattr(self, name)(x) | |
| if relu: | |
| x = self.relu(x) | |
| if self.conf.use_bn: | |
| x = getattr(self, name + "_bn")(x) | |
| return x | |
| def load_default_weights(self, weights_path: str) -> None: | |
| """Load and rename superpoint_v6_from_tf.pth weights.""" | |
| checkpoint = torch.load(weights_path, map_location="cpu") | |
| state_dict = convert_v6_state_dict(checkpoint) | |
| missing, unexpected = self.load_state_dict(state_dict, strict=False) | |
| missing = [k for k in missing if not k.endswith("num_batches_tracked")] | |
| if missing or unexpected: | |
| raise RuntimeError( | |
| f"Failed to load weights from {weights_path}. missing={missing} unexpected={unexpected}" | |
| ) | |
| print(f"Loaded default weights from {weights_path}") | |
| def load_pruned_weights(self, checkpoint_path: str) -> None: | |
| """Load pruned SP weights from checkpoint""" | |
| checkpoint = torch.load(checkpoint_path, map_location="cpu")["state_dict"] | |
| checkpoint = {k.replace("model.", ""): v for k, v in checkpoint.items()} | |
| self.load_state_dict(checkpoint, strict=True) | |
| def _bn_name(self, conv_name: str) -> str: | |
| return conv_name + "_bn" | |
| def _previous_backbone_layer(self, layer: str) -> str: | |
| parts = layer.split("_") | |
| stage, index = int(parts[1]), int(parts[2]) | |
| if index == 1: | |
| return f"backbone_{stage}_0" | |
| if index == 0 and stage > 0: | |
| return f"backbone_{stage - 1}_1" | |
| raise ValueError(f"Invalid layer: {layer}") | |
| def _prune_bn(self, conv_name: str, keep_idx: torch.Tensor) -> None: | |
| """Resize the BN that follows a pruned conv so channel counts still match.""" | |
| bn_name = self._bn_name(conv_name) | |
| old_bn = getattr(self, bn_name) | |
| new_bn = nn.BatchNorm2d( | |
| len(keep_idx), | |
| eps=old_bn.eps, | |
| momentum=old_bn.momentum, | |
| affine=old_bn.affine, | |
| track_running_stats=old_bn.track_running_stats, | |
| ) | |
| with torch.no_grad(): | |
| if old_bn.affine: | |
| new_bn.weight.copy_(old_bn.weight[keep_idx]) | |
| new_bn.bias.copy_(old_bn.bias[keep_idx]) | |
| if old_bn.track_running_stats: | |
| new_bn.running_mean.copy_(old_bn.running_mean[keep_idx]) | |
| new_bn.running_var.copy_(old_bn.running_var[keep_idx]) | |
| new_bn.num_batches_tracked.copy_(old_bn.num_batches_tracked) | |
| setattr(self, bn_name, new_bn) | |
| def prune_backbone(self, config: dict): | |
| for layer, channel in config.items(): | |
| previous_layer = self._previous_backbone_layer(layer) | |
| old1 = getattr(self, previous_layer) | |
| old2 = getattr(self, layer) | |
| magnitude = old2.weight.abs().mean(dim=(0, 2, 3)) | |
| keep_idx = ( | |
| torch.topk(magnitude, k=channel, largest=True).indices.sort().values | |
| ) | |
| new1 = torch.nn.Conv2d( | |
| old1.in_channels, | |
| channel, | |
| kernel_size=old1.kernel_size, | |
| stride=1, | |
| padding=1, | |
| ) | |
| new2 = torch.nn.Conv2d( | |
| channel, | |
| old2.out_channels, | |
| kernel_size=old2.kernel_size, | |
| stride=1, | |
| padding=1, | |
| ) | |
| with torch.no_grad(): | |
| new1.weight.copy_(old1.weight[keep_idx, :, :, :]) | |
| new1.bias.copy_(old1.bias[keep_idx]) | |
| new2.weight.copy_(old2.weight[:, keep_idx, :, :]) | |
| new2.bias.copy_(old2.bias) | |
| setattr(self, previous_layer, new1) | |
| setattr(self, layer, new2) | |
| if self.conf.use_bn: | |
| self._prune_bn(previous_layer, keep_idx) | |
| print("Pruned model structure:") | |
| print(self) | |
| def dense_head( | |
| self, scores: torch.Tensor, descriptors_dense: torch.Tensor | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| descriptors_dense = torch.nn.functional.normalize(descriptors_dense, p=2, dim=1) | |
| scores = torch.nn.functional.softmax(scores, 1)[:, :-1] | |
| b, _, h, w = scores.shape | |
| scores = scores.permute(0, 2, 3, 1).reshape(b, h, w, self.stride, self.stride) | |
| scores = scores.permute(0, 1, 3, 2, 4).reshape( | |
| b, h * self.stride, w * self.stride | |
| ) | |
| scores = batched_nms( | |
| scores, self.conf.nms_radius, skip_refinement=self.conf.skip_refinement | |
| ) | |
| # Discard keypoints near the image borders | |
| if self.conf.remove_borders: | |
| pad = self.conf.remove_borders | |
| scores[:, :pad] = -1 | |
| scores[:, :, :pad] = -1 | |
| scores[:, -pad:] = -1 | |
| scores[:, :, -pad:] = -1 | |
| if self.conf.hierarchical_topk: | |
| top_scores, top_indices = hierarchical_topk( | |
| scores, self.conf.hierarchical_tile_size, self.conf.num_keypoints | |
| ) | |
| else: | |
| top_scores, top_indices = scores.reshape( | |
| b, h * self.stride * w * self.stride | |
| ).topk(self.conf.num_keypoints) | |
| y_idx = torch.div(top_indices, w * self.stride, rounding_mode="floor") | |
| x_idx = torch.remainder(top_indices, w * self.stride) | |
| top_keypoints = torch.stack((x_idx, y_idx), dim=-1).to(dtype=torch.float32) | |
| top_descriptors = sample_descriptors( | |
| top_keypoints, descriptors_dense, self.stride | |
| ) | |
| return ( | |
| top_keypoints, | |
| top_scores, | |
| top_descriptors, | |
| ) | |
| def forward( | |
| self, image: torch.Tensor | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| x = self._forward_conv("backbone_0_0", image) | |
| x = self._forward_conv("backbone_0_1", x) | |
| x = self.pool(x) | |
| x = self._forward_conv("backbone_1_0", x) | |
| x = self._forward_conv("backbone_1_1", x) | |
| x = self.pool(x) | |
| x = self._forward_conv("backbone_2_0", x) | |
| x = self._forward_conv("backbone_2_1", x) | |
| x = self.pool(x) | |
| x = self._forward_conv("backbone_3_0", x) | |
| x = self._forward_conv("backbone_3_1", x) | |
| descriptors_dense = self._forward_conv( | |
| "descriptor_1", self._forward_conv("descriptor_0", x), relu=False | |
| ) | |
| scores = self._forward_conv( | |
| "detector_1", self._forward_conv("detector_0", x), relu=False | |
| ) | |
| if self.conf.return_dense: | |
| return scores, descriptors_dense | |
| return self.dense_head(scores, descriptors_dense) | |
| if __name__ == "__main__": | |
| sp = SuperPoint(num_keypoints=512) | |
| inputs = torch.zeros(1, 1, 768, 1024) | |
| torch.onnx.export( | |
| sp.cpu(), | |
| inputs, | |
| f"SP.onnx", | |
| input_names=["inputs"], | |
| output_names=["keypoints", "scores", "descriptors"], | |
| # opset_version=opset, | |
| # dynamic_axes=dynamic_axes, | |
| # dynamic_shapes=dynamic_shapes, | |
| # dynamo=True, | |
| ) | |