"""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, )