Initialize checkpoint manager. Args: work_dir: Working directory for checkpoint file. checkpoint_name: Name of checkpoint file. logger: Logger instance.
(
self, work_dir: str, checkpoint_name: str = "checkpoint.json", logger=None
)
| 15 | """ |
| 16 | |
| 17 | def __init__( |
| 18 | self, work_dir: str, checkpoint_name: str = "checkpoint.json", logger=None |
| 19 | ): |
| 20 | """Initialize checkpoint manager. |
| 21 | |
| 22 | Args: |
| 23 | work_dir: Working directory for checkpoint file. |
| 24 | checkpoint_name: Name of checkpoint file. |
| 25 | logger: Logger instance. |
| 26 | """ |
| 27 | self.work_dir = work_dir |
| 28 | self.checkpoint_file = os.path.join(work_dir, checkpoint_name) |
| 29 | self.logger = logger |
| 30 | self.checkpoints: Dict[str, dict] = {} |
| 31 | |
| 32 | os.makedirs(work_dir, exist_ok=True) |
| 33 | self._load() |
| 34 | |
| 35 | def _load(self) -> None: |
| 36 | """Load checkpoint from file.""" |