(self, save_name, save_path)
| 303 | 'eval/precision': precision, 'eval/recall': recall, 'eval/F1': F1, 'eval/AUC': AUC} |
| 304 | |
| 305 | def save_model(self, save_name, save_path): |
| 306 | save_filename = os.path.join(save_path, save_name) |
| 307 | # copy EMA parameters to ema_model for saving with model as temp |
| 308 | self.model.eval() |
| 309 | self.ema.apply_shadow() |
| 310 | ema_model = self.model.state_dict() |
| 311 | self.ema.restore() |
| 312 | self.model.train() |
| 313 | |
| 314 | torch.save({'model': self.model.state_dict(), |
| 315 | 'optimizer': self.optimizer.state_dict(), |
| 316 | 'scheduler': self.scheduler.state_dict(), |
| 317 | 'it': self.it, |
| 318 | 'ema_model': ema_model}, |
| 319 | save_filename) |
| 320 | |
| 321 | self.print_fn(f"model saved: {save_filename}") |
| 322 | |
| 323 | def load_model(self, load_path): |
| 324 | checkpoint = torch.load(load_path) |
no test coverage detected