mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
MP-CNN: add trainer and evaluator for WikiQA (#80)
This commit is contained in:
+3
-3
@@ -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_):
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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']
|
||||
@@ -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']
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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))
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user