| 164 | |
| 165 | |
| 166 | def save_checkpoint(save_dir, accelerator, epoch, max_acc, config, lr_scheduler=None): |
| 167 | # let accelerator handle the model and optimizer state for ddp and deepspeed. |
| 168 | accelerator.save_state(save_dir) |
| 169 | |
| 170 | if accelerator.is_main_process: |
| 171 | save_state = { |
| 172 | 'lr_scheduler': lr_scheduler.state_dict(), |
| 173 | 'max_acc': max_acc, |
| 174 | 'epoch': epoch, |
| 175 | 'config': config |
| 176 | } |
| 177 | torch.save(save_state, os.path.join(save_dir, 'additional_state.pth')) |
| 178 | |
| 179 | |
| 180 | def load_checkpoint_if_needed(accelerator, config, lr_scheduler=None): |