"""
training/plot_curves.py — Generate reward curve HTML from training results
==========================================================================
Reads checkpoints/cpu_run/reward_curve.csv and produces a self-contained
HTML file with interactive Plotly charts — no matplotlib, no GPU.
Run:
python -m training.plot_curves --input checkpoints/cpu_run/reward_curve.csv
# Opens reward_curves.html in browser
"""
from __future__ import annotations
import argparse
import csv
import json
import os
import webbrowser
def load_csv(path: str) -> dict[str, list]:
data: dict[str, list] = {}
with open(path) as f:
reader = csv.DictReader(f)
for row in reader:
for k, v in row.items():
data.setdefault(k, []).append(float(v))
return data
def generate_html(data: dict[str, list], output_path: str) -> str:
steps = data.get("step", [])
f1 = data.get("overseer_f1", [])
det = data.get("detection_rate", [])
fp = data.get("false_positive_rate", [])
# Compute improvement annotations
f1_before = f1[0] if f1 else 0
f1_after = f1[-1] if f1 else 0
gap = f1_after - f1_before
html = f"""
NegotiArena Training Curves
🏛️ NegotiArena — Training Results
REINFORCE on CPU → same reward signal used for GRPO on GPU
Overseer F1 (Before)
{f1_before:.3f}
Random baseline
Overseer F1 (After)
{f1_after:.3f}
After {int(steps[-1]) if steps else 0} training steps
Improvement
+{gap:.3f}
{gap/f1_before*100:.0f}% relative gain
FP Rate (After)
{fp[-1]:.1%}
Was {fp[0]:.1%}
"""
with open(output_path, "w", encoding="utf-8") as f:
f.write(html)
return output_path
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--input", default="checkpoints/cpu_run/reward_curve.csv")
parser.add_argument("--output", default="reward_curves.html")
parser.add_argument("--open", action="store_true", default=True,
help="Open in browser after generating")
args = parser.parse_args()
if not os.path.exists(args.input):
print(f"❌ File not found: {args.input}")
print(" Run training first: python -m training.train_cpu --steps 200")
return
data = load_csv(args.input)
out = generate_html(data, args.output)
print(f"✅ Chart saved to {out}")
if args.open:
webbrowser.open(f"file://{os.path.abspath(out)}")
if __name__ == "__main__":
main()