diff --git a/.gitignore b/.gitignore index 2318c3e..e371d91 100644 --- a/.gitignore +++ b/.gitignore @@ -10,3 +10,5 @@ trec_eval-9.0.5 *.pt text/ kim_cnn/data +.results +.qrel diff --git a/common/evaluation.py b/common/evaluation.py index d3d80e7..18248ee 100644 --- a/common/evaluation.py +++ b/common/evaluation.py @@ -22,7 +22,7 @@ class EvaluatorFactory(object): } @staticmethod - def get_evaluator(dataset_cls, model, embedding, data_loader, batch_size, device, nce=False): + def get_evaluator(dataset_cls, model, embedding, data_loader, batch_size, device, nce=False, keep_results=False): if data_loader is None: return None @@ -38,5 +38,5 @@ class EvaluatorFactory(object): raise ValueError('{} is not implemented.'.format(dataset_cls)) return evaluator_map[dataset_cls.NAME]( - dataset_cls, model, embedding, data_loader, batch_size, device + dataset_cls, model, embedding, data_loader, batch_size, device, keep_results ) diff --git a/common/evaluators/evaluator.py b/common/evaluators/evaluator.py index 7318bec..b739c6b 100644 --- a/common/evaluators/evaluator.py +++ b/common/evaluators/evaluator.py @@ -3,13 +3,14 @@ class Evaluator(object): Evaluates a model on a Dataset, using metrics specific to the Dataset. """ - def __init__(self, dataset_cls, model, embedding, data_loader, batch_size, device): + def __init__(self, dataset_cls, model, embedding, data_loader, batch_size, device, keep_results=False): self.dataset_cls = dataset_cls self.model = model self.embedding = embedding self.data_loader = data_loader self.batch_size = batch_size self.device = device + self.keep_results = keep_results def get_sentence_embeddings(self, batch): sent1 = self.embedding(batch.sentence_1).transpose(1, 2) diff --git a/common/evaluators/qa_evaluator.py b/common/evaluators/qa_evaluator.py index 46133eb..f5cd186 100644 --- a/common/evaluators/qa_evaluator.py +++ b/common/evaluators/qa_evaluator.py @@ -28,7 +28,9 @@ class QAEvaluator(Evaluator): qids = list(map(lambda n: int(round(n * 10, 0)) / 10, qids)) - mean_average_precision, mean_reciprocal_rank = get_map_mrr(qids, predictions, true_labels, self.data_loader.device) + mean_average_precision, mean_reciprocal_rank = get_map_mrr(qids, predictions, true_labels, + self.data_loader.device, + keep_results=self.keep_results) test_cross_entropy_loss /= len(batch.dataset.examples) return [mean_average_precision, mean_reciprocal_rank, test_cross_entropy_loss], ['map', 'mrr', 'cross entropy loss'] diff --git a/mp_cnn/__main__.py b/mp_cnn/__main__.py index 2c39516..ddbf9e6 100644 --- a/mp_cnn/__main__.py +++ b/mp_cnn/__main__.py @@ -29,8 +29,9 @@ def get_logger(): return logger -def evaluate_dataset(split_name, dataset_cls, model, embedding, loader, batch_size, device): - saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, loader, batch_size, device) +def evaluate_dataset(split_name, dataset_cls, model, embedding, loader, batch_size, device, keep_results=False): + saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, loader, batch_size, device, + keep_results=keep_results) scores, metric_names = saved_model_evaluator.get_scores() logger.info('Evaluation metrics for {}'.format(split_name)) logger.info('\t'.join([' '] + metric_names)) @@ -79,6 +80,9 @@ if __name__ == '__main__': parser.add_argument('--tensorboard', action='store_true', default=False, help='use TensorBoard to visualize training (default: false)') parser.add_argument('--run-label', type=str, help='label to describe run') + parser.add_argument('--keep-results', action='store_true', + help='store the output score and qrel files into disk for the test set') + args = parser.parse_args() device = torch.device(f'cuda:{args.device}' if torch.cuda.is_available() and args.device >= 0 else 'cpu') @@ -114,9 +118,12 @@ if __name__ == '__main__': else: raise ValueError('optimizer not recognized: it should be either adam or sgd') - train_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, train_loader, args.batch_size, args.device) - test_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, test_loader, args.batch_size, args.device) - dev_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, dev_loader, args.batch_size, args.device) + train_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, train_loader, args.batch_size, + args.device) + test_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, test_loader, args.batch_size, + args.device) + dev_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, embedding, dev_loader, args.batch_size, + args.device) trainer_config = { 'optimizer': optimizer, @@ -147,4 +154,4 @@ if __name__ == '__main__': model.load_state_dict(state_dict) if dev_loader: evaluate_dataset('dev', dataset_cls, model, embedding, dev_loader, args.batch_size, args.device) - evaluate_dataset('test', dataset_cls, model, embedding, test_loader, args.batch_size, args.device) + evaluate_dataset('test', dataset_cls, model, embedding, test_loader, args.batch_size, args.device, args.keep_results) diff --git a/nce/nce_pairwise_mp/evaluators/trecqa_evaluator.py b/nce/nce_pairwise_mp/evaluators/trecqa_evaluator.py index ba08b4f..2351da7 100644 --- a/nce/nce_pairwise_mp/evaluators/trecqa_evaluator.py +++ b/nce/nce_pairwise_mp/evaluators/trecqa_evaluator.py @@ -3,5 +3,5 @@ from nce.nce_pairwise_mp.evaluators.qa_evaluator import QAEvaluator class TRECQAEvaluatorNCE(QAEvaluator): - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(TRECQAEvaluatorNCE, self).__init__(dataset_cls, model, data_loader, batch_size, device) + def __init__(self, dataset_cls, model, data_loader, batch_size, device, keep_results=False): + super(TRECQAEvaluatorNCE, self).__init__(dataset_cls, model, data_loader, batch_size, device, keep_results) diff --git a/nce/nce_pairwise_mp/evaluators/wikiqa_evaluator.py b/nce/nce_pairwise_mp/evaluators/wikiqa_evaluator.py index 0263e28..03db889 100644 --- a/nce/nce_pairwise_mp/evaluators/wikiqa_evaluator.py +++ b/nce/nce_pairwise_mp/evaluators/wikiqa_evaluator.py @@ -3,5 +3,5 @@ from nce.nce_pairwise_mp.evaluators.qa_evaluator import QAEvaluator class WikiQAEvaluatorNCE(QAEvaluator): - def __init__(self, dataset_cls, model, data_loader, batch_size, device): - super(WikiQAEvaluatorNCE, self).__init__(dataset_cls, model, data_loader, batch_size, device) + def __init__(self, dataset_cls, model, data_loader, batch_size, device, keep_results=False): + super(WikiQAEvaluatorNCE, self).__init__(dataset_cls, model, data_loader, batch_size, device, keep_results) diff --git a/utils/relevancy_metrics.py b/utils/relevancy_metrics.py index f617ffb..429ae0b 100644 --- a/utils/relevancy_metrics.py +++ b/utils/relevancy_metrics.py @@ -3,7 +3,7 @@ import subprocess import time -def get_map_mrr(qids, predictions, labels, device=0): +def get_map_mrr(qids, predictions, labels, device=0, keep_results=False): """ Get the map and mrr using the trec_eval utility. qids, predictions, labels should have the same length. @@ -30,7 +30,11 @@ def get_map_mrr(qids, predictions, labels, device=0): mean_average_precision = float(trec_out_lines[0].split('\t')[-1]) mean_reciprocal_rank = float(trec_out_lines[1].split('\t')[-1]) - os.remove(qrel_fname) - os.remove(results_fname) + if keep_results: + print("Saving prediction file to {}".format(results_fname)) + print("Saving qrel file to {}".format(qrel_fname)) + else: + os.remove(results_fname) + os.remove(qrel_fname) return mean_average_precision, mean_reciprocal_rank