QuarksLab / Down /plot_dual_down_antidown.py
TwoQuarks's picture
Update Down/plot_dual_down_antidown.py
612a416 verified
Raw
History Blame Contribute Delete
4.92 kB
"""
Overlay plot: Down vs AntiDown (tabular experiments).
Reads CSVs from a shared results directory and writes a PNG overlay plot.
This is the canonical comparative result for the Down vs AntiDown experiment.
"""
from __future__ import annotations
# Headless-safe plotting (Hugging Face / Linux without DISPLAY)
import matplotlib
matplotlib.use("Agg")
import argparse
import csv
import os
from collections import defaultdict
from pathlib import Path
from typing import Dict, List, Tuple
import matplotlib.pyplot as plt
# ----------------------------
# Utilities
# ----------------------------
def _read_csv(path: Path) -> List[Dict[str, str]]:
with path.open("r", newline="") as f:
return list(csv.DictReader(f))
def _mean(xs: List[float]) -> float:
if not xs:
return float("nan")
return sum(xs) / len(xs)
# ----------------------------
# Series builders
# ----------------------------
def _down_series(rows: List[Dict[str, str]]) -> Tuple[List[int], List[float]]:
"""
Down CSV fields (expected):
phase, episode, agent, episode_reward, [global_episode], ...
Aggregate by (phase, episode), averaging across agents.
"""
bucket: Dict[Tuple[int, int], List[float]] = defaultdict(list)
for r in rows:
try:
phase = int(r["phase"])
ep = int(r["episode"])
rew = float(r["episode_reward"])
except Exception:
continue
bucket[(phase, ep)].append(rew)
xs, ys = [], []
g = 0
for key in sorted(bucket.keys()):
xs.append(g)
ys.append(_mean(bucket[key]))
g += 1
return xs, ys
def _antidown_series(rows: List[Dict[str, str]]) -> Tuple[List[int], List[float]]:
"""
AntiDown CSV fields (expected):
phase, episode, agent, total_reward, ...
Aggregate by (phase, episode), averaging across agents.
"""
bucket: Dict[Tuple[int, int], List[float]] = defaultdict(list)
for r in rows:
try:
phase = int(r["phase"])
ep = int(r["episode"])
rew = float(r["total_reward"])
except Exception:
continue
bucket[(phase, ep)].append(rew)
xs, ys = [], []
g = 0
for key in sorted(bucket.keys()):
xs.append(g)
ys.append(_mean(bucket[key]))
g += 1
return xs, ys
# ----------------------------
# CLI
# ----------------------------
def _build_parser() -> argparse.ArgumentParser:
p = argparse.ArgumentParser(description="Plot Down vs AntiDown (dual overlay).")
p.add_argument(
"--results_dir",
type=str,
default="",
help="Directory containing both CSVs. Defaults to TWOQUARKS_RESULTS_DIR or ./results.",
)
p.add_argument(
"--out",
type=str,
default="",
help="Output directory for PNG. Defaults to TWOQUARKS_GRAPHICS_DIR or ./graphics.",
)
p.add_argument(
"--prefix",
type=str,
default="dual_down_antidown",
help="Filename prefix for output image.",
)
return p
# ----------------------------
# Main
# ----------------------------
def main() -> None:
args = _build_parser().parse_args()
root = Path(__file__).resolve().parent
results_dir = Path(
args.results_dir
or os.environ.get("TWOQUARKS_RESULTS_DIR", (root / "results").as_posix())
)
graphics_dir = Path(
args.out
or os.environ.get("TWOQUARKS_GRAPHICS_DIR", (root / "graphics").as_posix())
)
graphics_dir.mkdir(parents=True, exist_ok=True)
down_csv = results_dir / "down_paradox_tabular_results.csv"
anti_csv = results_dir / "antidown_corrupted_valley_tabular.csv"
if not down_csv.exists():
raise FileNotFoundError(f"Missing Down CSV: {down_csv}")
if not anti_csv.exists():
raise FileNotFoundError(f"Missing AntiDown CSV: {anti_csv}")
down_rows = _read_csv(down_csv)
anti_rows = _read_csv(anti_csv)
if not down_rows:
raise RuntimeError("Down CSV is empty.")
if not anti_rows:
raise RuntimeError("AntiDown CSV is empty.")
x1, y1 = _down_series(down_rows)
x2, y2 = _antidown_series(anti_rows)
if len(x1) != len(x2):
print(
f"[dual] WARNING: series length mismatch "
f"(Down={len(x1)}, AntiDown={len(x2)})"
)
plt.figure(figsize=(10, 5))
plt.plot(x1, y1, label="Down (mean across agents)")
plt.plot(x2, y2, label="AntiDown (mean across agents)")
plt.xlabel("Global episode (phase-concatenated)")
plt.ylabel("Episode return")
plt.title("Down vs AntiDown — Tabular returns (mean across agents)")
plt.legend()
plt.tight_layout()
out_path = graphics_dir / f"{args.prefix}_episode_return.png"
plt.savefig(out_path, dpi=160)
plt.close()
print(f"[dual] saved {out_path}")
if __name__ == "__main__":
main()