| 6 | |
| 7 | |
| 8 | class 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 |
nothing calls this directly
no outgoing calls
no test coverage detected