(self, checkpoint_dir, log_dir, log_dir_link, infor="", metric=None)
| 134 | link_file(source, target) |
| 135 | |
| 136 | def save_and_link_checkpoint(self, checkpoint_dir, log_dir, log_dir_link, infor="", metric=None): |
| 137 | assert metric is not None |
| 138 | ensure_dir(checkpoint_dir) |
| 139 | if not osp.exists(log_dir_link): |
| 140 | link_file(log_dir, log_dir_link) |
| 141 | self.checkpoint_state.append({"epoch": self.state.epoch, "metric": metric}) |
| 142 | self.checkpoint_state.sort(key=lambda x: x["metric"], reverse=True) |
| 143 | if len(self.checkpoint_state) > 5: |
| 144 | try: |
| 145 | os.remove( |
| 146 | osp.join( |
| 147 | checkpoint_dir, |
| 148 | f"epoch-{self.checkpoint_state[-1]['epoch']}_miou_{self.checkpoint_state[-1]['metric']}.pth", |
| 149 | ) |
| 150 | ) |
| 151 | logger.info(f"remove inferior checkpoint: {self.checkpoint_state[-1]}") |
| 152 | except: |
| 153 | pass |
| 154 | self.checkpoint_state.pop() |
| 155 | checkpoint = osp.join(checkpoint_dir, f"epoch-{self.state.epoch}{infor}.pth") |
| 156 | self.save_checkpoint(checkpoint) |
| 157 | |
| 158 | def restore_checkpoint(self): |
| 159 | t_start = time.time() |
no test coverage detected