MCPcopy Create free account
hub / github.com/NJUNLP/GTS / train

Function train

code/BertModel/main.py:16–60  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

14
15
16def 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
63def eval(model, dataset, args):

Callers 1

main.pyFile · 0.70

Calls 5

get_batchMethod · 0.95
load_data_instancesFunction · 0.90
DataIteratorClass · 0.90
MultiInferBertClass · 0.90
evalFunction · 0.70

Tested by

no test coverage detected