From fbd8629ca2d81a1f948a67188694170706568370 Mon Sep 17 00:00:00 2001 From: Ralph Tang Date: Thu, 24 May 2018 18:02:37 -0400 Subject: [PATCH] Move MPCNN API to common module (#105) --- common/__init__.py | 2 ++ {mp_cnn => common}/dataset.py | 2 +- {mp_cnn => common}/evaluation.py | 14 ++++++------- {mp_cnn => common}/evaluators/__init__.py | 0 {mp_cnn => common}/evaluators/evaluator.py | 0 .../evaluators/msrvid_evaluator.py | 2 +- {mp_cnn => common}/evaluators/qa_evaluator.py | 2 +- .../evaluators/sick_evaluator.py | 2 +- .../evaluators/trecqa_evaluator.py | 2 +- .../evaluators/wikiqa_evaluator.py | 2 +- {mp_cnn => common}/train.py | 14 ++++++------- {mp_cnn => common}/trainers/__init__.py | 0 {mp_cnn => common}/trainers/msrvid_trainer.py | 2 +- {mp_cnn => common}/trainers/qa_trainer.py | 2 +- {mp_cnn => common}/trainers/sick_trainer.py | 2 +- {mp_cnn => common}/trainers/trainer.py | 0 {mp_cnn => common}/trainers/trecqa_trainer.py | 2 +- {mp_cnn => common}/trainers/wikiqa_trainer.py | 2 +- mp_cnn/README.md | 16 +++++++-------- mp_cnn/{main.py => __main__.py} | 20 +++++++++---------- mp_cnn/model.py | 6 +++--- .../evaluators/qa_evaluator.py | 2 +- nce/nce_pairwise_mp/main.py | 6 +++--- nce/nce_pairwise_mp/trainers/qa_trainer.py | 2 +- 24 files changed, 53 insertions(+), 51 deletions(-) create mode 100644 common/__init__.py rename {mp_cnn => common}/dataset.py (99%) rename {mp_cnn => common}/evaluation.py (72%) rename {mp_cnn => common}/evaluators/__init__.py (100%) rename {mp_cnn => common}/evaluators/evaluator.py (100%) rename {mp_cnn => common}/evaluators/msrvid_evaluator.py (97%) rename {mp_cnn => common}/evaluators/qa_evaluator.py (96%) rename {mp_cnn => common}/evaluators/sick_evaluator.py (97%) rename {mp_cnn => common}/evaluators/trecqa_evaluator.py (79%) rename {mp_cnn => common}/evaluators/wikiqa_evaluator.py (79%) rename {mp_cnn => common}/train.py (71%) rename {mp_cnn => common}/trainers/__init__.py (100%) rename {mp_cnn => common}/trainers/msrvid_trainer.py (99%) rename {mp_cnn => common}/trainers/qa_trainer.py (98%) rename {mp_cnn => common}/trainers/sick_trainer.py (98%) rename {mp_cnn => common}/trainers/trainer.py (100%) rename {mp_cnn => common}/trainers/trecqa_trainer.py (85%) rename {mp_cnn => common}/trainers/wikiqa_trainer.py (85%) rename mp_cnn/{main.py => __main__.py} (85%) diff --git a/common/__init__.py b/common/__init__.py new file mode 100644 index 0000000..1c4e196 --- /dev/null +++ b/common/__init__.py @@ -0,0 +1,2 @@ +from .evaluators import * +from .trainers import * \ No newline at end of file diff --git a/mp_cnn/dataset.py b/common/dataset.py similarity index 99% rename from mp_cnn/dataset.py rename to common/dataset.py index 67d9b0c..32a76cd 100644 --- a/mp_cnn/dataset.py +++ b/common/dataset.py @@ -24,7 +24,7 @@ class UnknownWordVecCache(object): return cls.cache[size_tup] -class MPCNNDatasetFactory(object): +class DatasetFactory(object): """ Get the corresponding Dataset class for a particular dataset. """ diff --git a/mp_cnn/evaluation.py b/common/evaluation.py similarity index 72% rename from mp_cnn/evaluation.py rename to common/evaluation.py index 9163ce6..aa34dee 100644 --- a/mp_cnn/evaluation.py +++ b/common/evaluation.py @@ -1,11 +1,11 @@ -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 +from .evaluators.sick_evaluator import SICKEvaluator +from .evaluators.msrvid_evaluator import MSRVIDEvaluator +from .evaluators.trecqa_evaluator import TRECQAEvaluator +from .evaluators.wikiqa_evaluator import WikiQAEvaluator from nce.nce_pairwise_mp.evaluators.trecqa_evaluator import TRECQAEvaluatorNCE from nce.nce_pairwise_mp.evaluators.wikiqa_evaluator import WikiQAEvaluatorNCE -class MPCNNEvaluatorFactory(object): +class EvaluatorFactory(object): """ Get the corresponding Evaluator class for a particular dataset. """ @@ -27,9 +27,9 @@ class MPCNNEvaluatorFactory(object): return None if nce: - evaluator_map = MPCNNEvaluatorFactory.evaluator_map_nce + evaluator_map = EvaluatorFactory.evaluator_map_nce else: - evaluator_map = MPCNNEvaluatorFactory.evaluator_map + evaluator_map = EvaluatorFactory.evaluator_map if not hasattr(dataset_cls, 'NAME'): raise ValueError('Invalid dataset. Dataset should have NAME attribute.') diff --git a/mp_cnn/evaluators/__init__.py b/common/evaluators/__init__.py similarity index 100% rename from mp_cnn/evaluators/__init__.py rename to common/evaluators/__init__.py diff --git a/mp_cnn/evaluators/evaluator.py b/common/evaluators/evaluator.py similarity index 100% rename from mp_cnn/evaluators/evaluator.py rename to common/evaluators/evaluator.py diff --git a/mp_cnn/evaluators/msrvid_evaluator.py b/common/evaluators/msrvid_evaluator.py similarity index 97% rename from mp_cnn/evaluators/msrvid_evaluator.py rename to common/evaluators/msrvid_evaluator.py index 3af565f..975638e 100644 --- a/mp_cnn/evaluators/msrvid_evaluator.py +++ b/common/evaluators/msrvid_evaluator.py @@ -2,7 +2,7 @@ from scipy.stats import pearsonr import torch import torch.nn.functional as F -from mp_cnn.evaluators.evaluator import Evaluator +from .evaluator import Evaluator class MSRVIDEvaluator(Evaluator): diff --git a/mp_cnn/evaluators/qa_evaluator.py b/common/evaluators/qa_evaluator.py similarity index 96% rename from mp_cnn/evaluators/qa_evaluator.py rename to common/evaluators/qa_evaluator.py index a42266f..dc61297 100644 --- a/mp_cnn/evaluators/qa_evaluator.py +++ b/common/evaluators/qa_evaluator.py @@ -1,6 +1,6 @@ import torch.nn.functional as F -from mp_cnn.evaluators.evaluator import Evaluator +from .evaluator import Evaluator from utils.relevancy_metrics import get_map_mrr diff --git a/mp_cnn/evaluators/sick_evaluator.py b/common/evaluators/sick_evaluator.py similarity index 97% rename from mp_cnn/evaluators/sick_evaluator.py rename to common/evaluators/sick_evaluator.py index 1101b07..d0f2280 100644 --- a/mp_cnn/evaluators/sick_evaluator.py +++ b/common/evaluators/sick_evaluator.py @@ -2,7 +2,7 @@ from scipy.stats import pearsonr, spearmanr import torch import torch.nn.functional as F -from mp_cnn.evaluators.evaluator import Evaluator +from .evaluator import Evaluator class SICKEvaluator(Evaluator): diff --git a/mp_cnn/evaluators/trecqa_evaluator.py b/common/evaluators/trecqa_evaluator.py similarity index 79% rename from mp_cnn/evaluators/trecqa_evaluator.py rename to common/evaluators/trecqa_evaluator.py index baae709..93b1d4a 100644 --- a/mp_cnn/evaluators/trecqa_evaluator.py +++ b/common/evaluators/trecqa_evaluator.py @@ -1,4 +1,4 @@ -from mp_cnn.evaluators.qa_evaluator import QAEvaluator +from .qa_evaluator import QAEvaluator class TRECQAEvaluator(QAEvaluator): diff --git a/mp_cnn/evaluators/wikiqa_evaluator.py b/common/evaluators/wikiqa_evaluator.py similarity index 79% rename from mp_cnn/evaluators/wikiqa_evaluator.py rename to common/evaluators/wikiqa_evaluator.py index 82463e4..810bacb 100644 --- a/mp_cnn/evaluators/wikiqa_evaluator.py +++ b/common/evaluators/wikiqa_evaluator.py @@ -1,4 +1,4 @@ -from mp_cnn.evaluators.qa_evaluator import QAEvaluator +from .qa_evaluator import QAEvaluator class WikiQAEvaluator(QAEvaluator): diff --git a/mp_cnn/train.py b/common/train.py similarity index 71% rename from mp_cnn/train.py rename to common/train.py index 428574f..96810ec 100644 --- a/mp_cnn/train.py +++ b/common/train.py @@ -1,12 +1,12 @@ -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 +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 nce.nce_pairwise_mp.trainers.trecqa_trainer import TRECQATrainerNCE from nce.nce_pairwise_mp.trainers.wikiqa_trainer import WikiQATrainerNCE -class MPCNNTrainerFactory(object): +class TrainerFactory(object): """ Get the corresponding Trainer class for a particular dataset. """ @@ -25,9 +25,9 @@ class MPCNNTrainerFactory(object): @staticmethod def get_trainer(dataset_name, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None, nce=False): if nce: - trainer_map = MPCNNTrainerFactory.trainer_map_nce + trainer_map = TrainerFactory.trainer_map_nce else: - trainer_map = MPCNNTrainerFactory.trainer_map + trainer_map = TrainerFactory.trainer_map if dataset_name not in trainer_map: raise ValueError('{} is not implemented.'.format(dataset_name)) diff --git a/mp_cnn/trainers/__init__.py b/common/trainers/__init__.py similarity index 100% rename from mp_cnn/trainers/__init__.py rename to common/trainers/__init__.py diff --git a/mp_cnn/trainers/msrvid_trainer.py b/common/trainers/msrvid_trainer.py similarity index 99% rename from mp_cnn/trainers/msrvid_trainer.py rename to common/trainers/msrvid_trainer.py index de0807d..7390c77 100644 --- a/mp_cnn/trainers/msrvid_trainer.py +++ b/common/trainers/msrvid_trainer.py @@ -6,7 +6,7 @@ import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau from scipy.stats import pearsonr -from mp_cnn.trainers.trainer import Trainer +from .trainer import Trainer class MSRVIDTrainer(Trainer): diff --git a/mp_cnn/trainers/qa_trainer.py b/common/trainers/qa_trainer.py similarity index 98% rename from mp_cnn/trainers/qa_trainer.py rename to common/trainers/qa_trainer.py index 44ea911..7960b8c 100644 --- a/mp_cnn/trainers/qa_trainer.py +++ b/common/trainers/qa_trainer.py @@ -4,7 +4,7 @@ import torch import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau -from mp_cnn.trainers.trainer import Trainer +from .trainer import Trainer class QATrainer(Trainer): diff --git a/mp_cnn/trainers/sick_trainer.py b/common/trainers/sick_trainer.py similarity index 98% rename from mp_cnn/trainers/sick_trainer.py rename to common/trainers/sick_trainer.py index 5ff8870..c0e5d96 100644 --- a/mp_cnn/trainers/sick_trainer.py +++ b/common/trainers/sick_trainer.py @@ -4,7 +4,7 @@ import torch import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau -from mp_cnn.trainers.trainer import Trainer +from .trainer import Trainer class SICKTrainer(Trainer): diff --git a/mp_cnn/trainers/trainer.py b/common/trainers/trainer.py similarity index 100% rename from mp_cnn/trainers/trainer.py rename to common/trainers/trainer.py diff --git a/mp_cnn/trainers/trecqa_trainer.py b/common/trainers/trecqa_trainer.py similarity index 85% rename from mp_cnn/trainers/trecqa_trainer.py rename to common/trainers/trecqa_trainer.py index c49785c..03a4194 100644 --- a/mp_cnn/trainers/trecqa_trainer.py +++ b/common/trainers/trecqa_trainer.py @@ -1,4 +1,4 @@ -from mp_cnn.trainers.qa_trainer import QATrainer +from .qa_trainer import QATrainer class TRECQATrainer(QATrainer): diff --git a/mp_cnn/trainers/wikiqa_trainer.py b/common/trainers/wikiqa_trainer.py similarity index 85% rename from mp_cnn/trainers/wikiqa_trainer.py rename to common/trainers/wikiqa_trainer.py index 801ae3e..dfc24a7 100644 --- a/mp_cnn/trainers/wikiqa_trainer.py +++ b/common/trainers/wikiqa_trainer.py @@ -1,4 +1,4 @@ -from mp_cnn.trainers.qa_trainer import QATrainer +from .qa_trainer import QATrainer class WikiQATrainer(QATrainer): diff --git a/mp_cnn/README.md b/mp_cnn/README.md index d6abd15..3b00917 100644 --- a/mp_cnn/README.md +++ b/mp_cnn/README.md @@ -11,7 +11,7 @@ Please ensure you have followed instructions in the main [README](../README.md) To run MP-CNN on the SICK dataset, use the following command. `--dropout 0` is for mimicking the original paper, although adding dropout can improve performance. If you have any problems running it check the Troubleshooting section below. ``` -python main.py mpcnn.sick.model.castor --dataset sick --epochs 19 --epsilon 1e-7 --dropout 0 +python -m mp_cnn mpcnn.sick.model.castor --dataset sick --epochs 19 --epsilon 1e-7 --dropout 0 ``` | Implementation and config | Pearson's r | Spearman's p | @@ -23,7 +23,7 @@ python main.py mpcnn.sick.model.castor --dataset sick --epochs 19 --epsilon 1e-7 To run MP-CNN on the MSRVID dataset, use the following command: ``` -python main.py mpcnn.msrvid.model.castor --dataset msrvid --batch-size 16 --epsilon 1e-7 --epochs 32 --dropout 0 --regularization 0.0025 +python -m mp_cnn mpcnn.msrvid.model.castor --dataset msrvid --batch-size 16 --epsilon 1e-7 --epochs 32 --dropout 0 --regularization 0.0025 ``` | Implementation and config | Pearson's r | @@ -37,7 +37,7 @@ To run MP-CNN on (Raw) TrecQA, you first need to run `./get_trec_eval.sh` in `ut Then, you can run: ``` -python main.py mpcnn.trecqa.model --dataset trecqa --epochs 5 --regularization 0.0005 --dropout 0.5 --eps 0.1 +python -m mp_cnn mpcnn.trecqa.model --dataset trecqa --epochs 5 --regularization 0.0005 --dropout 0.5 --eps 0.1 ``` | Implementation and config | map | mrr | @@ -53,7 +53,7 @@ You also need `trec_eval` for this dataset, similar to TrecQA. Then, you can run: ``` -python main.py mpcnn.wikiqa.model --epochs 10 --dataset wikiqa --batch-size 64 --lr 0.0004 --regularization 0.02 +python -m mp_cnn mpcnn.wikiqa.model --epochs 10 --dataset wikiqa --batch-size 64 --lr 0.0004 --regularization 0.02 ``` | Implementation and config | map | mrr | | -------------------------------- |:------:|:------:| @@ -67,7 +67,7 @@ These are not the optimal hyperparameters but they are decent. This README will To see all options available, use ``` -python main.py --help +python -m mp_cnn --help ``` ## Troubleshooting @@ -75,9 +75,9 @@ python main.py --help ### ModuleNotFoundError: datasets ``` Traceback (most recent call last): - File "main.py", line 9, in - from dataset import MPCNNDatasetFactory - File "/u/z3tu/castorini/Castor/mp_cnn/dataset.py", line 12, in + File "__main__.py", line 9, in + from common.dataset import DatasetFactory + File "/u/z3tu/castorini/Castor/common/dataset.py", line 12, in from datasets.sick import SICK ModuleNotFoundError: No module named 'datasets' ``` diff --git a/mp_cnn/main.py b/mp_cnn/__main__.py similarity index 85% rename from mp_cnn/main.py rename to mp_cnn/__main__.py index c89fe80..cf75446 100644 --- a/mp_cnn/main.py +++ b/mp_cnn/__main__.py @@ -8,10 +8,10 @@ import numpy as np import torch import torch.optim as optim -from mp_cnn.dataset import MPCNNDatasetFactory -from mp_cnn.evaluation import MPCNNEvaluatorFactory -from mp_cnn.model import MPCNN -from mp_cnn.train import MPCNNTrainerFactory +from common.dataset import DatasetFactory +from common.evaluation import EvaluatorFactory +from common.train import TrainerFactory +from .model import MPCNN if __name__ == '__main__': @@ -62,7 +62,7 @@ if __name__ == '__main__': logger.info(pprint.pformat(vars(args))) dataset_cls, embedding, train_loader, test_loader, dev_loader \ - = MPCNNDatasetFactory.get_dataset(args.dataset, args.word_vectors_dir, args.word_vectors_file, args.batch_size, args.device) + = DatasetFactory.get_dataset(args.dataset, args.word_vectors_dir, args.word_vectors_file, args.batch_size, args.device) filter_widths = list(range(1, args.max_window_size + 1)) + [np.inf] model = MPCNN(embedding, args.holistic_filters, args.per_dim_filters, filter_widths, @@ -80,9 +80,9 @@ if __name__ == '__main__': else: raise ValueError('optimizer not recognized: it should be either adam or sgd') - train_evaluator = MPCNNEvaluatorFactory.get_evaluator(dataset_cls, model, train_loader, args.batch_size, args.device) - test_evaluator = MPCNNEvaluatorFactory.get_evaluator(dataset_cls, model, test_loader, args.batch_size, args.device) - dev_evaluator = MPCNNEvaluatorFactory.get_evaluator(dataset_cls, model, dev_loader, args.batch_size, args.device) + train_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, train_loader, args.batch_size, args.device) + test_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, test_loader, args.batch_size, args.device) + dev_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, dev_loader, args.batch_size, args.device) trainer_config = { 'optimizer': optimizer, @@ -95,7 +95,7 @@ if __name__ == '__main__': 'run_label': args.run_label, 'logger': logger } - trainer = MPCNNTrainerFactory.get_trainer(args.dataset, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) + trainer = TrainerFactory.get_trainer(args.dataset, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator) if not args.skip_training: total_params = 0 @@ -106,7 +106,7 @@ if __name__ == '__main__': trainer.train(args.epochs) model = torch.load(args.model_outfile) - saved_model_evaluator = MPCNNEvaluatorFactory.get_evaluator(dataset_cls, model, test_loader, args.batch_size, args.device) + saved_model_evaluator = EvaluatorFactory.get_evaluator(dataset_cls, model, test_loader, args.batch_size, args.device) scores, metric_names = saved_model_evaluator.get_scores() logger.info('Evaluation metrics for test') logger.info('\t'.join([' '] + metric_names)) diff --git a/mp_cnn/model.py b/mp_cnn/model.py index 651cdd2..520ad66 100644 --- a/mp_cnn/model.py +++ b/mp_cnn/model.py @@ -91,7 +91,7 @@ class MPCNN(nn.Module): x2 = sent2_block_a[ws][pool] batch_size = x1.size()[0] comparison_feats.append(F.cosine_similarity(x1, x2).contiguous().view(batch_size, 1)) - comparison_feats.append(F.pairwise_distance(x1, x2)) + comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(-1)) return torch.cat(comparison_feats, dim=1) def _algo_2_vert_comp(self, sent1_block_a, sent2_block_a, sent1_block_b, sent2_block_b): @@ -105,7 +105,7 @@ class MPCNN(nn.Module): x2 = sent2_block_a[ws2][pool] if (not np.isinf(ws1) and not np.isinf(ws2)) or (np.isinf(ws1) and np.isinf(ws2)): comparison_feats.append(F.cosine_similarity(x1, x2).contiguous().view(batch_size, 1)) - comparison_feats.append(F.pairwise_distance(x1, x2)) + comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(-1)) comparison_feats.append(torch.abs(x1 - x2)) for pool in ('max', 'min'): @@ -117,7 +117,7 @@ class MPCNN(nn.Module): x2 = oG_2B[:, :, i] batch_size = x1.size()[0] comparison_feats.append(F.cosine_similarity(x1, x2).contiguous().view(batch_size, 1)) - comparison_feats.append(F.pairwise_distance(x1, x2)) + comparison_feats.append(F.pairwise_distance(x1, x2).unsqueeze(-1)) comparison_feats.append(torch.abs(x1 - x2)) return torch.cat(comparison_feats, dim=1) diff --git a/nce/nce_pairwise_mp/evaluators/qa_evaluator.py b/nce/nce_pairwise_mp/evaluators/qa_evaluator.py index d8db2bb..c546181 100644 --- a/nce/nce_pairwise_mp/evaluators/qa_evaluator.py +++ b/nce/nce_pairwise_mp/evaluators/qa_evaluator.py @@ -1,6 +1,6 @@ import numpy as np -from mp_cnn.evaluators.evaluator import Evaluator +from common.evaluators.evaluator import Evaluator from utils.relevancy_metrics import get_map_mrr diff --git a/nce/nce_pairwise_mp/main.py b/nce/nce_pairwise_mp/main.py index 32b39f8..f63132f 100644 --- a/nce/nce_pairwise_mp/main.py +++ b/nce/nce_pairwise_mp/main.py @@ -8,10 +8,10 @@ import numpy as np import torch import torch.optim as optim -from mp_cnn.dataset import MPCNNDatasetFactory -from mp_cnn.evaluation import MPCNNEvaluatorFactory +from common.dataset import MPCNNDatasetFactory +from common.evaluation import MPCNNEvaluatorFactory from nce.nce_pairwise_mp.model import MPCNN4NCE, PairwiseConv -from mp_cnn.train import MPCNNTrainerFactory +from common.train import MPCNNTrainerFactory if __name__ == '__main__': diff --git a/nce/nce_pairwise_mp/trainers/qa_trainer.py b/nce/nce_pairwise_mp/trainers/qa_trainer.py index df9a04b..b935898 100644 --- a/nce/nce_pairwise_mp/trainers/qa_trainer.py +++ b/nce/nce_pairwise_mp/trainers/qa_trainer.py @@ -5,7 +5,7 @@ import torch import torch.nn.functional as F from torch.optim.lr_scheduler import ReduceLROnPlateau -from mp_cnn.trainers.trainer import Trainer +from common.trainers.trainer import Trainer from utils.nce_neighbors import get_nearest_neg_id, get_random_neg_id, get_batch class QATrainer(Trainer):