dashboard / server.py
3v324v23's picture
fix score trend
13d819f
Raw
History Blame Contribute Delete
8.55 kB
#!/usr/bin/env python3
"""Serve the dashboard and periodically refresh its public W&B data."""
from __future__ import annotations
import http.server
import importlib
import json
import os
import socketserver
import sys
import threading
import time
from pathlib import Path
from urllib.parse import unquote, urlsplit
ROOT = Path(__file__).resolve().parent
sys.path.insert(0, str(ROOT / "scripts"))
PORT = int(os.environ.get("PORT", "7860"))
PROJECT = os.environ.get("WANDB_PROJECT", "").strip()
ENTITY = os.environ.get("WANDB_ENTITY") or None
REFRESH_SECONDS = int(os.environ.get("REFRESH_SECONDS", "1800"))
PROGRESS_REFRESH_SECONDS = max(
15, int(os.environ.get("PROGRESS_REFRESH_SECONDS", "15"))
)
DATA_PATH = ROOT / "data.json"
DATA_URL_PATH = "/data.json"
PUBLIC_PATHS = {"/", "/index.html", "/assets/app.js", "/assets/styles.css"}
_data_lock = threading.Lock()
_full_refresh_lock = threading.Lock()
_data_bytes = DATA_PATH.read_bytes()
def _set_data(payload: dict) -> None:
global _data_bytes
encoded = (json.dumps(payload, indent=2) + "\n").encode("utf-8")
with _data_lock:
_data_bytes = encoded
def _get_data() -> bytes:
with _data_lock:
return _data_bytes
def _get_data_payload() -> dict:
return json.loads(_get_data().decode("utf-8"))
def _set_progress(progress_by_validator: dict[str, dict]) -> None:
global _data_bytes
with _data_lock:
data = json.loads(_data_bytes.decode("utf-8"))
for validator in data.get("validators", []):
progress = progress_by_validator.get(validator.get("hotkey"))
if progress is None:
validator.pop("progress", None)
else:
validator["progress"] = progress
_data_bytes = (json.dumps(data, indent=2) + "\n").encode("utf-8")
def _progress_from_payload(payload: dict) -> dict[str, dict]:
return {
validator["hotkey"]: validator["progress"]
for validator in payload.get("validators", [])
if isinstance(validator.get("hotkey"), str)
and isinstance(validator.get("progress"), dict)
}
def _has_value(values, index: int) -> bool:
if not isinstance(values, list) or index >= len(values):
return False
value = values[index]
return isinstance(value, (int, float)) and value == value
def _history_has_scored_epoch(validator: dict, epoch) -> bool:
history = validator.get("history", {})
epochs = history.get("epochs", [])
if not isinstance(history, dict) or not isinstance(epochs, list):
return False
try:
index = epochs.index(epoch)
except ValueError:
return False
original = history.get("original", {})
if isinstance(original, dict):
for field in ("correctness_score", "completion_len", "score"):
if _has_value(original.get(field), index):
return True
miners = history.get("miners", {})
if isinstance(miners, dict):
for series in miners.values():
if not isinstance(series, dict):
continue
for field in ("correctness_score", "completion_len", "score"):
if _has_value(series.get(field), index):
return True
return False
def _progress_requires_data_refresh(
payload: dict, progress_by_validator: dict[str, dict]
) -> bool:
validators = {
validator.get("hotkey"): validator
for validator in payload.get("validators", [])
if isinstance(validator.get("hotkey"), str)
}
for hotkey, progress in progress_by_validator.items():
validator = validators.get(hotkey)
if validator is None:
return True
if progress.get("status") != "completed":
continue
progress_epoch = progress.get("epoch")
epochs = validator.get("history", {}).get("epochs", [])
history_epoch = epochs[-1] if epochs else None
if isinstance(progress_epoch, (int, float)) and (
history_epoch is None or progress_epoch > history_epoch
):
return True
if isinstance(progress_epoch, (int, float)) and not _history_has_scored_epoch(
validator, progress_epoch
):
return True
return False
def _build_progress(project: str, entity: str | None) -> dict[str, dict]:
build_data = importlib.import_module("build_data")
build_progress = getattr(build_data, "build_progress", None)
if callable(build_progress):
return build_progress(project, entity)
build = getattr(build_data, "build", None)
if not callable(build):
raise ImportError("build_data must define build_progress or build")
return _progress_from_payload(build(project, entity))
def _refresh_data(build, *, wait: bool = True) -> bool:
if not _full_refresh_lock.acquire(blocking=wait):
return False
try:
_set_data(build(PROJECT, ENTITY))
return True
finally:
_full_refresh_lock.release()
def _refresh_loop() -> None:
key = os.environ.get("WANDB_KEY") or os.environ.get("WANDB_API_KEY")
if not key:
print("[refresh] WANDB_KEY not set -- serving the bundled sample data.json only", flush=True)
return
if not PROJECT:
print("[refresh] WANDB_PROJECT not set -- serving the bundled sample data.json only", flush=True)
return
os.environ["WANDB_API_KEY"] = key
from build_data import build
while True:
try:
_refresh_data(build)
data = _get_data_payload()
print(f"[refresh] updated data: {len(data['validators'])} validators, "
f"{len(data['aggregate']['rankings'])} miners", flush=True)
except Exception as exc:
print(f"[refresh] failed ({type(exc).__name__}); keeping previous data", flush=True)
time.sleep(REFRESH_SECONDS)
def _progress_refresh_loop() -> None:
key = os.environ.get("WANDB_KEY") or os.environ.get("WANDB_API_KEY")
if not key or not PROJECT:
return
os.environ["WANDB_API_KEY"] = key
from build_data import build
time.sleep(min(5, PROGRESS_REFRESH_SECONDS))
while True:
try:
progress = _build_progress(PROJECT, ENTITY)
_set_progress(progress)
if _progress_requires_data_refresh(_get_data_payload(), progress):
if _refresh_data(build, wait=False):
print("[refresh] synced completed evaluation", flush=True)
except Exception as exc:
print(
f"[progress] refresh failed ({type(exc).__name__}); keeping previous data",
flush=True,
)
time.sleep(PROGRESS_REFRESH_SECONDS)
class Handler(http.server.SimpleHTTPRequestHandler):
server_version = "ThinkerDashboard"
sys_version = ""
def __init__(self, *args, **kwargs):
super().__init__(*args, directory=str(ROOT), **kwargs)
def _request_path(self) -> str:
return unquote(urlsplit(self.path).path)
def _send_data(self, include_body: bool) -> None:
payload = _get_data()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.send_header("Cache-Control", "no-store")
self.end_headers()
if include_body:
self.wfile.write(payload)
def send_head(self):
if self._request_path() not in PUBLIC_PATHS:
self.send_error(404)
return None
return super().send_head()
def do_GET(self):
if self._request_path() == DATA_URL_PATH:
self._send_data(include_body=True)
return
super().do_GET()
def do_HEAD(self):
if self._request_path() == DATA_URL_PATH:
self._send_data(include_body=False)
return
super().do_HEAD()
def log_message(self, *args):
pass
def main() -> int:
threading.Thread(target=_refresh_loop, daemon=True).start()
threading.Thread(target=_progress_refresh_loop, daemon=True).start()
socketserver.ThreadingTCPServer.allow_reuse_address = True
with socketserver.ThreadingTCPServer(("0.0.0.0", PORT), Handler) as httpd:
print(
f"[serve] dashboard on http://0.0.0.0:{PORT} "
f"(data {REFRESH_SECONDS}s, progress {PROGRESS_REFRESH_SECONDS}s)",
flush=True,
)
httpd.serve_forever()
return 0
if __name__ == "__main__":
raise SystemExit(main())