(self, batch)
| 107 | param_group["lr"] = 0 |
| 108 | |
| 109 | def forward(self, batch): |
| 110 | doc_embedding = self.token_embedding( |
| 111 | batch[cDataset.DOC_TOKEN].to(self.config.device), |
| 112 | batch[cDataset.DOC_TOKEN_OFFSET].to(self.config.device)) |
| 113 | length = batch[cDataset.DOC_TOKEN_LEN].to(self.config.device) |
| 114 | if self.config.feature.token_ngram > 1: |
| 115 | doc_embedding += self.token_ngram_embedding( |
| 116 | batch[cDataset.DOC_TOKEN_NGRAM].to(self.config.device), |
| 117 | batch[cDataset.DOC_TOKEN_NGRAM_OFFSET].to(self.config.device)) |
| 118 | length += batch[cDataset.DOC_TOKEN_NGRAM_LEN].to(self.config.device) |
| 119 | if "keyword" in self.config.feature.feature_names: |
| 120 | doc_embedding += self.keyword_embedding( |
| 121 | batch[cDataset.DOC_KEYWORD].to(self.config.device), |
| 122 | batch[cDataset.DOC_KEYWORD_OFFSET].to(self.config.device)) |
| 123 | length += batch[cDataset.DOC_KEYWORD_LEN].to(self.config.device) |
| 124 | if "topic" in self.config.feature.feature_names: |
| 125 | doc_embedding += self.topic_embedding( |
| 126 | batch[cDataset.DOC_TOPIC].to(self.config.device), |
| 127 | batch[cDataset.DOC_TOPIC_OFFSET].to(self.config.device)) |
| 128 | length += batch[cDataset.DOC_TOPIC_LEN].to(self.config.device) |
| 129 | |
| 130 | doc_embedding /= length.resize_(doc_embedding.size()[0], 1) |
| 131 | doc_embedding = self.dropout(doc_embedding) |
| 132 | return self.linear(doc_embedding) |
nothing calls this directly
no outgoing calls
no test coverage detected