(self, save_name, save_path)
| 225 | return {'eval/loss': total_loss / total_num, 'eval/top-1-acc': top1, 'eval/top-5-acc': top5} |
| 226 | |
| 227 | def save_model(self, save_name, save_path): |
| 228 | if self.it < 1000000: |
| 229 | return |
| 230 | save_filename = os.path.join(save_path, save_name) |
| 231 | # copy EMA parameters to ema_model for saving with model as temp |
| 232 | self.model.eval() |
| 233 | self.ema.apply_shadow() |
| 234 | ema_model = deepcopy(self.model) |
| 235 | self.ema.restore() |
| 236 | self.model.train() |
| 237 | |
| 238 | torch.save({'model': self.model.state_dict(), |
| 239 | 'optimizer': self.optimizer.state_dict(), |
| 240 | 'scheduler': self.scheduler.state_dict(), |
| 241 | 'it': self.it + 1, |
| 242 | 'ema_model': ema_model.state_dict()}, |
| 243 | save_filename) |
| 244 | |
| 245 | self.print_fn(f"model saved: {save_filename}") |
| 246 | |
| 247 | def load_model(self, load_path): |
| 248 | checkpoint = torch.load(load_path) |
no test coverage detected