HaM_World / scripts /replot_main_curves.py
haodongcui's picture
first commit (part 4)
0838417 verified
Raw History Blame Contribute Delete
10.5 kB
from __future__ import annotations
import re
from pathlib import Path
import sys
import matplotlib
matplotlib.use("Agg")
import matplotlib.image as mpimg
import matplotlib.pyplot as plt
from matplotlib.ticker import FuncFormatter, MaxNLocator
import numpy as np
SCRIPT_ROOT = Path(__file__).resolve().parents[1]
if str(SCRIPT_ROOT) not in sys.path:
sys.path.insert(0, str(SCRIPT_ROOT))
from common import RESULTS_MAIN_FIG_ROOT, RUNS_MAIN_ROOT
LOG_DIR = RUNS_MAIN_ROOT / "model_based" / "batch" / "jobs"
LOW_BUDGET_LOG_DIR = RUNS_MAIN_ROOT / "model_free" / "batch" / "jobs"
OUT_DIR = RESULTS_MAIN_FIG_ROOT
OUT_DIR.mkdir(parents=True, exist_ok=True)
ASSEMBLED_STEM = RESULTS_MAIN_FIG_ROOT / "main_results_curves_assembled"
ALGOS = ["hamworld", "tdmpc2", "dreamerv3", "ppo_low", "sac_low"]
ALGO_LABELS = {
"hamworld": "HaM-World",
"tdmpc2": "TD-MPC2",
"dreamerv3": "DreamerV3",
"ppo_low": "PPO",
"sac_low": "SAC",
}
ALGO_COLOR = {
"hamworld": "#c63d2f",
"tdmpc2": "#2563eb",
"dreamerv3": "#0f766e",
"ppo_low": "#94a3b8",
"sac_low": "#b7791f",
}
LINEWIDTH = {algo: (2.4 if algo == "hamworld" else 1.7) for algo in ALGOS}
ZORDER = {"hamworld": 6, "tdmpc2": 5, "dreamerv3": 4, "sac_low": 3, "ppo_low": 2}
LINESTYLE = {
"hamworld": "solid",
"tdmpc2": "solid",
"dreamerv3": "solid",
"ppo_low": (0, (5.0, 2.2)),
"sac_low": (0, (3.2, 1.6, 1.2, 1.6)),
}
BAND_ALPHA = {
"hamworld": 0.16,
"tdmpc2": 0.12,
"dreamerv3": 0.12,
"ppo_low": 0.08,
"sac_low": 0.08,
}
TASKS = ["cartpole_swingup", "cheetah_run", "finger_spin", "reacher_easy"]
TASK_TITLE = {
"cartpole_swingup": "Cartpole Swingup",
"cheetah_run": "Cheetah Run",
"finger_spin": "Finger Spin",
"reacher_easy": "Reacher Easy",
}
SEEDS = [7, 8, 9]
STEP_RE = re.compile(r"^\[step\s+(\d+)\]\s+(.*)$")
EVAL_RE = re.compile(r"eval/return_mean=([-+]?\d+(?:\.\d+)?(?:[eE][-+]?\d+)?)")
plt.rcParams.update(
{
"font.family": "DejaVu Sans",
"font.size": 10.5,
"axes.titlesize": 12,
"axes.titleweight": "semibold",
"axes.labelsize": 10.5,
"axes.labelcolor": "#334155",
"axes.edgecolor": "#cbd5e1",
"axes.linewidth": 1.0,
"axes.spines.top": False,
"axes.spines.right": False,
"xtick.color": "#475569",
"ytick.color": "#475569",
"xtick.labelsize": 9.5,
"ytick.labelsize": 9.5,
"xtick.major.size": 3.0,
"ytick.major.size": 3.0,
"xtick.major.width": 0.8,
"ytick.major.width": 0.8,
"legend.fontsize": 10.5,
"legend.frameon": False,
"figure.dpi": 220,
"savefig.dpi": 300,
"savefig.bbox": "tight",
"savefig.pad_inches": 0.03,
}
)
def _parse_eval_curve(path: Path) -> dict[int, float]:
out: dict[int, float] = {}
if not path.exists():
return out
for line in path.read_text(errors="ignore").splitlines():
match = STEP_RE.match(line)
if not match:
continue
eval_match = EVAL_RE.search(match.group(2))
if eval_match:
out[int(match.group(1))] = float(eval_match.group(1))
return out
def _curve_for_algo_task(algo: str, task: str) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
per_seed = []
for seed in SEEDS:
if algo in {"ppo_low", "sac_low"}:
path = LOW_BUDGET_LOG_DIR / f"{algo.replace('_low', '')}_{task}_seed{seed}.log"
else:
path = LOG_DIR / f"{algo}_{task}_seed{seed}.log"
parsed = _parse_eval_curve(path)
if parsed:
per_seed.append(parsed)
if not per_seed:
return np.array([]), np.array([]), np.array([]), np.array([])
steps = sorted({step for parsed in per_seed for step in parsed})
xs, means, lows, highs = [], [], [], []
for step in steps:
values = np.array([parsed[step] for parsed in per_seed if step in parsed], dtype=float)
if values.size == 0:
continue
xs.append(step)
means.append(float(values.mean()))
lows.append(float(values.min()))
highs.append(float(values.max()))
return np.array(xs), np.array(means), np.array(lows), np.array(highs)
def _smooth(y: np.ndarray, width: int = 3) -> np.ndarray:
if y.size < width:
return y
kernel = np.ones(width) / width
pad = width // 2
padded = np.pad(y, (pad, pad), mode="edge")
return np.convolve(padded, kernel, mode="valid")
def _fmt_steps(x, _pos) -> str:
if x == 0:
return "0"
if x >= 1000:
value = x / 1000
return f"{value:.0f}k" if abs(value - round(value)) < 1e-6 else f"{value:.1f}k"
return f"{x:.0f}"
def _xticks_for(x_max: float) -> list[int]:
if x_max >= 99_000:
return [0, 25_000, 50_000, 75_000, 100_000]
return [0, int(round(x_max / 3)), int(round(2 * x_max / 3)), int(round(x_max))]
def _save_pair(fig: plt.Figure, stem: Path) -> None:
fig.savefig(stem.with_suffix(".png"))
fig.savefig(stem.with_suffix(".pdf"))
def plot_task_curve(task: str) -> None:
fig, ax = plt.subplots(figsize=(3.05, 2.35))
fig.patch.set_facecolor("white")
ax.set_facecolor("#fbfcfe")
has_any = False
x_max = 0.0
for algo in ALGOS:
x, mean, low, high = _curve_for_algo_task(algo, task)
if x.size == 0:
continue
color = ALGO_COLOR[algo]
zorder = ZORDER[algo]
mean_s = _smooth(mean)
low_s = _smooth(low)
high_s = _smooth(high)
ax.fill_between(x, low_s, high_s, color=color, alpha=BAND_ALPHA[algo], linewidth=0, zorder=zorder - 0.5)
ax.plot(
x,
mean_s,
color=color,
linewidth=LINEWIDTH[algo],
linestyle=LINESTYLE[algo],
solid_capstyle="round",
solid_joinstyle="round",
dash_capstyle="round",
label=ALGO_LABELS[algo],
zorder=zorder,
)
ax.scatter(
x[-1],
mean_s[-1],
s=22 if algo == "hamworld" else 16,
color=color,
edgecolors="white",
linewidths=0.9,
zorder=zorder + 0.2,
)
x_max = max(x_max, float(x.max()))
has_any = True
if not has_any:
plt.close(fig)
return
ax.set_title(TASK_TITLE[task], color="#0f172a", pad=6, loc="left")
ax.set_xlabel("env steps")
ax.set_ylabel("eval return")
ax.set_xlim(0, x_max)
ax.margins(y=0.08)
ax.spines["left"].set_color("#cbd5e1")
ax.spines["bottom"].set_color("#cbd5e1")
ax.set_xticks(_xticks_for(x_max))
ax.yaxis.set_major_locator(MaxNLocator(nbins=4, integer=True))
ax.xaxis.set_major_formatter(FuncFormatter(_fmt_steps))
ax.grid(True, axis="y", linestyle="-", linewidth=0.6, color="#e5e7eb", zorder=0)
ax.set_axisbelow(True)
ax.tick_params(axis="x", which="major", pad=2)
ax.tick_params(axis="y", which="major", pad=2)
fig.tight_layout(pad=0.3)
_save_pair(fig, OUT_DIR / f"curve_{task}")
plt.close(fig)
def plot_legend_strip() -> None:
fig, ax = plt.subplots(figsize=(7.4, 0.5))
fig.patch.set_facecolor("white")
ax.axis("off")
handles = [
plt.Line2D(
[0],
[0],
color=ALGO_COLOR[algo],
linewidth=2.4 if algo == "hamworld" else 2.0,
linestyle=LINESTYLE[algo],
solid_capstyle="round",
dash_capstyle="round",
marker="o",
markersize=4.2 if algo == "hamworld" else 3.6,
markerfacecolor=ALGO_COLOR[algo],
markeredgewidth=0.0,
label=ALGO_LABELS[algo],
)
for algo in ALGOS
]
ax.legend(
handles=handles,
loc="center",
ncol=5,
frameon=False,
handlelength=2.0,
handletextpad=0.55,
columnspacing=1.45,
borderaxespad=0.0,
)
_save_pair(fig, OUT_DIR / "legend")
plt.close(fig)
def build_assembled_figure() -> None:
curve_order = [
"finger_spin",
"reacher_easy",
"cheetah_run",
"cartpole_swingup",
]
def _crop_curve_image(image: np.ndarray) -> np.ndarray:
height, width = image.shape[:2]
top = int(height * 0.01)
bottom = int(height * 0.93)
left = int(width * 0.015)
right = int(width * 0.99)
return image[top:bottom, left:right]
curve_images = [_crop_curve_image(mpimg.imread(OUT_DIR / f"curve_{task}.png")) for task in curve_order]
fig = plt.figure(figsize=(12.6, 3.52), facecolor="white")
grid = fig.add_gridspec(1, 4, wspace=0.08)
for idx, image in enumerate(curve_images):
ax = fig.add_subplot(grid[0, idx])
ax.imshow(image)
ax.axis("off")
handles = [
plt.Line2D(
[0],
[0],
color=ALGO_COLOR[algo],
linewidth=2.4 if algo == "hamworld" else 2.0,
linestyle=LINESTYLE[algo],
solid_capstyle="round",
dash_capstyle="round",
marker="o",
markersize=4.4 if algo == "hamworld" else 3.8,
markerfacecolor=ALGO_COLOR[algo],
markeredgewidth=0.0,
label=ALGO_LABELS[algo],
)
for algo in ALGOS
]
fig.legend(
handles=handles,
loc="lower center",
ncol=5,
frameon=False,
bbox_to_anchor=(0.5, 0.05),
handlelength=1.9,
handletextpad=0.45,
columnspacing=1.1,
fontsize=15.0,
borderaxespad=0.0,
)
fig.subplots_adjust(top=0.99, bottom=0.045, left=0.02, right=0.98)
fig.savefig(ASSEMBLED_STEM.with_suffix(".png"))
fig.savefig(ASSEMBLED_STEM.with_suffix(".pdf"))
plt.close(fig)
def main() -> int:
for task in TASKS:
plot_task_curve(task)
print(f"[curve] {task}")
plot_legend_strip()
build_assembled_figure()
print(f"[assembled] {ASSEMBLED_STEM.with_suffix('.png')}")
print(f"[done] {OUT_DIR}")
return 0
if __name__ == "__main__":
raise SystemExit(main())