diff --git a/torchtitan/components/checkpointer/dcp.py b/torchtitan/components/checkpointer/dcp.py index f1277397a6..c64e0229ba 100644 --- a/torchtitan/components/checkpointer/dcp.py +++ b/torchtitan/components/checkpointer/dcp.py @@ -65,6 +65,11 @@ class SaveDone: pass +def _parse_checkpoint_step(dirname: str) -> int | None: + match = re.fullmatch(r"step-(\d+)", dirname) + return int(match.group(1)) if match else None + + class _FilesystemCheckpointStorage: """``CheckpointStorage`` backed by ``torchtitan.tools.filesystem``. @@ -688,12 +693,10 @@ def _find_load_step(self, folder: str = "") -> int: if not self._storage.isdir(folder): return -1 - pattern = r"step-(\d+)" valid_steps = [] for filename in self._storage.listdir(folder): - match = re.search(pattern, filename) - if not match: + if (step := _parse_checkpoint_step(filename)) is None: continue # A checkpoint is valid only if it contains core metadata @@ -704,7 +707,7 @@ def _find_load_step(self, folder: str = "") -> int: ) if is_dcp or is_hf: - valid_steps.append(int(match.group(1))) + valid_steps.append(step) return max(valid_steps) if valid_steps else -1 @@ -840,10 +843,9 @@ def _purge_stale_checkpoints(self): if self._should_purge(): discovered_checkpoints = [] for filename in self._storage.listdir(self.folder): - match = re.search(r"step-(\d+)", filename) - if match: + if (step := _parse_checkpoint_step(filename)) is not None: path = filesystem.join(self.folder, filename) - discovered_checkpoints.append((int(match.group(1)), path)) + discovered_checkpoints.append((step, path)) discovered_checkpoints.sort() to_delete = discovered_checkpoints[: -1 * self.keep_latest_k]