| from __future__ import annotations |
|
|
| import json |
| import logging |
| import math |
| import re |
| import threading |
| from datetime import datetime, timezone |
| from pathlib import Path |
| from typing import Any |
|
|
| from PySide6.QtCore import QObject, QThread, Signal |
|
|
| from adam.executor import ToolCancelled, ToolExecutionError, ToolExecutor |
| from adam.assets import AssetRegistry |
| from adam.atlas import AtlasSupervisor |
| from adam.models import ExecutionPlan, Job, JobStatus, StepStatus, utc_now |
| from adam.nova import evaluate_job_output |
|
|
|
|
| class JobWorker(QThread): |
| event = Signal(object) |
|
|
| def __init__(self, job: Job, executor: ToolExecutor) -> None: |
| super().__init__() |
| self.job = job |
| self.executor = executor |
| self.cancel_event = threading.Event() |
| self.run_event = threading.Event() |
| self.run_event.set() |
|
|
| def pause(self) -> None: |
| self.run_event.clear() |
|
|
| def resume(self) -> None: |
| self.run_event.set() |
|
|
| def cancel(self) -> None: |
| self.cancel_event.set() |
| self.run_event.set() |
|
|
| def run(self) -> None: |
| total_steps = len(self.job.plan.steps) |
| try: |
| for index, step in enumerate(self.job.plan.steps): |
| if self.cancel_event.is_set(): |
| raise ToolCancelled("Job cancelled by user.") |
| self.event.emit( |
| { |
| "type": "step_started", |
| "index": index, |
| "message": step.title, |
| } |
| ) |
|
|
| preview_state = {"epoch": 0, "path": ""} |
|
|
| def on_progress(percent: int, message: str, step_index: int = index) -> None: |
| overall = int(((step_index + percent / 100) / total_steps) * 100) |
| self.event.emit( |
| { |
| "type": "progress", |
| "step_percent": percent, |
| "overall": overall, |
| "message": message, |
| } |
| ) |
| self._discover_external_preview(step, message, preview_state) |
|
|
| def on_log(message: str) -> None: |
| self.event.emit({"type": "log", "message": message}) |
|
|
| def on_preview(preview: dict[str, Any]) -> None: |
| preview_state["epoch"] = int(preview.get("epoch", 0) or 0) |
| preview_state["path"] = str(preview.get("path", "")) |
| self.event.emit({"type": "preview", **preview}) |
|
|
| result = self.executor.execute( |
| step.tool_id, |
| step.arguments, |
| job_id=self.job.id, |
| cancel_event=self.cancel_event, |
| run_event=self.run_event, |
| progress_callback=on_progress, |
| log_callback=on_log, |
| preview_callback=on_preview, |
| ) |
| self.event.emit( |
| { |
| "type": "step_finished", |
| "index": index, |
| "result": result, |
| } |
| ) |
| self.event.emit({"type": "completed"}) |
| except ToolCancelled as exc: |
| self.event.emit({"type": "cancelled", "message": str(exc)}) |
| except Exception as exc: |
| self.event.emit( |
| { |
| "type": "failed", |
| "message": str(exc), |
| "exception": type(exc).__name__, |
| } |
| ) |
|
|
| def _discover_external_preview( |
| self, step: Any, message: str, state: dict[str, Any] |
| ) -> None: |
| """Discover conventional preview files from any registered trainer.""" |
| if not step.tool_id.endswith("_trainer") or not step.arguments.get("preview_enabled", False): |
| return |
| match = re.search(r"\bepoch\s+(\d+)(?:\s+(?:of|/)|/)?", message, re.I) |
| if not match: |
| return |
| epoch = int(match.group(1)) |
| interval = max(1, int(step.arguments.get("preview_every", 5) or 5)) |
| if epoch % interval or epoch == int(state.get("epoch", 0)): |
| return |
| output = Path(str(step.arguments.get("output_dir", ""))).expanduser() |
| if not output.is_dir(): |
| return |
| try: |
| candidates = [ |
| path for path in output.rglob("*") |
| if path.is_file() and path.suffix.lower() in {".png", ".jpg", ".jpeg", ".webp", ".bmp"} |
| and any(token in path.name.lower() for token in ("preview", "sample", "epoch")) |
| ] |
| latest = max(candidates, key=lambda path: path.stat().st_mtime) if candidates else None |
| except OSError: |
| latest = None |
| if latest and str(latest) != state.get("path"): |
| state["epoch"] = epoch |
| state["path"] = str(latest) |
| self.event.emit({ |
| "type": "preview", "path": str(latest), "epoch": epoch, |
| "next_epoch": min(int(step.arguments.get("epochs", epoch + interval)), epoch + interval), |
| "prompt": str(step.arguments.get("preview_prompt", "")), |
| "seed": step.arguments.get("preview_seed"), |
| "steps": int(step.arguments.get("preview_steps", 0) or 0), |
| }) |
|
|
|
|
| class JobManager(QObject): |
| job_created = Signal(object) |
| job_updated = Signal(object) |
| log_added = Signal(str, str) |
| active_changed = Signal(object) |
| notification = Signal(str, str) |
|
|
| def __init__( |
| self, |
| root: Path, |
| executor: ToolExecutor, |
| logger: logging.Logger, |
| config: Any | None = None, |
| ) -> None: |
| super().__init__() |
| self.root = root |
| self.executor = executor |
| self.logger = logger |
| self.jobs_path = root / "data" / "jobs.json" |
| self.assets = AssetRegistry(root) |
| self.atlas = AtlasSupervisor(config) |
| self.jobs: list[Job] = [] |
| self._queue: list[str] = [] |
| self._worker: JobWorker | None = None |
| self._active_job: Job | None = None |
| self._load() |
|
|
| @property |
| def active_job(self) -> Job | None: |
| return self._active_job |
|
|
| def submit(self, plan: ExecutionPlan) -> Job: |
| status = ( |
| JobStatus.AWAITING_CONFIRMATION |
| if plan.requires_confirmation |
| else JobStatus.QUEUED |
| ) |
| job = Job(plan=plan, status=status) |
| self.jobs.insert(0, job) |
| self._append_log(job, f"Plan created: {plan.summary}") |
| if plan.requires_confirmation: |
| self._append_log(job, "Waiting for user confirmation.") |
| else: |
| self._queue.append(job.id) |
| self._save() |
| self.job_created.emit(job) |
| self.job_updated.emit(job) |
| if not plan.requires_confirmation: |
| self._start_next() |
| return job |
|
|
| def confirm(self, job_id: str) -> None: |
| job = self.get(job_id) |
| if job.status != JobStatus.AWAITING_CONFIRMATION: |
| return |
| job.status = JobStatus.QUEUED |
| self._append_log(job, "Plan approved by user.") |
| self._queue.append(job.id) |
| self._save() |
| self.job_updated.emit(job) |
| self._start_next() |
|
|
| def reject(self, job_id: str) -> None: |
| job = self.get(job_id) |
| if job.status != JobStatus.AWAITING_CONFIRMATION: |
| return |
| job.status = JobStatus.CANCELLED |
| job.ended_at = utc_now() |
| self._append_log(job, "Plan cancelled before execution.") |
| self._save() |
| self.job_updated.emit(job) |
|
|
| def pause(self, job_id: str) -> None: |
| job = self.get(job_id) |
| if job is self._active_job and job.status == JobStatus.RUNNING and self._worker: |
| self._worker.pause() |
| job.status = JobStatus.PAUSED |
| self._append_log(job, "Job paused.") |
| self._save() |
| self.job_updated.emit(job) |
|
|
| def resume(self, job_id: str) -> None: |
| job = self.get(job_id) |
| if job is self._active_job and job.status == JobStatus.PAUSED and self._worker: |
| self._worker.resume() |
| job.status = JobStatus.RUNNING |
| self._append_log(job, "Job resumed.") |
| self._save() |
| self.job_updated.emit(job) |
|
|
| def cancel(self, job_id: str) -> None: |
| job = self.get(job_id) |
| if job is self._active_job and self._worker: |
| self._append_log(job, "Cancellation requested.") |
| self._worker.cancel() |
| return |
| if job.id in self._queue: |
| self._queue.remove(job.id) |
| if job.status in { |
| JobStatus.QUEUED, |
| JobStatus.AWAITING_CONFIRMATION, |
| JobStatus.DRAFT, |
| }: |
| job.status = JobStatus.CANCELLED |
| job.ended_at = utc_now() |
| self._append_log(job, "Job cancelled.") |
| self._save() |
| self.job_updated.emit(job) |
|
|
| def get(self, job_id: str) -> Job: |
| for job in self.jobs: |
| if job.id == job_id: |
| return job |
| raise KeyError(f"Unknown job: {job_id}") |
|
|
| def retry(self, job_id: str) -> Job: |
| """Create an approval-gated retry, resuming interrupted DDPM work when possible.""" |
| original = self.get(job_id) |
| plan = ExecutionPlan.from_dict(original.to_dict()["plan"]) |
| resume_note = "" |
| if original.status == JobStatus.INTERRUPTED: |
| start_index = max(0, min(original.current_step, len(plan.steps) - 1)) |
| plan.steps = plan.steps[start_index:] |
| if plan.steps: |
| step = plan.steps[0] |
| resume_note = self._prepare_ddpm_resume(step.arguments, step.tool_id) |
| for step in plan.steps: |
| step.status = StepStatus.PENDING |
| plan.id = original.plan.id + "-retry" |
| plan.created_at = utc_now() |
| plan.requires_confirmation = True |
| plan.confirmation_reason = ( |
| "This retries a previous job. Review paths, checkpoints, and settings " |
| "because files or available resources may have changed." |
| ) |
| retry = self.submit(plan) |
| if resume_note: |
| self._append_log(retry, resume_note) |
| return retry |
|
|
| @staticmethod |
| def _prepare_ddpm_resume(arguments: dict[str, Any], tool_id: str) -> str: |
| """Attach the newest complete Accelerate checkpoint to a DDPM retry.""" |
| if tool_id != "ddpm_trainer": |
| return "" |
| output = Path(str(arguments.get("output_dir", ""))).expanduser() |
| dataset = Path(str(arguments.get("dataset_dir", ""))).expanduser() |
| try: |
| checkpoints = sorted( |
| (path for path in output.glob("checkpoint-*") if path.is_dir() |
| and (path / "unet" / "diffusion_pytorch_model.safetensors").is_file() |
| and (path / "optimizer.bin").is_file() and (path / "scheduler.bin").is_file()), |
| key=lambda path: int(path.name.rsplit("-", 1)[-1]), |
| ) |
| image_count = sum(1 for item in dataset.iterdir() if item.is_file() and item.suffix.lower() in {".png", ".jpg", ".jpeg", ".webp", ".bmp"}) |
| batch_size = max(1, int(arguments.get("batch_size", 1))) |
| completed_epochs = int(checkpoints[-1].name.rsplit("-", 1)[-1]) // max(1, math.ceil(image_count / batch_size)) |
| remaining_epochs = int(arguments.get("epochs", 0)) - completed_epochs |
| except (IndexError, OSError, ValueError, TypeError): |
| return "" |
| if remaining_epochs <= 0: |
| return "" |
| arguments["resume_from"] = str(checkpoints[-1]) |
| arguments["epochs"] = remaining_epochs |
| return f"Resuming interrupted DDPM training from {checkpoints[-1].name} (about epoch {completed_epochs}; {remaining_epochs} epochs remaining)." |
|
|
| def end_task(self, job_id: str) -> bool: |
| """Acknowledge an interrupted job and leave it inactive in history. |
| |
| This is intentionally limited to interrupted jobs: active work must still use |
| ``cancel`` so its worker receives the cancellation signal. |
| """ |
| job = self.get(job_id) |
| if job.status != JobStatus.INTERRUPTED: |
| return False |
| job.status = JobStatus.CANCELLED |
| job.ended_at = utc_now() |
| job.logs = [ |
| line |
| for line in job.logs |
| if not line.startswith("[startup] Previous session ended before this job.") |
| ] |
| self._append_log(job, "Interrupted job ended by user; no retry is pending.") |
| self._save() |
| self.job_updated.emit(job) |
| return True |
|
|
| def remove_completed_or_failed(self) -> int: |
| """Remove completed and failed history records without touching output files.""" |
| removable = {JobStatus.FINISHED, JobStatus.FAILED} |
| before = len(self.jobs) |
| self.jobs = [job for job in self.jobs if job.status not in removable] |
| removed = before - len(self.jobs) |
| if removed: |
| self._save() |
| return removed |
|
|
| def _start_next(self) -> None: |
| if self._worker is not None and self._worker.isRunning(): |
| return |
| while self._queue: |
| job_id = self._queue.pop(0) |
| job = self.get(job_id) |
| if job.status != JobStatus.QUEUED: |
| continue |
| self._active_job = job |
| job.status = JobStatus.RUNNING |
| job.started_at = utc_now() |
| self._append_log(job, f"Job {job.id} started.") |
| self._worker = JobWorker(job, self.executor) |
| self._worker.event.connect(self._handle_event) |
| self._worker.finished.connect(self._worker_finished) |
| self._save() |
| self.active_changed.emit(job) |
| self.job_updated.emit(job) |
| self.notification.emit("Training started" if self._has_training(job) else "Job started", job.plan.project_name) |
| self._worker.start() |
| return |
| self._active_job = None |
| self.active_changed.emit(None) |
|
|
| def _handle_event(self, event: dict[str, Any]) -> None: |
| job = self._active_job |
| if job is None: |
| return |
| event_type = event.get("type") |
| if event_type == "step_started": |
| index = int(event["index"]) |
| job.current_step = index |
| job.plan.steps[index].status = StepStatus.RUNNING |
| if job.plan.steps[index].tool_id.endswith("_trainer"): |
| job.preview_path = None |
| job.preview_epoch = 0 |
| job.preview_next_epoch = 0 |
| job.preview_prompt = "" |
| job.preview_seed = None |
| job.preview_steps = 0 |
| job.preview_kind = "training" |
| job.preview_current = 0 |
| job.preview_total = 0 |
| job.preview_image_index = 0 |
| job.preview_image_count = 0 |
| self._append_log(job, f"Starting: {job.plan.steps[index].title}") |
| elif event_type == "progress": |
| job.progress = int(event["overall"]) |
| message = str(event["message"]) |
| if message and (not job.logs or message not in job.logs[-1]): |
| self._append_log(job, message) |
| elif event_type == "log": |
| self._append_log(job, str(event["message"])) |
| elif event_type == "preview": |
| job.preview_path = str(event.get("path", "")) or None |
| job.preview_epoch = int(event.get("epoch", 0) or 0) |
| job.preview_next_epoch = int(event.get("next_epoch", 0) or 0) |
| job.preview_prompt = str(event.get("prompt", "")) |
| seed = event.get("seed") |
| job.preview_seed = int(seed) if seed is not None else None |
| job.preview_steps = int(event.get("steps", 0) or 0) |
| job.preview_kind = str(event.get("kind", "training")) |
| job.preview_current = int(event.get("current", 0) or 0) |
| job.preview_total = int(event.get("total", 0) or 0) |
| job.preview_image_index = int(event.get("image_index", 0) or 0) |
| job.preview_image_count = int(event.get("image_count", 0) or 0) |
| label = "Denoising" if job.preview_kind == "generation" else "Training" |
| position = f" step {job.preview_current}" if job.preview_current else f" epoch {job.preview_epoch}" |
| self._append_log(job, f"{label} preview updated at{position}.") |
| elif event_type == "step_finished": |
| index = int(event["index"]) |
| job.plan.steps[index].status = StepStatus.FINISHED |
| result = event.get("result") or {} |
| self.assets.ingest_result(result) |
| if result.get("output_folder"): |
| job.output_folder = str(result["output_folder"]) |
| if job.plan.steps[index].tool_id.endswith("_trainer"): |
| evaluation = evaluate_job_output(job) |
| if evaluation: |
| evaluation["step"] = job.plan.steps[index].title |
| evaluation["model_name"] = str(result.get("model_name", "")) |
| reports = list(job.nova_report.get("evaluations", [])) |
| reports.append(evaluation) |
| job.nova_report = { |
| "agent": "NOVA", "evaluations": reports, "latest": evaluation, |
| } |
| self._append_log( |
| job, |
| f"NOVA — {evaluation['status']}: {evaluation['summary']}", |
| ) |
| self._append_log(job, f"Finished: {job.plan.steps[index].title}") |
| elif event_type == "completed": |
| job.status = JobStatus.FINISHED |
| job.progress = 100 |
| job.ended_at = utc_now() |
| self._append_log(job, "Job finished successfully.") |
| self.notification.emit("Job complete", job.plan.project_name) |
| elif event_type == "cancelled": |
| job.status = JobStatus.CANCELLED |
| job.ended_at = utc_now() |
| self._append_log(job, str(event.get("message", "Job cancelled."))) |
| for step in job.plan.steps: |
| if step.status == StepStatus.RUNNING: |
| step.status = StepStatus.SKIPPED |
| self.notification.emit("Job cancelled", job.plan.project_name) |
| elif event_type == "failed": |
| job.status = JobStatus.FAILED |
| job.ended_at = utc_now() |
| job.error = str(event.get("message", "Unknown error")) |
| if 0 <= job.current_step < len(job.plan.steps): |
| job.plan.steps[job.current_step].status = StepStatus.FAILED |
| self._append_log(job, f"Failed: {job.error}") |
| self.logger.error( |
| "Job %s failed (%s): %s", |
| job.id, |
| event.get("exception"), |
| job.error, |
| ) |
| self.notification.emit("Job failed", job.error) |
| self._save() |
| self.job_updated.emit(job) |
|
|
| def _worker_finished(self) -> None: |
| self._worker = None |
| self._active_job = None |
| self.active_changed.emit(None) |
| self._start_next() |
|
|
| def supervise(self, snapshot: Any) -> None: |
| """Let ATLAS inspect the active training run and apply critical pauses.""" |
| job = self._active_job |
| if job is None or job.status != JobStatus.RUNNING: |
| return |
| decision = self.atlas.observe(job, snapshot) |
| previous = str(job.atlas_report.get("message", "")) |
| unchanged = ( |
| previous == decision.message |
| and job.atlas_report.get("severity") == decision.severity |
| and job.atlas_report.get("action") == decision.action |
| ) |
| if unchanged: |
| return |
| job.atlas_report = { |
| "agent": "ATLAS", |
| "severity": decision.severity, |
| "message": decision.message, |
| "action": decision.action, |
| "updated_at": utc_now(), |
| } |
| if decision.message != previous and decision.severity in {"warning", "critical"}: |
| self._append_log(job, f"ATLAS {decision.severity.upper()} — {decision.message}") |
| self.notification.emit(f"ATLAS {decision.severity}", decision.message) |
| self._save() |
| self.job_updated.emit(job) |
| if decision.action == "pause" and job.status == JobStatus.RUNNING: |
| self.pause(job.id) |
|
|
| def _append_log(self, job: Job, message: str) -> None: |
| timestamp = datetime.now().strftime("%H:%M:%S") |
| line = f"[{timestamp}] {message}" |
| job.logs.append(line) |
| job.logs = job.logs[-1000:] |
| self.log_added.emit(job.id, line) |
| self.logger.info("Job %s | %s", job.id, message) |
|
|
| @staticmethod |
| def _has_training(job: Job) -> bool: |
| return any(step.tool_id.endswith("trainer") for step in job.plan.steps) |
|
|
| def _load(self) -> None: |
| if not self.jobs_path.exists(): |
| return |
| try: |
| payload = json.loads(self.jobs_path.read_text(encoding="utf-8")) |
| self.jobs = [Job.from_dict(item) for item in payload.get("jobs", [])] |
| except (OSError, ValueError, TypeError, KeyError, json.JSONDecodeError): |
| self.jobs = [] |
| return |
| for job in self.jobs: |
| if job.status in {JobStatus.RUNNING, JobStatus.PAUSED, JobStatus.QUEUED}: |
| job.status = JobStatus.INTERRUPTED |
| job.ended_at = utc_now() |
| job.logs.append( |
| "[startup] Previous session ended before this job. " |
| "Review it before retrying." |
| ) |
| self._save() |
|
|
| def _save(self) -> None: |
| self.jobs_path.parent.mkdir(parents=True, exist_ok=True) |
| temporary = self.jobs_path.with_suffix(".tmp") |
| temporary.write_text( |
| json.dumps({"jobs": [job.to_dict() for job in self.jobs]}, indent=2), |
| encoding="utf-8", |
| ) |
| temporary.replace(self.jobs_path) |
|
|
| def shutdown(self) -> None: |
| if self._worker is not None and self._worker.isRunning(): |
| self._worker.cancel() |
| self._worker.wait(2500) |
| if self._active_job and self._active_job.status in { |
| JobStatus.RUNNING, |
| JobStatus.PAUSED, |
| }: |
| self._active_job.status = JobStatus.CANCELLED |
| self._active_job.ended_at = utc_now() |
| self._append_log(self._active_job, "ADAM closed; the active job was stopped.") |
| self._save() |
|
|