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