395 lines
15 KiB
Python
395 lines
15 KiB
Python
#!/usr/bin/env python3
|
|
"""Plot SAC training progress from live campaign logs.
|
|
|
|
Pure-stdlib SVG output (matplotlib not available on this box).
|
|
Generates:
|
|
docs/graph_eval_winrates.svg - two stacked panels: Run 3 trend (raw dots +
|
|
rolling-10 thick line) per opponent, and
|
|
Run 1 vs Run 3 Corners on a fraction-of-run axis
|
|
docs/graph_losses.svg - critic/actor/alpha loss curves from training_metrics.jsonl
|
|
|
|
Usage:
|
|
python3 tools/plot_progress.py [campaign_log] [metrics_jsonl] [v1_log] [outdir]
|
|
All args optional; defaults relative to the SAC_LSTM_Bot/ root (parent of tools/).
|
|
python3 tools/plot_progress.py --selftest # tiny built-in sanity check
|
|
"""
|
|
import json
|
|
import math
|
|
import re
|
|
import sys
|
|
import tempfile
|
|
from pathlib import Path
|
|
|
|
ROOT = Path(__file__).resolve().parent.parent
|
|
EVAL_RE = re.compile(r">>> \[eval\] win rate: (\d+)/(\d+) \(([\d.]+)%\) vs (\S+)")
|
|
TREND_WINDOW = 10 # rolling mean shown as the thick trend line
|
|
BUCKETS = 40 # v1 downsampling for the Run1-vs-Run3 panel
|
|
COLORS = {"Corners": "#d62728", "Crazy": "#1f77b4", "Target": "#2ca02c"}
|
|
V1_COLOR = "#999999"
|
|
W, M_L, M_R, M_T, M_B = 1000, 70, 20, 40, 45 # graph 2 geometry (unchanged)
|
|
|
|
|
|
def parse_eval_series(path):
|
|
"""Return {opponent: [win% per eval, in file order]}."""
|
|
series = {}
|
|
if not path.is_file():
|
|
print(f"[skip] eval log not found: {path}")
|
|
return series
|
|
for line in path.read_text(errors="replace").splitlines():
|
|
m = EVAL_RE.search(line)
|
|
if not m:
|
|
continue
|
|
try:
|
|
pct = float(m.group(3))
|
|
except ValueError:
|
|
continue
|
|
series.setdefault(m.group(4), []).append(pct)
|
|
return series
|
|
|
|
|
|
def parse_metrics(path):
|
|
"""Return list of metric dicts, skipping malformed lines."""
|
|
rows = []
|
|
if not path.is_file():
|
|
print(f"[skip] metrics file not found: {path}")
|
|
return rows
|
|
for line in path.read_text(errors="replace").splitlines():
|
|
line = line.strip()
|
|
if not line:
|
|
continue
|
|
try:
|
|
rows.append(json.loads(line))
|
|
except json.JSONDecodeError:
|
|
continue
|
|
return rows
|
|
|
|
|
|
def rolling(vals, w=TREND_WINDOW):
|
|
out, s = [], 0.0
|
|
for i, v in enumerate(vals):
|
|
s += v
|
|
if i >= w:
|
|
s -= vals[i - w]
|
|
out.append(s / min(i + 1, w))
|
|
return out
|
|
|
|
|
|
# ---------- tiny SVG helpers ----------
|
|
|
|
def esc(s):
|
|
return str(s).replace("&", "&").replace("<", "<").replace(">", ">")
|
|
|
|
|
|
def svg_open(w, h, title):
|
|
return (f'<svg xmlns="http://www.w3.org/2000/svg" width="{w}" height="{h}" '
|
|
f'viewBox="0 0 {w} {h}" font-family="sans-serif">\n'
|
|
f'<rect width="{w}" height="{h}" fill="white"/>\n'
|
|
f'<text x="{w // 2}" y="22" text-anchor="middle" font-size="17" font-weight="bold">'
|
|
f"{esc(title)}</text>\n")
|
|
|
|
|
|
def polyline(pts, color, width=1.5, dash=None, opacity=1.0):
|
|
if len(pts) < 2:
|
|
return ""
|
|
d = f' stroke-dasharray="{dash}"' if dash else ""
|
|
p = " ".join(f"{x:.1f},{y:.1f}" for x, y in pts)
|
|
return (f'<polyline points="{p}" fill="none" stroke="{color}" '
|
|
f'stroke-width="{width}" opacity="{opacity}"{d}/>\n')
|
|
|
|
|
|
def dots(pts, color, r=2, opacity=0.25):
|
|
return "".join(f'<circle cx="{x:.1f}" cy="{y:.1f}" r="{r}" '
|
|
f'fill="{color}" opacity="{opacity}"/>\n' for x, y in pts)
|
|
|
|
|
|
def hgrid(x0, x1, ys):
|
|
return "".join(f'<line x1="{x0}" y1="{y:.1f}" x2="{x1}" y2="{y:.1f}" '
|
|
f'stroke="#dddddd"/>\n' for y in ys)
|
|
|
|
|
|
def bucket_means(vals, n):
|
|
"""Split vals into <=n contiguous buckets of near-equal size; per-bucket mean."""
|
|
if not vals:
|
|
return []
|
|
n = min(n, len(vals))
|
|
k, rem = divmod(len(vals), n)
|
|
out, i = [], 0
|
|
for b in range(n):
|
|
size = k + (1 if b < rem else 0)
|
|
out.append(sum(vals[i:i + size]) / size)
|
|
i += size
|
|
return out
|
|
|
|
|
|
def axis(x0, y0, x1, y1, xt, yt, xlabel, ylabel, ylog=False):
|
|
"""Draw axes + ticks + labels. xt/yt are (value, px) tick lists."""
|
|
s = (f'<line x1="{x0}" y1="{y0}" x2="{x1}" y2="{y0}" stroke="black"/>\n'
|
|
f'<line x1="{x0}" y1="{y0}" x2="{x0}" y2="{y1}" stroke="black"/>\n')
|
|
for v, px in xt:
|
|
s += (f'<line x1="{px:.1f}" y1="{y0}" x2="{px:.1f}" y2="{y0 + 4}" stroke="black"/>\n'
|
|
f'<text x="{px:.1f}" y="{y0 + 17}" text-anchor="middle" font-size="11">'
|
|
f"{esc(v)}</text>\n")
|
|
for v, py in yt:
|
|
s += (f'<line x1="{x0 - 4}" y1="{py:.1f}" x2="{x0}" y2="{py:.1f}" stroke="black"/>\n'
|
|
f'<text x="{x0 - 7}" y="{py + 4:.1f}" text-anchor="end" font-size="11">'
|
|
f"{esc(v)}</text>\n")
|
|
s += (f'<text x="{(x0 + x1) // 2}" y="{y0 + 33}" text-anchor="middle" font-size="12">'
|
|
f"{esc(xlabel)}</text>\n"
|
|
f'<text x="16" y="{(y0 + y1) // 2}" text-anchor="middle" font-size="12" '
|
|
f'transform="rotate(-90 16 {(y0 + y1) // 2})">{esc(ylabel)}</text>\n')
|
|
return s
|
|
|
|
|
|
def ticks_linear(vmin, vmax, p0, p1, n=5, fmt="{:.0f}"):
|
|
return [(fmt.format(vmin + (vmax - vmin) * i / (n - 1)),
|
|
p0 + (p1 - p0) * i / (n - 1)) for i in range(n)]
|
|
|
|
|
|
def ticks_log(vmin, vmax, p0, p1, n=5):
|
|
vals = [10 ** (vmin + (vmax - vmin) * i / (n - 1)) for i in range(n)]
|
|
return [("{:.3g}".format(v), p0 + (p1 - p0) * i / (n - 1)) for i, v in enumerate(vals)]
|
|
|
|
|
|
def legend(items, x, y):
|
|
"""items: [(color, label)]"""
|
|
s = ""
|
|
for i, (color, label) in enumerate(items):
|
|
yy = y + i * 18
|
|
s += (f'<line x1="{x}" y1="{yy}" x2="{x + 28}" y2="{yy}" stroke="{color}" '
|
|
f'stroke-width="3"/>\n'
|
|
f'<text x="{x + 34}" y="{yy + 4}" font-size="12">{esc(label)}</text>\n')
|
|
return s
|
|
|
|
|
|
def map_fn(p0, p1, vmin, vmax, log=False):
|
|
def f(v):
|
|
t = (math.log10(v) - vmin) / (vmax - vmin) if log else (v - vmin) / (vmax - vmin)
|
|
return p1 - max(0.0, min(1.0, t)) * (p1 - p0)
|
|
return f
|
|
|
|
|
|
# ---------- graph 1: eval win rates ----------
|
|
|
|
def graph_eval(series_v2, series_v1, out):
|
|
if not series_v2 and not series_v1:
|
|
print("[skip] no eval data at all -> no win-rate graph")
|
|
return 0
|
|
# panel geometry (local: graph 2 keeps the module-level constants)
|
|
CW, ML, MR = 1200, 64, 24
|
|
TOP, PH, GAP, MB = 78, 310, 84, 56
|
|
H = TOP + PH + GAP + PH + MB # 838 >= 800
|
|
x0, x1 = ML, CW - MR
|
|
s = svg_open(CW, H, "SAC-LSTM campaign: eval win rate vs opponents (Run 3 live, Run 1 archived)")
|
|
|
|
# ---- top panel: run 3 recent trend ----
|
|
y1t, y0t = TOP, TOP + PH
|
|
xmax = max(max((len(v) for v in series_v2.values()), default=1), 2)
|
|
xm = map_fn(x0, x1, 1, xmax)
|
|
ymt = map_fn(y0t, y1t, 0, 100)
|
|
s += hgrid(x0, x1, [ymt(v) for v in range(0, 101, 20)])
|
|
s += (f'<text x="{x0 + 6}" y="{TOP - 34}" font-size="14" font-weight="bold">'
|
|
f"Run 3 — recent trend</text>\n"
|
|
f'<text x="{x0 + 6}" y="{TOP - 16}" font-size="11" fill="#555">'
|
|
f"win % per test match — thick line = trend</text>\n")
|
|
xt = [(str(v), xm(v)) for v in range(10, xmax + 1, 10)] or \
|
|
[("1", xm(1)), (str(xmax), xm(xmax))]
|
|
s += axis(x0, y0t, x1, y1t, xt,
|
|
ticks_linear(0, 100, y0t, y1t, n=6),
|
|
"eval number", "win rate (%)")
|
|
items = []
|
|
for name in ("Corners", "Crazy", "Target"):
|
|
vals = series_v2.get(name, [])
|
|
if not vals:
|
|
continue
|
|
c = COLORS[name]
|
|
s += dots([(xm(i + 1), ymt(v)) for i, v in enumerate(vals)], c)
|
|
s += polyline([(xm(i + 1), ymt(v))
|
|
for i, v in enumerate(rolling(vals, TREND_WINDOW))], c, 3.5)
|
|
items.append((c, f"{name} — {len(vals)} evals"))
|
|
s += legend(items, x0 + 12, y1t + 14)
|
|
|
|
# ---- bottom panel: run 1 vs run 3 on the same enemy ----
|
|
tb = TOP + PH + GAP
|
|
y1b, y0b = tb, tb + PH
|
|
ymb = map_fn(y0b, y1b, 0, 100)
|
|
|
|
def xf(f):
|
|
return x0 + (x1 - x0) * f
|
|
|
|
s += hgrid(x0, x1, [ymb(v) for v in range(0, 101, 20)])
|
|
s += (f'<text x="{x0 + 6}" y="{tb - 22}" font-size="14" font-weight="bold">'
|
|
f"Run 1 vs Run 3 (Corners) — same enemy — did learning stick?</text>\n")
|
|
s += axis(x0, y0b, x1, y1b,
|
|
[(("0", "0.2", "0.4", "0.6", "0.8", "1")[i], xf(i / 5)) for i in range(6)],
|
|
ticks_linear(0, 100, y0b, y1b, n=6),
|
|
"fraction of run (start → end)", "win rate (%)")
|
|
items_b = []
|
|
v1 = series_v1.get("Corners", [])
|
|
bm = bucket_means(v1, BUCKETS)
|
|
if bm:
|
|
fb = [xf((i + 0.5) / len(bm)) for i in range(len(bm))]
|
|
s += polyline(list(zip(fb, map(ymb, bm))), V1_COLOR, 1.5)
|
|
items_b.append((V1_COLOR, f"Run 1 (old) — {len(v1)} evals in {len(bm)} buckets"))
|
|
c3 = series_v2.get("Corners", [])
|
|
if c3:
|
|
rm = rolling(c3, TREND_WINDOW)
|
|
fr = [xf((i + 1) / len(rm)) for i in range(len(rm))]
|
|
s += polyline(list(zip(fr, map(ymb, rm))), COLORS["Corners"], 3.5)
|
|
items_b.append((COLORS["Corners"], f"Run 3 trend (rolling-{TREND_WINDOW})"))
|
|
s += legend(items_b, x0 + 12, y1b + 14)
|
|
|
|
out.write_text(s)
|
|
return 1
|
|
|
|
|
|
# ---------- graph 2: loss curves ----------
|
|
|
|
def graph_losses(rows, out):
|
|
if not rows:
|
|
print("[skip] no metric rows -> no loss graph")
|
|
return 0
|
|
|
|
def col(key, log=False):
|
|
vals = []
|
|
for r in rows:
|
|
try:
|
|
v = abs(float(r[key])) # abs: log panels plot magnitude
|
|
except (KeyError, TypeError, ValueError):
|
|
continue
|
|
if log and v <= 0:
|
|
continue
|
|
vals.append(v)
|
|
return vals
|
|
|
|
critic = col("critic_loss", log=True)
|
|
actor = col("actor_loss", log=True)
|
|
alpha = col("alpha")
|
|
if not (critic or actor or alpha):
|
|
print("[skip] metric rows lack critic_loss/actor_loss/alpha -> no loss graph")
|
|
return 0
|
|
|
|
PH, GAP = 210, 55
|
|
H = M_T + 3 * PH + 2 * GAP + M_B
|
|
x0, x1 = M_L, W - M_R
|
|
n = len(rows)
|
|
s = svg_open(W, H, "SAC-LSTM training losses (x = metric line number)")
|
|
|
|
def panel(top, vals, title, log, ylab, fixed_range=None):
|
|
nonlocal s
|
|
y0, y1 = top + PH, top
|
|
if not vals:
|
|
s += (f'<text x="{x0 + 10}" y="{top + 20}" font-size="12" fill="#a00">'
|
|
f"{esc(title)}: no valid points</text>\n")
|
|
return
|
|
if log:
|
|
lo = math.floor(math.log10(min(vals)))
|
|
hi = math.ceil(math.log10(max(vals)))
|
|
if lo == hi:
|
|
hi = lo + 1
|
|
yt = ticks_log(lo, hi, y0, y1)
|
|
else:
|
|
lo, hi = fixed_range or (min(vals), max(vals))
|
|
if lo == hi:
|
|
hi = lo + 1
|
|
yt = ticks_linear(lo, hi, y0, y1, n=5, fmt="{:.3g}")
|
|
s += (f'<text x="{x0 + 6}" y="{top + 16}" font-size="13" font-weight="bold">'
|
|
f"{esc(title)}</text>\n")
|
|
s += axis(x0, y0, x1, y1, ticks_linear(1, n, x0, x1, fmt="{:.0f}"), yt,
|
|
"", ylab)
|
|
ym = map_fn(y0, y1, lo, hi, log=log)
|
|
pts = []
|
|
for i, r in enumerate(rows):
|
|
try:
|
|
v = abs(float(r[title_key(title)]))
|
|
except (KeyError, TypeError, ValueError):
|
|
continue
|
|
if log and v <= 0:
|
|
continue
|
|
pts.append((x0 + (x1 - x0) * i / max(n - 1, 1), ym(v)))
|
|
s += polyline(pts, "#1f77b4", 1.6)
|
|
|
|
def title_key(title):
|
|
return {"critic_loss (log scale)": "critic_loss",
|
|
"actor_loss |abs| (log scale)": "actor_loss",
|
|
"alpha (temperature)": "alpha"}[title]
|
|
|
|
panel(M_T, critic, "critic_loss (log scale)", True, "critic_loss")
|
|
panel(M_T + PH + GAP, actor, "actor_loss |abs| (log scale)", True, "|actor_loss|")
|
|
panel(M_T + 2 * (PH + GAP), alpha, "alpha (temperature)", False, "alpha",
|
|
fixed_range=(0, max(1.0, max(alpha))))
|
|
s += (f'<text x="{x0 + (x1 - x0) // 2}" y="{H - 8}" text-anchor="middle" '
|
|
f'font-size="12">metric line number ({n} rows)</text>\n</svg>\n')
|
|
out.write_text(s)
|
|
return 1
|
|
|
|
|
|
# ---------- selftest ----------
|
|
|
|
def selftest():
|
|
with tempfile.TemporaryDirectory() as td:
|
|
td = Path(td)
|
|
(td / "log").write_text(
|
|
">>> [eval] win rate: 3/10 (30%) vs Corners\n"
|
|
"garbage line\n"
|
|
">>> [eval] win rate: 7/10 (70%) vs Crazy\n"
|
|
">>> [eval] win rate: broken\n"
|
|
">>> [eval] win rate: 5/10 (50%) vs Corners\n")
|
|
ser = parse_eval_series(td / "log")
|
|
assert ser == {"Corners": [30.0, 50.0], "Crazy": [70.0]}, ser
|
|
assert rolling([10] * 25, 20)[-1] == 10.0
|
|
assert rolling([1, 2, 3], 20) == [1.0, 1.5, 2.0]
|
|
assert bucket_means(list(range(1287)), BUCKETS) is not None
|
|
assert len(bucket_means(list(range(1287)), BUCKETS)) == BUCKETS
|
|
bm = bucket_means([0, 10], BUCKETS)
|
|
assert bm == [0.0, 10.0], bm # fewer points than buckets -> no empty buckets
|
|
(td / "m.jsonl").write_text(
|
|
'{"critic_loss": 10, "actor_loss": -2, "alpha": 0.5}\n'
|
|
"not json\n"
|
|
'{"critic_loss": 100, "actor_loss": -4, "alpha": 0.25}\n')
|
|
rows = parse_metrics(td / "m.jsonl")
|
|
assert len(rows) == 2 and rows[1]["critic_loss"] == 100
|
|
ok = graph_losses(rows, td / "g2.svg") and graph_eval(ser, {"Corners": [0, 10]}, td / "g1.svg")
|
|
assert ok and (td / "g1.svg").stat().st_size > 500
|
|
g1 = (td / "g1.svg").read_text()
|
|
assert 'width="1200"' in g1, "canvas must be >=1200 wide"
|
|
assert g1.count("<circle") == sum(len(v) for v in ser.values()) # one dot per raw eval
|
|
n_trend = sum(1 for v in ser.values() if len(v) >= 2) + 2 # trends + v1 buckets + v3 trend
|
|
assert g1.count("<polyline") == n_trend
|
|
assert graph_eval({}, {}, td / "none.svg") == 0 # missing data handled
|
|
print("selftest OK")
|
|
|
|
|
|
def main():
|
|
if "--selftest" in sys.argv:
|
|
selftest()
|
|
return
|
|
args = [a for a in sys.argv[1:] if not a.startswith("-")]
|
|
campaign = Path(args[0]) if len(args) > 0 else ROOT / "campaign_v2_stdout.log"
|
|
metrics = Path(args[1]) if len(args) > 1 else ROOT / "training_metrics.jsonl"
|
|
v1log = Path(args[2]) if len(args) > 2 else ROOT / "weights_v1_archive" / "campaign_stdout.log"
|
|
outdir = Path(args[3]) if len(args) > 3 else ROOT / "docs"
|
|
outdir.mkdir(parents=True, exist_ok=True)
|
|
|
|
made = 0
|
|
series_v2 = parse_eval_series(campaign)
|
|
print(f"[info] v2 evals parsed: " +
|
|
", ".join(f"{k}={len(v)}" for k, v in sorted(series_v2.items())) or "(none)")
|
|
series_v1 = parse_eval_series(v1log)
|
|
print(f"[info] v1 evals parsed: " +
|
|
", ".join(f"{k}={len(v)}" for k, v in sorted(series_v1.items())) or "(none)")
|
|
made += graph_eval(series_v2, series_v1, outdir / "graph_eval_winrates.svg")
|
|
|
|
rows = parse_metrics(metrics)
|
|
print(f"[info] metric rows parsed: {len(rows)}")
|
|
made += graph_losses(rows, outdir / "graph_losses.svg")
|
|
|
|
if made:
|
|
print(f"[done] {made} graph(s) written to {outdir}")
|
|
else:
|
|
print("[error] nothing plotted - check paths above")
|
|
sys.exit(1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|