tiny-agent-112m / code /scripts /dashboard.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
8.53 kB
"""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: <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()