MCPcopy Create free account
hub / github.com/Tencent/NeuralNLP-NeuralClassifier / TextCNN

Class TextCNN

model/classification/textcnn.py:21–73  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

19
20
21class 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))

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected