| 3 | import torch.nn.functional as F |
| 4 | |
| 5 | class LocalDynamics(nn.Module): |
| 6 | def __init__(self, dims_in): |
| 7 | |
| 8 | super(LocalDynamics, self).__init__() |
| 9 | self.linear = nn.Conv1d(dims_in*2, dims_in, kernel_size=1, stride=1, padding=0) |
| 10 | self.softmax = nn.Softmax(dim = -1) |
| 11 | self.k = 1 |
| 12 | def forward(self, x, mask): |
| 13 | mask = F.interpolate(mask.detach(), size=x.size()[2:], mode='nearest') |
| 14 | B, C, H, W = x.shape |
| 15 | |
| 16 | flattened_features_query = x * mask |
| 17 | flattened_features_support = x * (1-mask) |
| 18 | |
| 19 | flattened_features_query = flattened_features_query.view(B, C, -1) |
| 20 | flattened_features_support = flattened_features_support.view(B, C, -1) |
| 21 | |
| 22 | masked_features = [] |
| 23 | for i in range(B): |
| 24 | |
| 25 | query_features = flattened_features_query[i].view(C, -1) |
| 26 | support_features = flattened_features_support[i].view(C, -1) |
| 27 | |
| 28 | query = query_features.unsqueeze(0) |
| 29 | support = support_features.unsqueeze(0) |
| 30 | |
| 31 | simi_matrix = torch.matmul(query.permute(0, 2, 1), support) |
| 32 | |
| 33 | weights_value, index = torch.topk(simi_matrix, dim=2, k=self.k) |
| 34 | |
| 35 | views = [support.shape[0]] + [1 if i != 2 else -1 for i in range(1, len(support.shape))] |
| 36 | expanse = list(support.shape) |
| 37 | expanse[0] = -1 |
| 38 | expanse[2] = -1 |
| 39 | index = index.view(views) |
| 40 | index = index.expand(expanse) |
| 41 | weights_value = weights_value.view(views) |
| 42 | weights_value = weights_value.expand(expanse) |
| 43 | select_value = torch.gather(support, 2, index) |
| 44 | |
| 45 | select_value = select_value.view(1, C, self.k, -1) |
| 46 | weights_value = weights_value.view(1, C, self.k, -1) |
| 47 | weights_value = self.softmax(weights_value) |
| 48 | fuse_tensor = weights_value * select_value |
| 49 | fuse_tensor = torch.sum(fuse_tensor, -2) |
| 50 | |
| 51 | hybrid_feat = torch.cat((fuse_tensor, query), 1) |
| 52 | hybrid_feat = self.linear(hybrid_feat) |
| 53 | masked_features.append(hybrid_feat) |
| 54 | |
| 55 | refined_feat = torch.cat(masked_features, 0) |
| 56 | refined_feat = refined_feat.view(B, C, H, W) |
| 57 | |
| 58 | return refined_feat * mask + x * (1-mask) |
| 59 | |
| 60 | |
| 61 | |