Move MPCNN API to common module (#105)

This commit is contained in:
Ralph Tang
2018-05-24 18:02:37 -04:00
committed by GitHub
parent 4ece3c7ade
commit fbd8629ca2
24 changed files with 53 additions and 51 deletions
+2
View File
@@ -0,0 +1,2 @@
from .evaluators import *
from .trainers import *
+1 -1
View File
@@ -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.
"""
@@ -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.')
@@ -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):
@@ -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
@@ -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):
@@ -1,4 +1,4 @@
from mp_cnn.evaluators.qa_evaluator import QAEvaluator
from .qa_evaluator import QAEvaluator
class TRECQAEvaluator(QAEvaluator):
@@ -1,4 +1,4 @@
from mp_cnn.evaluators.qa_evaluator import QAEvaluator
from .qa_evaluator import QAEvaluator
class WikiQAEvaluator(QAEvaluator):
+7 -7
View File
@@ -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))
@@ -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):
@@ -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):
@@ -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):
@@ -1,4 +1,4 @@
from mp_cnn.trainers.qa_trainer import QATrainer
from .qa_trainer import QATrainer
class TRECQATrainer(QATrainer):
@@ -1,4 +1,4 @@
from mp_cnn.trainers.qa_trainer import QATrainer
from .qa_trainer import QATrainer
class WikiQATrainer(QATrainer):
+8 -8
View File
@@ -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 <module>
from dataset import MPCNNDatasetFactory
File "/u/z3tu/castorini/Castor/mp_cnn/dataset.py", line 12, in <module>
File "__main__.py", line 9, in <module>
from common.dataset import DatasetFactory
File "/u/z3tu/castorini/Castor/common/dataset.py", line 12, in <module>
from datasets.sick import SICK
ModuleNotFoundError: No module named 'datasets'
```
+10 -10
View File
@@ -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))
+3 -3
View File
@@ -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)
@@ -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
+3 -3
View File
@@ -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__':
+1 -1
View File
@@ -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):