"""Live progress dashboard: pretraining/decay, teacher data, the eval gate and GRPO, read straight from the logs on every refresh (stdlib only, nothing to install). setsid nohup $TA_PY scripts/dashboard.py >/dev/null 2>&1 < /dev/null & # http://127.0.0.1:8765 $TA_PY scripts/dashboard.py --host 0.0.0.0 # reachable from other devices on the LAN """ import argparse import ast import glob import json import os import re import subprocess import time from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from tiny_agent.tasks import HELDOUT_KINDS DATA = os.environ.get("TA_DATA", "data") DECAY_FRAC = 0.2 # train.py --decay_frac: the forced decay lasts frac/(1-frac) of the steps done def jsonl(path, max_points=None): try: rows = [json.loads(l) for l in open(path) if l.strip()] except (OSError, ValueError): return [] if max_points and len(rows) > max_points: k = len(rows) / max_points rows = [rows[int(i * k)] for i in range(max_points)] + [rows[-1]] return rows def tail(path, n): try: with open(path, "rb") as f: f.seek(0, 2) f.seek(max(0, f.tell() - 64_000)) return f.read().decode(errors="replace").splitlines()[-n:] except OSError: return [] def procs(): out = subprocess.run(["ps", "-eo", "pid,etimes,args"], capture_output=True, text=True).stdout found = {} for line in out.splitlines()[1:]: parts = line.split(None, 2) if len(parts) < 3: continue pid, et, args = parts for name in ("train.py", "teacher.py", "eval_agent.py", "grpo.py"): if f"scripts/{name}" in args.split(" ")[1:3]: found[name] = {"pid": int(pid), "elapsed_s": int(et), "args": args[:1000]} return found def pretrain(): rows = [r for r in jsonl(f"{DATA}/runs/main/log.jsonl") if "loss" in r] if not rows: return None pts = rows if len(rows) <= 800 else rows[:: len(rows) // 800] + [rows[-1]] series = [{k: r[k] for k in ("step", "tokens", "loss", "lr_mult", "tok_s")} for r in pts] val = [{"step": r["step"], "tokens": r["tokens"], **r["val"]} for r in rows if "val" in r] last = rows[-1] decay = None # the decay starts at the last drop from lr_mult 1.0 (warmup also runs below 1.0) for i in range(len(rows) - 1, 0, -1): r = rows[i] if r["lr_mult"] < 1.0 and rows[i - 1]["lr_mult"] >= 1.0: start = rows[i - 1]["step"] steps = max(1, int(start * DECAY_FRAC / (1 - DECAY_FRAC))) done = last["step"] - start recent = [x for x in rows[-20:] if x["lr_mult"] < 1.0] rate = None if len(recent) >= 2 and recent[-1]["elapsed_min"] > recent[0]["elapsed_min"]: rate = (recent[-1]["step"] - recent[0]["step"]) / (recent[-1]["elapsed_min"] - recent[0]["elapsed_min"]) decay = {"start_step": start, "steps": steps, "done": done, "frac": min(1.0, done / steps), "eta_min": round((steps - done) / rate, 1) if rate else None} break return {"series": series, "val": val, "last": last, "decay": decay, "final": os.path.exists(f"{DATA}/runs/main/final.pt")} def teacher(): out = {"stats": None, "build": None} for line in reversed(tail(f"{DATA}/logs/teacher.log", 400)): if line.startswith("{") and '"episodes"' in line: try: out["stats"] = json.loads(line) except ValueError: pass break for line in tail(f"{DATA}/logs/teacher_window.log", 400): if line.startswith("teacher {"): try: out["build"] = ast.literal_eval(line[len("teacher "):]) except (ValueError, SyntaxError): pass return out def eval_gate(run_dir=None): """The pre-RL baseline of the shown run: /baseline.json if present (e.g. r2 starts from the SFT warm start), else the eval gate on final.pt.""" for p in ([os.path.join(run_dir, "baseline.json")] if run_dir else []) + [f"{DATA}/logs/eval_final.json"]: try: return json.load(open(p)) except (OSError, ValueError): continue return None def rl(): runs = [] for d in sorted(glob.glob(f"{DATA}/rl/*/")): p = os.path.join(d, "log.jsonl") if os.path.exists(p): runs.append((os.path.getmtime(p), d)) if not runs: return None _, d = max(runs) rows = jsonl(os.path.join(d, "log.jsonl")) steps = [r for r in rows if "reward" in r] evals = [{"step": r["step"], **r["eval"]} for r in rows if "eval" in r] keep = ("step", "reward", "acc", "grounded", "gen_tokens", "tool_tokens", "turns", "flagged_turns", "repeat_calls", "parse_err", "truncated", "groups", "no_signal", "stale_dropped", "staleness", "invented", "nf_answerable", "wrong_copied", "mismatch", "clip_frac", "grad_norm", "rollout_s", "step_s", "loss") notes = "" for p in (os.path.join(d, "NOTES.md"), f"{DATA}/rl/NOTES.md"): if os.path.exists(p): notes = open(p).read()[-20000:] break refs = {} for _, other in runs: if other != d: ev = [{"step": r["step"], **r["eval"]} for r in jsonl(os.path.join(other, "log.jsonl")) if "eval" in r] if ev: refs[os.path.basename(other.rstrip("/"))] = ev stdout_tail = [l for l in tail(os.path.join(d, "stdout.log"), 200) if re.search(r"Traceback|Error|OOM|out of memory|warn|stall", l, re.I)][-8:] return {"run": os.path.basename(d.rstrip("/")), "all_runs": [os.path.basename(x.rstrip("/")) for _, x in sorted(runs)], "steps": [{k: r.get(k) for k in keep} for r in steps], "last": steps[-1] if steps else None, "evals": evals, "notes": notes, "problems": stdout_tail, "refs": refs, "dir": d} def state(): p = procs() pt = pretrain() r = rl() if "grpo.py" in p: stage = "rl" elif "eval_agent.py" in p: stage = "eval" elif "teacher.py" in p: stage = "teacher" elif "train.py" in p and "/runs/sft" in p["train.py"]["args"]: stage = "sft" elif "train.py" in p: stage = "decay" if pt and pt["decay"] else "pretrain" elif r and r["steps"]: stage = "rl_done" else: stage = "idle" events = [l for l in tail(f"{DATA}/logs/teacher_window.log", 60) if l.startswith("[")][-6:] + \ [l for l in tail(f"{DATA}/logs/rl_after_decay.log", 60) if l.startswith("[")][-6:] + \ [l for l in tail(f"{DATA}/logs/rl_pipeline.log", 60) if l.startswith("[")][-8:] + \ [l for l in tail(f"{DATA}/logs/rl_watchdog.log", 60) if l.startswith("[")][-8:] events.sort() return {"now": time.strftime("%Y-%m-%d %H:%M:%S"), "stage": stage, "procs": p, "pretrain": pt, "teacher": teacher(), "eval_gate": eval_gate(r["dir"] if r else None), "rl": r, "events": events[-14:], "heldout_kinds": sorted(HELDOUT_KINDS)} HTML = os.path.join(os.path.dirname(os.path.abspath(__file__)), "dashboard.html") def PAGE(): with open(HTML) as f: # re-read per request, so page edits show up without a restart return f.read() class H(BaseHTTPRequestHandler): def log_message(self, *a): pass def do_GET(self): if self.path.startswith("/api/state"): try: body, ctype = json.dumps(state()).encode(), "application/json" except Exception as e: # never take the page down over one bad log line body, ctype = json.dumps({"error": repr(e)}).encode(), "application/json" elif self.path in ("/", "/index.html"): body, ctype = PAGE().encode(), "text/html; charset=utf-8" else: self.send_error(404) return self.send_response(200) self.send_header("Content-Type", ctype) self.send_header("Content-Length", str(len(body))) self.send_header("Cache-Control", "no-store") self.end_headers() self.wfile.write(body) if __name__ == "__main__": ap = argparse.ArgumentParser() ap.add_argument("--host", default="127.0.0.1") ap.add_argument("--port", type=int, default=8765) a = ap.parse_args() print(f"dashboard on http://{a.host}:{a.port}", flush=True) ThreadingHTTPServer((a.host, a.port), H).serve_forever()