mirror of
https://github.com/wassname/Castor.git
synced 2026-09-10 11:40:44 +08:00
Move MPCNN API to common module (#105)
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
from .evaluators import *
|
||||
from .trainers import *
|
||||
@@ -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):
|
||||
@@ -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
@@ -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'
|
||||
```
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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__':
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user