Base model.
| 11 | |
| 12 | |
| 13 | class BaseModel(): |
| 14 | """Base model.""" |
| 15 | |
| 16 | def __init__(self, opt): |
| 17 | self.opt = opt |
| 18 | self.device = torch.device('cuda' if opt['num_gpu'] != 0 else 'cpu') |
| 19 | self.is_train = opt['is_train'] |
| 20 | self.schedulers = [] |
| 21 | self.optimizers = [] |
| 22 | |
| 23 | def feed_data(self, data): |
| 24 | pass |
| 25 | |
| 26 | def optimize_parameters(self): |
| 27 | pass |
| 28 | |
| 29 | def get_current_visuals(self): |
| 30 | pass |
| 31 | |
| 32 | def save(self, epoch, current_iter): |
| 33 | """Save networks and training state.""" |
| 34 | pass |
| 35 | |
| 36 | def validation(self, dataloader, current_iter, tb_logger, save_img=False,test_num=-1,save_num=-1): |
| 37 | """Validation function. |
| 38 | |
| 39 | Args: |
| 40 | dataloader (torch.utils.data.DataLoader): Validation dataloader. |
| 41 | current_iter (int): Current iteration. |
| 42 | tb_logger (tensorboard logger): Tensorboard logger. |
| 43 | save_img (bool): Whether to save images. Default: False. |
| 44 | """ |
| 45 | if self.opt['dist']: |
| 46 | self.dist_validation(dataloader, current_iter, tb_logger, save_img,test_num,save_num) |
| 47 | else: |
| 48 | self.nondist_validation(dataloader, current_iter, tb_logger, save_img,test_num,save_num) |
| 49 | |
| 50 | def _initialize_best_metric_results(self, dataset_name): |
| 51 | """Initialize the best metric results dict for recording the best metric value and iteration.""" |
| 52 | if hasattr(self, 'best_metric_results') and dataset_name in self.best_metric_results: |
| 53 | return |
| 54 | elif not hasattr(self, 'best_metric_results'): |
| 55 | self.best_metric_results = dict() |
| 56 | |
| 57 | # add a dataset record |
| 58 | record = dict() |
| 59 | for metric, content in self.opt['val']['metrics'].items(): |
| 60 | better = content.get('better', 'higher') |
| 61 | init_val = float('-inf') if better == 'higher' else float('inf') |
| 62 | record[metric] = dict(better=better, val=init_val, iter=-1) |
| 63 | self.best_metric_results[dataset_name] = record |
| 64 | |
| 65 | def _update_best_metric_result(self, dataset_name, metric, val, current_iter): |
| 66 | print(val) |
| 67 | if self.best_metric_results[dataset_name][metric]['better'] == 'higher': |
| 68 | if val >= self.best_metric_results[dataset_name][metric]['val']: |
| 69 | self.best_metric_results[dataset_name][metric]['val'] = val |
| 70 | self.best_metric_results[dataset_name][metric]['iter'] = current_iter |
nothing calls this directly
no outgoing calls
no test coverage detected