Files
Castor/mp_cnn/train.py
Victor Yang 51d8e29525 add NCE to MP-CNN (#84)
* update nce-sm

* refactor code, update torchtext

* use shared evaluation

* refactor code, use shared data loader

* refactor code

* refactor code

* refactor code according to Michael's great suggestions

* update readme and requirement

* update datasets and readme

* update data loader

* add space between +

* update refactor code

* add nce-mp

* remove duplicate files

* update readme, refactor code according to mp_cnn and delete duplicate code, follow PEP8 standard

* refactor code, add/delete comments

* import exit from sys
2018-01-03 18:12:57 -05:00

38 lines
1.3 KiB
Python

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 nce.nce_pairwise_mp.trainers.trecqa_trainer import TRECQATrainerNCE
from nce.nce_pairwise_mp.trainers.wikiqa_trainer import WikiQATrainerNCE
class MPCNNTrainerFactory(object):
"""
Get the corresponding Trainer class for a particular dataset.
"""
trainer_map = {
'sick': SICKTrainer,
'msrvid': MSRVIDTrainer,
'trecqa': TRECQATrainer,
'wikiqa': WikiQATrainer
}
trainer_map_nce = {
'trecqa': TRECQATrainerNCE,
'wikiqa': WikiQATrainerNCE
}
@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
else:
trainer_map = MPCNNTrainerFactory.trainer_map
if dataset_name not in trainer_map:
raise ValueError('{} is not implemented.'.format(dataset_name))
return trainer_map[dataset_name](
model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator
)