from __future__ import annotations import math import re import time from dataclasses import dataclass from datetime import datetime from typing import Any from adam.models import Job, SystemSnapshot @dataclass(slots=True) class AtlasDecision: severity: str message: str action: str = "none" class AtlasSupervisor: """Stateful, conservative runtime guard for active training jobs.""" def __init__(self, config: Any | None = None) -> None: getter = config.get if config is not None else lambda _key, default: default self.warning_temp = float(getter("atlas_warning_temperature", 82)) self.critical_temp = float(getter("atlas_critical_temperature", 90)) self.warning_disk_gb = float(getter("atlas_warning_disk_gb", 10)) self.critical_disk_gb = float(getter("atlas_critical_disk_gb", 1)) self.stall_minutes = float(getter("atlas_stall_minutes", 30)) self._state: dict[str, dict[str, Any]] = {} @staticmethod def _is_training(job: Job) -> bool: return 0 <= job.current_step < len(job.plan.steps) and job.plan.steps[job.current_step].tool_id.endswith("_trainer") def observe(self, job: Job, snapshot: SystemSnapshot, *, now: float | None = None) -> AtlasDecision: if not self._is_training(job): return AtlasDecision("info", "ATLAS is standing by; the active step is not training.") moment = time.monotonic() if now is None else now state = self._state.setdefault(job.id, { "progress": job.progress, "changed_at": moment, "hot_samples": 0, "disk_samples": 0, "log_index": 0, }) if job.progress != state["progress"]: state["progress"] = job.progress state["changed_at"] = moment new_logs = job.logs[int(state["log_index"]):] state["log_index"] = len(job.logs) recent = "\n".join(new_logs) if re.search(r"\bloss\s*[:=]?\s*(?:nan|[+-]?inf)\b", recent, re.I): return AtlasDecision("critical", "Non-finite loss detected. ATLAS paused the job for review.", "pause") temperature = snapshot.gpu_temperature state["hot_samples"] = state["hot_samples"] + 1 if temperature is not None and temperature >= self.critical_temp else 0 if state["hot_samples"] >= 3: return AtlasDecision("critical", f"GPU temperature remained at {temperature:.0f}°C. ATLAS paused the job.", "pause") free_disk = max(0.0, snapshot.disk_total_gb - snapshot.disk_used_gb) state["disk_samples"] = state["disk_samples"] + 1 if snapshot.disk_total_gb and free_disk <= self.critical_disk_gb else 0 if state["disk_samples"] >= 2: return AtlasDecision("critical", f"Only {free_disk:.1f} GB remains on the output drive. ATLAS paused the job.", "pause") stalled_minutes = (moment - state["changed_at"]) / 60 if stalled_minutes >= self.stall_minutes and snapshot.gpu_percent < 5: return AtlasDecision("warning", f"No recorded progress and little GPU activity for {stalled_minutes:.0f} minutes. Check the trainer.") if temperature is not None and temperature >= self.warning_temp: return AtlasDecision("warning", f"GPU temperature is elevated at {temperature:.0f}°C; ATLAS is watching it closely.") if snapshot.disk_total_gb and free_disk <= self.warning_disk_gb: return AtlasDecision("warning", f"Output drive space is getting low ({free_disk:.1f} GB free).") if snapshot.memory_percent >= 95: return AtlasDecision("warning", f"System memory usage is very high at {snapshot.memory_percent:.0f}%.") try: started = datetime.fromisoformat(job.started_at) if job.started_at else None elapsed_minutes = (datetime.now(started.tzinfo) - started).total_seconds() / 60 if started else 0 except ValueError: elapsed_minutes = 0 expected = float(job.plan.orion_review.get("estimated_high_minutes", 0) or 0) if expected and elapsed_minutes > expected * 2 and job.progress < 90: return AtlasDecision("warning", "Runtime is now more than twice ORION's broad estimate. The job is still running.") return AtlasDecision("healthy", "Training behavior is within the current safety limits.") def forget(self, job_id: str) -> None: self._state.pop(job_id, None)