| 65 | |
| 66 | |
| 67 | def save_checkpoint(save_dir, save_name, model, optim, log, def_noise=None): |
| 68 | torch.save({ |
| 69 | 'model_state_dict': utils.get_model_state(model), |
| 70 | 'optim_state_dict': optim.state_dict(), |
| 71 | }, os.path.join(save_dir, '{}-model.pkl'.format(save_name))) |
| 72 | with open(os.path.join(save_dir, '{}-log.pkl'.format(save_name)), 'wb') as f: |
| 73 | pickle.dump(log, f) |
| 74 | if def_noise is not None: |
| 75 | # def_noise = (def_noise * 255).round() |
| 76 | # def_noise = def_noise * 255 |
| 77 | # def_noise = def_noise.round() |
| 78 | # assert (def_noise.max()<=127 and def_noise.min()>=-128) |
| 79 | # def_noise = def_noise.astype(np.int8) |
| 80 | with open(os.path.join(save_dir, '{}-def-noise.pkl'.format(save_name)), 'wb') as f: |
| 81 | pickle.dump(def_noise, f) |
| 82 | |
| 83 | |
| 84 | def main(args, logger): |