(model, dataset, args)
| 61 | |
| 62 | |
| 63 | def 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 | |
| 101 | def test(args): |
no test coverage detected