| 7 | |
| 8 | |
| 9 | class 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): |