| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| from copy import deepcopy |
| from pathlib import Path |
| import torch |
| from torch import nn |
| |
| from torch_scatter import scatter as super_pixel_pooling |
| import argparse |
|
|
| def MLP(channels: list, do_bn=True): |
| """ Multi-layer perceptron """ |
| n = len(channels) |
| layers = [] |
| for i in range(1, n): |
| layers.append( |
| nn.Conv1d(channels[i - 1], channels[i], kernel_size=1, bias=True)) |
| if i < (n-1): |
| if do_bn: |
| |
| layers.append(nn.InstanceNorm1d(channels[i])) |
| layers.append(nn.ReLU()) |
| return nn.Sequential(*layers) |
|
|
|
|
| def normalize_keypoints(kpts, image_shape): |
| """ Normalize keypoints locations based on image image_shape""" |
| _, _, height, width = image_shape |
| one = kpts.new_tensor(1) |
| size = torch.stack([one*width, one*width, one*height, one*height])[None] |
| center = size / 2 |
| scaling = size.max(1, keepdim=True).values * 0.7 |
| |
| return (kpts - center[:, None, :]) / scaling[:, None, :] |
|
|
| class ThreeLayerDecoder(nn.Module): |
| """ Joint encoding of visual appearance and location using MLPs""" |
| def __init__(self, enc_dim): |
| super().__init__() |
| |
| self.layer1 = nn.Conv2d(enc_dim, enc_dim, 3, padding=1) |
| self.non_linear1 = nn.ReLU() |
| self.layer2 = nn.Conv2d(enc_dim, enc_dim, 3, padding=1) |
| self.non_linear2 = nn.ReLU() |
| self.layer3 = nn.Conv2d(enc_dim, enc_dim, 1) |
|
|
| self.norm1 = nn.InstanceNorm2d(enc_dim) |
| self.norm2 = nn.InstanceNorm2d(enc_dim) |
| self.norm3 = nn.InstanceNorm2d(enc_dim) |
|
|
| for m in self.modules(): |
| if isinstance(m, nn.Conv2d): |
| nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') |
| nn.init.constant_(m.bias, 0.0) |
|
|
| def forward(self, img): |
| x = self.non_linear1(self.norm1(self.layer1(img))) |
| x = self.non_linear2(self.norm2(self.layer2(x))) |
| x = self.norm3(self.layer3(x)) |
| |
| |
| |
| return x |
|
|
|
|
| class SegmentDescriptor(nn.Module): |
| """ Joint encoding of visual appearance and location using MLPs""" |
| def __init__(self, enc_dim): |
| super().__init__() |
| |
| |
| |
| |
|
|
| def forward(self, img, seg): |
| |
| n, c, h, w = x.size() |
| assert((h, w) == img.size()[2:4]) |
| return super_pixel_pooling(x.view(n, c, -1), seg.view(-1).long(), reduce='mean') |
| |
|
|
| class SegmentAdaptor(nn.Module): |
| """ Joint encoding of visual appearance and location using MLPs""" |
| def __init__(self, enc_dim): |
| super().__init__() |
| self.encoder = ThreeLayerDecoder(enc_dim) |
| |
| |
| |
|
|
| def forward(self, desc, seg): |
| x = desc[seg] |
| x = self.decoder(x) |
| return x |
| |
|
|
| class KeypointEncoder(nn.Module): |
| """ Joint encoding of visual appearance and location using MLPs""" |
| def __init__(self, feature_dim, layers): |
| super().__init__() |
| self.encoder = MLP([4] + layers + [feature_dim]) |
| |
| |
| |
| |
| nn.init.constant_(self.encoder[-1].bias, 0.0) |
|
|
| def forward(self, kpts): |
| inputs = kpts.transpose(1, 2) |
| |
| x = self.encoder(inputs) |
| |
| return x |
|
|
|
|
| def attention(query, key, value): |
| dim = query.shape[1] |
| scores = torch.einsum('bdhn,bdhm->bhnm', query, key) / dim**.5 |
| prob = torch.nn.functional.softmax(scores, dim=-1) |
| return torch.einsum('bhnm,bdhm->bdhn', prob, value), prob |
|
|
|
|
| class MultiHeadedAttention(nn.Module): |
| """ Multi-head attention to increase model expressivitiy """ |
| def __init__(self, num_heads: int, d_model: int): |
| super().__init__() |
| assert d_model % num_heads == 0 |
| self.dim = d_model // num_heads |
| self.num_heads = num_heads |
| self.merge = nn.Conv1d(d_model, d_model, kernel_size=1) |
| self.proj = nn.ModuleList([deepcopy(self.merge) for _ in range(3)]) |
|
|
| def forward(self, query, key, value): |
| batch_dim = query.size(0) |
| query, key, value = [l(x).view(batch_dim, self.dim, self.num_heads, -1) |
| for l, x in zip(self.proj, (query, key, value))] |
| x, prob = attention(query, key, value) |
| self.prob.append(prob) |
| return self.merge(x.contiguous().view(batch_dim, self.dim*self.num_heads, -1)) |
|
|
|
|
| class AttentionalPropagation(nn.Module): |
| def __init__(self, feature_dim: int, num_heads: int): |
| super().__init__() |
| self.attn = MultiHeadedAttention(num_heads, feature_dim) |
| self.mlp = MLP([feature_dim*2, feature_dim*2, feature_dim]) |
| nn.init.constant_(self.mlp[-1].bias, 0.0) |
|
|
| def forward(self, x, source): |
| message = self.attn(x, source, source) |
| return self.mlp(torch.cat([x, message], dim=1)) |
|
|
|
|
| class AttentionalGNN(nn.Module): |
| def __init__(self, feature_dim: int, layer_names: list): |
| super().__init__() |
| self.layers = nn.ModuleList([ |
| AttentionalPropagation(feature_dim, 4) |
| for _ in range(len(layer_names))]) |
| self.names = layer_names |
|
|
| def forward(self, desc0, desc1): |
| for layer, name in zip(self.layers, self.names): |
| layer.attn.prob = [] |
| if name == 'cross': |
| src0, src1 = desc1, desc0 |
| else: |
| src0, src1 = desc0, desc1 |
| delta0, delta1 = layer(desc0, src0), layer(desc1, src1) |
| desc0, desc1 = (desc0 + delta0), (desc1 + delta1) |
| return desc0, desc1 |
|
|
|
|
| def log_sinkhorn_iterations(Z, log_mu, log_nu, iters: int): |
| """ Perform Sinkhorn Normalization in Log-space for stability""" |
| u, v = torch.zeros_like(log_mu), torch.zeros_like(log_nu) |
| for _ in range(iters): |
| u = log_mu - torch.logsumexp(Z + v.unsqueeze(1), dim=2) |
| v = log_nu - torch.logsumexp(Z + u.unsqueeze(2), dim=1) |
| return Z + u.unsqueeze(2) + v.unsqueeze(1) |
|
|
|
|
| def log_optimal_transport(scores, alpha, iters: int): |
| """ Perform Differentiable Optimal Transport in Log-space for stability""" |
| b, m, n = scores.shape |
| one = scores.new_tensor(1) |
| ms, ns = (m*one).to(scores), (n*one).to(scores) |
|
|
| bins0 = alpha.expand(b, m, 1) |
| bins1 = alpha.expand(b, 1, n) |
| alpha = alpha.expand(b, 1, 1) |
|
|
| couplings = torch.cat([torch.cat([scores, bins0], -1), |
| torch.cat([bins1, alpha], -1)], 1) |
|
|
| norm = - (ms + ns).log() |
| log_mu = torch.cat([norm.expand(m), ns.log()[None] + norm]) |
| log_nu = torch.cat([norm.expand(n), ms.log()[None] + norm]) |
| log_mu, log_nu = log_mu[None].expand(b, -1), log_nu[None].expand(b, -1) |
|
|
| Z = log_sinkhorn_iterations(couplings, log_mu, log_nu, iters) |
| Z = Z - norm |
| return Z |
|
|
|
|
| def arange_like(x, dim: int): |
| return x.new_ones(x.shape[dim]).cumsum(0) - 1 |
|
|
|
|
| class Augmentor(nn.Module): |
| """SuperGlue feature matching middle-end |
| |
| Given two sets of keypoints and locations, we determine the |
| correspondences by: |
| 1. Keypoint Encoding (normalization + visual feature and location fusion) |
| 2. Graph Neural Network with multiple self and cross-attention layers |
| 3. Final projection layer |
| 4. Optimal Transport Layer (a differentiable Hungarian matching algorithm) |
| 5. Thresholding matrix based on mutual exclusivity and a match_threshold |
| |
| The correspondence ids use -1 to indicate non-matching points. |
| |
| Paul-Edouard Sarlin, Daniel DeTone, Tomasz Malisiewicz, and Andrew |
| Rabinovich. SuperGlue: Learning Feature Matching with Graph Neural |
| Networks. In CVPR, 2020. https://arxiv.org/abs/1911.11763 |
| |
| """ |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| def __init__(self, config=None): |
| super().__init__() |
|
|
| default_config = argparse.Namespace() |
| default_config.descriptor_dim = 128 |
| |
| |
| default_config.GNN_layers = ['self', 'cross'] * 9 |
| default_config.sinkhorn_iterations = 100 |
| default_config.match_threshold = 0.2 |
| |
|
|
| if config is None: |
| self.config = default_config |
| else: |
| self.config = config |
| self.config.GNN_layers = ['self', 'cross'] * self.config.GNN_layer_num |
| |
|
|
| self.kenc = KeypointEncoder( |
| self.config.descriptor_dim, self.config.keypoint_encoder) |
|
|
| self.gnn = AttentionalGNN( |
| self.config.descriptor_dim, self.config.GNN_layers) |
|
|
| self.final_proj = nn.Conv1d( |
| self.config.descriptor_dim, self.config.descriptor_dim, |
| kernel_size=1, bias=True) |
|
|
| bin_score = torch.nn.Parameter(torch.tensor(1.)) |
| self.register_parameter('bin_score', bin_score) |
| self.segment_adaptor = SegmentAdaptor(self.config.descriptor_dim) |
|
|
| |
| |
| |
| |
| |
| |
|
|
| def forward(self, data): |
| """Run SuperGlue on a pair of keypoints and descriptors""" |
| |
| |
| desc0, desc1 = self.segment_adaptor(data['feat0'], data['segment0']), self.segment_adaptor(data['feat1'], data['segment1']) |
| |
| kpts0, kpts1 = data['keypoints0'].float(), data['keypoints1'].float() |
|
|
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| kpts0 = normalize_keypoints(kpts0, data['image0'].shape) |
| kpts1 = normalize_keypoints(kpts1, data['image1'].shape) |
|
|
| |
| |
| |
| |
| pos0 = self.kenc(kpts0) |
| pos1 = self.kenc(kpts1) |
| |
| desc0 = desc0 + pos0 |
| desc1 = desc1 + pos1 |
|
|
| |
| desc0, desc1 = self.gnn(desc0, desc1) |
|
|
| |
| mdesc0, mdesc1 = self.final_proj(desc0), self.final_proj(desc1) |
| return mdesc0, mdesc1, None |
| |
|
|
|
|
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| if __name__ == '__main__': |
| from anime_seg_mat_dataset import fetch_dataloader |
| args = argparse.Namespace() |
| ss = SuperGlue() |
| args.subset = 'trytry' |
| args.batch_size = 1 |
| args.stage = 'anime' |
| args.image_size = (368, 368) |
| loader = fetch_dataloader(args) |
| |
| for data in loader: |
| |
| dict1 = data |
|
|
| kp1 = dict1['keypoints0'] |
| kp2 = dict1['keypoints1'] |
| p1 = dict1['image0'] |
| p2 = dict1['image1'] |
| s1 = dict1['segment0'] |
| s2 = dict1['segment1'] |
| |
| |
| mi = dict1['all_matches'] |
| fname = dict1['file_name'] |
| |
| |
|
|
| a = ss(data) |
| |
| |
| |
| |
| |