File size: 8,525 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 | """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()
|