| 19 | |
| 20 | |
| 21 | class TextCNN(Classifier): |
| 22 | def __init__(self, dataset, config): |
| 23 | super(TextCNN, self).__init__(dataset, config) |
| 24 | |
| 25 | self.kernel_sizes = config.TextCNN.kernel_sizes |
| 26 | self.convs = torch.nn.ModuleList() |
| 27 | for kernel_size in self.kernel_sizes: |
| 28 | self.convs.append(torch.nn.Conv1d( |
| 29 | config.embedding.dimension, config.TextCNN.num_kernels, |
| 30 | kernel_size, padding=kernel_size - 1)) |
| 31 | |
| 32 | self.top_k = self.config.TextCNN.top_k_max_pooling |
| 33 | hidden_size = len(config.TextCNN.kernel_sizes) * \ |
| 34 | config.TextCNN.num_kernels * self.top_k |
| 35 | self.linear = torch.nn.Linear(hidden_size, len(dataset.label_map)) |
| 36 | self.dropout = torch.nn.Dropout(p=config.train.hidden_layer_dropout) |
| 37 | |
| 38 | def get_parameter_optimizer_dict(self): |
| 39 | params = list() |
| 40 | params.append({'params': self.token_embedding.parameters()}) |
| 41 | params.append({'params': self.char_embedding.parameters()}) |
| 42 | params.append({'params': self.convs.parameters()}) |
| 43 | params.append({'params': self.linear.parameters()}) |
| 44 | return params |
| 45 | |
| 46 | def update_lr(self, optimizer, epoch): |
| 47 | """Update lr |
| 48 | """ |
| 49 | if epoch > self.config.train.num_epochs_static_embedding: |
| 50 | for param_group in optimizer.param_groups[:2]: |
| 51 | param_group["lr"] = self.config.optimizer.learning_rate |
| 52 | else: |
| 53 | for param_group in optimizer.param_groups[:2]: |
| 54 | param_group["lr"] = 0 |
| 55 | |
| 56 | def forward(self, batch): |
| 57 | if self.config.feature.feature_names[0] == "token": |
| 58 | embedding = self.token_embedding( |
| 59 | batch[cDataset.DOC_TOKEN].to(self.config.device)) |
| 60 | else: |
| 61 | embedding = self.char_embedding( |
| 62 | batch[cDataset.DOC_CHAR].to(self.config.device)) |
| 63 | embedding = embedding.transpose(1, 2) |
| 64 | pooled_outputs = [] |
| 65 | for i, conv in enumerate(self.convs): |
| 66 | #convolution = torch.nn.ReLU(conv(embedding)) |
| 67 | convolution = torch.nn.functional.relu(conv(embedding)) |
| 68 | pooled = torch.topk(convolution, self.top_k)[0].view( |
| 69 | convolution.size(0), -1) |
| 70 | pooled_outputs.append(pooled) |
| 71 | |
| 72 | doc_embedding = torch.cat(pooled_outputs, 1) |
| 73 | return self.dropout(self.linear(doc_embedding)) |
nothing calls this directly
no outgoing calls
no test coverage detected