diff --git a/datasets/wikiqa.py b/datasets/wikiqa.py index 9de669b..2bf98e3 100644 --- a/datasets/wikiqa.py +++ b/datasets/wikiqa.py @@ -10,7 +10,7 @@ from datasets.castor_dataset import CastorPairDataset from datasets.idf_utils import get_pairwise_word_to_doc_freq, get_pairwise_overlap_features -class WIKIQA(CastorPairDataset): +class WikiQA(CastorPairDataset): NAME = 'wikiqa' NUM_CLASSES = 2 ID_FIELD = Field(sequential=False, tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True) @@ -26,11 +26,11 @@ class WIKIQA(CastorPairDataset): """ Create a WIKIQA dataset instance """ - super(WIKIQA, self).__init__(path) + super(WikiQA, self).__init__(path) @classmethod def splits(cls, path, train='train', validation='dev', test='test', **kwargs): - return super(WIKIQA, cls).splits(path, train=train, validation=validation, test=test, **kwargs) + return super(WikiQA, cls).splits(path, train=train, validation=validation, test=test, **kwargs) @classmethod def iters(cls, path, vectors_name, vectors_cache, batch_size=64, shuffle=True, device=0, vectors=None, unk_init=torch.Tensor.zero_): diff --git a/mp_cnn/dataset.py b/mp_cnn/dataset.py index a1239df..ebe3d04 100644 --- a/mp_cnn/dataset.py +++ b/mp_cnn/dataset.py @@ -6,6 +6,7 @@ import torch.nn as nn from datasets.sick import SICK from datasets.msrvid import MSRVID from datasets.trecqa import TRECQA +from datasets.wikiqa import WikiQA class UnknownWordVecCache(object): @@ -53,6 +54,15 @@ class MPCNNDatasetFactory(object): embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) embedding.weight = nn.Parameter(TRECQA.TEXT_FIELD.vocab.vectors) return TRECQA, embedding, train_loader, test_loader, dev_loader + elif dataset_name == 'wikiqa': + if not os.path.exists('../utils/trec_eval-9.0.5/trec_eval'): + raise FileNotFoundError('TrecQA requires the trec_eval tool to run. Please run get_trec_eval.sh inside Castor/utils (as working directory) before continuing.') + dataset_root = os.path.join(os.pardir, os.pardir, 'data', 'WikiQA/') + train_loader, dev_loader, test_loader = WikiQA.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) + embedding_dim = WikiQA.TEXT_FIELD.vocab.vectors.size() + embedding = nn.Embedding(embedding_dim[0], embedding_dim[1]) + embedding.weight = nn.Parameter(WikiQA.TEXT_FIELD.vocab.vectors) + return WikiQA, embedding, train_loader, test_loader, dev_loader else: raise ValueError('{} is not a valid dataset.'.format(dataset_name)) diff --git a/mp_cnn/evaluation.py b/mp_cnn/evaluation.py index 47fc05d..1dfa9c1 100644 --- a/mp_cnn/evaluation.py +++ b/mp_cnn/evaluation.py @@ -1,6 +1,7 @@ from mp_cnn.evaluators.sick_evaluator import SICKEvaluator from mp_cnn.evaluators.msrvid_evaluator import MSRVIDEvaluator from mp_cnn.evaluators.trecqa_evaluator import TRECQAEvaluator +from mp_cnn.evaluators.wikiqa_evaluator import WikiQAEvaluator class MPCNNEvaluatorFactory(object): @@ -10,7 +11,8 @@ class MPCNNEvaluatorFactory(object): evaluator_map = { 'sick': SICKEvaluator, 'msrvid': MSRVIDEvaluator, - 'trecqa': TRECQAEvaluator + 'trecqa': TRECQAEvaluator, + 'wikiqa': WikiQAEvaluator } @staticmethod diff --git a/mp_cnn/evaluators/qa_evaluator.py b/mp_cnn/evaluators/qa_evaluator.py new file mode 100644 index 0000000..a42266f --- /dev/null +++ b/mp_cnn/evaluators/qa_evaluator.py @@ -0,0 +1,34 @@ +import torch.nn.functional as F + +from mp_cnn.evaluators.evaluator import Evaluator +from utils.relevancy_metrics import get_map_mrr + + +class QAEvaluator(Evaluator): + + def __init__(self, dataset_cls, model, data_loader, batch_size, device): + super(QAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) + + def get_scores(self): + self.model.eval() + test_cross_entropy_loss = 0 + qids = [] + true_labels = [] + predictions = [] + + for batch in self.data_loader: + qids.extend(batch.id.data.cpu().numpy()) + output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) + test_cross_entropy_loss += F.cross_entropy(output, batch.label, size_average=False).data[0] + + true_labels.extend(batch.label.data.cpu().numpy()) + predictions.extend(output.data.exp()[:, 1].cpu().numpy()) + + del output + + 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) + test_cross_entropy_loss /= len(batch.dataset.examples) + + return [test_cross_entropy_loss, mean_average_precision, mean_reciprocal_rank], ['cross entropy loss', 'map', 'mrr'] diff --git a/mp_cnn/evaluators/trecqa_evaluator.py b/mp_cnn/evaluators/trecqa_evaluator.py index 9391745..baae709 100644 --- a/mp_cnn/evaluators/trecqa_evaluator.py +++ b/mp_cnn/evaluators/trecqa_evaluator.py @@ -1,34 +1,7 @@ -import torch.nn.functional as F - -from mp_cnn.evaluators.evaluator import Evaluator -from utils.relevancy_metrics import get_map_mrr +from mp_cnn.evaluators.qa_evaluator import QAEvaluator -class TRECQAEvaluator(Evaluator): +class TRECQAEvaluator(QAEvaluator): def __init__(self, dataset_cls, model, data_loader, batch_size, device): super(TRECQAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) - - def get_scores(self): - self.model.eval() - test_cross_entropy_loss = 0 - qids = [] - true_labels = [] - predictions = [] - - for batch in self.data_loader: - qids.extend(batch.id.data.cpu().numpy()) - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - test_cross_entropy_loss += F.cross_entropy(output, batch.label, size_average=False).data[0] - - true_labels.extend(batch.label.data.cpu().numpy()) - predictions.extend(output.data.exp()[:, 1].cpu().numpy()) - - del output - - 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) - test_cross_entropy_loss /= len(batch.dataset.examples) - - return [test_cross_entropy_loss, mean_average_precision, mean_reciprocal_rank], ['cross entropy loss', 'map', 'mrr'] diff --git a/mp_cnn/evaluators/wikiqa_evaluator.py b/mp_cnn/evaluators/wikiqa_evaluator.py new file mode 100644 index 0000000..82463e4 --- /dev/null +++ b/mp_cnn/evaluators/wikiqa_evaluator.py @@ -0,0 +1,7 @@ +from mp_cnn.evaluators.qa_evaluator import QAEvaluator + + +class WikiQAEvaluator(QAEvaluator): + + def __init__(self, dataset_cls, model, data_loader, batch_size, device): + super(WikiQAEvaluator, self).__init__(dataset_cls, model, data_loader, batch_size, device) diff --git a/mp_cnn/main.py b/mp_cnn/main.py index 53e33d9..a91a300 100644 --- a/mp_cnn/main.py +++ b/mp_cnn/main.py @@ -17,7 +17,7 @@ from mp_cnn.train import MPCNNTrainerFactory if __name__ == '__main__': parser = argparse.ArgumentParser(description='PyTorch implementation of Multi-Perspective CNN') parser.add_argument('model_outfile', help='file to save final model') - parser.add_argument('--dataset', help='dataset to use, one of [sick, msrvid, trecqa]', default='sick') + parser.add_argument('--dataset', help='dataset to use, one of [sick, msrvid, trecqa, wikiqa]', default='sick') parser.add_argument('--word-vectors-dir', help='word vectors directory', default=os.path.join(os.pardir, os.pardir, 'data', 'GloVe')) parser.add_argument('--word-vectors-file', help='word vectors filename', default='glove.840B.300d.txt') parser.add_argument('--skip-training', help='will load pre-trained model', action='store_true') diff --git a/mp_cnn/train.py b/mp_cnn/train.py index 039297b..be22982 100644 --- a/mp_cnn/train.py +++ b/mp_cnn/train.py @@ -1,6 +1,7 @@ from mp_cnn.trainers.sick_trainer import SICKTrainer from mp_cnn.trainers.msrvid_trainer import MSRVIDTrainer from mp_cnn.trainers.trecqa_trainer import TRECQATrainer +from mp_cnn.trainers.wikiqa_trainer import WikiQATrainer class MPCNNTrainerFactory(object): @@ -10,7 +11,8 @@ class MPCNNTrainerFactory(object): trainer_map = { 'sick': SICKTrainer, 'msrvid': MSRVIDTrainer, - 'trecqa': TRECQATrainer + 'trecqa': TRECQATrainer, + 'wikiqa': WikiQATrainer } @staticmethod diff --git a/mp_cnn/trainers/qa_trainer.py b/mp_cnn/trainers/qa_trainer.py new file mode 100644 index 0000000..44ea911 --- /dev/null +++ b/mp_cnn/trainers/qa_trainer.py @@ -0,0 +1,76 @@ +import time + +import torch +import torch.nn.functional as F +from torch.optim.lr_scheduler import ReduceLROnPlateau + +from mp_cnn.trainers.trainer import Trainer + + +class QATrainer(Trainer): + + def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): + super(QATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + + def train_epoch(self, epoch): + self.model.train() + total_loss = 0 + for batch_idx, batch in enumerate(self.train_loader): + self.optimizer.zero_grad() + output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) + loss = F.cross_entropy(output, batch.label, size_average=False) + total_loss += loss.data[0] + loss.backward() + self.optimizer.step() + if batch_idx % self.log_interval == 0: + self.logger.info('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( + epoch, min(batch_idx * self.batch_size, len(batch.dataset.examples)), + len(batch.dataset.examples), + 100. * batch_idx / (len(self.train_loader)), loss.data[0]) + ) + + average_loss, mean_average_precision, mean_reciprocal_rank = self.evaluate(self.train_evaluator, 'train') + + if self.use_tensorboard: + self.writer.add_scalar('{}/train/cross_entropy_loss'.format(self.train_loader.dataset.NAME), average_loss, epoch) + self.writer.add_scalar('{}/train/map'.format(self.train_loader.dataset.NAME), mean_average_precision, epoch) + self.writer.add_scalar('{}/train/mrr'.format(self.train_loader.dataset.NAME), mean_reciprocal_rank, epoch) + + return total_loss + + def train(self, epochs): + scheduler = ReduceLROnPlateau(self.optimizer, mode='max', factor=self.lr_reduce_factor, patience=self.patience) + epoch_times = [] + prev_loss = -1 + best_dev_score = -1 + for epoch in range(1, epochs + 1): + start = time.time() + self.logger.info('Epoch {} started...'.format(epoch)) + self.train_epoch(epoch) + + dev_scores = self.evaluate(self.dev_evaluator, 'dev') + new_loss, mean_average_precision, mean_reciprocal_rank = dev_scores + + if self.use_tensorboard: + self.writer.add_scalar('{}/lr'.format(self.train_loader.dataset.NAME), self.optimizer.param_groups[0]['lr'], epoch) + self.writer.add_scalar('{}/dev/cross_entropy_loss'.format(self.train_loader.dataset.NAME), new_loss, epoch) + self.writer.add_scalar('{}/dev/map'.format(self.train_loader.dataset.NAME), mean_average_precision, epoch) + self.writer.add_scalar('{}/dev/mrr'.format(self.train_loader.dataset.NAME), mean_reciprocal_rank, epoch) + + end = time.time() + duration = end - start + self.logger.info('Epoch {} finished in {:.2f} minutes'.format(epoch, duration / 60)) + epoch_times.append(duration) + + if dev_scores[0] > best_dev_score: + best_dev_score = dev_scores[0] + torch.save(self.model, self.model_outfile) + + if abs(prev_loss - new_loss) <= 0.0002: + self.logger.info('Early stopping. Loss changed by less than 0.0002.') + break + + prev_loss = new_loss + scheduler.step(dev_scores[0]) + + self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) diff --git a/mp_cnn/trainers/trecqa_trainer.py b/mp_cnn/trainers/trecqa_trainer.py index 69e6c5c..c49785c 100644 --- a/mp_cnn/trainers/trecqa_trainer.py +++ b/mp_cnn/trainers/trecqa_trainer.py @@ -1,76 +1,7 @@ -import time - -import torch -import torch.nn.functional as F -from torch.optim.lr_scheduler import ReduceLROnPlateau - -from mp_cnn.trainers.trainer import Trainer +from mp_cnn.trainers.qa_trainer import QATrainer -class TRECQATrainer(Trainer): +class TRECQATrainer(QATrainer): def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): super(TRECQATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) - - def train_epoch(self, epoch): - self.model.train() - total_loss = 0 - for batch_idx, batch in enumerate(self.train_loader): - self.optimizer.zero_grad() - output = self.model(batch.sentence_1, batch.sentence_2, batch.ext_feats) - loss = F.cross_entropy(output, batch.label, size_average=False) - total_loss += loss.data[0] - loss.backward() - self.optimizer.step() - if batch_idx % self.log_interval == 0: - self.logger.info('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( - epoch, min(batch_idx * self.batch_size, len(batch.dataset.examples)), - len(batch.dataset.examples), - 100. * batch_idx / (len(self.train_loader)), loss.data[0]) - ) - - average_loss, mean_average_precision, mean_reciprocal_rank = self.evaluate(self.train_evaluator, 'train') - - if self.use_tensorboard: - self.writer.add_scalar('trecqa/train/cross_entropy_loss', average_loss, epoch) - self.writer.add_scalar('trecqa/train/map', mean_average_precision, epoch) - self.writer.add_scalar('trecqa/train/mrr', mean_reciprocal_rank, epoch) - - return total_loss - - def train(self, epochs): - scheduler = ReduceLROnPlateau(self.optimizer, mode='max', factor=self.lr_reduce_factor, patience=self.patience) - epoch_times = [] - prev_loss = -1 - best_dev_score = -1 - for epoch in range(1, epochs + 1): - start = time.time() - self.logger.info('Epoch {} started...'.format(epoch)) - self.train_epoch(epoch) - - dev_scores = self.evaluate(self.dev_evaluator, 'dev') - new_loss, mean_average_precision, mean_reciprocal_rank = dev_scores - - if self.use_tensorboard: - self.writer.add_scalar('trecqa/lr', self.optimizer.param_groups[0]['lr'], epoch) - self.writer.add_scalar('trecqa/dev/cross_entropy_loss', new_loss, epoch) - self.writer.add_scalar('trecqa/dev/map', mean_average_precision, epoch) - self.writer.add_scalar('trecqa/dev/mrr', mean_reciprocal_rank, epoch) - - end = time.time() - duration = end - start - self.logger.info('Epoch {} finished in {:.2f} minutes'.format(epoch, duration / 60)) - epoch_times.append(duration) - - if dev_scores[0] > best_dev_score: - best_dev_score = dev_scores[0] - torch.save(self.model, self.model_outfile) - - if abs(prev_loss - new_loss) <= 0.0002: - self.logger.info('Early stopping. Loss changed by less than 0.0002.') - break - - prev_loss = new_loss - scheduler.step(dev_scores[0]) - - self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) diff --git a/mp_cnn/trainers/wikiqa_trainer.py b/mp_cnn/trainers/wikiqa_trainer.py new file mode 100644 index 0000000..801ae3e --- /dev/null +++ b/mp_cnn/trainers/wikiqa_trainer.py @@ -0,0 +1,7 @@ +from mp_cnn.trainers.qa_trainer import QATrainer + + +class WikiQATrainer(QATrainer): + + def __init__(self, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None): + super(WikiQATrainer, self).__init__(model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator)