mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
* 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
38 lines
1.3 KiB
Python
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
|
|
)
|