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()