(args)
| 14 | |
| 15 | |
| 16 | def train(args): |
| 17 | |
| 18 | # load dataset |
| 19 | train_sentence_packs = json.load(open(args.prefix + args.dataset + '/train.json')) |
| 20 | random.shuffle(train_sentence_packs) |
| 21 | dev_sentence_packs = json.load(open(args.prefix + args.dataset + '/dev.json')) |
| 22 | instances_train = load_data_instances(train_sentence_packs, args) |
| 23 | instances_dev = load_data_instances(dev_sentence_packs, args) |
| 24 | random.shuffle(instances_train) |
| 25 | trainset = DataIterator(instances_train, args) |
| 26 | devset = DataIterator(instances_dev, args) |
| 27 | |
| 28 | if not os.path.exists(args.model_dir): |
| 29 | os.makedirs(args.model_dir) |
| 30 | model = MultiInferBert(args).to(args.device) |
| 31 | |
| 32 | optimizer = torch.optim.Adam([ |
| 33 | {'params': model.bert.parameters(), 'lr': 5e-5}, |
| 34 | {'params': model.cls_linear.parameters()} |
| 35 | ], lr=5e-5) |
| 36 | |
| 37 | best_joint_f1 = 0 |
| 38 | best_joint_epoch = 0 |
| 39 | for i in range(args.epochs): |
| 40 | print('Epoch:{}'.format(i)) |
| 41 | for j in trange(trainset.batch_count): |
| 42 | _, tokens, lengths, masks, _, _, aspect_tags, tags = trainset.get_batch(j) |
| 43 | preds = model(tokens, masks) |
| 44 | |
| 45 | preds_flatten = preds.reshape([-1, preds.shape[3]]) |
| 46 | tags_flatten = tags.reshape([-1]) |
| 47 | loss = F.cross_entropy(preds_flatten, tags_flatten, ignore_index=-1) |
| 48 | |
| 49 | optimizer.zero_grad() |
| 50 | loss.backward() |
| 51 | optimizer.step() |
| 52 | |
| 53 | joint_precision, joint_recall, joint_f1 = eval(model, devset, args) |
| 54 | |
| 55 | if joint_f1 > best_joint_f1: |
| 56 | model_path = args.model_dir + 'bert' + args.task + '.pt' |
| 57 | torch.save(model, model_path) |
| 58 | best_joint_f1 = joint_f1 |
| 59 | best_joint_epoch = i |
| 60 | print('best epoch: {}\tbest dev {} f1: {:.5f}\n\n'.format(best_joint_epoch, args.task, best_joint_f1)) |
| 61 | |
| 62 | |
| 63 | def eval(model, dataset, args): |
no test coverage detected