File size: 2,664 Bytes
0838417 | 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 | 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())
|