MP-CNN: add trainer and evaluator for WikiQA (#80)

This commit is contained in:
Michael Tu
2017-11-06 20:37:40 -05:00
committed by GitHub
parent 4132874bad
commit 2344354cbf
11 changed files with 148 additions and 106 deletions
+3 -3
View File
@@ -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_):
+10
View File
@@ -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))
+3 -1
View File
@@ -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
+34
View File
@@ -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']
+2 -29
View File
@@ -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']
+7
View File
@@ -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)
+1 -1
View File
@@ -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')
+3 -1
View File
@@ -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
+76
View File
@@ -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))
+2 -71
View File
@@ -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))
+7
View File
@@ -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)