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

Function train

code/NNModel/main.py:18–78  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

16
17
18def 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)

Callers 1

main.pyFile · 0.70

Calls 6

get_batchMethod · 0.95
load_data_instancesFunction · 0.90
DataIteratorClass · 0.90
MultiInferRNNModelClass · 0.90
MultiInferCNNModelClass · 0.90
evalFunction · 0.70

Tested by

no test coverage detected