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()