diff --git a/mostlyai/sdk/_local/execution/jobs.py b/mostlyai/sdk/_local/execution/jobs.py index f374ced6..1aca9bc4 100644 --- a/mostlyai/sdk/_local/execution/jobs.py +++ b/mostlyai/sdk/_local/execution/jobs.py @@ -122,20 +122,25 @@ def _move_generation_artefacts(synthetic_dataset_dir: Path, job_workspace_dir: P def _mark_in_progress(resource: Generator | SyntheticDataset, resource_dir: Path): if isinstance(resource, Generator): resource.training_status = ProgressStatus.in_progress + write_resource_to_json = write_generator_to_json else: resource.generation_status = ProgressStatus.in_progress + write_resource_to_json = write_synthetic_dataset_to_json job_progress = read_job_progress_from_json(resource_dir) job_progress.status = ProgressStatus.in_progress for step in job_progress.steps: step.status = ProgressStatus.in_progress + write_resource_to_json(resource_dir, resource) write_job_progress_to_json(resource_dir, job_progress) def _mark_done(resource: Generator | SyntheticDataset, resource_dir: Path): if isinstance(resource, Generator): resource.training_status = ProgressStatus.done + write_resource_to_json = write_generator_to_json else: resource.generation_status = ProgressStatus.done + write_resource_to_json = write_synthetic_dataset_to_json now = get_current_utc_time() job_progress = read_job_progress_from_json(resource_dir) job_progress.status = ProgressStatus.done @@ -147,19 +152,23 @@ def _mark_done(resource: Generator | SyntheticDataset, resource_dir: Path): step.start_date = step.start_date or now step.end_date = step.end_date or now step.progress.value = step.progress.max + write_resource_to_json(resource_dir, resource) write_job_progress_to_json(resource_dir, job_progress) def _mark_failed(resource: Generator | SyntheticDataset, resource_dir: Path): if isinstance(resource, Generator): resource.training_status = ProgressStatus.failed + write_resource_to_json = write_generator_to_json else: resource.generation_status = ProgressStatus.failed + write_resource_to_json = write_synthetic_dataset_to_json job_progress = read_job_progress_from_json(resource_dir) job_progress.status = ProgressStatus.failed for step in job_progress.steps: if step.status != ProgressStatus.done: step.status = ProgressStatus.failed + write_resource_to_json(resource_dir, resource) write_job_progress_to_json(resource_dir, job_progress) @@ -607,7 +616,6 @@ def execute_training_job(generator_id: str, home_dir: Path): finally: execution.clear_job_workspace() execution.clear_file_upload_connectors() - write_generator_to_json(generator_dir, generator) try: _probe_random_samples(home_dir=home_dir, generator=generator) @@ -652,7 +660,6 @@ def execute_generation_job(synthetic_dataset_id: str, home_dir: Path): finally: execution.clear_job_workspace() execution.clear_file_upload_connectors() - write_synthetic_dataset_to_json(synthetic_dataset_dir, synthetic_dataset) def execute_probing_job(synthetic_dataset_id: str, home_dir: Path) -> list[Probe]: diff --git a/mostlyai/sdk/_local/progress.py b/mostlyai/sdk/_local/progress.py index 2f3016fd..eb8f784e 100644 --- a/mostlyai/sdk/_local/progress.py +++ b/mostlyai/sdk/_local/progress.py @@ -13,14 +13,13 @@ # limitations under the License. import datetime -import json import math from pathlib import Path from pydantic import BaseModel -from mostlyai.sdk._local.storage import write_to_json -from mostlyai.sdk.domain import JobProgress, ProgressStatus, StepCode +from mostlyai.sdk._local.storage import read_job_progress_from_json, write_job_progress_to_json +from mostlyai.sdk.domain import ProgressStatus, StepCode def get_current_utc_time() -> datetime.datetime: @@ -43,8 +42,7 @@ def __init__(self, resource_path: Path, model_label: str | None, step_code: Step self._last_send_progress_time = None self._total = None - self.progress_file = self.resource_path / "job_progress.json" - self.job_progress = JobProgress(**json.loads(self.progress_file.read_text())) + self.job_progress = read_job_progress_from_json(self.resource_path) def _check_elapsed_interval(self): now = get_current_utc_time() @@ -119,9 +117,8 @@ def __call__( self.job_progress.progress.max = len(self.job_progress.steps) # of steps if self.job_progress.start_date is None: self.job_progress.start_date = now - if self.job_progress.progress.value >= self.job_progress.progress.max: - self.job_progress.end_date = now - self.job_progress.status = ProgressStatus.done + # NOTE: do not set job status to DONE here, that should happen as the very last thing + # in order for `job_wait` not to finish polling prematurely # send progress if we are DONE, or if we have a message to pass, # or if enough time has passed since last progress update @@ -133,6 +130,6 @@ def __call__( or (completed is not None and elapsed_enough_time) or (increase_by > 0 and elapsed_enough_time) ): - write_to_json(self.progress_file, self.job_progress) + write_job_progress_to_json(self.resource_path, self.job_progress) return {} diff --git a/mostlyai/sdk/client/_utils.py b/mostlyai/sdk/client/_utils.py index a020f0e9..3b2f04b4 100644 --- a/mostlyai/sdk/client/_utils.py +++ b/mostlyai/sdk/client/_utils.py @@ -288,18 +288,16 @@ def _rich_table_insert_row(*renderables: RenderableType | None, table: Table, id if step.status in (ProgressStatus.failed, ProgressStatus.canceled): rich.print(f"[red]Step {step.model_label} {step.step_code.value} {step.status.lower()}") return - # check whether we are done - if job.progress.value >= job.progress.max: - live.refresh() - time.sleep(1) # give the system a moment to update the status - return - else: - if job.end_date or job.progress in ( - ProgressStatus.failed, - ProgressStatus.canceled, - ): - rich.print(f"Job {job.status.lower()}") - return + + # check whether we are done + if job.status in ( + ProgressStatus.done, + ProgressStatus.failed, + ProgressStatus.canceled, + ): + live.refresh() + time.sleep(1) # give the system a moment to update the status + return except KeyboardInterrupt: rich.print(f"[red]Step {step.model_label} {step.step_code.value} {step.status.lower()}") return