| 51 | self.init_test() |
| 52 | |
| 53 | def init_test(self,): |
| 54 | self.test = True |
| 55 | save_dir = os.path.join(self.logdir, "test") |
| 56 | if 'ckpt' in self.test_args: |
| 57 | ckpt_name = os.path.basename(self.test_args.ckpt).split('.ckpt')[0] + f'_epoch{self._cur_epoch}' |
| 58 | self.root = os.path.join(save_dir, ckpt_name) |
| 59 | else: |
| 60 | self.root = save_dir |
| 61 | if 'test_subdir' in self.test_args: |
| 62 | self.root = os.path.join(save_dir, self.test_args.test_subdir) |
| 63 | |
| 64 | self.root_zs = os.path.join(self.root, "zs") |
| 65 | self.root_dec = os.path.join(self.root, "reconstructions") |
| 66 | self.root_inputs = os.path.join(self.root, "inputs") |
| 67 | os.makedirs(self.root, exist_ok=True) |
| 68 | |
| 69 | if self.test_args.save_z: |
| 70 | os.makedirs(self.root_zs, exist_ok=True) |
| 71 | if self.test_args.save_reconstruction: |
| 72 | os.makedirs(self.root_dec, exist_ok=True) |
| 73 | if self.test_args.save_input: |
| 74 | os.makedirs(self.root_inputs, exist_ok=True) |
| 75 | assert(self.test_args is not None) |
| 76 | self.test_maximum = getattr(self.test_args, 'test_maximum', None) |
| 77 | self.count = 0 |
| 78 | self.eval_metrics = {} |
| 79 | self.decodes = [] |
| 80 | self.save_decode_samples = 2048 |
| 81 | |
| 82 | def init_from_ckpt(self, path, ignore_keys=list()): |
| 83 | sd = torch.load(path, map_location="cpu") |