MCPcopy Create free account
hub / github.com/ICTMCG/FakeSV / main

Method main

code/run.py:459–503  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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")

Callers 1

main.pyFile · 0.80

Calls 6

get_modelMethod · 0.95
get_dataloaderMethod · 0.95
trainMethod · 0.95
TrainerClass · 0.90
Trainer3Class · 0.90

Tested by

no test coverage detected