Download scripts/replot_phase_portrait.py from haodongcui/HaM_World: direct link, hf CLI and curl.
- Browser
- Download file 2.66 kB
-
https://huggingface.co/haodongcui/HaM_World/resolve/main/scripts/replot_phase_portrait.py
- Command line
-
hf download hf://haodongcui/HaM_World/scripts/replot_phase_portrait.py
-
curl -L -o replot_phase_portrait.py https://huggingface.co/haodongcui/HaM_World/resolve/main/scripts/replot_phase_portrait.py
2.66 kB
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| import sys | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| 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 ANALYSIS_MECHANISM_FREERUN_MANIFEST, RESULTS_MECHANISM_FIG_ROOT, read_csv_rows, resolve_repo_path | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Replot simple q/p phase portrait from paper free-run traces.") | |
| parser.add_argument( | |
| "--manifest", | |
| type=Path, | |
| default=ANALYSIS_MECHANISM_FREERUN_MANIFEST, | |
| ) | |
| return parser.parse_args() | |
| def _load_qp(path: Path) -> tuple[np.ndarray, np.ndarray]: | |
| payload = np.load(path, allow_pickle=True) | |
| q = np.asarray(payload["Q"], dtype=np.float32) | |
| p = np.asarray(payload["P"], dtype=np.float32) | |
| valid = np.asarray(payload["valid_mask"], dtype=bool) | |
| episode = 0 | |
| steps = np.where(valid[episode])[0] | |
| if steps.size == 0: | |
| return np.asarray([]), np.asarray([]) | |
| q_series = q[episode, steps, 0] | |
| p_series = p[episode, steps, 0] | |
| return q_series, p_series | |
| def main() -> int: | |
| args = parse_args() | |
| rows = [row for row in read_csv_rows(args.manifest) if row["paper_group"] == "phase_portrait"] | |
| rows.sort(key=lambda row: row["task"]) | |
| fig, axes = plt.subplots(1, 2, figsize=(5.0, 2.2), dpi=250) | |
| titles = {"finger_spin": "Finger Spin", "cheetah_run": "Cheetah Run"} | |
| for ax, row in zip(axes, rows): | |
| q_series, p_series = _load_qp(resolve_repo_path(row["trace_path"])) | |
| if q_series.size == 0: | |
| continue | |
| color = "#1f4ea8" if row["task"] == "cheetah_run" else "#c3324b" | |
| ax.plot(q_series, p_series, color=color, linewidth=1.4, alpha=0.95) | |
| ax.scatter(q_series[0], p_series[0], color="#111827", s=10, zorder=3) | |
| ax.set_title(titles.get(row["task"], row["task"])) | |
| ax.set_xlabel("q[0]") | |
| ax.set_ylabel("p[0]") | |
| ax.grid(True, linestyle=":", linewidth=0.5, alpha=0.45) | |
| ax.spines["top"].set_visible(False) | |
| ax.spines["right"].set_visible(False) | |
| fig.tight_layout(pad=0.5) | |
| out_stem = RESULTS_MECHANISM_FIG_ROOT / "phase_qp_no_action_combined" | |
| fig.savefig(out_stem.with_suffix(".pdf")) | |
| fig.savefig(out_stem.with_suffix(".png")) | |
| plt.close(fig) | |
| print(out_stem.with_suffix(".pdf")) | |
| print(out_stem.with_suffix(".png")) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |