oskarkuuse-pruna's picture
commit repo mirror
6979012
Raw History Blame Contribute Delete
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,
)