| 41 | |
| 42 | |
| 43 | class Model(nn.Module): |
| 44 | def __init__(self, config): |
| 45 | super(Model, self).__init__() |
| 46 | if config.embedding_pretrained is not None: |
| 47 | self.embedding = nn.Embedding.from_pretrained(config.embedding_pretrained, freeze=False) |
| 48 | else: |
| 49 | self.embedding = nn.Embedding(config.n_vocab, config.embed, padding_idx=config.n_vocab - 1) |
| 50 | self.embedding_ngram2 = nn.Embedding(config.n_gram_vocab, config.embed) |
| 51 | self.embedding_ngram3 = nn.Embedding(config.n_gram_vocab, config.embed) |
| 52 | self.dropout = nn.Dropout(config.dropout) |
| 53 | self.fc1 = nn.Linear(config.embed * 3, config.hidden_size) |
| 54 | # self.dropout2 = nn.Dropout(config.dropout) |
| 55 | self.fc2 = nn.Linear(config.hidden_size, config.num_classes) |
| 56 | |
| 57 | def forward(self, x): |
| 58 | |
| 59 | out_word = self.embedding(x[0]) |
| 60 | out_bigram = self.embedding_ngram2(x[2]) |
| 61 | out_trigram = self.embedding_ngram3(x[3]) |
| 62 | out = torch.cat((out_word, out_bigram, out_trigram), -1) |
| 63 | |
| 64 | out = out.mean(dim=1) |
| 65 | out = self.dropout(out) |
| 66 | out = self.fc1(out) |
| 67 | out = F.relu(out) |
| 68 | out = self.fc2(out) |
| 69 | return out |
nothing calls this directly
no outgoing calls
no test coverage detected