MCPcopy Create free account
hub / github.com/42dot/VFDepth / BaseModel

Class BaseModel

models/base_model.py:8–93  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class BaseModel:
9 def __init__(self, cfg):
10 self._dataloaders = {}
11 self.mode = None
12 self.models = None
13 self.optimizer = None
14 self.lr_scheduler = None
15 self.ddp_enable = False
16
17 def read_config(self, cfg):
18 raise NotImplementedError('Not implemented for BaseModel')
19
20 def prepare_dataset(self):
21 raise NotImplementedError('Not implemented for BaseModel')
22
23 def set_optimizer(self):
24 raise NotImplementedError('Not implemented for BaseModel')
25
26 def train_dataloader(self):
27 return self._dataloaders['train']
28
29 def val_dataloader(self):
30 return self._dataloaders['val']
31
32 def eval_dataloader(self):
33 return self._dataloaders['eval']
34
35 def set_train(self):
36 self.mode = 'train'
37 for m in self.models.values():
38 m.train()
39
40 def set_val(self):
41 self.mode = 'val'
42 for m in self.models.values():
43 m.eval()
44
45 def save_model(self, epoch):
46 curr_model_weights_dir = os.path.join(self.save_weights_root, f'weights_{epoch}')
47 os.makedirs(curr_model_weights_dir, exist_ok=True)
48
49 for model_name, model in self.models.items():
50 model_file_path = os.path.join(curr_model_weights_dir, f'{model_name}.pth')
51 to_save = model.state_dict()
52 torch.save(to_save, model_file_path)
53
54 # save optimizer
55 optim_file_path = os.path.join(curr_model_weights_dir, f'{_OPTIMIZER_NAME}.pth')
56 torch.save(self.optimizer.state_dict(), optim_file_path)
57
58 def load_weights(self):
59 assert os.path.isdir(self.load_weights_dir), f'\tCannot find {self.load_weights_dir}'
60 print(f'Loading a model from {self.load_weights_dir}')
61
62 # to retrain
63 if self.pretrain and self.ddp_enable:
64 map_location = {'cuda:%d' % 0: 'cuda:%d' % (self.world_size-1)}
65

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected