MCPcopy Create free account
hub / github.com/NJUNLP/GTS / multi_hops

Method multi_hops

code/BertModel/model.py:19–51  ·  view source on GitHub ↗

generate mask

(self, features, mask, k)

Source from the content-addressed store, hash-verified

17 self.dropout_output = torch.nn.Dropout(0.1)
18
19 def multi_hops(self, features, mask, k):
20 '''generate mask'''
21 max_length = features.shape[1]
22 mask = mask[:, :max_length]
23 mask_a = mask.unsqueeze(1).expand([-1, max_length, -1])
24 mask_b = mask.unsqueeze(2).expand([-1, -1, max_length])
25 mask = mask_a * mask_b
26 mask = torch.triu(mask).unsqueeze(3).expand([-1, -1, -1, self.args.class_num])
27
28 '''save all logits'''
29 logits_list = []
30 logits = self.cls_linear(features)
31 logits_list.append(logits)
32
33 for i in range(k):
34 #probs = torch.softmax(logits, dim=3)
35 probs = logits
36 logits = probs * mask
37
38 logits_a = torch.max(logits, dim=1)[0]
39 logits_b = torch.max(logits, dim=2)[0]
40 logits = torch.cat([logits_a.unsqueeze(3), logits_b.unsqueeze(3)], dim=3)
41 logits = torch.max(logits, dim=3)[0]
42
43 logits = logits.unsqueeze(2).expand([-1,-1, max_length, -1])
44 logits_T = logits.transpose(1, 2)
45 logits = torch.cat([logits, logits_T], dim=3)
46
47 new_features = torch.cat([features, logits, probs], dim=3)
48 features = self.feature_linear(new_features)
49 logits = self.cls_linear(features)
50 logits_list.append(logits)
51 return logits_list
52
53 def forward(self, tokens, masks):
54 bert_feature, _ = self.bert(tokens, masks)

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected