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())