armnet-eval / runner.py
villekuosmanen's picture
deploy demo space from 9889c231ec4c
a191eda verified
Raw History Blame Contribute Delete
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"
@dataclass
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,
)
@staticmethod
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
@staticmethod
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,
}
@staticmethod
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
@staticmethod
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()