Files
Castor/mp_cnn/train.py
T
Michael Tu 4470983504 TrecQA for MP-CNN (#77)
* Add TrecQA dataset and modularize MP-CNN infra

* Stylistic improvements

* Fix and warn about trec_eval path issue

* Update README for MP-CNN

* Update incorrect map/mrr

* MP-CNN: address code review comments

* Create common Castor pair Dataset class

* Move map and mrr computation to Castor utils

* Make map mrr utility trec_eval path more general
2017-11-04 18:22:40 -04:00

24 lines
859 B
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
class MPCNNTrainerFactory(object):
"""
Get the corresponding Trainer class for a particular dataset.
"""
trainer_map = {
'sick': SICKTrainer,
'msrvid': MSRVIDTrainer,
'trecqa': TRECQATrainer
}
@staticmethod
def get_trainer(dataset_name, model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator=None):
if dataset_name not in MPCNNTrainerFactory.trainer_map:
raise ValueError('{} is not implemented.'.format(dataset_name))
return MPCNNTrainerFactory.trainer_map[dataset_name](
model, train_loader, trainer_config, train_evaluator, test_evaluator, dev_evaluator
)