#!/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 os
import re
import sys
import tempfile
import xml.etree.ElementTree as ET
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 write_svg(path, text):
"""Validate the finished SVG, then atomically swap it into place.
Readers never see partial output; an invalid render aborts without
touching the previous good file."""
try:
ET.fromstring(text)
except ET.ParseError as e:
print(f"[error] {path.name}: generated SVG invalid, keeping old file ({e})")
return False
tmp = path.with_name(path.name + ".tmp")
tmp.write_text(text)
os.replace(tmp, path)
return True
def svg_open(w, h, title):
return (f'\n"
return 1 if write_svg(out, s) else 0
# ---------- 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''
f"{esc(title)}: no valid points\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''
f"{esc(title)}\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'metric line number ({n} rows)\n\n')
return 1 if write_svg(out, s) else 0
# ---------- 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()
ET.fromstring(g1) # whole doc must parse -> closing tag present
ET.fromstring((td / "g2.svg").read_text())
assert 'width="1200"' in g1, "canvas must be >=1200 wide"
assert g1.count("= 2) + 2 # trends + v1 buckets + v3 trend
assert g1.count(" 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)")
try:
made += graph_eval(series_v2, series_v1, outdir / "graph_eval_winrates.svg")
except Exception as e:
print(f"[error] win-rate graph failed (loss graph still attempted): {e}")
rows = parse_metrics(metrics)
print(f"[info] metric rows parsed: {len(rows)}")
try:
made += graph_losses(rows, outdir / "graph_losses.svg")
except Exception as e:
print(f"[error] loss graph failed: {e}")
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()