diff --git a/common/dataset.py b/common/dataset.py index 4faf4bc..5cffeeb 100644 --- a/common/dataset.py +++ b/common/dataset.py @@ -7,6 +7,7 @@ from datasets.sick import SICK from datasets.msrvid import MSRVID from datasets.trecqa import TRECQA from datasets.wikiqa import WikiQA +from datasets.pit2015 import PIT2015 class UnknownWordVecCache(object): @@ -55,6 +56,13 @@ class DatasetFactory(object): 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 = nn.Embedding.from_pretrained(WikiQA.TEXT_FIELD.vocab.vectors) return WikiQA, embedding, train_loader, test_loader, dev_loader + elif dataset_name == 'pit2015': + if not os.path.exists(os.path.join(castor_dir, utils_trecqa)): + 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(castor_dir, os.pardir, 'Castor-data', 'datasets', 'SemEval-PIT2015/') + train_loader, dev_loader, test_loader = PIT2015.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk) + embedding = nn.Embedding.from_pretrained(PIT2015.TEXT_FIELD.vocab.vectors) + return PIT2015, embedding, train_loader, test_loader, dev_loader else: raise ValueError('{} is not a valid dataset.'.format(dataset_name)) diff --git a/common/evaluation.py b/common/evaluation.py index e5eee09..2fa467d 100644 --- a/common/evaluation.py +++ b/common/evaluation.py @@ -3,6 +3,7 @@ from .evaluators.msrvid_evaluator import MSRVIDEvaluator from .evaluators.sst_evaluator import SSTEvaluator from .evaluators.trecqa_evaluator import TRECQAEvaluator from .evaluators.wikiqa_evaluator import WikiQAEvaluator +from .evaluators.pit2015_evaluator import PIT2015Evaluator from nce.nce_pairwise_mp.evaluators.trecqa_evaluator import TRECQAEvaluatorNCE from nce.nce_pairwise_mp.evaluators.wikiqa_evaluator import WikiQAEvaluatorNCE @@ -17,7 +18,8 @@ class EvaluatorFactory(object): 'SST-1': SSTEvaluator, 'SST-2': SSTEvaluator, 'trecqa': TRECQAEvaluator, - 'wikiqa': WikiQAEvaluator + 'wikiqa': WikiQAEvaluator, + 'pit2015': PIT2015Evaluator } evaluator_map_nce = { diff --git a/common/evaluators/pit2015_evaluator.py b/common/evaluators/pit2015_evaluator.py new file mode 100644 index 0000000..5ae2186 --- /dev/null +++ b/common/evaluators/pit2015_evaluator.py @@ -0,0 +1,32 @@ +import torch +import torch.nn.functional as F + +from .evaluator import Evaluator + +class PIT2015Evaluator(Evaluator): + + def get_scores(self): + self.model.eval() + self.data_loader.init_epoch() + n_dev_correct = 0 + total_loss = 0 + acc_total = 0 + rel_total = 0 + pre_total = 0 + for batch_idx, batch in enumerate(self.data_loader): + sent1, sent2 = self.get_sentence_embeddings(batch) + scores = self.model(sent1, sent2, batch.ext_feats, batch.dataset.word_to_doc_cnt, batch.sentence_1_raw, batch.sentence_2_raw) + prediction = torch.max(scores, 1)[1].view(batch.label.size()).data + gold_label = batch.label.data + n_dev_correct += (prediction == gold_label).sum().item() + acc_total += ((prediction == batch.label.data) * (prediction == 1)).sum().item() + total_loss += F.nll_loss(scores, batch.label, size_average=False).item() + rel_total += batch.label.data.sum().item() + pre_total += torch.max(scores, 1)[1].view(batch.label.size()).data.sum().item() + + precision = acc_total / pre_total + recall = acc_total / rel_total + f1 = 2 * precision * recall / (precision + recall) + accuracy = 100. * n_dev_correct / len(self.data_loader.dataset.examples) + avg_loss = total_loss / len(self.data_loader.dataset.examples) + return [accuracy, avg_loss, precision, recall, f1], ['accuracy', 'cross_entropy_loss', 'precision', 'recall', 'f1'] diff --git a/common/train.py b/common/train.py index 3eb05c6..bf8a54f 100644 --- a/common/train.py +++ b/common/train.py @@ -2,6 +2,7 @@ from .trainers.sick_trainer import SICKTrainer from .trainers.msrvid_trainer import MSRVIDTrainer from .trainers.trecqa_trainer import TRECQATrainer from .trainers.wikiqa_trainer import WikiQATrainer +from .trainers.pit2015_trainer import PIT2015Trainer from .trainers.sst_trainer import SSTTrainer from nce.nce_pairwise_mp.trainers.trecqa_trainer import TRECQATrainerNCE from nce.nce_pairwise_mp.trainers.wikiqa_trainer import WikiQATrainerNCE @@ -17,7 +18,8 @@ class TrainerFactory(object): 'SST-1': SSTTrainer, 'SST-2': SSTTrainer, 'trecqa': TRECQATrainer, - 'wikiqa': WikiQATrainer + 'wikiqa': WikiQATrainer, + 'pit2015': PIT2015Trainer } trainer_map_nce = { diff --git a/common/trainers/pit2015_trainer.py b/common/trainers/pit2015_trainer.py new file mode 100644 index 0000000..9551736 --- /dev/null +++ b/common/trainers/pit2015_trainer.py @@ -0,0 +1,85 @@ +import time + +import torch.nn.functional as F +from torch.optim.lr_scheduler import ReduceLROnPlateau + +from .trainer import Trainer +from utils.serialization import save_checkpoint + + +class PIT2015Trainer(Trainer): + + def train_epoch(self, epoch): + self.model.train() + total_loss = 0 + for batch_idx, batch in enumerate(self.train_loader): + self.optimizer.zero_grad() + + # Select embedding + sent1, sent2 = self.get_sentence_embeddings(batch) + + output = self.model(sent1, sent2, batch.ext_feats, batch.dataset.word_to_doc_cnt, batch.sentence_1_raw, batch.sentence_2_raw) + loss = F.nll_loss(output, batch.label, size_average=False) + total_loss += loss.item() + 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.item() / len(batch)) + ) + + accuracy, avg_loss, precision, recall, f1 = self.evaluate(self.train_evaluator, 'train') + + if self.use_tensorboard: + self.writer.add_scalar('{}/train/cross_entropy_loss'.format(self.train_loader.dataset.NAME), avg_loss, epoch) + self.writer.add_scalar('{}/train/accuracy'.format(self.train_loader.dataset.NAME), accuracy, epoch) + self.writer.add_scalar('{}/train/precision'.format(self.train_loader.dataset.NAME), precision, epoch) + self.writer.add_scalar('{}/train/recall'.format(self.train_loader.dataset.NAME), recall, epoch) + self.writer.add_scalar('{}/train/f1'.format(self.train_loader.dataset.NAME), f1, epoch) + + return total_loss + + def train(self, epochs): + scheduler = None + if self.lr_reduce_factor != 1 and self.lr_reduce_factor != None: + 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') + accuracy, avg_loss, precision, recall, f1 = dev_scores + + test_scores = self.evaluate(self.test_evaluator, 'test') + 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), avg_loss, epoch) + self.writer.add_scalar('{}/dev/accuracy'.format(self.train_loader.dataset.NAME), accuracy, epoch) + self.writer.add_scalar('{}/dev/precision'.format(self.train_loader.dataset.NAME), precision, epoch) + self.writer.add_scalar('{}/dev/recall'.format(self.train_loader.dataset.NAME), recall, epoch) + self.writer.add_scalar('{}/dev/f1'.format(self.train_loader.dataset.NAME), f1, epoch) + + end = time.time() + duration = end - start + self.logger.info('Epoch {} finished in {:.2f} minutes'.format(epoch, duration / 60)) + epoch_times.append(duration) + + if f1 > best_dev_score: + best_dev_score = f1 + save_checkpoint(epoch, self.model.arch, self.model.state_dict(), self.optimizer.state_dict(), best_dev_score, self.model_outfile) + + if abs(prev_loss - avg_loss) <= 0.0002: + self.logger.info('Early stopping. Loss changed by less than 0.0002.') + break + + prev_loss = avg_loss + if scheduler is not None: + scheduler.step(f1) + + self.logger.info('Training took {:.2f} minutes overall...'.format(sum(epoch_times) / 60)) diff --git a/datasets/pit2015.py b/datasets/pit2015.py new file mode 100644 index 0000000..0cede2b --- /dev/null +++ b/datasets/pit2015.py @@ -0,0 +1,65 @@ +import os + +import torch +from torchtext.data.field import Field, RawField +from torchtext.data.iterator import BucketIterator +from torchtext.vocab import Vectors + +from datasets.castor_dataset import CastorPairDataset + + +class PIT2015(CastorPairDataset): + NAME = 'pit2015' + NUM_CLASSES = 2 + ID_FIELD = Field(sequential=False, tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True) + AID_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + TEXT_FIELD = Field(batch_first=True, tokenize=lambda x: x) # tokenizer is identity since we already tokenized it to compute external features + EXT_FEATS_FIELD = Field(tensor_type=torch.FloatTensor, use_vocab=False, batch_first=True, tokenize=lambda x: x) + LABEL_FIELD = Field(sequential=False, use_vocab=False, batch_first=True) + RAW_TEXT_FIELD = RawField() + VOCAB_SIZE = 0 + + @staticmethod + def sort_key(ex): + return len(ex.sentence_1) + + def __init__(self, path): + """ + Create a PIT2015 dataset instance + """ + super(PIT2015, self).__init__(path) + + @classmethod + def splits(cls, path, train='train', validation='dev', test='test', **kwargs): + return super(PIT2015, cls).splits(path, train=train, validation=validation, test=test, **kwargs) + + @classmethod + def iters(cls, path, vectors_name, vectors_dir, batch_size=64, shuffle=True, device=0, pt_file=False, vectors=None, + unk_init=torch.Tensor.zero_): + """ + :param path: directory containing train, test, dev files + :param vectors_name: name of word vectors file + :param vectors_dir: directory containing word vectors file + :param batch_size: batch size + :param device: GPU device + :param vectors: custom vectors - either predefined torchtext vectors or your own custom Vector classes + :param pt_file: load cached embedding file from disk if it is true + :param unk_init: function used to generate vector for OOV words + :return: + """ + + train, validation, test = cls.splits(path) + if not pt_file: + if vectors is None: + vectors = Vectors(name=vectors_name, cache=vectors_dir, unk_init=unk_init) + cls.TEXT_FIELD.build_vocab(train, validation, test, vectors=vectors) + else: + cls.TEXT_FIELD.build_vocab(train, validation, test) + cls.TEXT_FIELD = cls.set_vectors(cls.TEXT_FIELD, os.path.join(vectors_dir, vectors_name)) + + cls.LABEL_FIELD.build_vocab(train, validation, test) + + cls.VOCAB_SIZE = len(cls.TEXT_FIELD.vocab) + + return BucketIterator.splits((train, validation, test), batch_size=batch_size, repeat=False, shuffle=shuffle, + sort_within_batch=True, device=device) diff --git a/utils/relevancy_metrics.py b/utils/relevancy_metrics.py index 429ae0b..8154ec2 100644 --- a/utils/relevancy_metrics.py +++ b/utils/relevancy_metrics.py @@ -17,7 +17,7 @@ def get_map_mrr(qids, predictions, labels, device=0, keep_results=False): qrel_fname = 'trecqa_{}_{}.qrel'.format(time.time(), device) results_fname = 'trecqa_{}_{}.results'.format(time.time(), device) qrel_template = '{qid} 0 {docno} {rel}\n' - results_template = '{qid} 0 {docno} 0 {sim} mpcnn\n' + results_template = '{qid} 0 {docno} 0 {sim} castor-model\n' with open(qrel_fname, 'w') as f1, open(results_fname, 'w') as f2: docnos = range(len(qids)) for qid, docno, predicted, actual in zip(qids, docnos, predictions, labels): diff --git a/vdpwi/__main__.py b/vdpwi/__main__.py index c8522ef..e1befcc 100644 --- a/vdpwi/__main__.py +++ b/vdpwi/__main__.py @@ -26,7 +26,7 @@ if __name__ == '__main__': parser = argparse.ArgumentParser(description='PyTorch implementation of VDPWI') parser.add_argument('model_outfile', help='file to save final model') 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, 'Castor-data', 'embeddings', 'GloVe')) + parser.add_argument('--word-vectors-dir', help='word vectors directory', default=os.path.join(os.pardir, 'Castor-data', 'embeddings', 'GloVe')) parser.add_argument('--word-vectors-file', help='word vectors filename', default='glove.840B.300d.txt') parser.add_argument('--word-vectors-dim', type=int, default=300, help='number of dimensions of word vectors (default: 300)') diff --git a/vdpwi/model.py b/vdpwi/model.py index 4c0dd7f..3203fd0 100644 --- a/vdpwi/model.py +++ b/vdpwi/model.py @@ -40,7 +40,7 @@ class VDPWIConvNet(nn.Module): def make_conv(n_in, n_out): conv = nn.Conv2d(n_in, n_out, 3, padding=1) conv.bias.data.zero_() - nn.init.xavier_normal(conv.weight) + nn.init.xavier_normal_(conv.weight) return conv self.conv1 = make_conv(12, 128) self.conv2 = make_conv(128, 164)