| 55 | |
| 56 | |
| 57 | def save_checkpoint(save_dir, save_name, model, optim, log, def_noise=None): |
| 58 | torch.save({ |
| 59 | 'model_state_dict': utils.get_model_state(model), |
| 60 | 'optim_state_dict': optim.state_dict(), |
| 61 | }, os.path.join(save_dir, '{}-model.pkl'.format(save_name))) |
| 62 | with open(os.path.join(save_dir, '{}-log.pkl'.format(save_name)), 'wb') as f: |
| 63 | pickle.dump(log, f) |
| 64 | if def_noise is not None: |
| 65 | def_noise = (def_noise * 255).round() |
| 66 | assert (def_noise.max()<=127 and def_noise.min()>=-128) |
| 67 | def_noise = def_noise.astype(np.int8) |
| 68 | with open(os.path.join(save_dir, '{}-def-noise.pkl'.format(save_name)), 'wb') as f: |
| 69 | pickle.dump(def_noise, f) |
| 70 | |
| 71 | |
| 72 | def main(args, logger): |