MCPcopy Create free account
hub / github.com/chenhaoxing/HDNet / LocalDynamics

Class LocalDynamics

models/att.py:5–58  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

3import torch.nn.functional as F
4
5class 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

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected