save checkpoint for future restore
(self, name)
| 139 | self.scheduler = optim.lr_scheduler.ExponentialLR(self.optimizer, lr_decay) |
| 140 | |
| 141 | def save_ckpt(self, name): |
| 142 | """save checkpoint for future restore""" |
| 143 | save_path = os.path.join(self.log_dir, f"ckpt_{name}.pth") |
| 144 | |
| 145 | save_dict = { |
| 146 | "net": self.net.cpu().state_dict(), |
| 147 | "optimizer": self.optimizer.state_dict(), |
| 148 | "scheduler": self.scheduler.state_dict(), |
| 149 | "Ka": self.Ka, |
| 150 | "Kd": self.Kd, |
| 151 | "Ks": self.Ks, |
| 152 | "Ns": self.Ns, |
| 153 | "aabb": self.aabb.tolist(), |
| 154 | "featmap_size": self.featmap_size, |
| 155 | } |
| 156 | torch.save(save_dict, save_path) |
| 157 | self.net.to(self.device) |
| 158 | |
| 159 | def load_ckpt(self, name): |
| 160 | """load saved checkpoint""" |