(self)
| 457 | |
| 458 | |
| 459 | def main(self): |
| 460 | if self.mode_eval == "nocv": |
| 461 | self.model = self.get_model() |
| 462 | dataloaders = self.get_dataloader(data_type=self.data_type, data_fold=self.fold) |
| 463 | trainer = Trainer(model=self.model, device = self.device, lr = self.lr, dataloaders = dataloaders, epoches = self.epoches, dropout = self.dropout, weight_decay = self.weight_decay, mode = self.mode, model_name = self.model_name, event_num = self.event_num, |
| 464 | epoch_stop = self.epoch_stop, save_param_path = self.save_param_dir+self.data_type+"/"+self.model_name+"/", writer = SummaryWriter(self.path_tensorboard)) |
| 465 | result=trainer.train() |
| 466 | |
| 467 | elif self.mode_eval == "temporal": |
| 468 | self.model = self.get_model() |
| 469 | dataloaders = self.get_dataloader_temporal(data_type=self.data_type) |
| 470 | trainer = Trainer3(model=self.model, device = self.device, lr = self.lr, dataloaders = dataloaders, epoches = self.epoches, dropout = self.dropout, weight_decay = self.weight_decay, mode = self.mode, model_name = self.model_name, event_num = self.event_num, |
| 471 | epoch_stop = self.epoch_stop, save_param_path = self.save_param_dir+self.data_type+"/"+self.model_name+"/", writer = SummaryWriter(self.path_tensorboard)) |
| 472 | result=trainer.train() |
| 473 | return result |
| 474 | |
| 475 | elif self.mode_eval == "cv": |
| 476 | collate_fn=None |
| 477 | if self.model_name == 'TextCNN': |
| 478 | wv_from_text = KeyedVectors.load_word2vec_format("./stores/tencent-ailab-embedding-zh-d100-v0.2.0-s/tencent-ailab-embedding-zh-d100-v0.2.0-s.txt", binary=False) |
| 479 | |
| 480 | history = collections.defaultdict(list) |
| 481 | for fold in range(1, 6): |
| 482 | print('-' * 50) |
| 483 | print ('fold %d:' % fold) |
| 484 | print('-' * 50) |
| 485 | self.model = self.get_model() |
| 486 | dataloaders = self.get_dataloader(data_type=self.data_type, data_fold=fold) |
| 487 | trainer = Trainer(model = self.model, device = self.device, lr = self.lr, dataloaders = dataloaders, epoches = self.epoches, dropout = self.dropout, weight_decay = self.weight_decay, mode = self.mode, model_name = self.model_name, event_num = self.event_num, |
| 488 | epoch_stop = self.epoch_stop, save_param_path = self.save_param_dir+self.data_type+"/"+self.model_name+"/", writer = SummaryWriter(self.path_tensorboard+"fold_"+str(fold)+"/")) |
| 489 | |
| 490 | result = trainer.train() |
| 491 | |
| 492 | history['auc'].append(result['auc']) |
| 493 | history['f1'].append(result['f1']) |
| 494 | history['recall'].append(result['recall']) |
| 495 | history['precision'].append(result['precision']) |
| 496 | history['acc'].append(result['acc']) |
| 497 | |
| 498 | print ('results on 5-fold cross-validation: ') |
| 499 | for metric in ['acc', 'f1', 'precision', 'recall', 'auc']: |
| 500 | print ('%s : %.4f +/- %.4f' % (metric, np.mean(history[metric]), np.std(history[metric]))) |
| 501 | |
| 502 | else: |
| 503 | print ("Not Available") |
no test coverage detected