Save a model checkpoint.
(iteration, model, lr_scheduler, args, use_ema = False)
| 196 | |
| 197 | |
| 198 | def save_ds_checkpoint(iteration, model, lr_scheduler, args, use_ema = False): |
| 199 | """Save a model checkpoint.""" |
| 200 | |
| 201 | sd = {} |
| 202 | sd['iteration'] = iteration |
| 203 | if lr_scheduler is not None: |
| 204 | sd['client_lr_scheduler'] = lr_scheduler.state_dict() |
| 205 | # rng states. |
| 206 | if not args.no_save_rng: |
| 207 | sd['random_rng_state'] = random.getstate() |
| 208 | sd['np_rng_state'] = np.random.get_state() |
| 209 | sd['torch_rng_state'] = torch.get_rng_state() |
| 210 | sd['cuda_rng_state'] = torch.cuda.get_rng_state() |
| 211 | if not use_ema: |
| 212 | save_ds_checkpoint_no_optim(model, args.save, str(iteration), client_state=sd) |
| 213 | else: |
| 214 | save_ds_checkpoint_no_optim(model, args.save, str(iteration)+'-ema', client_state=sd) |
| 215 | |
| 216 | |
| 217 | def save_ds_checkpoint_no_optim(model, save_dir, tag=None, client_state={}, save_latest=True): |
no test coverage detected