Spaces:
Runtime error
Runtime error
File size: 3,789 Bytes
b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 b5285b3 ccba775 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 | """
Plot results produced by exp/run_paradox_tabular.py.
Creates PNG plots under down/graphics/ by default.
Now supports:
--csv path to results csv
--out output directory for pngs
--prefix filename prefix for plots
"""
from __future__ import annotations
import argparse
import csv
import os
from pathlib import Path
from typing import Dict, List
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
def _load_csv(path: Path) -> List[Dict[str, str]]:
with path.open("r", newline="") as f:
return list(csv.DictReader(f))
def _to_float(x: str) -> float:
try:
return float(x)
except Exception:
return float("nan")
def _build_parser() -> argparse.ArgumentParser:
p = argparse.ArgumentParser(description="Plot DOWN paradox tabular results.")
p.add_argument(
"--csv",
type=str,
default="",
help="Path to results CSV. If omitted, uses TWOQUARKS_RESULTS_DIR/down_paradox_tabular_results.csv or down/results/...",
)
p.add_argument(
"--out",
type=str,
default="",
help="Output directory for PNGs. If omitted, uses TWOQUARKS_GRAPHICS_DIR or down/graphics.",
)
p.add_argument("--prefix", type=str, default="Down", help="Prefix for output filenames.")
return p
def main() -> None:
args = _build_parser().parse_args()
quark_dir = Path(__file__).resolve().parents[1] # .../down
# Default dirs (env overrides)
results_dir = Path(os.environ.get("TWOQUARKS_RESULTS_DIR", (quark_dir / "results").as_posix()))
graphics_dir = Path(os.environ.get("TWOQUARKS_GRAPHICS_DIR", (quark_dir / "graphics").as_posix()))
# Respect CLI overrides
if args.out.strip():
graphics_dir = Path(args.out)
graphics_dir.mkdir(parents=True, exist_ok=True)
if args.csv.strip():
results_csv = Path(args.csv)
else:
results_dir.mkdir(parents=True, exist_ok=True)
results_csv = results_dir / "down_paradox_tabular_results.csv"
if not results_csv.exists():
raise FileNotFoundError(f"Missing results CSV: {results_csv}")
rows = _load_csv(results_csv)
if not rows:
raise RuntimeError(f"CSV is empty: {results_csv}")
prefix = args.prefix.strip() or "Down"
# Agents present in CSV
agents = sorted({r.get("agent", "unknown") for r in rows})
# Reward plot
plt.figure()
for a in agents:
xs = [int(r.get("global_episode", r.get("episode", 0)) or 0) for r in rows if r.get("agent") == a]
ys = [_to_float(r.get("episode_reward", r.get("return", "nan"))) for r in rows if r.get("agent") == a]
if xs:
plt.plot(xs, ys, label=a)
plt.xlabel("Global episode")
plt.ylabel("Episode reward")
plt.title(f"{prefix}: reward over training")
plt.legend()
out1 = graphics_dir / f"{prefix}_paradox_tabular_reward.png"
plt.savefig(out1, dpi=160, bbox_inches="tight")
plt.close()
# Rho plot (if present)
has_rho = any(("rho_state" in r) for r in rows) or any(("rho" in r) for r in rows)
if has_rho:
plt.figure()
for a in agents:
xs = [int(r.get("global_episode", r.get("episode", 0)) or 0) for r in rows if r.get("agent") == a]
ys = [_to_float(r.get("rho_state", r.get("rho", "nan"))) for r in rows if r.get("agent") == a]
if xs:
plt.plot(xs, ys, label=a)
plt.xlabel("Global episode")
plt.ylabel("rho")
plt.title(f"{prefix}: rho over training")
plt.legend()
out2 = graphics_dir / f"{prefix}_paradox_tabular_rho.png"
plt.savefig(out2, dpi=160, bbox_inches="tight")
plt.close()
print(f"Saved plots to {graphics_dir}")
if __name__ == "__main__":
main()
|