feat(SAC_LSTM_Bot): progress graph tooling + first graphs
This commit is contained in:
@@ -318,3 +318,7 @@ SACLSTM_LR_CRITIC=1e-4 \
|
|||||||
| Crash banners | 0 |
|
| Crash banners | 0 |
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
### Progress graphs
|
||||||
|
|
||||||
|
Graphs live in `docs/` (`graph_eval_winrates.svg`, `graph_losses.svg`). Regenerate anytime: `python3 tools/plot_progress.py` (pure stdlib, ~0.05 s; paths overridable via argv, `--selftest` for sanity check).
|
||||||
|
|||||||
File diff suppressed because one or more lines are too long
|
After Width: | Height: | Size: 34 KiB |
@@ -0,0 +1,83 @@
|
|||||||
|
<svg xmlns="http://www.w3.org/2000/svg" width="1000" height="825" viewBox="0 0 1000 825" font-family="sans-serif">
|
||||||
|
<rect width="1000" height="825" fill="white"/>
|
||||||
|
<text x="500" y="22" text-anchor="middle" font-size="17" font-weight="bold">SAC-LSTM training losses (x = metric line number)</text>
|
||||||
|
<text x="76" y="56" font-size="13" font-weight="bold">critic_loss (log scale)</text>
|
||||||
|
<line x1="70" y1="250" x2="980" y2="250" stroke="black"/>
|
||||||
|
<line x1="70" y1="250" x2="70" y2="40" stroke="black"/>
|
||||||
|
<line x1="70.0" y1="250" x2="70.0" y2="254" stroke="black"/>
|
||||||
|
<text x="70.0" y="267" text-anchor="middle" font-size="11">1</text>
|
||||||
|
<line x1="297.5" y1="250" x2="297.5" y2="254" stroke="black"/>
|
||||||
|
<text x="297.5" y="267" text-anchor="middle" font-size="11">25</text>
|
||||||
|
<line x1="525.0" y1="250" x2="525.0" y2="254" stroke="black"/>
|
||||||
|
<text x="525.0" y="267" text-anchor="middle" font-size="11">48</text>
|
||||||
|
<line x1="752.5" y1="250" x2="752.5" y2="254" stroke="black"/>
|
||||||
|
<text x="752.5" y="267" text-anchor="middle" font-size="11">72</text>
|
||||||
|
<line x1="980.0" y1="250" x2="980.0" y2="254" stroke="black"/>
|
||||||
|
<text x="980.0" y="267" text-anchor="middle" font-size="11">96</text>
|
||||||
|
<line x1="66" y1="250.0" x2="70" y2="250.0" stroke="black"/>
|
||||||
|
<text x="63" y="254.0" text-anchor="end" font-size="11">0.1</text>
|
||||||
|
<line x1="66" y1="197.5" x2="70" y2="197.5" stroke="black"/>
|
||||||
|
<text x="63" y="201.5" text-anchor="end" font-size="11">3.16e+03</text>
|
||||||
|
<line x1="66" y1="145.0" x2="70" y2="145.0" stroke="black"/>
|
||||||
|
<text x="63" y="149.0" text-anchor="end" font-size="11">1e+08</text>
|
||||||
|
<line x1="66" y1="92.5" x2="70" y2="92.5" stroke="black"/>
|
||||||
|
<text x="63" y="96.5" text-anchor="end" font-size="11">3.16e+12</text>
|
||||||
|
<line x1="66" y1="40.0" x2="70" y2="40.0" stroke="black"/>
|
||||||
|
<text x="63" y="44.0" text-anchor="end" font-size="11">1e+17</text>
|
||||||
|
<text x="525" y="283" text-anchor="middle" font-size="12"></text>
|
||||||
|
<text x="16" y="145" text-anchor="middle" font-size="12" transform="rotate(-90 16 145)">critic_loss</text>
|
||||||
|
<polyline points="70.0,67.6 79.6,66.5 89.2,67.9 98.7,70.0 108.3,72.6 117.9,74.1 127.5,73.1 137.1,71.4 146.6,240.6 156.2,65.5 165.8,55.6 175.4,61.7 184.9,67.4 194.5,71.4 204.1,74.2 213.7,74.9 223.3,75.7 232.8,75.9 242.4,76.9 252.0,76.5 261.6,75.5 271.2,74.8 280.7,185.2 290.3,70.5 299.9,68.6 309.5,67.4 319.1,202.4 328.6,60.3 338.2,60.4 347.8,59.0 357.4,59.7 366.9,58.6 376.5,56.8 386.1,55.7 395.7,56.0 405.3,55.9 414.8,50.4 424.4,50.9 434.0,53.5 443.6,48.3 453.2,45.7 462.7,240.6 472.3,240.6 481.9,56.2 491.5,59.1 501.1,63.5 510.6,65.0 520.2,240.6 529.8,68.4 539.4,70.5 548.9,185.8 558.5,71.7 568.1,72.8 577.7,69.0 587.3,62.6 596.8,56.9 606.4,63.6 616.0,240.6 625.6,240.6 635.2,73.5 644.7,74.7 654.3,73.7 663.9,73.7 673.5,240.6 683.1,73.7 692.6,74.6 702.2,72.8 711.8,70.5 721.4,72.5 730.9,70.7 740.5,73.1 750.1,69.5 759.7,67.9 769.3,67.1 778.8,64.0 788.4,53.5 798.0,184.6 807.6,63.5 817.2,114.6 826.7,161.5 836.3,197.0 845.9,204.7 855.5,205.0 865.1,205.3 874.6,205.5 884.2,205.6 893.8,205.7 903.4,205.6 912.9,206.0 922.5,206.2 932.1,206.4 941.7,206.5 951.3,206.8 960.8,207.0 970.4,240.5 980.0,240.5" fill="none" stroke="#1f77b4" stroke-width="1.6" opacity="1.0"/>
|
||||||
|
<text x="76" y="321" font-size="13" font-weight="bold">actor_loss |abs| (log scale)</text>
|
||||||
|
<line x1="70" y1="515" x2="980" y2="515" stroke="black"/>
|
||||||
|
<line x1="70" y1="515" x2="70" y2="305" stroke="black"/>
|
||||||
|
<line x1="70.0" y1="515" x2="70.0" y2="519" stroke="black"/>
|
||||||
|
<text x="70.0" y="532" text-anchor="middle" font-size="11">1</text>
|
||||||
|
<line x1="297.5" y1="515" x2="297.5" y2="519" stroke="black"/>
|
||||||
|
<text x="297.5" y="532" text-anchor="middle" font-size="11">25</text>
|
||||||
|
<line x1="525.0" y1="515" x2="525.0" y2="519" stroke="black"/>
|
||||||
|
<text x="525.0" y="532" text-anchor="middle" font-size="11">48</text>
|
||||||
|
<line x1="752.5" y1="515" x2="752.5" y2="519" stroke="black"/>
|
||||||
|
<text x="752.5" y="532" text-anchor="middle" font-size="11">72</text>
|
||||||
|
<line x1="980.0" y1="515" x2="980.0" y2="519" stroke="black"/>
|
||||||
|
<text x="980.0" y="532" text-anchor="middle" font-size="11">96</text>
|
||||||
|
<line x1="66" y1="515.0" x2="70" y2="515.0" stroke="black"/>
|
||||||
|
<text x="63" y="519.0" text-anchor="end" font-size="11">1</text>
|
||||||
|
<line x1="66" y1="462.5" x2="70" y2="462.5" stroke="black"/>
|
||||||
|
<text x="63" y="466.5" text-anchor="end" font-size="11">56.2</text>
|
||||||
|
<line x1="66" y1="410.0" x2="70" y2="410.0" stroke="black"/>
|
||||||
|
<text x="63" y="414.0" text-anchor="end" font-size="11">3.16e+03</text>
|
||||||
|
<line x1="66" y1="357.5" x2="70" y2="357.5" stroke="black"/>
|
||||||
|
<text x="63" y="361.5" text-anchor="end" font-size="11">1.78e+05</text>
|
||||||
|
<line x1="66" y1="305.0" x2="70" y2="305.0" stroke="black"/>
|
||||||
|
<text x="63" y="309.0" text-anchor="end" font-size="11">1e+07</text>
|
||||||
|
<text x="525" y="548" text-anchor="middle" font-size="12"></text>
|
||||||
|
<text x="16" y="410" text-anchor="middle" font-size="12" transform="rotate(-90 16 410)">|actor_loss|</text>
|
||||||
|
<polyline points="70.0,320.3 79.6,325.2 89.2,327.4 98.7,330.1 108.3,333.3 117.9,335.0 127.5,333.9 137.1,332.2 146.6,328.2 156.2,325.3 165.8,318.1 175.4,313.9 184.9,325.2 194.5,331.1 204.1,334.8 213.7,337.2 223.3,338.1 232.8,338.6 242.4,339.7 252.0,340.5 261.6,340.3 271.2,340.1 280.7,339.5 290.3,339.3 299.9,338.2 309.5,337.7 319.1,335.9 328.6,335.6 338.2,332.5 347.8,331.0 357.4,329.6 366.9,329.0 376.5,329.3 386.1,329.7 395.7,327.8 405.3,327.5 414.8,329.1 424.4,329.6 434.0,327.7 443.6,330.3 453.2,331.0 462.7,332.4 472.3,335.8 481.9,338.6 491.5,339.7 501.1,341.4 510.6,342.5 520.2,344.1 529.8,345.3 539.4,346.6 548.9,347.7 558.5,348.6 568.1,349.3 577.7,348.9 587.3,348.7 596.8,347.9 606.4,346.7 616.0,344.9 625.6,342.6 635.2,341.6 644.7,340.5 654.3,339.6 663.9,339.5 673.5,340.3 683.1,341.3 692.6,340.4 702.2,341.3 711.8,342.4 721.4,341.3 730.9,341.7 740.5,340.4 750.1,342.0 759.7,342.5 769.3,342.2 778.8,342.7 788.4,344.5 798.0,347.4 807.6,347.7 817.2,385.3 826.7,446.3 836.3,492.0 845.9,501.9 855.5,502.3 865.1,502.7 874.6,502.9 884.2,503.0 893.8,503.2 903.4,503.4 912.9,503.6 922.5,503.8 932.1,504.0 941.7,504.2 951.3,504.6 960.8,504.9 970.4,505.0 980.0,505.2" fill="none" stroke="#1f77b4" stroke-width="1.6" opacity="1.0"/>
|
||||||
|
<text x="76" y="586" font-size="13" font-weight="bold">alpha (temperature)</text>
|
||||||
|
<line x1="70" y1="780" x2="980" y2="780" stroke="black"/>
|
||||||
|
<line x1="70" y1="780" x2="70" y2="570" stroke="black"/>
|
||||||
|
<line x1="70.0" y1="780" x2="70.0" y2="784" stroke="black"/>
|
||||||
|
<text x="70.0" y="797" text-anchor="middle" font-size="11">1</text>
|
||||||
|
<line x1="297.5" y1="780" x2="297.5" y2="784" stroke="black"/>
|
||||||
|
<text x="297.5" y="797" text-anchor="middle" font-size="11">25</text>
|
||||||
|
<line x1="525.0" y1="780" x2="525.0" y2="784" stroke="black"/>
|
||||||
|
<text x="525.0" y="797" text-anchor="middle" font-size="11">48</text>
|
||||||
|
<line x1="752.5" y1="780" x2="752.5" y2="784" stroke="black"/>
|
||||||
|
<text x="752.5" y="797" text-anchor="middle" font-size="11">72</text>
|
||||||
|
<line x1="980.0" y1="780" x2="980.0" y2="784" stroke="black"/>
|
||||||
|
<text x="980.0" y="797" text-anchor="middle" font-size="11">96</text>
|
||||||
|
<line x1="66" y1="780.0" x2="70" y2="780.0" stroke="black"/>
|
||||||
|
<text x="63" y="784.0" text-anchor="end" font-size="11">0</text>
|
||||||
|
<line x1="66" y1="727.5" x2="70" y2="727.5" stroke="black"/>
|
||||||
|
<text x="63" y="731.5" text-anchor="end" font-size="11">0.252</text>
|
||||||
|
<line x1="66" y1="675.0" x2="70" y2="675.0" stroke="black"/>
|
||||||
|
<text x="63" y="679.0" text-anchor="end" font-size="11">0.504</text>
|
||||||
|
<line x1="66" y1="622.5" x2="70" y2="622.5" stroke="black"/>
|
||||||
|
<text x="63" y="626.5" text-anchor="end" font-size="11">0.756</text>
|
||||||
|
<line x1="66" y1="570.0" x2="70" y2="570.0" stroke="black"/>
|
||||||
|
<text x="63" y="574.0" text-anchor="end" font-size="11">1.01</text>
|
||||||
|
<text x="525" y="813" text-anchor="middle" font-size="12"></text>
|
||||||
|
<text x="16" y="675" text-anchor="middle" font-size="12" transform="rotate(-90 16 675)">alpha</text>
|
||||||
|
<polyline points="70.0,778.3 79.6,778.1 89.2,778.0 98.7,777.8 108.3,777.5 117.9,777.1 127.5,777.0 137.1,776.8 146.6,776.6 156.2,776.5 165.8,776.2 175.4,775.8 184.9,775.7 194.5,776.0 204.1,776.2 213.7,776.5 223.3,776.7 232.8,777.0 242.4,777.3 252.0,777.5 261.6,777.8 271.2,778.0 280.7,778.3 290.3,778.5 299.9,778.8 309.5,778.9 319.1,779.3 328.6,779.4 338.2,779.2 347.8,779.1 357.4,778.9 366.9,778.8 376.5,778.6 386.1,778.4 395.7,778.2 405.3,778.0 414.8,777.6 424.4,777.4 434.0,777.1 443.6,776.9 453.2,776.6 462.7,776.4 472.3,776.3 481.9,776.5 491.5,776.6 501.1,776.8 510.6,776.9 520.2,777.1 529.8,777.3 539.4,777.5 548.9,777.8 558.5,778.0 568.1,778.2 577.7,778.4 587.3,778.8 596.8,779.1 606.4,779.3 616.0,779.8 625.6,780.0 635.2,779.8 644.7,779.6 654.3,779.4 663.9,779.1 673.5,778.9 683.1,778.4 692.6,778.2 702.2,777.9 711.8,777.7 721.4,777.4 730.9,777.3 740.5,777.1 750.1,776.8 759.7,776.6 769.3,776.5 778.8,776.5 788.4,776.7 798.0,777.2 807.6,777.4 817.2,777.5 826.7,777.3 836.3,777.0 845.9,776.9 855.5,776.6 865.1,776.2 874.6,776.1 884.2,776.0 893.8,775.9 903.4,775.6 912.9,775.5 922.5,775.3 932.1,775.0 941.7,774.8 951.3,774.5 960.8,774.2 970.4,774.1 980.0,774.0" fill="none" stroke="#1f77b4" stroke-width="1.6" opacity="1.0"/>
|
||||||
|
<text x="525" y="817" text-anchor="middle" font-size="12">metric line number (96 rows)</text>
|
||||||
|
</svg>
|
||||||
|
After Width: | Height: | Size: 8.9 KiB |
@@ -0,0 +1,325 @@
|
|||||||
|
#!/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 - eval win rate per opponent (+ rolling-20 mean, v1 overlay)
|
||||||
|
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+)")
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
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=ROLL_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 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, dash, label)]"""
|
||||||
|
s = ""
|
||||||
|
for i, (color, dash, 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'<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
|
||||||
|
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)
|
||||||
|
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 = []
|
||||||
|
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)"))
|
||||||
|
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')
|
||||||
|
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]
|
||||||
|
(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
|
||||||
|
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()
|
||||||
Reference in New Issue
Block a user