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

Class MultiInferRNNModel

code/NNModel/model.py:9–100  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class MultiInferRNNModel(torch.nn.Module):
10 def __init__(self, gen_emb, domain_emb, args):
11 '''double embedding + lstm encoder + dot self attention'''
12 super(MultiInferRNNModel, self).__init__()
13
14 self.args = args
15 self.gen_embedding = torch.nn.Embedding(gen_emb.shape[0], gen_emb.shape[1])
16 self.gen_embedding.weight.data.copy_(gen_emb)
17 self.gen_embedding.weight.requires_grad = False
18
19 self.domain_embedding = torch.nn.Embedding(domain_emb.shape[0], domain_emb.shape[1])
20 self.domain_embedding.weight.data.copy_(domain_emb)
21 self.domain_embedding.weight.requires_grad = False
22
23 self.dropout1 = torch.nn.Dropout(0.5)
24 self.dropout2 = torch.nn.Dropout(0)
25
26 self.bilstm = torch.nn.LSTM(300+100, args.lstm_dim,
27 num_layers=1, batch_first=True, bidirectional=True)
28 self.attention_layer = SelfAttention(args)
29
30 self.feature_linear = torch.nn.Linear(args.lstm_dim*4 + args.class_num*3, args.lstm_dim*4)
31 self.cls_linear = torch.nn.Linear(args.lstm_dim*4, args.class_num)
32
33 def _get_embedding(self, sentence_tokens, mask):
34 gen_embed = self.gen_embedding(sentence_tokens)
35 domain_embed = self.domain_embedding(sentence_tokens)
36 embedding = torch.cat([gen_embed, domain_embed], dim=2)
37 embedding = self.dropout1(embedding)
38 embedding = embedding * mask.unsqueeze(2).float().expand_as(embedding)
39 return embedding
40
41 def _lstm_feature(self, embedding, lengths):
42 embedding = pack_padded_sequence(embedding, lengths, batch_first=True)
43 context, _ = self.bilstm(embedding)
44 context, _ = pad_packed_sequence(context, batch_first=True)
45 return context
46
47 def _cls_logits(self, features):
48 # features = self.dropout2(features)
49 tags = self.cls_linear(features)
50 return tags
51
52 def multi_hops(self, features, lengths, mask, k):
53 '''generate mask'''
54 max_length = features.shape[1]
55 mask = mask[:, :max_length]
56 mask_a = mask.unsqueeze(1).expand([-1, max_length, -1])
57 mask_b = mask.unsqueeze(2).expand([-1, -1, max_length])
58 mask = mask_a * mask_b
59 mask = torch.triu(mask).unsqueeze(3).expand([-1, -1, -1, self.args.class_num])
60
61 '''save all logits'''
62 logits_list = []
63 logits = self._cls_logits(features)
64 logits_list.append(logits)
65
66 for i in range(k):

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected