sundea/text-classification
0
1# coding: UTF-82import time3import torch4import numpy as np5from train_eval import train, init_network6from importlib import import_module7import argparse8 9parser = argparse.ArgumentParser(description='Chinese Text Classification')10parser.add_argument('--model', type=str, required=True, help='choose a model: TextCNN, TextRNN, FastText, TextRCNN, TextRNN_Att, DPCNN, Transformer')11parser.add_argument('--embedding', default='pre_trained', type=str, help='random or pre_trained')12parser.add_argument('--word', default=False, type=bool, help='True for word, False for char')13args = parser.parse_args()14 15 16if __name__ == '__main__':17 dataset = 'THUCNews' # 数据集18 19 # 搜狗新闻:embedding_SougouNews.npz, 腾讯:embedding_Tencent.npz, 随机初始化:random20 embedding = 'embedding_SougouNews.npz'21 if args.embedding == 'random':22 embedding = 'random'23 model_name = args.model # 'TextRCNN' # TextCNN, TextRNN, FastText, TextRCNN, TextRNN_Att, DPCNN, Transformer24 if model_name == 'FastText':25 from utils_fasttext import build_dataset, build_iterator, get_time_dif26 embedding = 'random'27 else:28 from utils import build_dataset, build_iterator, get_time_dif29 30 x = import_module('models.' + model_name)31 config = x.Config(dataset, embedding)32 np.random.seed(1)33 torch.manual_seed(1)34 torch.cuda.manual_seed_all(1)35 torch.backends.cudnn.deterministic = True # 保证每次结果一样36 37 start_time = time.time()38 print("Loading data...")39 vocab, train_data, dev_data, test_data = build_dataset(config, args.word)40 train_iter = build_iterator(train_data, config)41 dev_iter = build_iterator(dev_data, config)42 test_iter = build_iterator(test_data, config)43 time_dif = get_time_dif(start_time)44 print("Time usage:", time_dif)45 46 # train47 config.n_vocab = len(vocab)48 model = x.Model(config).to(config.device)49 if model_name != 'Transformer':50 init_network(model)51 print(model.parameters)52 train(config, model, train_iter, dev_iter, test_iter)53 