Download scripts/replot_main_curves.py from haodongcui/HaM_World: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/haodongcui/HaM_World/resolve/main/scripts/replot_main_curves.py
- Command line
-
hf download hf://haodongcui/HaM_World/scripts/replot_main_curves.py
-
curl -L -o replot_main_curves.py https://huggingface.co/haodongcui/HaM_World/resolve/main/scripts/replot_main_curves.py
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()) | |