diff --git a/src/spindle/control_plane/service.py b/src/spindle/control_plane/service.py index e1186df..e61af5d 100644 --- a/src/spindle/control_plane/service.py +++ b/src/spindle/control_plane/service.py @@ -315,6 +315,8 @@ async def training_runs(self) -> list[dict[str, object]]: async def training_run(self, training_run_id: str) -> dict[str, object]: entries = await self.checkpoints(training_run_id) + if not entries and training_run_id.endswith(":train:0"): + entries = await self.checkpoints(training_run_id.removesuffix(":train:0")) if not entries: raise RecordNotFound("training run", training_run_id) return self._training_run(entries) diff --git a/tests/control_plane/test_models.py b/tests/control_plane/test_models.py index 44f50eb..de1d5ea 100644 --- a/tests/control_plane/test_models.py +++ b/tests/control_plane/test_models.py @@ -160,6 +160,8 @@ async def run() -> None: assert run_a["is_lora"] is True assert run_a["lora_rank"] == 32 assert run_a["last_checkpoint"]["tinker_path"] == "tinker://run-a/weights/newer" + sampler_run = await plane.training_run("run-a:train:0") + assert sampler_run["training_run_id"] == "run-a" runs = await plane.training_runs() assert [run["training_run_id"] for run in runs] == ["run-a", "run-b"] assert runs[1]["is_lora"] is False