(args)
| 16 | |
| 17 | |
| 18 | def train(args): |
| 19 | |
| 20 | # load double embedding |
| 21 | word2index = json.load(open(args.prefix + 'doubleembedding/word_idx.json')) |
| 22 | general_embedding = numpy.load(args.prefix +'doubleembedding/gen.vec.npy') |
| 23 | general_embedding = torch.from_numpy(general_embedding) |
| 24 | domain_embedding = numpy.load(args.prefix +'doubleembedding/'+args.dataset+'_emb.vec.npy') |
| 25 | domain_embedding = torch.from_numpy(domain_embedding) |
| 26 | |
| 27 | # load dataset |
| 28 | train_sentence_packs = json.load(open(args.prefix + args.dataset + '/train.json')) |
| 29 | random.shuffle(train_sentence_packs) |
| 30 | dev_sentence_packs = json.load(open(args.prefix + args.dataset + '/dev.json')) |
| 31 | |
| 32 | instances_train = load_data_instances(train_sentence_packs, word2index, args) |
| 33 | instances_dev = load_data_instances(dev_sentence_packs, word2index, args) |
| 34 | |
| 35 | random.shuffle(instances_train) |
| 36 | trainset = DataIterator(instances_train, args) |
| 37 | devset = DataIterator(instances_dev, args) |
| 38 | |
| 39 | if not os.path.exists(args.model_dir): |
| 40 | os.makedirs(args.model_dir) |
| 41 | |
| 42 | # build model |
| 43 | if args.model == 'bilstm': |
| 44 | model = MultiInferRNNModel(general_embedding, domain_embedding, args).to(args.device) |
| 45 | elif args.model == 'cnn': |
| 46 | model = MultiInferCNNModel(general_embedding, domain_embedding, args).to(args.device) |
| 47 | |
| 48 | parameters = list(model.parameters()) |
| 49 | parameters = filter(lambda x: x.requires_grad, parameters) |
| 50 | optimizer = torch.optim.Adam(parameters, lr=args.lr) |
| 51 | |
| 52 | # training |
| 53 | best_joint_f1 = 0 |
| 54 | best_joint_epoch = 0 |
| 55 | for i in range(args.epochs): |
| 56 | print('Epoch:{}'.format(i)) |
| 57 | for j in trange(trainset.batch_count): |
| 58 | _, sentence_tokens, lengths, masks, aspect_tags, _, tags = trainset.get_batch(j) |
| 59 | predictions = model(sentence_tokens, lengths, masks) |
| 60 | |
| 61 | loss = 0. |
| 62 | tags_flatten = tags[:, :lengths[0], :lengths[0]].reshape([-1]) |
| 63 | for k in range(len(predictions)): |
| 64 | prediction_flatten = predictions[k].reshape([-1, predictions[k].shape[3]]) |
| 65 | loss = loss + F.cross_entropy(prediction_flatten, tags_flatten, ignore_index=-1) |
| 66 | |
| 67 | optimizer.zero_grad() |
| 68 | loss.backward() |
| 69 | optimizer.step() |
| 70 | |
| 71 | joint_precision, joint_recall, joint_f1 = eval(model, devset, args) |
| 72 | |
| 73 | if joint_f1 > best_joint_f1: |
| 74 | model_path = args.model_dir + args.model + args.task + '.pt' |
| 75 | torch.save(model, model_path) |
no test coverage detected