Spaces:
Running
Running
Download runner.py from armnet/armnet-eval: direct link, hf CLI and curl.
- Browser
- Download file 56.4 kB
-
https://huggingface.co/spaces/armnet/armnet-eval/resolve/main/runner.py
- Command line
-
hf download hf://spaces/armnet/armnet-eval/runner.py
-
curl -L -o runner.py https://huggingface.co/spaces/armnet/armnet-eval/resolve/main/runner.py
56.4 kB
| """Submit and follow shared demo evals, live. | |
| Jobs route to any free BusyBox cell; several can run at once. Process-global | |
| `STATE` tracks each live eval in a per-cell `RunSlot` so every Gradio session | |
| can watch the same fleet: pick a cell to see that job's video, rollouts, and | |
| status. A background worker submits to the orchestrator and, for each dispatched | |
| job, two threads consume its websockets: one decodes the Rerun video feed into | |
| a frame, the other parses per-rollout results from the log stream. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import logging | |
| import os | |
| import re | |
| import ssl | |
| import threading | |
| import time | |
| from collections import deque | |
| from dataclasses import dataclass, field | |
| from typing import Any, Optional | |
| import certifi | |
| import websocket | |
| from armnet_core import API_KEY_HEADER, Job, JobSpec, JobStatus, TerminalStatus | |
| from armnet_client._config import api_key, orchestrator_url | |
| from armnet_client.client import OrchestratorClient, OrchestratorError | |
| from armnet_client.execute import _ws_url | |
| import config | |
| import leaderboard | |
| from frames import decode_frame | |
| from leaderboard import LEADERBOARD | |
| logger = logging.getLogger(__name__) | |
| # "[armnet:progress] episode 3/10: complete=True success=False ticks=420 scored_by=busybox" | |
| _EPISODE_RE = re.compile( | |
| r"episode\s+(\d+)/(\d+):\s+complete=(\w+)\s+success=(\w+)\s+ticks=(\d+)\s+scored_by=(\w+)" | |
| ) | |
| # The HTTP poll (not the websockets) is the authoritative source of terminal | |
| # state, so it must stay resilient to transient orchestrator errors. | |
| _POLL_INTERVAL_S = 3.0 | |
| _POLL_BACKOFF_MAX_S = 15.0 | |
| # Stop following + unlock the Space after this long with no orchestrator contact, | |
| # instead of freezing for the whole job timeout. Kept short so a wedged Space | |
| # frees up quickly for the next user; the job may still be running on the cell, | |
| # we just stop watching. (Backoff between retries means the real unlock lands a | |
| # few seconds past this.) | |
| _LOST_CONTACT_GIVEUP_S = float(os.environ.get("ARMNET_DEMO_LOST_CONTACT_GIVEUP_S", "10")) | |
| # Cap on the websocket reconnect backoff so a broken/rate-limited orchestrator | |
| # isn't reconnected against every fraction of a second. | |
| _WS_RECONNECT_MAX_S = 10.0 | |
| # A cell-status call that does not return is treated as failed, and the HTTP | |
| # client is closed so the next poll cannot reuse a half-open connection. | |
| _STATUS_CALL_TIMEOUT_S = 12.0 | |
| # One blip must not paint the fleet offline. This many failed polls in a row | |
| # does, which is about a minute at the default interval. | |
| _STATUS_FAILURES_BEFORE_OFFLINE = 3 | |
| # How many undecodable packets to accept before saying so in the run log. A | |
| # handful can be lost legitimately (a truncated first packet, a dropped frame), | |
| # so this is high enough not to cry wolf and low enough to land in the first | |
| # second or two of a run. | |
| _UNDECODABLE_NOTICE_AFTER = 30 | |
| _PENDING_CELL_PREFIX = "pending:" | |
| def pending_cell_id(job_id: str) -> str: | |
| """Placeholder key until the orchestrator reports the job's assigned cell.""" | |
| return f"{_PENDING_CELL_PREFIX}{job_id}" | |
| def display_cell_id(cell_id: Optional[str]) -> Optional[str]: | |
| """A real orchestrator cell id, or None while routing is still pending.""" | |
| if not cell_id or cell_id.startswith(_PENDING_CELL_PREFIX): | |
| return None | |
| return cell_id | |
| def fleet_availability( | |
| cell_statuses: dict[str, str], *, aggregate: str = "offline" | |
| ) -> tuple[str, str]: | |
| """Dot colour and caption for the BusyBox fleet, not a named cell. | |
| Counts come from ``GET /cells/status`` ``cells[]``. The aggregate status is | |
| only used when that list is empty (first poll, or the orchestrator is | |
| unreachable). | |
| """ | |
| n_available = sum(1 for status in cell_statuses.values() if status == "available") | |
| n_busy = sum( | |
| 1 for status in cell_statuses.values() if status in {"occupied", "maintenance"} | |
| ) | |
| if cell_statuses: | |
| if n_available > 0: | |
| noun = "cell" if n_available == 1 else "cells" | |
| return "green", f"{n_available} {noun} available" | |
| if n_busy > 0: | |
| return "yellow", "0 cells available" | |
| return "red", "cells offline" | |
| if aggregate == "available": | |
| return "green", "cells available" | |
| if aggregate == "occupied": | |
| return "yellow", "0 cells available" | |
| return "red", "cells offline" | |
| class RunSlot: | |
| """Independent live state for one BusyBox cell's current/recent eval.""" | |
| cell_id: str | |
| job_id: str | |
| policy: str | |
| revision: Optional[str] | |
| task: config.Task | |
| status: str | |
| active: bool = True | |
| rollouts: list[dict[str, Any]] = field(default_factory=list) | |
| latest_frame: Any = None | |
| frame_version: int = 0 | |
| events: deque[str] = field(default_factory=lambda: deque(maxlen=200)) | |
| summary: Optional[dict] = None | |
| error: Optional[str] = None | |
| started_at: Optional[float] = None | |
| class DemoState: | |
| """Process-global, thread-safe state for shared demo evals.""" | |
| def __init__(self) -> None: | |
| self._lock = threading.Lock() | |
| self._reset(active=False) | |
| # Whether we've provisioned the HF token as an orchestrator secret this | |
| # process (done once; the token is stable for the Space's lifetime). | |
| self._eval_secret_ensured = False | |
| # Live availability per task slug ("available"/"occupied"/"offline"), | |
| # refreshed by one process-wide poller thread and shared by all viewers. | |
| # Per task because a cell serves an environment and the tasks in it can | |
| # in principle live on different cells. Starts "offline" so submission | |
| # stays blocked until the first good poll. | |
| self._cell_status = {task.slug: "offline" for task in config.TASKS.values()} | |
| self._status_failures = {task.slug: 0 for task in config.TASKS.values()} | |
| self._queued_eval_total = 0 | |
| self._queue_items: list[str] = [] | |
| self._queued_eval_policies: set[str] = set() | |
| self._pending_eval_policies: set[str] = set() | |
| self._cells_by_task: dict[str, dict[str, str]] = { | |
| task.slug: {} for task in config.TASKS.values() | |
| } | |
| self._slots_by_cell: dict[str, RunSlot] = {} | |
| self._following_job_ids: set[str] = set() | |
| self._default_cell_id: Optional[str] = None | |
| # Per Gradio session: which cell's feed that viewer last picked. | |
| self._view_by_session: dict[str, str] = {} | |
| threading.Thread(target=self._poll_cell_status_loop, daemon=True).start() | |
| def _reset(self, *, active: bool) -> None: | |
| self.active = active | |
| self.job_id: Optional[str] = None | |
| self.policy: Optional[str] = None | |
| # Which task the shared run is evaluating. Viewers can have a different | |
| # task selected while watching, so the Live tab names it explicitly. | |
| self.task: Optional[config.Task] = None | |
| # Resolved commit SHA of the policy being evaluated (for the leaderboard). | |
| self.revision: Optional[str] = None | |
| self.status: str = "idle" | |
| self.rollouts: list[dict[str, Any]] = [] | |
| self.latest_frame = None | |
| # Bumped on every new frame so the UI can skip re-sending an unchanged | |
| # image (heavy) while still pushing lightweight text updates every tick. | |
| self.frame_version = 0 | |
| self.events: deque[str] = deque(maxlen=200) | |
| self.summary: Optional[dict] = None | |
| self.error: Optional[str] = None | |
| self.started_at: Optional[float] = None | |
| # ---------------------------------------------------------------- mutations | |
| def _event(self, message: str) -> None: | |
| ts = time.strftime("%H:%M:%S") | |
| self.events.append(f"[{ts}] {message}") | |
| def set_status(self, status: str, *, job_id: Optional[str] = None) -> None: | |
| with self._lock: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| if status != slot.status: | |
| slot.status = status | |
| self._slot_event(slot, f"status: {status}") | |
| if self.job_id == job_id and status != self.status: | |
| self.status = status | |
| return | |
| if job_id is not None and self.job_id != job_id: | |
| return | |
| if status != self.status: | |
| self.status = status | |
| self._event(f"status: {status}") | |
| def set_frame( # noqa: ANN001 - numpy array | |
| self, frame, *, job_id: Optional[str] = None | |
| ) -> None: | |
| with self._lock: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| slot.latest_frame = frame | |
| slot.frame_version += 1 | |
| if self.job_id == job_id: | |
| self.latest_frame = frame | |
| self.frame_version += 1 | |
| return | |
| if job_id is not None and self.job_id != job_id: | |
| return | |
| self.latest_frame = frame | |
| self.frame_version += 1 | |
| def record_rollout( | |
| self, | |
| episode: int, | |
| total: int, | |
| success: bool, | |
| scored_by: str, | |
| *, | |
| job_id: Optional[str] = None, | |
| ) -> None: | |
| with self._lock: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| for rollout in slot.rollouts: | |
| if rollout["episode"] == episode: | |
| return | |
| slot.rollouts.append( | |
| { | |
| "episode": episode, | |
| "success": success, | |
| "scored_by": scored_by, | |
| } | |
| ) | |
| slot.rollouts.sort(key=lambda rollout: rollout["episode"]) | |
| self._slot_event( | |
| slot, | |
| f"rollout {episode}/{total}: " | |
| f"{'SUCCESS' if success else 'fail'} (scored by {scored_by})", | |
| ) | |
| if self.job_id == job_id: | |
| self.rollouts = list(slot.rollouts) | |
| return | |
| if job_id is not None and self.job_id != job_id: | |
| return | |
| for r in self.rollouts: | |
| if r["episode"] == episode: | |
| return # already recorded (duplicate log line) | |
| self.rollouts.append( | |
| {"episode": episode, "success": success, "scored_by": scored_by} | |
| ) | |
| self.rollouts.sort(key=lambda r: r["episode"]) | |
| self._event( | |
| f"rollout {episode}/{total}: {'SUCCESS' if success else 'fail'} " | |
| f"(scored by {scored_by})" | |
| ) | |
| def log_line(self, line: str, *, job_id: Optional[str] = None) -> None: | |
| text = line.rstrip("\n") | |
| if not text: | |
| return | |
| with self._lock: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| slot.events.append(text[:400]) | |
| return | |
| if job_id is not None and self.job_id != job_id: | |
| return | |
| self.events.append(text[:400]) | |
| def _slot_for_job_unlocked(self, job_id: Optional[str]) -> Optional[RunSlot]: | |
| if job_id is None: | |
| return None | |
| return next( | |
| ( | |
| slot | |
| for slot in self._slots_by_cell.values() | |
| if slot.job_id == job_id | |
| ), | |
| None, | |
| ) | |
| def _slot_event(slot: RunSlot, message: str) -> None: | |
| slot.events.append(f"[{time.strftime('%H:%M:%S')}] {message}") | |
| def set_view(self, session_id: Optional[str], cell_id: Optional[str]) -> None: | |
| """Remember which cell a Gradio session last asked to watch.""" | |
| if not session_id or not cell_id: | |
| return | |
| with self._lock: | |
| self._view_by_session[session_id] = cell_id | |
| def _apply_slot_to_global_unlocked(self, slot: RunSlot) -> None: | |
| self.active = slot.active | |
| self.job_id = slot.job_id | |
| self.policy = slot.policy | |
| self.revision = slot.revision | |
| self.task = slot.task | |
| self.status = slot.status | |
| self.rollouts = slot.rollouts | |
| self.latest_frame = slot.latest_frame | |
| self.frame_version = slot.frame_version | |
| self.events = slot.events | |
| self.summary = slot.summary | |
| self.error = slot.error | |
| self.started_at = slot.started_at | |
| def _create_slot_unlocked( | |
| self, | |
| *, | |
| job_id: str, | |
| cell_id: Optional[str], | |
| policy: str, | |
| revision: Optional[str], | |
| task: config.Task, | |
| status: str, | |
| ) -> RunSlot: | |
| key = display_cell_id(cell_id) or pending_cell_id(job_id) | |
| slot = RunSlot( | |
| cell_id=key, | |
| job_id=job_id, | |
| policy=policy, | |
| revision=revision, | |
| task=task, | |
| status=status, | |
| active=True, | |
| started_at=time.time(), | |
| ) | |
| self._slot_event(slot, f"following eval {job_id}: {policy}") | |
| self._slots_by_cell[key] = slot | |
| self._default_cell_id = key | |
| return slot | |
| def _bind_slot_cell_unlocked(self, job_id: str, cell_id: str) -> None: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is None or not cell_id or slot.cell_id == cell_id: | |
| return | |
| old = slot.cell_id | |
| if self._slots_by_cell.get(old) is slot: | |
| del self._slots_by_cell[old] | |
| slot.cell_id = cell_id | |
| self._slots_by_cell[cell_id] = slot | |
| if self._default_cell_id == old: | |
| self._default_cell_id = cell_id | |
| for session_id, viewed in list(self._view_by_session.items()): | |
| if viewed == old: | |
| self._view_by_session[session_id] = cell_id | |
| def _running_policies_unlocked(self) -> set[str]: | |
| policies = {slot.policy for slot in self._slots_by_cell.values() if slot.active} | |
| if self.active and self.policy: | |
| policies.add(self.policy) | |
| return policies | |
| def _preferred_cell_unlocked(self, choices: list[str]) -> Optional[str]: | |
| active = [ | |
| slot | |
| for slot in self._slots_by_cell.values() | |
| if slot.active and slot.cell_id in choices | |
| ] | |
| if active: | |
| return max(active, key=lambda slot: slot.started_at or 0).cell_id | |
| if self._default_cell_id in choices: | |
| return self._default_cell_id | |
| return choices[0] if choices else None | |
| def _ingest_cells_unlocked(self, task_slug: str, cells: list[Any]) -> None: | |
| self._cells_by_task[task_slug] = { | |
| cell.cell_id: cell.status.value for cell in cells | |
| } | |
| for cell in cells: | |
| if cell.current_job_id: | |
| self._bind_slot_cell_unlocked(cell.current_job_id, cell.cell_id) | |
| # ---------------------------------------------------------------- accessors | |
| def snapshot( | |
| self, | |
| cell_id: Optional[str] = None, | |
| *, | |
| task_slug: Optional[str] = None, | |
| session_id: Optional[str] = None, | |
| ) -> dict[str, Any]: | |
| with self._lock: | |
| snap = self._snapshot_unlocked() | |
| cell_statuses = ( | |
| dict(self._cells_by_task.get(task_slug, {})) | |
| if task_slug | |
| else { | |
| cell: status | |
| for statuses in self._cells_by_task.values() | |
| for cell, status in statuses.items() | |
| } | |
| ) | |
| live_feeds = [ | |
| { | |
| "cell_id": slot.cell_id, | |
| "job_id": slot.job_id, | |
| "policy": slot.policy, | |
| "status": slot.status, | |
| "active": slot.active, | |
| "task_label": slot.task.label if slot.task is not None else None, | |
| } | |
| for slot in sorted( | |
| self._slots_by_cell.values(), | |
| key=lambda slot: (not slot.active, slot.cell_id), | |
| ) | |
| ] | |
| feed_ids = [feed["cell_id"] for feed in live_feeds] | |
| stored = self._view_by_session.get(session_id) if session_id else None | |
| selected = None | |
| if cell_id in feed_ids: | |
| selected = cell_id | |
| elif stored in feed_ids: | |
| selected = stored | |
| else: | |
| selected = self._preferred_cell_unlocked(feed_ids) | |
| if session_id and selected: | |
| self._view_by_session[session_id] = selected | |
| slot = self._slots_by_cell.get(selected or "") | |
| if slot is not None: | |
| snap.update(self._slot_snapshot_unlocked(slot)) | |
| elif selected is not None: | |
| snap.update(self._empty_run_snapshot()) | |
| color, caption = fleet_availability( | |
| cell_statuses, | |
| aggregate=snap["cell_status"].get(task_slug or "", "offline") | |
| if task_slug | |
| else ( | |
| "available" | |
| if any(s == "available" for s in snap["cell_status"].values()) | |
| else ( | |
| "occupied" | |
| if any(s == "occupied" for s in snap["cell_status"].values()) | |
| else "offline" | |
| ) | |
| ), | |
| ) | |
| snap.update( | |
| { | |
| "selected_cell": selected, | |
| "cell_choices": feed_ids, | |
| "cell_statuses": cell_statuses, | |
| "live_feeds": live_feeds, | |
| "availability_color": color, | |
| "availability_caption": caption, | |
| } | |
| ) | |
| return snap | |
| def _slot_snapshot_unlocked(slot: RunSlot) -> dict[str, Any]: | |
| scored = list(slot.rollouts) | |
| return { | |
| "active": slot.active, | |
| "job_id": slot.job_id, | |
| "cell_id": slot.cell_id, | |
| "policy": slot.policy, | |
| "revision": slot.revision, | |
| "status": slot.status, | |
| "rollouts": scored, | |
| "latest_frame": slot.latest_frame, | |
| "frame_version": slot.frame_version, | |
| "events": list(slot.events), | |
| "summary": slot.summary, | |
| "error": slot.error, | |
| "n_scored": len(scored), | |
| "n_success": sum(1 for rollout in scored if rollout["success"]), | |
| "started_at": slot.started_at, | |
| "task": slot.task, | |
| } | |
| def _empty_run_snapshot() -> dict[str, Any]: | |
| return { | |
| "active": False, | |
| "job_id": None, | |
| "policy": None, | |
| "revision": None, | |
| "status": "idle", | |
| "rollouts": [], | |
| "latest_frame": None, | |
| "frame_version": 0, | |
| "events": [], | |
| "summary": None, | |
| "error": None, | |
| "n_scored": 0, | |
| "n_success": 0, | |
| "started_at": None, | |
| "task": None, | |
| "cell_id": None, | |
| } | |
| def _snapshot_unlocked(self) -> dict[str, Any]: | |
| """Copy shared state while the caller owns ``_lock``.""" | |
| scored = list(self.rollouts) | |
| n_success = sum(1 for r in scored if r["success"]) | |
| return { | |
| "active": self.active, | |
| "job_id": self.job_id, | |
| "cell_id": None, | |
| "policy": self.policy, | |
| "revision": self.revision, | |
| "status": self.status, | |
| "rollouts": scored, | |
| "latest_frame": self.latest_frame, | |
| "frame_version": self.frame_version, | |
| "events": list(self.events), | |
| "summary": self.summary, | |
| "error": self.error, | |
| "n_scored": len(scored), | |
| "n_success": n_success, | |
| "started_at": self.started_at, | |
| "cell_status": dict(self._cell_status), | |
| "queue_total": len(self._queue_items), | |
| "queued_eval_total": self._queued_eval_total, | |
| "queue_items": list(self._queue_items), | |
| "task": self.task, | |
| } | |
| # ------------------------------------------------------------- cell status | |
| def _poll_cell_status_loop(self) -> None: | |
| """Refresh each task's availability light, forever, in one thread. | |
| A single process-wide poll (not per Gradio session) keeps the lights | |
| live for every viewer without adding to the per-tab request budget that | |
| HF rate-limits. One request per task per interval is affordable at this | |
| scale and stays correct if the tasks ever move to different cells. | |
| Any failure — orchestrator unreachable, key rejected, or no cell able to | |
| run the task — is treated as ``offline`` (red, blocks submission), which | |
| is the safe default, and the client is rebuilt so a wedged connection | |
| self-heals. Those causes need different fixes and are indistinguishable | |
| in the UI, so the reason is logged when it changes. | |
| """ | |
| client: Optional[OrchestratorClient] = None | |
| previous: dict[str, str] = {} | |
| while True: | |
| for task in config.TASKS.values(): | |
| status, reason, client, cells = self._poll_one(client, task) | |
| # Logged on change rather than every poll: at a few seconds per | |
| # poll this would otherwise be the loudest thing in the Space | |
| # log, but without it a red light gives an operator nothing to | |
| # act on. | |
| shown = self._apply_poll(task.slug, status, cells) | |
| if shown != previous.get(task.slug) and reason and shown == "offline": | |
| logger.warning("%s shown as offline: %s", task.label, reason) | |
| previous[task.slug] = shown | |
| try: | |
| self._refresh_queue(client) | |
| self._adopt_dispatched_eval(client) | |
| except Exception as exc: # noqa: BLE001 - keep last good queue view | |
| logger.warning("queue status poll failed: %r", exc) | |
| time.sleep(config.CELL_STATUS_POLL_S) | |
| def _apply_poll( | |
| self, task_slug: str, status: str, cells: Optional[list[Any]] | |
| ) -> str: | |
| """Record one poll. A failed poll keeps the last good fleet for a while.""" | |
| with self._lock: | |
| if cells is None: | |
| fails = self._status_failures.get(task_slug, 0) + 1 | |
| self._status_failures[task_slug] = fails | |
| if fails < _STATUS_FAILURES_BEFORE_OFFLINE: | |
| return self._cell_status.get(task_slug, "offline") | |
| self._cell_status[task_slug] = "offline" | |
| self._cells_by_task[task_slug] = {} | |
| return "offline" | |
| self._status_failures[task_slug] = 0 | |
| self._cell_status[task_slug] = status | |
| self._ingest_cells_unlocked(task_slug, cells) | |
| return status | |
| def _poll_one( | |
| self, client: Optional[OrchestratorClient], task: config.Task | |
| ) -> tuple[str, str, Optional[OrchestratorClient], Optional[list[Any]]]: | |
| """One task's availability, plus the client to reuse for the next poll.""" | |
| try: | |
| if client is None: | |
| client = OrchestratorClient( | |
| orchestrator_url(), timeout=_STATUS_CALL_TIMEOUT_S | |
| ) | |
| summary = _call_bounded( | |
| lambda: client.get_cell_status( | |
| embodiment=config.EMBODIMENT, task=task.slug | |
| ), | |
| _STATUS_CALL_TIMEOUT_S, | |
| ) | |
| if summary.status.value == "offline" and not summary.cells: | |
| # The orchestrator answered, so it is not a connectivity or key | |
| # problem: no cell can run this task right now. A cell is | |
| # assigned an *environment* in the FMS and serves every task in | |
| # it, so a healthy cell set to the wrong environment lands here | |
| # looking identical to one that is switched off. | |
| return summary.status.value, ( | |
| f"no cell is serving task {task.slug!r} for embodiment " | |
| f"{config.EMBODIMENT!r} (check the environment assigned to " | |
| "the cell in the FMS, and that the task belongs to it)" | |
| ), client, [] | |
| return summary.status.value, "", client, list(summary.cells) | |
| except Exception as exc: # noqa: BLE001 - transient; treat as offline + rebuild | |
| if client is not None: | |
| try: | |
| client.close() | |
| except Exception: # noqa: BLE001 | |
| pass | |
| return "offline", f"cell status poll failed: {exc!r}", None, None | |
| def _refresh_queue(self, client: OrchestratorClient) -> None: | |
| """Read only jobs owned by the Space's demo API user.""" | |
| response = client._client.get( # noqa: SLF001 - shared auth/session | |
| "/jobs", params={"limit": 100, "queued_only": True} | |
| ) | |
| response.raise_for_status() | |
| waiting = [ | |
| job | |
| for item in response.json() | |
| if (job := Job.model_validate(item)).status | |
| in (JobStatus.SUBMITTED, JobStatus.QUEUED) | |
| ] | |
| labels: list[str] = [] | |
| eval_total = 0 | |
| eval_policies: set[str] = set() | |
| for job in reversed(waiting): # API is newest-first; queue is oldest-first. | |
| policy = self._space_eval_policy(job) | |
| if policy is not None: | |
| eval_total += 1 | |
| eval_policies.add(policy) | |
| labels.append(policy) | |
| else: | |
| labels.append("other job") | |
| with self._lock: | |
| self._queued_eval_total = eval_total | |
| self._queued_eval_policies = eval_policies | |
| self._queue_items = labels | |
| def _space_eval_policy(job: Job) -> Optional[str]: | |
| """Policy repo for a Space eval, including jobs queued before tagging.""" | |
| policy = job.spec.args.get("policy_path") | |
| if not isinstance(policy, str) or not policy: | |
| return None | |
| if ( | |
| job.spec.args.get("submission_source") == config.SUBMISSION_SOURCE | |
| or "lerobot-eval" in job.spec.image | |
| ): | |
| return policy | |
| return None | |
| def _adopt_dispatched_eval(self, client: OrchestratorClient) -> None: | |
| """Follow every previously queued demo job once the scheduler starts it.""" | |
| if client is None: | |
| return | |
| candidates = client.list_jobs(limit=50) | |
| for job in candidates: | |
| if job.status not in (JobStatus.DISPATCHED, JobStatus.RUNNING): | |
| continue | |
| policy = self._space_eval_policy(job) | |
| task = next( | |
| ( | |
| configured | |
| for configured in config.TASKS.values() | |
| if configured.slug == job.spec.task | |
| ), | |
| None, | |
| ) | |
| if policy is None or task is None: | |
| continue | |
| revision = job.spec.args.get("policy_revision") | |
| if not isinstance(revision, str): | |
| revision = None | |
| with self._lock: | |
| if job.id in self._following_job_ids: | |
| if job.cell_id: | |
| self._bind_slot_cell_unlocked(job.id, job.cell_id) | |
| continue | |
| self._following_job_ids.add(job.id) | |
| slot = self._create_slot_unlocked( | |
| job_id=job.id, | |
| cell_id=job.cell_id, | |
| policy=policy, | |
| revision=revision, | |
| task=task, | |
| status=job.status.value, | |
| ) | |
| if not self.active: | |
| self._apply_slot_to_global_unlocked(slot) | |
| self._event(f"following queued eval {job.id}: {policy}") | |
| follower = OrchestratorClient(orchestrator_url()) | |
| threading.Thread( | |
| target=self._follow, | |
| args=(follower, job.id), | |
| daemon=True, | |
| ).start() | |
| # ------------------------------------------------------------------- run | |
| def start_run(self, policy_repo: str, task: config.Task) -> tuple[bool, str]: | |
| policy_repo = (policy_repo or "").strip() | |
| if not policy_repo: | |
| return False, "Enter or select a policy repo id." | |
| if not config.eval_image_configured(): | |
| return False, ( | |
| "Not configured: ARMNET_DEMO_EVAL_IMAGE is unset. An operator " | |
| "must build + push the lerobot-eval runtime image once and set it." | |
| ) | |
| with self._lock: | |
| if policy_repo in self._running_policies_unlocked(): | |
| return False, ( | |
| f"{policy_repo} is already running. You can submit it again " | |
| "after this run finishes." | |
| ) | |
| if policy_repo in self._queued_eval_policies: | |
| return False, ( | |
| f"{policy_repo} is already queued. You can submit it again " | |
| "after that run finishes." | |
| ) | |
| if policy_repo in self._pending_eval_policies: | |
| return False, f"{policy_repo} is already being submitted." | |
| self._pending_eval_policies.add(policy_repo) | |
| any_live = self.active or any( | |
| slot.active for slot in self._slots_by_cell.values() | |
| ) | |
| owns_placeholder = not any_live | |
| if owns_placeholder: | |
| self._reset(active=True) | |
| self.policy = policy_repo | |
| self.task = task | |
| self.status = "submitting" | |
| self.started_at = time.time() | |
| self._event(f"submitting eval of {policy_repo} on {task.label!r}") | |
| else: | |
| self._event(f"submitting eval of {policy_repo}") | |
| def release_submission() -> None: | |
| with self._lock: | |
| self._pending_eval_policies.discard(policy_repo) | |
| # Recheck this API user's jobs synchronously rather than trusting the | |
| # five-second UI snapshot. Only submissions tagged as coming from this | |
| # Space consume its client-side abuse cap. | |
| client: Optional[OrchestratorClient] = None | |
| try: | |
| client = OrchestratorClient(orchestrator_url()) | |
| self._refresh_queue(client) | |
| except Exception as exc: # noqa: BLE001 - fail closed on abuse guard | |
| if client is not None: | |
| client.close() | |
| with self._lock: | |
| if owns_placeholder: | |
| self._reset(active=False) | |
| release_submission() | |
| logger.warning("could not verify queue capacity", exc_info=True) | |
| return False, f"Cannot check queue capacity: {type(exc).__name__}" | |
| with self._lock: | |
| queue_full = self._queued_eval_total >= config.MAX_QUEUED_JOBS | |
| duplicate_queued = policy_repo in self._queued_eval_policies | |
| if queue_full or duplicate_queued: | |
| client.close() | |
| with self._lock: | |
| if owns_placeholder: | |
| self._reset(active=False) | |
| release_submission() | |
| if duplicate_queued: | |
| return False, ( | |
| f"{policy_repo} is already queued. You can submit it again " | |
| "after that run finishes." | |
| ) | |
| return False, ( | |
| f"This Space already has {config.MAX_QUEUED_JOBS} eval jobs queued. " | |
| "Please wait for one to start." | |
| ) | |
| # Pin the policy's current commit so the run is reproducible and pools | |
| # correctly on the leaderboard (a repo updated later is a new revision). | |
| revision = self._resolve_revision(policy_repo) | |
| if owns_placeholder: | |
| with self._lock: | |
| self.revision = revision | |
| self._ensure_eval_secret(client) | |
| try: | |
| job = client.submit( | |
| self._build_spec(policy_repo, revision, task), allow_queue=True | |
| ) | |
| except OrchestratorError as exc: | |
| client.close() | |
| with self._lock: | |
| if owns_placeholder: | |
| self._reset(active=False) | |
| release_submission() | |
| if exc.status_code == 409: | |
| return False, ( | |
| "Eval could not be queued because capacity changed. Please " | |
| "refresh and try again." | |
| ) | |
| return False, f"Eval failed to submit (HTTP {exc.status_code})." | |
| except Exception as exc: # noqa: BLE001 - surface submit failure to the UI | |
| client.close() | |
| with self._lock: | |
| if owns_placeholder: | |
| self._reset(active=False) | |
| release_submission() | |
| logger.exception("demo submit failed") | |
| return False, f"Eval failed to submit: {type(exc).__name__}" | |
| if not job.dispatched: | |
| client.close() | |
| with self._lock: | |
| if owns_placeholder: | |
| self._reset(active=False) | |
| self._queued_eval_total += 1 | |
| self._queued_eval_policies.add(policy_repo) | |
| self._queue_items.append(policy_repo) | |
| release_submission() | |
| return True, ( | |
| f"Queued {config.N_EPISODES}-rollout eval of {policy_repo}. " | |
| "It will start automatically when a cell is free." | |
| ) | |
| already_following = False | |
| routed = display_cell_id(job.cell_id) | |
| with self._lock: | |
| if job.id in self._following_job_ids: | |
| already_following = True | |
| if job.cell_id: | |
| self._bind_slot_cell_unlocked(job.id, job.cell_id) | |
| else: | |
| self._following_job_ids.add(job.id) | |
| slot = self._create_slot_unlocked( | |
| job_id=job.id, | |
| cell_id=job.cell_id, | |
| policy=policy_repo, | |
| revision=revision, | |
| task=task, | |
| status="running", | |
| ) | |
| if owns_placeholder or not self.active: | |
| self._apply_slot_to_global_unlocked(slot) | |
| else: | |
| self._event( | |
| f"job {job.id} dispatched" | |
| + (f" to {routed}" if routed else "") | |
| ) | |
| release_submission() | |
| if already_following: | |
| client.close() | |
| else: | |
| threading.Thread( | |
| target=self._follow, args=(client, job.id), daemon=True | |
| ).start() | |
| where = f" ({routed})" if routed else "" | |
| return True, ( | |
| f"Running {config.N_EPISODES}-rollout eval of {policy_repo} " | |
| f"on {task.label.lower()}{where}…" | |
| ) | |
| def _follow_after_current( | |
| self, | |
| job_id: str, | |
| policy: str, | |
| revision: Optional[str], | |
| task: config.Task, | |
| cell_id: Optional[str] = None, | |
| ) -> None: | |
| """Attach to an unexpectedly immediate dispatch after the old run closes.""" | |
| while True: | |
| with self._lock: | |
| if job_id in self._following_job_ids and self.job_id == job_id: | |
| return # The queue poller already adopted it. | |
| if not self.active: | |
| if job_id not in self._following_job_ids: | |
| self._following_job_ids.add(job_id) | |
| slot = self._create_slot_unlocked( | |
| job_id=job_id, | |
| cell_id=cell_id, | |
| policy=policy, | |
| revision=revision, | |
| task=task, | |
| status="running", | |
| ) | |
| self._apply_slot_to_global_unlocked(slot) | |
| else: | |
| self._reset(active=True) | |
| self.job_id = job_id | |
| self.policy = policy | |
| self.revision = revision | |
| self.task = task | |
| self.status = "running" | |
| self.started_at = time.time() | |
| self._event(f"following newly dispatched eval {job_id}: {policy}") | |
| break | |
| time.sleep(0.2) | |
| self._follow(OrchestratorClient(orchestrator_url()), job_id) | |
| def _ensure_eval_secret(self, client: OrchestratorClient) -> None: | |
| """Store the HF token as a named orchestrator secret (once per process). | |
| The eval references it by name ({"HF_TOKEN": EVAL_HF_SECRET_NAME}); the | |
| orchestrator resolves the name to its value from the demo user's secret | |
| store. Best-effort: if it fails, a gated policy simply won't load. | |
| """ | |
| if not config.HF_TOKEN or self._eval_secret_ensured: | |
| return | |
| try: | |
| client.create_secret(config.EVAL_HF_SECRET_NAME, config.HF_TOKEN) | |
| self._eval_secret_ensured = True | |
| logger.info("provisioned orchestrator secret %r for the eval", config.EVAL_HF_SECRET_NAME) | |
| except Exception: # noqa: BLE001 - proceed; gated policies may then fail to load | |
| logger.warning("could not provision HF secret for the eval", exc_info=True) | |
| def _resolve_revision(self, policy_repo: str) -> Optional[str]: | |
| """The policy repo's current commit SHA, or None if it can't be resolved.""" | |
| try: | |
| from huggingface_hub import HfApi | |
| return HfApi(token=config.HF_TOKEN).model_info(policy_repo).sha | |
| except Exception: # noqa: BLE001 - run unpinned rather than fail the demo | |
| logger.info("could not resolve revision for %s; running unpinned", policy_repo) | |
| return None | |
| def _build_spec( | |
| self, policy_repo: str, revision: Optional[str], task: config.Task | |
| ) -> JobSpec: | |
| return JobSpec( | |
| image=config.EVAL_IMAGE, | |
| embodiment=config.EMBODIMENT, | |
| task=task.slug, | |
| args={ | |
| "policy_path": policy_repo, | |
| "submission_source": config.SUBMISSION_SOURCE, | |
| # whoami() is the token's user, not the org. Name the namespace | |
| # so a recorded dataset lands in the org even when the token | |
| # belongs to a person. | |
| "hf_user": config.DATASET_NAMESPACE, | |
| **({"policy_revision": revision} if revision else {}), | |
| "n_episodes": config.N_EPISODES, | |
| "episode_time_s": config.EPISODE_TIME_S, | |
| "policy_fps": config.POLICY_FPS, | |
| "start_seed": config.START_SEED, | |
| # Fixed seed, no submitter control: every entry on the | |
| # leaderboard faces the same sequence of scenes, so the ranking | |
| # compares policies rather than how kind the setup was. | |
| "variation": True, | |
| "variation_seed": config.VARIATION_SEED, | |
| "policy_device": "cuda", | |
| "policy_use_amp": False, | |
| "streaming_encoding": True, | |
| "vcodec": "libsvtav1", | |
| "encoder_threads": 2, | |
| # Turns on the live rollout video feed, streamed small (see config). | |
| "use_rerun": True, | |
| "rerun_image_scale": config.RERUN_IMAGE_SCALE, | |
| "rerun_jpeg_quality": config.RERUN_JPEG_QUALITY, | |
| # The demo doesn't need the recorded dataset; skipping the push | |
| # avoids needing Hub write access and is faster. | |
| "no_push_to_hub": not config.PUSH_DATASET, | |
| }, | |
| # Reference the named orchestrator secret (provisioned in | |
| # _ensure_eval_secret) — {ENV_VAR: secret_name}, NOT an inline value. | |
| # Resolved to $HF_TOKEN in the container so it can pull a gated policy. | |
| secrets={"HF_TOKEN": config.EVAL_HF_SECRET_NAME} if config.HF_TOKEN else {}, | |
| timeout_seconds=config.TIMEOUT_SECONDS, | |
| ) | |
| def _follow(self, client: OrchestratorClient, job_id: str) -> None: | |
| """Stream the (already-dispatched) job's video + logs and poll to done.""" | |
| stop = threading.Event() | |
| threads: list[threading.Thread] = [] | |
| leaderboard_snapshot: Optional[dict[str, Any]] = None | |
| try: | |
| for target in (self._stream_rerun, self._stream_logs): | |
| t = threading.Thread(target=target, args=(job_id, stop), daemon=True) | |
| t.start() | |
| threads.append(t) | |
| final = self._poll_until_terminal(client, job_id) | |
| self._finish(final, job_id=job_id) | |
| with self._lock: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| leaderboard_snapshot = self._slot_snapshot_unlocked(slot) | |
| if final is not None and final.cell_id: | |
| leaderboard_snapshot["cell_id"] = final.cell_id | |
| elif self.job_id == job_id: | |
| leaderboard_snapshot = self._snapshot_unlocked() | |
| if final is not None and final.cell_id: | |
| leaderboard_snapshot["cell_id"] = final.cell_id | |
| except Exception as exc: # noqa: BLE001 - surface any failure to the UI | |
| logger.exception("demo run failed") | |
| with self._lock: | |
| message = f"{type(exc).__name__}: {exc}" | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| slot.error = message | |
| slot.status = "failed" | |
| self._slot_event(slot, message) | |
| if self.job_id == job_id: | |
| self.error = message | |
| self.status = "failed" | |
| self._event(message) | |
| finally: | |
| stop.set() | |
| for t in threads: | |
| t.join(timeout=2) | |
| try: | |
| client.close() | |
| except Exception: # noqa: BLE001 | |
| pass | |
| with self._lock: | |
| self._following_job_ids.discard(job_id) | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| slot.active = False | |
| if self.job_id == job_id: | |
| other = next( | |
| ( | |
| remaining | |
| for remaining in self._slots_by_cell.values() | |
| if remaining.active | |
| ), | |
| None, | |
| ) | |
| if other is not None: | |
| self._apply_slot_to_global_unlocked(other) | |
| else: | |
| self.active = False | |
| # Hub metadata/download/upload work must never delay stream cleanup or | |
| # keep the Space marked active. Persist an immutable snapshot in the | |
| # background after the current run has been released. | |
| if leaderboard_snapshot is not None: | |
| threading.Thread( | |
| target=self._record_to_leaderboard, | |
| args=(leaderboard_snapshot,), | |
| daemon=True, | |
| ).start() | |
| def _poll_until_terminal(self, client: OrchestratorClient, job_id: str): | |
| """Poll the job to a terminal state — the reliable completion signal. | |
| Independent of the log/rerun websockets: even if those drop, this keeps | |
| checking and will detect success/fail and unlock. If the orchestrator | |
| itself becomes unreachable (network, rate limit), it backs off and, after | |
| ``_LOST_CONTACT_GIVEUP_S`` with no contact, gives up and unlocks rather | |
| than freezing the Space for the whole job timeout. | |
| """ | |
| hard_deadline = time.monotonic() + config.TIMEOUT_SECONDS + 120 | |
| last_ok = time.monotonic() | |
| failures = 0 | |
| while True: | |
| try: | |
| job = client.get(job_id) | |
| except Exception as exc: # noqa: BLE001 - transient; back off + retry | |
| failures += 1 | |
| gone = time.monotonic() - last_ok | |
| if gone > _LOST_CONTACT_GIVEUP_S: | |
| logger.warning("giving up on job %s after %.0fs no contact: %r", job_id, gone, exc) | |
| self.set_status( | |
| "lost contact with the orchestrator — unlocking " | |
| "(the job may still be running on the cell)", | |
| job_id=job_id, | |
| ) | |
| return None | |
| # Stable message (no changing numbers) so it doesn't spam events. | |
| self.set_status("reconnecting to the orchestrator…", job_id=job_id) | |
| if time.monotonic() > hard_deadline: | |
| return None | |
| time.sleep(min(_POLL_INTERVAL_S * failures, _POLL_BACKOFF_MAX_S)) | |
| continue | |
| failures = 0 | |
| last_ok = time.monotonic() | |
| if job.cell_id: | |
| with self._lock: | |
| self._bind_slot_cell_unlocked(job_id, job.cell_id) | |
| if job.status in TerminalStatus: | |
| return job | |
| # Keep the live badge in sync (submitted -> dispatched -> running). | |
| self.set_status(job.status.value, job_id=job_id) | |
| if time.monotonic() > hard_deadline: | |
| self.set_status("timed out (demo budget)", job_id=job_id) | |
| return None | |
| time.sleep(_POLL_INTERVAL_S) | |
| def _record_to_leaderboard(self, snap: dict[str, Any]) -> None: | |
| """Append this run's pooled result to the persistent leaderboard.""" | |
| rollouts = snap["rollouts"] | |
| task = snap["task"] | |
| if not rollouts or not snap["policy"] or task is None: | |
| return | |
| n_success = sum(1 for r in rollouts if r["success"]) | |
| model_type = leaderboard.model_type_for( | |
| snap["policy"], snap.get("revision"), config.HF_TOKEN | |
| ) | |
| try: | |
| LEADERBOARD.record_run( | |
| repo_id=snap["policy"], | |
| revision=snap.get("revision") or "unknown", | |
| n_rollouts=len(rollouts), | |
| n_success=n_success, | |
| model_type=model_type, | |
| # The demo evaluates LeRobot policies via armnet-lerobot-eval. | |
| training_framework="lerobot", | |
| source="demo_space", | |
| cell_id=( | |
| display_cell_id(snap.get("cell_id")) or config.CELL_ID | |
| ), | |
| embodiment=config.EMBODIMENT, | |
| task=task.slug, | |
| ) | |
| except Exception: # noqa: BLE001 - never fail the run over leaderboard I/O | |
| logger.warning("failed to record run to leaderboard", exc_info=True) | |
| def _finish( # noqa: ANN001 - Job | None | |
| self, final, *, job_id: Optional[str] = None | |
| ) -> None: | |
| with self._lock: | |
| slot = self._slot_for_job_unlocked(job_id) | |
| if slot is not None: | |
| if final is not None and final.cell_id: | |
| self._bind_slot_cell_unlocked(job_id, final.cell_id) | |
| slot = self._slot_for_job_unlocked(job_id) or slot | |
| self._finish_target_unlocked(slot, final) | |
| if self.job_id == job_id: | |
| self._apply_slot_to_global_unlocked(slot) | |
| return | |
| if job_id is not None and self.job_id != job_id: | |
| return | |
| self._finish_target_unlocked(self, final) | |
| def _finish_target_unlocked(self, target: Any, final: Any) -> None: | |
| if final is None: | |
| return | |
| target.status = final.status.value | |
| result = final.result | |
| if result is not None: | |
| target.summary = result.return_value | |
| if result.error: | |
| target.error = result.error | |
| # Recover per-rollout results from the terminal job even if we missed | |
| # the live log lines (e.g. the log websocket dropped mid-run): the | |
| # eval's return_value carries every episode's outcome. This is what | |
| # lets a run that "finished but we weren't watching" still show its | |
| # full results instead of a half-filled table. | |
| if isinstance(target.summary, dict) and target.summary.get("per_episode"): | |
| target.rollouts = [ | |
| { | |
| "episode": int(item.get("episode_ix", i)) + 1, | |
| "success": bool(item.get("success")), | |
| "scored_by": item.get("scored_by", "?"), | |
| } | |
| for i, item in enumerate(target.summary["per_episode"]) | |
| ] | |
| if isinstance(target, RunSlot): | |
| self._slot_event(target, f"job finished: {target.status}") | |
| else: | |
| self._event(f"job finished: {target.status}") | |
| def _call_bounded(fn, timeout: float): | |
| """Run ``fn`` and raise ``TimeoutError`` if it has not finished in time. | |
| Closing the HTTP client unblocks a half-open read that ignores its own | |
| timeout. The caller does that when this raises. | |
| """ | |
| box: dict[str, Any] = {} | |
| def run() -> None: | |
| try: | |
| box["value"] = fn() | |
| except Exception as exc: # noqa: BLE001 - handed back to the caller | |
| box["error"] = exc | |
| worker = threading.Thread(target=run, name="armnet-demo-bounded", daemon=True) | |
| worker.start() | |
| worker.join(timeout) | |
| if worker.is_alive(): | |
| raise TimeoutError(f"did not finish within {timeout:.0f}s") | |
| if "error" in box: | |
| raise box["error"] | |
| return box.get("value") | |
| def _connect_ws(path_job_id: str, suffix: str): | |
| """Open an orchestrator websocket for a job with the API key header.""" | |
| key = api_key() | |
| url = _ws_url(orchestrator_url(), f"/jobs/{path_job_id}/{suffix}") | |
| return websocket.create_connection( | |
| url, | |
| header=[f"{API_KEY_HEADER}: {key}"], | |
| timeout=8, | |
| # websocket-client does not use certifi. The published armnet-client | |
| # the Space installs may not export a helper for this. | |
| sslopt={"context": ssl.create_default_context(cafile=certifi.where())}, | |
| ) | |
| def _release_ws_when_stopped(ws: Any, stop: threading.Event, done: threading.Event) -> None: | |
| """Close ``ws`` once the follow is over, so a blocked ``recv`` returns.""" | |
| while not done.is_set(): | |
| if stop.wait(0.25): | |
| try: | |
| ws.close() | |
| except Exception: # noqa: BLE001 | |
| pass | |
| return | |
| # Bound the two stream methods to the state instance below (defined as methods | |
| # for access to `self`, but kept out of the class body above for readability). | |
| def _stream_rerun(self: DemoState, job_id: str, stop: threading.Event) -> None: | |
| connect_backoff = 1.0 | |
| # Packets that arrived vs. packets we could turn into a frame. A run where | |
| # these diverge is the one failure mode the viewer cannot see for themselves: | |
| # the feed looks identical to a cell that simply isn't sending video. | |
| received = 0 | |
| decoded = 0 | |
| while not stop.is_set(): | |
| try: | |
| ws = _connect_ws(job_id, "rerun") | |
| except Exception: # noqa: BLE001 - back off + retry until stop | |
| if stop.wait(connect_backoff): | |
| return | |
| connect_backoff = min(connect_backoff * 2, _WS_RECONNECT_MAX_S) | |
| continue | |
| connect_backoff = 1.0 | |
| finished = threading.Event() | |
| threading.Thread( | |
| target=_release_ws_when_stopped, | |
| args=(ws, stop, finished), | |
| name="armnet-demo-rerun-stop", | |
| daemon=True, | |
| ).start() | |
| try: | |
| ws.settimeout(0.5) | |
| while not stop.is_set(): | |
| try: | |
| frame = ws.recv() | |
| except (TimeoutError, websocket.WebSocketTimeoutException): | |
| continue | |
| except Exception: # noqa: BLE001 - dropped; reconnect | |
| break | |
| if isinstance(frame, bytes) and frame: | |
| received += 1 | |
| try: | |
| img = decode_frame(frame) | |
| except Exception: # noqa: BLE001 - one frame, not the feed | |
| # decode_frame is written not to raise; if it ever does, | |
| # losing this thread would cost the rest of the run's | |
| # video, which is how a missing decoder once presented as | |
| # a cell that simply wasn't streaming. | |
| logger.warning("failed to decode a video packet", exc_info=True) | |
| img = None | |
| if img is not None: | |
| decoded += 1 | |
| self.set_frame(img, job_id=job_id) | |
| elif decoded == 0 and received == _UNDECODABLE_NOTICE_AFTER: | |
| self.log_line( | |
| "[armnet:demo] live video unavailable: the cell is " | |
| "sending frames but this Space cannot decode them " | |
| "(see the Space logs for why)", | |
| job_id=job_id, | |
| ) | |
| elif isinstance(frame, str) and frame: | |
| try: | |
| if json.loads(frame).get("type") == "terminal": | |
| return | |
| except json.JSONDecodeError: | |
| pass | |
| finally: | |
| finished.set() | |
| try: | |
| ws.close() | |
| except Exception: # noqa: BLE001 | |
| pass | |
| if stop.wait(0.5): | |
| return | |
| def _stream_logs(self: DemoState, job_id: str, stop: threading.Event) -> None: | |
| connect_backoff = 1.0 | |
| while not stop.is_set(): | |
| try: | |
| ws = _connect_ws(job_id, "logs") | |
| except Exception: # noqa: BLE001 - back off + retry until stop | |
| if stop.wait(connect_backoff): | |
| return | |
| connect_backoff = min(connect_backoff * 2, _WS_RECONNECT_MAX_S) | |
| continue | |
| connect_backoff = 1.0 | |
| finished = threading.Event() | |
| threading.Thread( | |
| target=_release_ws_when_stopped, | |
| args=(ws, stop, finished), | |
| name="armnet-demo-logs-stop", | |
| daemon=True, | |
| ).start() | |
| try: | |
| ws.settimeout(0.5) | |
| while not stop.is_set(): | |
| try: | |
| raw = ws.recv() | |
| except (TimeoutError, websocket.WebSocketTimeoutException): | |
| continue | |
| except Exception: # noqa: BLE001 - dropped; reconnect | |
| break | |
| try: | |
| payload = json.loads(raw) | |
| except (json.JSONDecodeError, TypeError): | |
| continue | |
| kind = payload.get("type") | |
| if kind == "terminal": | |
| return | |
| if kind != "log": | |
| continue | |
| line = payload.get("line", "") | |
| m = _EPISODE_RE.search(line) | |
| if m: | |
| self.record_rollout( | |
| episode=int(m.group(1)), | |
| total=int(m.group(2)), | |
| success=m.group(4) == "True", | |
| scored_by=m.group(6), | |
| job_id=job_id, | |
| ) | |
| elif "[armnet:progress]" in line: | |
| self.log_line(line, job_id=job_id) | |
| finally: | |
| finished.set() | |
| try: | |
| ws.close() | |
| except Exception: # noqa: BLE001 | |
| pass | |
| if stop.wait(0.5): | |
| return | |
| # Attach the stream helpers as methods and create the process-global state. | |
| DemoState._stream_rerun = _stream_rerun # type: ignore[attr-defined] | |
| DemoState._stream_logs = _stream_logs # type: ignore[attr-defined] | |
| STATE = DemoState() | |