SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
22.7 kB
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()