(self, sentence_tokens, lengths, mask)
| 84 | return logits_list |
| 85 | |
| 86 | def forward(self, sentence_tokens, lengths, mask): |
| 87 | embedding = self._get_embedding(sentence_tokens, mask) |
| 88 | lstm_feature = self._lstm_feature(embedding, lengths) |
| 89 | |
| 90 | # self attention |
| 91 | lstm_feature_attention = self.attention_layer(lstm_feature, lstm_feature, mask[:,:lengths[0]]) |
| 92 | #lstm_feature_attention = self.attention_layer.forward_perceptron(lstm_feature, lstm_feature, mask[:, :lengths[0]]) |
| 93 | lstm_feature = lstm_feature + lstm_feature_attention |
| 94 | |
| 95 | lstm_feature = lstm_feature.unsqueeze(2).expand([-1,-1, lengths[0], -1]) |
| 96 | lstm_feature_T = lstm_feature.transpose(1, 2) |
| 97 | features = torch.cat([lstm_feature, lstm_feature_T], dim=3) |
| 98 | |
| 99 | logits = self.multi_hops(features, lengths, mask, self.args.nhops) |
| 100 | return [logits[-1]] |
| 101 | |
| 102 | |
| 103 | class MultiInferCNNModel(torch.nn.Module): |
no test coverage detected