feat(SAC_LSTM_Bot): readable eval chart — trends primary, raw dots secondary, v1 comparison separated
This commit is contained in:
@@ -3,8 +3,10 @@
|
||||
|
||||
Pure-stdlib SVG output (matplotlib not available on this box).
|
||||
Generates:
|
||||
docs/graph_eval_winrates.svg - eval win rate per opponent (+ rolling-20 mean, v1 overlay)
|
||||
docs/graph_losses.svg - critic/actor/alpha loss curves from training_metrics.jsonl
|
||||
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]
|
||||
@@ -20,10 +22,11 @@ from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parent.parent
|
||||
EVAL_RE = re.compile(r">>> \[eval\] win rate: (\d+)/(\d+) \(([\d.]+)%\) vs (\S+)")
|
||||
ROLL_WINDOW = 20
|
||||
COLORS = {"Corners": "#1f77b4", "Crazy": "#2ca02c", "Target": "#d62728"}
|
||||
V1_COLOR = "#888888"
|
||||
W, M_L, M_R, M_T, M_B = 1000, 70, 20, 40, 45
|
||||
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):
|
||||
@@ -61,7 +64,7 @@ def parse_metrics(path):
|
||||
return rows
|
||||
|
||||
|
||||
def rolling(vals, w=ROLL_WINDOW):
|
||||
def rolling(vals, w=TREND_WINDOW):
|
||||
out, s = [], 0.0
|
||||
for i, v in enumerate(vals):
|
||||
s += v
|
||||
@@ -94,6 +97,30 @@ def polyline(pts, color, width=1.5, dash=None, opacity=1.0):
|
||||
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'
|
||||
@@ -124,13 +151,12 @@ def ticks_log(vmin, vmax, p0, p1, n=5):
|
||||
|
||||
|
||||
def legend(items, x, y):
|
||||
"""items: [(color, dash, label)]"""
|
||||
"""items: [(color, label)]"""
|
||||
s = ""
|
||||
for i, (color, dash, label) in enumerate(items):
|
||||
for i, (color, label) in enumerate(items):
|
||||
yy = y + i * 18
|
||||
d = f' stroke-dasharray="{dash}"' if dash else ""
|
||||
s += (f'<line x1="{x}" y1="{yy}" x2="{x + 28}" y2="{yy}" stroke="{color}" '
|
||||
f'stroke-width="3"{d}/>\n'
|
||||
f'stroke-width="3"/>\n'
|
||||
f'<text x="{x + 34}" y="{yy + 4}" font-size="12">{esc(label)}</text>\n')
|
||||
return s
|
||||
|
||||
@@ -148,36 +174,70 @@ 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
|
||||
H = 560
|
||||
x0, y0, x1, y1 = M_L, H - M_B, W - M_R, M_T
|
||||
xmax = max(max((len(v) for v in s.values()), default=1) for s in (series_v2, series_v1))
|
||||
xmax = max(xmax, 2)
|
||||
# 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)
|
||||
ym = map_fn(y0, y1, 0, 100)
|
||||
s = svg_open(W, H, "SAC-LSTM campaign: eval win rate vs opponents (v2 live, v1 archived)")
|
||||
s += axis(x0, y0, x1, y1, ticks_linear(1, xmax, x0, x1),
|
||||
ticks_linear(0, 100, y0, y1), f"eval number (per opponent)",
|
||||
"win rate (%)")
|
||||
counts = []
|
||||
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 += polyline([(xm(i + 1), ym(v)) for i, v in enumerate(vals)], c, 1, opacity=0.30)
|
||||
s += polyline([(xm(i + 1), ym(v)) for i, v in enumerate(rolling(vals))], c, 3)
|
||||
counts.append(f"{name}: {len(vals)} evals")
|
||||
items.append((c, None, f"{name} (thick = rolling-{ROLL_WINDOW} mean)"))
|
||||
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", [])
|
||||
if v1:
|
||||
s += polyline([(xm(i + 1), ym(v)) for i, v in enumerate(v1)], V1_COLOR, 1, dash="6,4")
|
||||
s += polyline([(xm(i + 1), ym(v)) for i, v in enumerate(rolling(v1))], V1_COLOR, 2, dash="6,4")
|
||||
counts.append(f"v1 Corners (old run): {len(v1)} evals")
|
||||
items.append((V1_COLOR, "6,4", "v1 Corners (old run)"))
|
||||
s += legend(items, x1 - 260, y1 + 14)
|
||||
s += (f'<text x="{x0 + 8}" y="{y1 + 14}" font-size="11" fill="#555">'
|
||||
f'{esc("; ".join(counts))}</text>\n</svg>\n')
|
||||
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
|
||||
|
||||
@@ -278,6 +338,10 @@ def selftest():
|
||||
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"
|
||||
@@ -286,6 +350,11 @@ def selftest():
|
||||
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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user