HaM_World / scripts /replot_pqc.py
haodongcui's picture
first commit (part 4)
0838417 verified
Raw History Blame Contribute Delete
7.67 kB
from __future__ import annotations
import argparse
from pathlib import Path
import sys
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.lines import Line2D
import numpy as np
from scipy.stats import gaussian_kde
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_MECHANISM_FIG_ROOT, RESULTS_MECHANISM_TRACE_EP10_ROOT, RESULTS_MECHANISM_TRACE_QUICK_ROOT
MAX_POINTS_PER_COMPONENT = 2000
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Replot 2-task P/Q/C UMAP figure from paper mechanism traces.")
parser.add_argument("--bundle", choices=["quick_seed7", "ep10_seed7"], default="ep10_seed7")
return parser.parse_args()
plt.rcParams.update(
{
"font.family": "DejaVu Sans",
"font.size": 10,
"axes.titlesize": 13.5,
"axes.titleweight": "semibold",
"axes.labelsize": 11.5,
"axes.labelcolor": "#334155",
"axes.edgecolor": "#cbd5e1",
"axes.linewidth": 1.0,
"xtick.color": "#64748b",
"ytick.color": "#64748b",
"legend.fontsize": 11,
"figure.facecolor": "white",
"axes.facecolor": "#fbfcfe",
"savefig.facecolor": "white",
"savefig.bbox": "tight",
"savefig.pad_inches": 0.05,
}
)
def _load_trace(task: str, bundle: str) -> dict[str, np.ndarray]:
traces_dir = (RESULTS_MECHANISM_TRACE_QUICK_ROOT if bundle == "quick_seed7" else RESULTS_MECHANISM_TRACE_EP10_ROOT) / "traces"
payload = np.load(traces_dir / f"{task}_seed7_teacher_forced.npz", allow_pickle=True)
return {
"valid_mask": np.asarray(payload["valid_mask"]),
"q": np.asarray(payload["q"]),
"p": np.asarray(payload["p"]),
"c": np.asarray(payload["c"]),
}
def _project(array: np.ndarray) -> np.ndarray:
try:
from umap import UMAP
return UMAP(n_neighbors=35, min_dist=0.12, random_state=0, n_components=2).fit_transform(array)
except Exception:
centered = array - np.mean(array, axis=0, keepdims=True)
_, _, vt = np.linalg.svd(centered, full_matrices=False)
basis = vt[:2].T
return centered @ basis
def _pad(values: np.ndarray, target_dim: int) -> np.ndarray:
if values.shape[1] >= target_dim:
return values[:, :target_dim]
return np.concatenate([values, np.zeros((values.shape[0], target_dim - values.shape[1]))], axis=1)
def _standardize(values: np.ndarray) -> np.ndarray:
mean = np.mean(values, axis=0, keepdims=True)
std = np.std(values, axis=0, keepdims=True)
return (values - mean) / np.clip(std, 1e-6, None)
def _prepare_component(values: np.ndarray, target_dim: int) -> np.ndarray:
standardized = _standardize(values)
padded = _pad(standardized, target_dim)
return padded / np.sqrt(float(values.shape[1]))
def _subsample_even(values: np.ndarray, count: int) -> np.ndarray:
if values.shape[0] <= count:
return values
indices = np.linspace(0, values.shape[0] - 1, count, dtype=np.int32)
return values[indices]
def _draw_density_contours(ax: plt.Axes, points: np.ndarray, color: str) -> None:
if points.shape[0] < 32:
return
x = points[:, 0]
y = points[:, 1]
try:
kde = gaussian_kde(np.vstack([x, y]))
except Exception:
return
x_lo, x_hi = np.quantile(x, [0.01, 0.99])
y_lo, y_hi = np.quantile(y, [0.01, 0.99])
x_pad = max(1e-3, 0.12 * (x_hi - x_lo))
y_pad = max(1e-3, 0.12 * (y_hi - y_lo))
xx, yy = np.meshgrid(
np.linspace(x_lo - x_pad, x_hi + x_pad, 160),
np.linspace(y_lo - y_pad, y_hi + y_pad, 160),
)
zz = kde(np.vstack([xx.ravel(), yy.ravel()])).reshape(xx.shape)
levels = np.quantile(zz[zz > 0], [0.72, 0.86, 0.95])
ax.contour(xx, yy, zz, levels=levels, colors=[color], linewidths=[0.9, 1.15, 1.45], alpha=0.7, zorder=3)
def _set_view_limits(ax: plt.Axes, embedding: np.ndarray) -> None:
x_lo, x_hi = np.quantile(embedding[:, 0], [0.01, 0.99])
y_lo, y_hi = np.quantile(embedding[:, 1], [0.01, 0.99])
x_pad = max(1e-3, 0.10 * (x_hi - x_lo))
y_pad = max(1e-3, 0.10 * (y_hi - y_lo))
ax.set_xlim(x_lo - x_pad, x_hi + x_pad)
ax.set_ylim(y_lo - y_pad, y_hi + y_pad)
def _save_pair(fig: plt.Figure, stem: Path) -> None:
stem.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(stem.with_suffix(".pdf"))
fig.savefig(stem.with_suffix(".png"))
def main() -> int:
args = parse_args()
fig, axes = plt.subplots(1, 2, figsize=(6.6, 2.9), sharex=False, sharey=False, dpi=220)
colors = {"q": "#2b6cb0", "p": "#dd6b20", "c": "#2f855a"}
titles = {"finger_spin": "Finger Spin", "cheetah_run": "Cheetah Run"}
legend_handles = [
Line2D([0], [0], marker="o", linestyle="", markersize=6.5, markerfacecolor=colors[key], markeredgecolor="white", markeredgewidth=0.6, label=fr"${key}$")
for key in ("q", "p", "c")
]
for ax, task in zip(axes, ["finger_spin", "cheetah_run"]):
trace = _load_trace(task, args.bundle)
mask = trace["valid_mask"].reshape(-1)
q = trace["q"].reshape(-1, trace["q"].shape[-1])[mask]
p = trace["p"].reshape(-1, trace["p"].shape[-1])[mask]
c = trace["c"].reshape(-1, trace["c"].shape[-1])[mask]
count = min(MAX_POINTS_PER_COMPONENT, q.shape[0], p.shape[0], c.shape[0])
q = _subsample_even(q, count)
p = _subsample_even(p, count)
c = _subsample_even(c, count)
target_dim = max(q.shape[1], p.shape[1], c.shape[1])
stacked = np.concatenate(
[
_prepare_component(q, target_dim),
_prepare_component(p, target_dim),
_prepare_component(c, target_dim),
],
axis=0,
)
labels = (["q"] * count) + (["p"] * count) + (["c"] * count)
embedding = _project(stacked)
for key in ("q", "p", "c"):
selected = np.array([label == key for label in labels])
points = embedding[selected]
ax.scatter(
points[:, 0],
points[:, 1],
s=4,
alpha=0.20,
color=colors[key],
label=fr"${key}$",
linewidths=0.0,
edgecolors="none",
rasterized=True,
zorder=2,
)
_draw_density_contours(ax, points, colors[key])
_set_view_limits(ax, embedding)
ax.set_title(titles[task], loc="left", pad=8, color="#0f172a")
ax.set_xticks([])
ax.set_yticks([])
ax.set_xlabel("Embedding 1", labelpad=8)
ax.set_ylabel("Embedding 2", labelpad=8)
for spine in ("left", "bottom"):
ax.spines[spine].set_color("#cbd5e1")
ax.spines["top"].set_visible(False)
ax.spines["right"].set_visible(False)
fig.legend(legend_handles, ["q", "p", "c"], loc="upper center", ncol=3, frameon=False, bbox_to_anchor=(0.5, 1.02), handletextpad=0.35, columnspacing=1.2)
fig.tight_layout(pad=0.7, w_pad=1.2, rect=(0, 0, 1, 0.92))
out_stem = RESULTS_MECHANISM_FIG_ROOT / "pqc"
_save_pair(fig, out_stem)
plt.close(fig)
print(out_stem.with_suffix(".pdf"))
print(out_stem.with_suffix(".png"))
return 0
if __name__ == "__main__":
raise SystemExit(main())