MCPcopy Create free account
hub / github.com/DragonisCV/RAM / BaseModel

Class BaseModel

ram/models/base_model.py:13–396  ·  view source on GitHub ↗

Base model.

Source from the content-addressed store, hash-verified

11
12
13class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected