Download scripts/plot_hamiltonian_slice_flow.py from haodongcui/HaM_World: direct link, hf CLI and curl.
- Browser
- Download file 18.6 kB
-
https://huggingface.co/haodongcui/HaM_World/resolve/main/scripts/plot_hamiltonian_slice_flow.py
- Command line
-
hf download hf://haodongcui/HaM_World/scripts/plot_hamiltonian_slice_flow.py
-
curl -L -o plot_hamiltonian_slice_flow.py https://huggingface.co/haodongcui/HaM_World/resolve/main/scripts/plot_hamiltonian_slice_flow.py
18.6 kB
| #!/usr/bin/env python3 | |
| from __future__ import annotations | |
| import argparse | |
| import copy | |
| import csv | |
| import json | |
| import os | |
| import sys | |
| from dataclasses import asdict, dataclass | |
| from pathlib import Path | |
| REPO_ROOT = Path(__file__).resolve().parents[1] | |
| if str(REPO_ROOT) not in sys.path: | |
| sys.path.insert(0, str(REPO_ROOT)) | |
| os.environ.setdefault("KMP_DUPLICATE_LIB_OK", "TRUE") | |
| MPLCONFIGDIR = REPO_ROOT / ".cache" / "matplotlib" | |
| MPLCONFIGDIR.mkdir(parents=True, exist_ok=True) | |
| os.environ.setdefault("MPLCONFIGDIR", str(MPLCONFIGDIR)) | |
| os.environ.setdefault("MUJOCO_GL", "egl") | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import torch | |
| from common import ensure_repo_on_path | |
| ensure_repo_on_path() | |
| from hamworld.runtime import make_env | |
| from hamworld.world_model import CanonicalDynamicsWorldModel, infer_checkpoint_step | |
| class PairScore: | |
| q_index: int | |
| p_index: int | |
| kind: str | |
| score: float | |
| coverage: float | |
| swirl: float | |
| speed: float | |
| occupancy: float | |
| q_span: float | |
| p_span: float | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Render 2D Hamiltonian slice flow maps from a HaM-World checkpoint.") | |
| parser.add_argument("--checkpoint", required=True, help="Path to checkpoint_*.pt.") | |
| parser.add_argument("--trace-path", required=True, help="Path to dynamics trace .npz file.") | |
| parser.add_argument("--output-dir", required=True, help="Directory for figures and pair summaries.") | |
| parser.add_argument("--device", default="cpu", help="Torch device, e.g. cpu or cuda:0.") | |
| parser.add_argument("--grid-size", type=int, default=81, help="Grid size for slice evaluation.") | |
| parser.add_argument("--same-k", type=int, default=4, help="How many same-index (q_i, p_i) pairs to render.") | |
| parser.add_argument("--cross-k", type=int, default=2, help="How many cross-index (q_i, p_j) pairs to render.") | |
| parser.add_argument("--trace-episodes", type=int, default=3, help="How many trace episodes to overlay.") | |
| parser.add_argument("--trace-quantile-lo", type=float, default=2.0, help="Low percentile for slice bounds.") | |
| parser.add_argument("--trace-quantile-hi", type=float, default=98.0, help="High percentile for slice bounds.") | |
| return parser.parse_args() | |
| def _load_checkpoint_world_model(checkpoint_path: Path, device: torch.device) -> tuple[CanonicalDynamicsWorldModel, dict]: | |
| payload = torch.load(checkpoint_path, map_location="cpu", weights_only=False) | |
| config = copy.deepcopy(payload["config"]) | |
| env, spec = make_env(config["task"], int(config["experiment"]["seed"]) + 1000) | |
| close_fn = getattr(env, "close", None) | |
| if callable(close_fn): | |
| close_fn() | |
| world_model = CanonicalDynamicsWorldModel( | |
| config=config, | |
| obs_dim=int(spec.observation_shape[0]), | |
| action_dim=int(spec.action_shape[0]), | |
| ).to(device) | |
| agent_state = payload.get("agent_state") | |
| if isinstance(agent_state, dict) and "model" in agent_state: | |
| state_dict = agent_state["model"] | |
| else: | |
| state_dict = payload.get("model") | |
| if state_dict is None: | |
| raise ValueError(f"Checkpoint does not contain a world model state: {checkpoint_path}") | |
| world_model.load_state_dict(state_dict) | |
| world_model.set_schedule_step(infer_checkpoint_step(payload, checkpoint_path)) | |
| world_model.eval() | |
| return world_model, config | |
| def _load_trace(trace_path: Path) -> dict[str, np.ndarray]: | |
| trace = np.load(trace_path, allow_pickle=True) | |
| return {key: trace[key] for key in trace.files} | |
| def _flatten_valid_vectors(values: np.ndarray, valid_mask: np.ndarray) -> np.ndarray: | |
| mask = np.asarray(valid_mask, dtype=bool) | |
| return np.asarray(values[mask], dtype=np.float32) | |
| def _robust_span(values: np.ndarray, lo: float = 5.0, hi: float = 95.0) -> float: | |
| low, high = np.percentile(values, [lo, hi]) | |
| return float(high - low) | |
| def _histogram_entropy(x: np.ndarray, y: np.ndarray, bins: int = 20) -> float: | |
| hist, _, _ = np.histogram2d(x, y, bins=bins) | |
| probs = hist.ravel() | |
| total = probs.sum() | |
| if total <= 0.0: | |
| return 0.0 | |
| probs = probs[probs > 0.0] / total | |
| entropy = -(probs * np.log(probs)).sum() | |
| max_entropy = np.log(float(bins * bins)) | |
| return float(entropy / max(1e-6, max_entropy)) | |
| def _normalize(values: list[float]) -> list[float]: | |
| maximum = max(values) if values else 1.0 | |
| if maximum <= 1e-6: | |
| return [0.0 for _ in values] | |
| return [float(value / maximum) for value in values] | |
| def score_pairs(trace: dict[str, np.ndarray]) -> list[PairScore]: | |
| valid_mask = np.asarray(trace["valid_mask"], dtype=bool) | |
| q = _flatten_valid_vectors(trace["q"], valid_mask) | |
| p = _flatten_valid_vectors(trace["p"], valid_mask) | |
| dH_dq = _flatten_valid_vectors(trace["dH_dq"], valid_mask) | |
| dH_dp = _flatten_valid_vectors(trace["dH_dp"], valid_mask) | |
| q_dim = int(np.asarray(trace["q_dim"]).item()) | |
| p_dim = int(np.asarray(trace["p_dim"]).item()) | |
| raw_rows: list[dict[str, float | int | str]] = [] | |
| for q_index in range(q_dim): | |
| for p_index in range(p_dim): | |
| qv = q[:, q_index] | |
| pv = p[:, p_index] | |
| dq = dH_dp[:, q_index] | |
| dp = -dH_dq[:, p_index] | |
| q_center = float(np.median(qv)) | |
| p_center = float(np.median(pv)) | |
| radial_norm = np.hypot(qv - q_center, pv - p_center) | |
| flow_norm = np.hypot(dq, dp) | |
| tangential = np.abs((qv - q_center) * dp - (pv - p_center) * dq) | |
| swirl = float(np.mean(tangential / np.maximum(1e-6, radial_norm * flow_norm))) | |
| q_span = _robust_span(qv) | |
| p_span = _robust_span(pv) | |
| coverage = float(q_span * p_span) | |
| speed = float(np.mean(flow_norm)) | |
| occupancy = _histogram_entropy(qv, pv) | |
| raw_rows.append( | |
| { | |
| "q_index": q_index, | |
| "p_index": p_index, | |
| "kind": "same" if q_index == p_index else "cross", | |
| "coverage": coverage, | |
| "swirl": swirl, | |
| "speed": speed, | |
| "occupancy": occupancy, | |
| "q_span": q_span, | |
| "p_span": p_span, | |
| } | |
| ) | |
| coverage_norm = _normalize([float(row["coverage"]) for row in raw_rows]) | |
| swirl_norm = _normalize([float(row["swirl"]) for row in raw_rows]) | |
| speed_norm = _normalize([float(row["speed"]) for row in raw_rows]) | |
| occupancy_norm = _normalize([float(row["occupancy"]) for row in raw_rows]) | |
| scored: list[PairScore] = [] | |
| for index, row in enumerate(raw_rows): | |
| score = ( | |
| 0.35 * coverage_norm[index] | |
| + 0.30 * swirl_norm[index] | |
| + 0.20 * speed_norm[index] | |
| + 0.15 * occupancy_norm[index] | |
| ) | |
| scored.append( | |
| PairScore( | |
| q_index=int(row["q_index"]), | |
| p_index=int(row["p_index"]), | |
| kind=str(row["kind"]), | |
| score=float(score), | |
| coverage=float(row["coverage"]), | |
| swirl=float(row["swirl"]), | |
| speed=float(row["speed"]), | |
| occupancy=float(row["occupancy"]), | |
| q_span=float(row["q_span"]), | |
| p_span=float(row["p_span"]), | |
| ) | |
| ) | |
| return sorted(scored, key=lambda item: item.score, reverse=True) | |
| def select_pairs(scores: list[PairScore], same_k: int, cross_k: int) -> list[PairScore]: | |
| same_pairs = [row for row in scores if row.kind == "same"][: max(0, int(same_k))] | |
| cross_pairs = [row for row in scores if row.kind == "cross"][: max(0, int(cross_k))] | |
| return same_pairs + cross_pairs | |
| def _slice_bounds(values: np.ndarray, quantile_lo: float, quantile_hi: float) -> tuple[float, float]: | |
| low, high = np.percentile(values, [quantile_lo, quantile_hi]) | |
| if not np.isfinite(low) or not np.isfinite(high) or abs(high - low) < 1e-5: | |
| center = float(np.median(values)) | |
| return center - 1.0, center + 1.0 | |
| margin = 0.18 * float(high - low) | |
| return float(low - margin), float(high + margin) | |
| def evaluate_slice( | |
| world_model: CanonicalDynamicsWorldModel, | |
| q_reference: np.ndarray, | |
| p_reference: np.ndarray, | |
| pair: PairScore, | |
| q_values: np.ndarray, | |
| p_values: np.ndarray, | |
| grid_size: int, | |
| quantile_lo: float, | |
| quantile_hi: float, | |
| device: torch.device, | |
| ) -> dict[str, np.ndarray]: | |
| x_min, x_max = _slice_bounds(q_values, quantile_lo, quantile_hi) | |
| y_min, y_max = _slice_bounds(p_values, quantile_lo, quantile_hi) | |
| q_axis = np.linspace(x_min, x_max, grid_size, dtype=np.float64) | |
| p_axis = np.linspace(y_min, y_max, grid_size, dtype=np.float64) | |
| q_grid, p_grid = np.meshgrid(q_axis, p_axis, indexing="xy") | |
| q_eval = np.repeat(q_reference[None, :], grid_size * grid_size, axis=0) | |
| p_eval = np.repeat(p_reference[None, :], grid_size * grid_size, axis=0) | |
| q_eval[:, pair.q_index] = q_grid.reshape(-1) | |
| p_eval[:, pair.p_index] = p_grid.reshape(-1) | |
| q_tensor = torch.as_tensor(q_eval, dtype=torch.float32, device=device).requires_grad_(True) | |
| p_tensor = torch.as_tensor(p_eval, dtype=torch.float32, device=device).requires_grad_(True) | |
| energy = world_model.energy_head(q_tensor, p_tensor) | |
| dH_dq, dH_dp = torch.autograd.grad(energy.sum(), (q_tensor, p_tensor), create_graph=False, retain_graph=False) | |
| H = energy.detach().cpu().numpy().reshape(grid_size, grid_size) | |
| U = dH_dp.detach().cpu().numpy()[:, pair.q_index].reshape(grid_size, grid_size) | |
| V = (-dH_dq.detach().cpu().numpy()[:, pair.p_index]).reshape(grid_size, grid_size) | |
| speed = np.hypot(U, V) | |
| return { | |
| "q_axis": q_axis, | |
| "p_axis": p_axis, | |
| "H": H, | |
| "U": U, | |
| "V": V, | |
| "speed": speed, | |
| "q_grid": q_grid, | |
| "p_grid": p_grid, | |
| } | |
| def _overlay_trace(axis: plt.Axes, trace: dict[str, np.ndarray], pair: PairScore, max_episodes: int) -> None: | |
| valid_mask = np.asarray(trace["valid_mask"], dtype=bool) | |
| q = np.asarray(trace["q"], dtype=np.float32) | |
| p = np.asarray(trace["p"], dtype=np.float32) | |
| episode_count = min(int(max_episodes), int(q.shape[0])) | |
| colors = ["#7dd3fc", "#38bdf8", "#0ea5e9"] | |
| for episode_index in range(episode_count): | |
| mask = valid_mask[episode_index] | |
| if not np.any(mask): | |
| continue | |
| x = q[episode_index, mask, pair.q_index] | |
| y = p[episode_index, mask, pair.p_index] | |
| axis.plot(x, y, color=colors[episode_index % len(colors)], linewidth=1.1, alpha=0.9) | |
| axis.scatter(x[0], y[0], color=colors[episode_index % len(colors)], s=14, marker="o", alpha=0.95) | |
| def draw_pair_figure( | |
| output_path: Path, | |
| slice_eval: dict[str, np.ndarray], | |
| trace: dict[str, np.ndarray], | |
| pair: PairScore, | |
| q_reference: np.ndarray, | |
| p_reference: np.ndarray, | |
| checkpoint_name: str, | |
| ) -> None: | |
| q_axis = slice_eval["q_axis"] | |
| p_axis = slice_eval["p_axis"] | |
| H = slice_eval["H"] | |
| U = slice_eval["U"] | |
| V = slice_eval["V"] | |
| speed = slice_eval["speed"] | |
| fig, axis = plt.subplots(figsize=(6.4, 5.2), dpi=220) | |
| heat = axis.contourf(q_axis, p_axis, H, levels=24, cmap="magma") | |
| axis.contour(q_axis, p_axis, H, levels=10, colors="white", linewidths=0.45, alpha=0.55) | |
| axis.streamplot( | |
| q_axis, | |
| p_axis, | |
| U, | |
| V, | |
| density=1.1, | |
| color=speed, | |
| cmap="viridis", | |
| linewidth=0.9, | |
| arrowsize=0.9, | |
| ) | |
| _overlay_trace(axis, trace, pair, max_episodes=3) | |
| axis.scatter( | |
| [q_reference[pair.q_index]], | |
| [p_reference[pair.p_index]], | |
| s=52, | |
| marker="x", | |
| linewidths=1.5, | |
| color="white", | |
| zorder=5, | |
| ) | |
| axis.set_xlabel(f"q[{pair.q_index}]") | |
| axis.set_ylabel(f"p[{pair.p_index}]") | |
| axis.set_title( | |
| f"{pair.kind} pair q[{pair.q_index}] / p[{pair.p_index}] score={pair.score:.3f}\n{checkpoint_name}", | |
| loc="left", | |
| ) | |
| cbar = fig.colorbar(heat, ax=axis, fraction=0.046, pad=0.04) | |
| cbar.set_label("H(q, p)") | |
| axis.spines["top"].set_visible(False) | |
| axis.spines["right"].set_visible(False) | |
| fig.tight_layout() | |
| fig.savefig(output_path, bbox_inches="tight") | |
| plt.close(fig) | |
| def draw_overview( | |
| output_path: Path, | |
| slices: list[tuple[PairScore, dict[str, np.ndarray]]], | |
| trace: dict[str, np.ndarray], | |
| q_reference: np.ndarray, | |
| p_reference: np.ndarray, | |
| checkpoint_name: str, | |
| ) -> None: | |
| columns = 3 | |
| rows = max(1, int(np.ceil(len(slices) / columns))) | |
| fig, axes = plt.subplots(rows, columns, figsize=(5.3 * columns, 4.5 * rows), dpi=220) | |
| axes_array = np.atleast_1d(axes).reshape(rows, columns) | |
| for axis, (pair, slice_eval) in zip(axes_array.flatten(), slices): | |
| heat = axis.contourf(slice_eval["q_axis"], slice_eval["p_axis"], slice_eval["H"], levels=20, cmap="magma") | |
| axis.contour(slice_eval["q_axis"], slice_eval["p_axis"], slice_eval["H"], levels=8, colors="white", linewidths=0.35, alpha=0.5) | |
| axis.streamplot( | |
| slice_eval["q_axis"], | |
| slice_eval["p_axis"], | |
| slice_eval["U"], | |
| slice_eval["V"], | |
| density=1.0, | |
| color=slice_eval["speed"], | |
| cmap="viridis", | |
| linewidth=0.8, | |
| arrowsize=0.8, | |
| ) | |
| _overlay_trace(axis, trace, pair, max_episodes=2) | |
| axis.scatter( | |
| [q_reference[pair.q_index]], | |
| [p_reference[pair.p_index]], | |
| s=32, | |
| marker="x", | |
| linewidths=1.2, | |
| color="white", | |
| zorder=5, | |
| ) | |
| axis.set_xlabel(f"q[{pair.q_index}]") | |
| axis.set_ylabel(f"p[{pair.p_index}]") | |
| axis.set_title(f"{pair.kind} q[{pair.q_index}] / p[{pair.p_index}] s={pair.score:.3f}", loc="left", fontsize=10) | |
| axis.spines["top"].set_visible(False) | |
| axis.spines["right"].set_visible(False) | |
| for axis in axes_array.flatten()[len(slices) :]: | |
| axis.axis("off") | |
| fig.suptitle(f"Hamiltonian slice flow candidates\n{checkpoint_name}", x=0.055, y=0.995, ha="left", fontsize=14) | |
| fig.tight_layout(rect=[0, 0, 1, 0.97]) | |
| fig.savefig(output_path, bbox_inches="tight") | |
| plt.close(fig) | |
| def write_scores_csv(output_path: Path, scores: list[PairScore]) -> None: | |
| fieldnames = list(asdict(scores[0]).keys()) if scores else [] | |
| with output_path.open("w", encoding="utf-8", newline="") as handle: | |
| writer = csv.DictWriter(handle, fieldnames=fieldnames) | |
| writer.writeheader() | |
| for row in scores: | |
| writer.writerow(asdict(row)) | |
| def main() -> int: | |
| args = parse_args() | |
| checkpoint_path = Path(args.checkpoint).expanduser().resolve() | |
| trace_path = Path(args.trace_path).expanduser().resolve() | |
| output_dir = Path(args.output_dir).expanduser().resolve() | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| device = torch.device(args.device) | |
| world_model, config = _load_checkpoint_world_model(checkpoint_path, device) | |
| trace = _load_trace(trace_path) | |
| valid_mask = np.asarray(trace["valid_mask"], dtype=bool) | |
| q_flat = _flatten_valid_vectors(trace["q"], valid_mask) | |
| p_flat = _flatten_valid_vectors(trace["p"], valid_mask) | |
| q_reference = np.median(q_flat, axis=0).astype(np.float32) | |
| p_reference = np.median(p_flat, axis=0).astype(np.float32) | |
| all_scores = score_pairs(trace) | |
| selected_pairs = select_pairs(all_scores, same_k=args.same_k, cross_k=args.cross_k) | |
| if not selected_pairs: | |
| raise RuntimeError("No q/p pairs were selected for rendering.") | |
| render_dir = output_dir / "pairs" | |
| render_dir.mkdir(parents=True, exist_ok=True) | |
| slice_bundle: list[tuple[PairScore, dict[str, np.ndarray]]] = [] | |
| for pair in selected_pairs: | |
| q_values = q_flat[:, pair.q_index] | |
| p_values = p_flat[:, pair.p_index] | |
| slice_eval = evaluate_slice( | |
| world_model=world_model, | |
| q_reference=q_reference, | |
| p_reference=p_reference, | |
| pair=pair, | |
| q_values=q_values, | |
| p_values=p_values, | |
| grid_size=int(args.grid_size), | |
| quantile_lo=float(args.trace_quantile_lo), | |
| quantile_hi=float(args.trace_quantile_hi), | |
| device=device, | |
| ) | |
| slice_bundle.append((pair, slice_eval)) | |
| filename = f"hamiltonian_slice_{pair.kind}_q{pair.q_index}_p{pair.p_index}.png" | |
| draw_pair_figure( | |
| output_path=render_dir / filename, | |
| slice_eval=slice_eval, | |
| trace=trace, | |
| pair=pair, | |
| q_reference=q_reference, | |
| p_reference=p_reference, | |
| checkpoint_name=checkpoint_path.parent.parent.parent.name, | |
| ) | |
| write_scores_csv(output_dir / "pair_scores.csv", all_scores) | |
| with (output_dir / "selected_pairs.json").open("w", encoding="utf-8") as handle: | |
| json.dump([asdict(pair) for pair in selected_pairs], handle, indent=2) | |
| draw_overview( | |
| output_path=output_dir / "hamiltonian_slice_overview.png", | |
| slices=slice_bundle, | |
| trace=trace, | |
| q_reference=q_reference, | |
| p_reference=p_reference, | |
| checkpoint_name=checkpoint_path.parent.parent.parent.name, | |
| ) | |
| summary = { | |
| "checkpoint": str(checkpoint_path), | |
| "trace_path": str(trace_path), | |
| "task": str(np.asarray(trace["task"]).item()), | |
| "seed": int(np.asarray(trace["seed"]).item()), | |
| "mode": str(np.asarray(trace["mode"]).item()), | |
| "q_reference": q_reference.tolist(), | |
| "p_reference": p_reference.tolist(), | |
| "selected_pairs": [asdict(pair) for pair in selected_pairs], | |
| } | |
| with (output_dir / "summary.json").open("w", encoding="utf-8") as handle: | |
| json.dump(summary, handle, indent=2) | |
| print(output_dir / "hamiltonian_slice_overview.png") | |
| for pair in selected_pairs: | |
| print(render_dir / f"hamiltonian_slice_{pair.kind}_q{pair.q_index}_p{pair.p_index}.png") | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |