| 30 | |
| 31 | |
| 32 | def save_checkpoint(save_dir, save_name, model, optim, log, def_noise=None): |
| 33 | torch.save({ |
| 34 | 'model_state_dict': utils.get_model_state(model), |
| 35 | 'optim_state_dict': optim.state_dict(), |
| 36 | }, os.path.join(save_dir, '{}-model.pkl'.format(save_name))) |
| 37 | with open(os.path.join(save_dir, '{}-log.pkl'.format(save_name)), 'wb') as f: |
| 38 | pickle.dump(log, f) |
| 39 | if def_noise is not None: |
| 40 | def_noise = (def_noise * 255).round() |
| 41 | assert (def_noise.max()<=127 and def_noise.min()>=-128) |
| 42 | def_noise = def_noise.astype(np.int8) |
| 43 | with open(os.path.join(save_dir, '{}-def-noise.pkl'.format(save_name)), 'wb') as f: |
| 44 | pickle.dump(def_noise, f) |
| 45 | |
| 46 | |
| 47 | def main(args, logger): |