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

Function eval

code/BertModel/main.py:63–98  ·  view source on GitHub ↗
(model, dataset, args)

Source from the content-addressed store, hash-verified

61
62
63def eval(model, dataset, args):
64 model.eval()
65 with torch.no_grad():
66 all_ids = []
67 all_preds = []
68 all_labels = []
69 all_lengths = []
70 all_sens_lengths = []
71 all_token_ranges = []
72 for i in range(dataset.batch_count):
73 sentence_ids, tokens, lengths, masks, sens_lens, token_ranges, aspect_tags, tags = dataset.get_batch(i)
74 preds = model(tokens, masks)
75 preds = torch.argmax(preds, dim=3)
76 all_preds.append(preds)
77 all_labels.append(tags)
78 all_lengths.append(lengths)
79 all_sens_lengths.extend(sens_lens)
80 all_token_ranges.extend(token_ranges)
81 all_ids.extend(sentence_ids)
82
83 all_preds = torch.cat(all_preds, dim=0).cpu().tolist()
84 all_labels = torch.cat(all_labels, dim=0).cpu().tolist()
85 all_lengths = torch.cat(all_lengths, dim=0).cpu().tolist()
86
87 metric = utils.Metric(args, all_preds, all_labels, all_lengths, all_sens_lengths, all_token_ranges, ignore_index=-1)
88 precision, recall, f1 = metric.score_uniontags()
89 aspect_results = metric.score_aspect()
90 opinion_results = metric.score_opinion()
91 print('Aspect term\tP:{:.5f}\tR:{:.5f}\tF1:{:.5f}'.format(aspect_results[0], aspect_results[1],
92 aspect_results[2]))
93 print('Opinion term\tP:{:.5f}\tR:{:.5f}\tF1:{:.5f}'.format(opinion_results[0], opinion_results[1],
94 opinion_results[2]))
95 print(args.task + '\tP:{:.5f}\tR:{:.5f}\tF1:{:.5f}\n'.format(precision, recall, f1))
96
97 model.train()
98 return precision, recall, f1
99
100
101def test(args):

Callers 2

trainFunction · 0.70
testFunction · 0.70

Calls 4

score_uniontagsMethod · 0.95
score_aspectMethod · 0.95
score_opinionMethod · 0.95
get_batchMethod · 0.45

Tested by

no test coverage detected