1textCNN_param = {
2 'vocab_size': len(word2ind) + 1,
3 'embed_dim': 128, # 1 x 128 vector
4 'class_num': len(label_w2n),
5 "kernel_num": 16,
6 "kernel_size": [3, 4, 5],
7 "dropout": 0.5,
8 }
9 dataLoader_param = {
10 'batch_size': 128,
11 'shuffle': True,
12 }
1class textCNN(nn.Module):
2 def __init__(self, param):
3 super(textCNN, self).__init__()
4 ci = 1 # input chanel size
5 # kernel 卷积核
6 kernel_num = param['kernel_num'] # output chanel size
7 kernel_size = param['kernel_size']
8 vocab_size = param['vocab_size']
9 embed_dim = param['embed_dim'] # embedding dimension
10 dropout = param['dropout']
11 class_num = param['class_num']
12 self.param = param
13 # 把token随机向量化
14 self.embed = nn.Embedding(vocab_size, embed_dim, padding_idx=1)
15
16 # 三个不同长度的卷积
17 self.conv11 = nn.Conv2d(ci, kernel_num, (kernel_size[0], embed_dim))
18 self.conv12 = nn.Conv2d(ci, kernel_num, (kernel_size[1], embed_dim))
19 self.conv13 = nn.Conv2d(ci, kernel_num, (kernel_size[2], embed_dim))
20 # 三个不同长度的卷积
21
22 # increasing the ability of calculation by dropout
23 self.dropout = nn.Dropout(dropout)
24 self.fc1 = nn.Linear(len(kernel_size) * kernel_num, class_num)
25
26 def init_embed(self, embed_matrix):
27 self.embed.weight = nn.Parameter(torch.Tensor(embed_matrix))
28
29 @staticmethod
30 def conv_and_pool(x, conv):
31 # x: (batch, 1, sentence_length, )
32 x = conv(x)
33 # x: (batch, kernel_num, H_out, 1)
34 x = F.relu(x.squeeze(3))
35 # x: (batch, kernel_num, H_out)
36 x = F.max_pool1d(x, x.size(2)).squeeze(2)
37 # (batch, kernel_num)
38 return x
39
40 def forward(self, x):
41 # x: (batch, sentence_length)
42 x = self.embed(x)
43 # x: (batch, sentence_length, embed_dim)
44 # TODO init embed matrix with pre-trained
45 x = x.unsqueeze(1)
46 # x: (batch, 1, sentence_length, embed_dim)
47 x1 = self.conv_and_pool(x, self.conv11) # (batch, kernel_num)
48 x2 = self.conv_and_pool(x, self.conv12) # (batch, kernel_num)
49 x3 = self.conv_and_pool(x, self.conv13) # (batch, kernel_num)
50 x = torch.cat((x1, x2, x3), 1) # (batch, 3 * kernel_num)
51 x = self.dropout(x)
52 logit = F.log_softmax(self.fc1(x), dim=1)
53 return logit
1 # set the seed for ensuring reproducibility
2 seed = 3407
3 torch.cuda.manual_seed(seed)
4 torch.manual_seed(seed)
5 torch.backends.cudnn.deterministic = True
6 torch.backends.cudnn.benchmark = False
7
8 word2ind, ind2word = sen2inds.get_worddict('wordLabel.txt')
9 label_w2n, label_n2w = sen2inds.read_labelFile('data/label.txt')
10
11 textCNN_param = {
12 'vocab_size': len(word2ind) + 1,
13 'embed_dim': 128, # 1 x 128 vector
14 'class_num': len(label_w2n),
15 "kernel_num": 16,
16 "kernel_size": [3, 4, 5],
17 "dropout": 0.5,
18 }
19 dataLoader_param = {
20 'batch_size': 128,
21 'shuffle': True,
22 }
23
24 # # device = 'cuda:0' if torch.cuda.is_available() else 'cpu'
25 device = 'cpu'
26
27 # init dataset
28 print('init dataset...')
29 trainDataFile = 'traindata_vec.txt'
30 valDataFile = 'devdata_vec.txt'
31 train_dataset = textCNN_data(trainDataFile)
32 train_dataLoader = DataLoader(train_dataset,
33 batch_size=dataLoader_param['batch_size'],
34 shuffle=True)
35
36 val_dataset = textCNN_data(valDataFile)
37 val_dataLoader = DataLoader(val_dataset,
38 batch_size=dataLoader_param['batch_size'], # batch size 128
39 shuffle=False)
40
41 # init net
42 print('init net...')
43 net = textCNN(textCNN_param)
44 print(net)
45 net.to(device)
46 optimizer = torch.optim.Adam(net.parameters(), lr=0.001)
47 criterion = nn.CrossEntropyLoss()
48
49 print("training...")
50 net.train()
51 best_dev_acc = 0
52 for epoch in range(100):
53 for i, (clas, sentences) in enumerate(train_dataLoader):
54 out = net(sentences)
55 loss = criterion(out, clas)
56 optimizer.zero_grad()
57 loss.backward()
58 optimizer.step()
59 if (i + 1) % 10 == 0:
60 print("epoch:", epoch + 1, "step:", i + 1, "loss:", loss.item())
61
62 dev_acc = validation(model=net, val_dataLoader=val_dataLoader,
63 device=device)
64 if best_dev_acc < dev_acc:
65 best_dev_acc = dev_acc
66 print("save model...")
67 torch.save(net.state_dict(), "textcnn.bin")
68 print("epoch:", epoch + 1, "step:", i + 1, "loss:", loss.item())
69 print("best dev acc %.4f dev acc %.4f" % (best_dev_acc, dev_acc))
70