From 39dc799083090145a9f444e0fc601c88338ccf4c Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:13:15 -0700 Subject: [PATCH 01/11] Add gamma sweep rollout data utilities --- cfm_mppi/visualization/gamma_sweep_data.py | 149 +++++++++++++++++++++ 1 file changed, 149 insertions(+) create mode 100644 cfm_mppi/visualization/gamma_sweep_data.py diff --git a/cfm_mppi/visualization/gamma_sweep_data.py b/cfm_mppi/visualization/gamma_sweep_data.py new file mode 100644 index 0000000..c53478c --- /dev/null +++ b/cfm_mppi/visualization/gamma_sweep_data.py @@ -0,0 +1,149 @@ +from __future__ import annotations + +import argparse +import time +from types import SimpleNamespace +from typing import Any, Dict, List, Optional, Sequence + +import numpy as np +import torch + +from cfm_mppi.evaluation.eval_benchmark import DEFAULTS, BenchmarkPolicies, _dynamics_step, _make_episode, _set_seed +from cfm_mppi.evaluation.metrics import compute_episode_metrics +from cfm_mppi.safegpc_adapter import SafeMPPIAdapter + + +def build_gamma_grid(values: Optional[Sequence[float]] = None, count: int = 21) -> List[float]: + raw = [float(v) for v in values] if values else np.linspace(0.0, 1.0, max(2, int(count))).tolist() + grid = sorted({round(float(np.clip(v, 0.0, 1.0)), 10) for v in raw}) + if 0.0 not in grid: + grid.insert(0, 0.0) + if 1.0 not in grid: + grid.append(1.0) + return grid + + +def jsonable(x: Any) -> Any: + if isinstance(x, np.generic): + return x.item() + if isinstance(x, np.ndarray): + return x.tolist() + if isinstance(x, torch.Tensor): + return x.detach().cpu().tolist() + if isinstance(x, dict): + return {k: jsonable(v) for k, v in x.items()} + if isinstance(x, (list, tuple)): + return [jsonable(v) for v in x] + if isinstance(x, (str, int, float, bool)) or x is None: + return x + return str(x) + + +def policy_args(args: argparse.Namespace) -> SimpleNamespace: + return SimpleNamespace( + mizuta_checkpoint=args.mizuta_checkpoint, + safe_cfm_checkpoint=args.safe_cfm_checkpoint, + drifting_checkpoint=args.drifting_checkpoint, + smoke=args.smoke, + seed=args.seed, + ) + + +def _finish(args: argparse.Namespace, method: str, episode: int, gamma: Optional[float], states, controls, obstacles, goal, times, infos): + states = np.asarray(states, dtype=np.float32) + controls = np.asarray(controls, dtype=np.float32) if controls else np.zeros((0, 2), dtype=np.float32) + min_h = None + violations = 0 + for info in infos: + if info.get("min_barrier_h") is not None: + h = float(info["min_barrier_h"]) + min_h = h if min_h is None else min(min_h, h) + violations += int(info.get("num_barrier_violations", 0) or 0) + rec = compute_episode_metrics( + states, + controls, + obstacles, + goal, + safety_margin=args.safety_margin, + success_threshold=args.success_threshold, + planning_times=times, + min_barrier_h=min_h, + num_barrier_violations=violations, + ) + calls = [int(i.get("model_calls_per_step", 0) or 0) for i in infos] + nfes = [int(i.get("nfe", 0) or 0) for i in infos] + rec.update( + method=method, + episode=int(episode), + seed=int(args.seed), + dataset=args.dataset, + dynamics=args.dynamics, + gamma=gamma, + safety_margin=float(args.safety_margin), + safety_guarantee_scope="linear_system_theorem_relevant" if args.dynamics == "doubleintegrator" else "empirical_only_unicycle", + model_calls_per_step=float(np.mean(calls)) if calls else 0.0, + nfe=float(np.mean(nfes)) if nfes else 0.0, + checkpoint_path=args.mizuta_checkpoint if method == "mizuta_cfm_mppi" else None, + states=states.astype(float).tolist(), + controls=controls.astype(float).tolist(), + obstacles=obstacles.astype(float).tolist(), + goal=goal.astype(float).tolist(), + planning_times=[float(t) for t in times], + step_infos=jsonable(infos), + ) + if method == "safemppi_gamma": + rec["safemppi_num_samples"] = int(args.safemppi_num_samples) + rec["safemppi_horizon"] = int(args.safemppi_horizon) + return jsonable(rec) + + +def rollout_mizuta(args: argparse.Namespace, episode: int) -> Dict[str, Any]: + policy = BenchmarkPolicies(policy_args(args), torch.device(args.device)) + state, goal, obstacles = _make_episode(args.seed + episode, args.dynamics, args.dataset) + states, controls, times, infos = [state.copy()], [], [], [] + for step in range(args.horizon): + action, info = policy.action("mizuta_cfm_mppi", state, goal, obstacles, controls, args.dynamics, 0.5, args.horizon) + times.append(float(info.get("planning_wall_time", 0.0))) + infos.append({"step": step, **jsonable(info)}) + controls.append(action.copy()) + state = _dynamics_step(state, action, args.dynamics, args.dt) + states.append(state.copy()) + if np.linalg.norm(state[:2] - goal[:2]) <= args.success_threshold: + break + return _finish(args, "mizuta_cfm_mppi", episode, None, states, controls, obstacles, goal, times, infos) + + +def rollout_safemppi(args: argparse.Namespace, episode: int, gamma: float) -> Dict[str, Any]: + device = torch.device(args.device) + state, goal, obstacles = _make_episode(args.seed + episode, args.dynamics, args.dataset) + planner = SafeMPPIAdapter( + horizon=min(args.horizon, args.safemppi_horizon), + dt=args.dt, + num_samples=args.safemppi_num_samples, + gamma=gamma, + dynamics_type=args.dynamics, + noise_sigma=args.safemppi_noise_sigma, + temperature=args.safemppi_temperature, + u_min=tuple(args.u_min), + u_max=tuple(args.u_max), + check_first_control_only=args.check_first_control_only, + ) + states, controls, times, infos = [state.copy()], [], [], [] + for step in range(args.horizon): + t0 = time.perf_counter() + action_t, info = planner.plan( + torch.tensor(state, dtype=torch.float32, device=device), + torch.tensor(goal, dtype=torch.float32, device=device), + torch.tensor(obstacles, dtype=torch.float32, device=device), + gamma=gamma, + seed=args.seed + 100000 * episode + step, + ) + action = action_t.detach().cpu().numpy().astype(np.float32) + times.append(float(info.get("solve_time", time.perf_counter() - t0))) + infos.append({"step": step, **jsonable(info)}) + controls.append(action.copy()) + state = _dynamics_step(state, action, args.dynamics, args.dt) + states.append(state.copy()) + if np.linalg.norm(state[:2] - goal[:2]) <= args.success_threshold: + break + return _finish(args, "safemppi_gamma", episode, gamma, states, controls, obstacles, goal, times, infos) From a4a951b5c9e364a23949484cd31b6e04902da533 Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:13:57 -0700 Subject: [PATCH 02/11] Add gamma sweep summary writers --- cfm_mppi/visualization/gamma_sweep_summary.py | 60 +++++++++++++++++++ 1 file changed, 60 insertions(+) create mode 100644 cfm_mppi/visualization/gamma_sweep_summary.py diff --git a/cfm_mppi/visualization/gamma_sweep_summary.py b/cfm_mppi/visualization/gamma_sweep_summary.py new file mode 100644 index 0000000..e1d56e6 --- /dev/null +++ b/cfm_mppi/visualization/gamma_sweep_summary.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +import csv +import json +from pathlib import Path +from typing import Any, Dict, Sequence + +import numpy as np + + +def aggregate(records: Sequence[Dict[str, Any]]) -> Dict[str, float]: + if not records: + return {"n": 0, "success_rate": 0.0, "collision_rate": 0.0, "mean_min_clearance": 0.0, "mean_final_goal_distance": 0.0, "mean_planning_time_ms": 0.0} + def mean(key: str) -> float: + return float(np.mean([float(r.get(key, 0.0) or 0.0) for r in records])) + def rate(key: str) -> float: + return float(np.mean([1.0 if r.get(key) else 0.0 for r in records])) + return dict( + n=len(records), + success_rate=rate("success"), + collision_rate=rate("collision"), + goal_reached_rate=rate("goal_reached"), + mean_min_clearance=mean("min_clearance"), + mean_final_goal_distance=mean("final_goal_distance"), + mean_control_effort=mean("control_effort"), + mean_control_smoothness=mean("control_smoothness"), + mean_planning_time_ms=1000.0 * mean("planning_wall_time_mean"), + p95_planning_time_ms=1000.0 * mean("planning_wall_time_p95"), + ) + + +def summarize_records(records: Sequence[Dict[str, Any]], gammas: Sequence[float]) -> Dict[str, Any]: + out: Dict[str, Any] = {"mizuta_cfm_mppi": aggregate([r for r in records if r["method"] == "mizuta_cfm_mppi"]), "safemppi_gamma": {}} + for gamma in gammas: + key = f"{gamma:.10g}" + rows = [r for r in records if r["method"] == "safemppi_gamma" and abs(float(r["gamma"]) - gamma) < 1e-9] + out["safemppi_gamma"][key] = {"gamma": float(gamma), **aggregate(rows)} + return out + + +def write_summary(root: Path, summary: Dict[str, Any], gammas: Sequence[float]) -> None: + rows = [{"method": "mizuta_cfm_mppi", "gamma": "", **summary["mizuta_cfm_mppi"]}] + rows += [{"method": "safemppi_gamma", **summary["safemppi_gamma"][f"{g:.10g}"]} for g in gammas] + with (root / "summary.json").open("w", encoding="utf-8") as f: + json.dump(summary, f, indent=2) + keys = sorted({k for r in rows for k in r}) + with (root / "summary.csv").open("w", encoding="utf-8", newline="") as f: + writer = csv.DictWriter(f, fieldnames=keys) + writer.writeheader() + writer.writerows(rows) + with (root / "summary.md").open("w", encoding="utf-8") as f: + f.write("# Mizuta CFM-MPPI vs online safeMPPI gamma sweep\n\n") + f.write("| method | gamma | success | collision | clearance | final dist | plan ms |\n") + f.write("|---|---:|---:|---:|---:|---:|---:|\n") + for r in rows: + f.write( + f"| {r.get('method','')} | {r.get('gamma','')} | {float(r.get('success_rate',0)):.3f} | " + f"{float(r.get('collision_rate',0)):.3f} | {float(r.get('mean_min_clearance',0)):.3f} | " + f"{float(r.get('mean_final_goal_distance',0)):.3f} | {float(r.get('mean_planning_time_ms',0)):.2f} |\n" + ) From 09cfdd3c7dd61cc86d338e118f3feb26d5337c79 Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:16:23 -0700 Subject: [PATCH 03/11] Probe branch file creation --- tmp_probe.txt | 1 + 1 file changed, 1 insertion(+) create mode 100644 tmp_probe.txt diff --git a/tmp_probe.txt b/tmp_probe.txt new file mode 100644 index 0000000..da0c4eb --- /dev/null +++ b/tmp_probe.txt @@ -0,0 +1 @@ +probe From 82dee082b39036090ae28a15b7ec2d46c5c9ba9f Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:16:30 -0700 Subject: [PATCH 04/11] Add gamma sweep renderer placeholder --- cfm_mppi/visualization/gamma_sweep_render.py | 3 +++ 1 file changed, 3 insertions(+) create mode 100644 cfm_mppi/visualization/gamma_sweep_render.py diff --git a/cfm_mppi/visualization/gamma_sweep_render.py b/cfm_mppi/visualization/gamma_sweep_render.py new file mode 100644 index 0000000..6f35c5f --- /dev/null +++ b/cfm_mppi/visualization/gamma_sweep_render.py @@ -0,0 +1,3 @@ +from __future__ import annotations + +# placeholder From a3eb209e55a629ef414f563ee61c5b7ee565cf7d Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:17:29 -0700 Subject: [PATCH 05/11] Add three-panel gamma sweep renderer --- cfm_mppi/visualization/gamma_sweep_render.py | 78 +++++++++++++++++++- 1 file changed, 77 insertions(+), 1 deletion(-) diff --git a/cfm_mppi/visualization/gamma_sweep_render.py b/cfm_mppi/visualization/gamma_sweep_render.py index 6f35c5f..cf08278 100644 --- a/cfm_mppi/visualization/gamma_sweep_render.py +++ b/cfm_mppi/visualization/gamma_sweep_render.py @@ -1,3 +1,79 @@ from __future__ import annotations -# placeholder +from pathlib import Path +from typing import Any, Dict, Sequence + +import numpy as np + + +def render_animation(root: Path, records: Sequence[Dict[str, Any]], summary: Dict[str, Any], gammas: Sequence[float], args) -> Dict[str, str]: + if not args.show_live: + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib import animation + from matplotlib.patches import Circle + + miz = next((r for r in records if r["method"] == "mizuta_cfm_mppi" and r["episode"] == args.video_episode), None) + safe_by_gamma = {f"{float(r['gamma']):.10g}": r for r in records if r["method"] == "safemppi_gamma" and r["episode"] == args.video_episode} + g = np.asarray(gammas, dtype=float) + s = summary["safemppi_gamma"] + success = np.asarray([s[f"{x:.10g}"]["success_rate"] for x in g]) + collision = np.asarray([s[f"{x:.10g}"]["collision_rate"] for x in g]) + clearance = np.asarray([s[f"{x:.10g}"]["mean_min_clearance"] for x in g]) + final_dist = np.asarray([s[f"{x:.10g}"]["mean_final_goal_distance"] for x in g]) + ms = np.asarray([s[f"{x:.10g}"]["mean_planning_time_ms"] for x in g]) + fig, axes = plt.subplots(1, 3, figsize=(18, 5.5)) + + def draw(frame: int): + gamma = float(g[frame]) + rec = safe_by_gamma[f"{gamma:.10g}"] + st = np.asarray(rec["states"], dtype=float) + goal = np.asarray(rec["goal"], dtype=float) + obs = np.asarray(rec["obstacles"], dtype=float) + for ax in axes: + ax.clear() + ax0, ax1, ax2 = axes + for o in obs: + ax0.add_patch(Circle((o[0], o[1]), o[2] + args.safety_margin, fill=False, linestyle="--")) + ax0.add_patch(Circle((o[0], o[1]), o[2], fill=False, alpha=0.5)) + if miz: + m = np.asarray(miz["states"], dtype=float) + ax0.plot(m[:, 0], m[:, 1], linestyle="--", label="Mizuta CFM-MPPI") + ax0.plot(st[:, 0], st[:, 1], label=f"safeMPPI gamma={gamma:.2f}") + ax0.scatter([st[0, 0]], [st[0, 1]], label="start") + ax0.scatter([goal[0]], [goal[1]], marker="*", s=120, label="goal") + ax0.axis("equal"); ax0.grid(True, alpha=0.3); ax0.set_title("1) trajectory"); ax0.legend(fontsize=8) + ax0.text(0.02, 0.02, f"success={int(rec['success'])} collision={int(rec['collision'])}\nmin_clearance={rec['min_clearance']:.3f}", transform=ax0.transAxes, fontsize=8, va="bottom") + ax1.plot(g, success, marker="o", label="safeMPPI success") + ax1.plot(g, collision, marker="x", label="safeMPPI collision") + ax1.axhline(summary["mizuta_cfm_mppi"].get("success_rate", 0.0), linestyle="--", label="Mizuta success") + ax1.axvline(gamma); ax1.set_xlim(-0.02, 1.02); ax1.set_ylim(-0.05, 1.05); ax1.grid(True, alpha=0.3); ax1.legend(fontsize=8); ax1.set_title("2) success/collision") + ax2.plot(g, clearance, marker="o", label="clearance") + ax2.plot(g, final_dist, marker="s", label="final distance") + ax2.plot(g, ms / max(float(np.max(ms)), 1e-9), marker="^", linestyle=":", label="normalized ms") + ax2.axhline(summary["mizuta_cfm_mppi"].get("mean_min_clearance", 0.0), linestyle="--", label="Mizuta clearance") + ax2.axvline(gamma); ax2.set_xlim(-0.02, 1.02); ax2.grid(True, alpha=0.3); ax2.legend(fontsize=8); ax2.set_title("3) margin/performance/compute") + fig.suptitle(f"{args.dataset}/{args.dynamics}: gamma={gamma:.2f}") + fig.tight_layout() + return [] + + anim = animation.FuncAnimation(fig, draw, frames=len(g), interval=int(1000 / max(args.fps, 1)), repeat=args.repeat) + draw(len(g) - 1) + png = root / "live_gamma_sweep_last_frame.png" + fig.savefig(png, dpi=args.dpi, bbox_inches="tight") + out = {"last_frame_png": str(png)} + if not args.no_video: + try: + mp4 = root / "live_gamma_sweep.mp4" + anim.save(mp4, writer="ffmpeg", fps=args.fps, dpi=args.dpi) + out["mp4"] = str(mp4) + except Exception as exc: + gif = root / "live_gamma_sweep.gif" + anim.save(gif, writer="pillow", fps=args.fps, dpi=args.dpi) + out["gif"] = str(gif) + out["render_note"] = str(exc) + if args.show_live: + plt.show() + plt.close(fig) + return out From c20a332ed6dac9c92bf75181de4977d587b32d89 Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:18:19 -0700 Subject: [PATCH 06/11] Add live gamma comparison CLI --- cfm_mppi/visualization/live_gamma_compare.py | 98 ++++++++++++++++++++ 1 file changed, 98 insertions(+) create mode 100644 cfm_mppi/visualization/live_gamma_compare.py diff --git a/cfm_mppi/visualization/live_gamma_compare.py b/cfm_mppi/visualization/live_gamma_compare.py new file mode 100644 index 0000000..7fc36f3 --- /dev/null +++ b/cfm_mppi/visualization/live_gamma_compare.py @@ -0,0 +1,98 @@ +from __future__ import annotations + +import argparse +import json +from datetime import datetime +from pathlib import Path +from typing import Any, Dict, List + +import torch + +from cfm_mppi.evaluation.eval_benchmark import DEFAULTS, _set_seed +from cfm_mppi.visualization.gamma_sweep_data import build_gamma_grid, rollout_mizuta, rollout_safemppi +from cfm_mppi.visualization.gamma_sweep_render import render_animation +from cfm_mppi.visualization.gamma_sweep_summary import summarize_records, write_summary + + +def run(args: argparse.Namespace) -> Dict[str, Any]: + _set_seed(args.seed) + args.u_min = tuple(float(x) for x in args.u_min) + args.u_max = tuple(float(x) for x in args.u_max) + gammas = build_gamma_grid(args.gamma_values, args.gamma_count) + if args.smoke: + args.horizon = min(args.horizon, 20) + args.num_episodes = min(args.num_episodes, 2) + args.safemppi_num_samples = min(args.safemppi_num_samples, 128) + args.no_video = True + root = Path(args.output_root) / datetime.now().strftime("%Y%m%d_%H%M%S") / args.dataset / args.dynamics + root.mkdir(parents=True, exist_ok=True) + records: List[Dict[str, Any]] = [] + jsonl = root / "gamma_sweep_records.jsonl" + with jsonl.open("w", encoding="utf-8", buffering=1) as f: + for ep in range(args.num_episodes): + rec = rollout_mizuta(args, ep) + records.append(rec) + f.write(json.dumps(rec) + "\n") + print(f"mizuta ep={ep+1}/{args.num_episodes} success={int(rec['success'])} collision={int(rec['collision'])} min_clearance={rec['min_clearance']:.3f}", flush=True) + for gamma in gammas: + for ep in range(args.num_episodes): + rec = rollout_safemppi(args, ep, gamma) + records.append(rec) + f.write(json.dumps(rec) + "\n") + print(f"safeMPPI gamma={gamma:.3f} ep={ep+1}/{args.num_episodes} success={int(rec['success'])} collision={int(rec['collision'])} min_clearance={rec['min_clearance']:.3f} plan_ms={1000*rec['planning_wall_time_mean']:.2f}", flush=True) + summary = summarize_records(records, gammas) + write_summary(root, summary, gammas) + artifacts = { + "output_dir": str(root), + "records_jsonl": str(jsonl), + "summary_json": str(root / "summary.json"), + "summary_csv": str(root / "summary.csv"), + "summary_md": str(root / "summary.md"), + } + artifacts.update(render_animation(root, records, summary, gammas, args)) + with (root / "artifacts.json").open("w", encoding="utf-8") as f: + json.dump(artifacts, f, indent=2) + print(json.dumps(artifacts, indent=2), flush=True) + return artifacts + + +def get_parser() -> argparse.ArgumentParser: + p = argparse.ArgumentParser(description="Three-panel live video comparing Mizuta CFM-MPPI and safeMPPI with gamma in [0,1].") + p.add_argument("--dataset", default="sfm", choices=["sfm", "ucy", "sdd"]) + p.add_argument("--dynamics", default="doubleintegrator", choices=["doubleintegrator", "unicycle"]) + p.add_argument("--num-episodes", type=int, default=10) + p.add_argument("--video-episode", type=int, default=0) + p.add_argument("--seed", type=int, default=0) + p.add_argument("--output-root", default="results/visualization/gamma_sweep") + p.add_argument("--gamma-count", type=int, default=21) + p.add_argument("--gamma-values", nargs="*", type=float, default=None) + p.add_argument("--horizon", type=int, default=DEFAULTS["horizon"]) + p.add_argument("--dt", type=float, default=DEFAULTS["dt"]) + p.add_argument("--safety-margin", type=float, default=DEFAULTS["safety_margin"]) + p.add_argument("--success-threshold", type=float, default=DEFAULTS["success_threshold"]) + p.add_argument("--u-min", nargs=2, type=float, default=list(DEFAULTS["u_min"])) + p.add_argument("--u-max", nargs=2, type=float, default=list(DEFAULTS["u_max"])) + p.add_argument("--safemppi-num-samples", type=int, default=1024) + p.add_argument("--safemppi-horizon", type=int, default=20) + p.add_argument("--safemppi-noise-sigma", type=float, default=0.6) + p.add_argument("--safemppi-temperature", type=float, default=1.0) + p.add_argument("--check-first-control-only", action="store_true") + p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") + p.add_argument("--mizuta-checkpoint", default="output_dir/cfm_transformer/checkpoint.pth") + p.add_argument("--safe-cfm-checkpoint", default="output_dir/safe_contextual_cfm/checkpoint_best.pth") + p.add_argument("--drifting-checkpoint", default="output_dir/drifting_generator/checkpoint_best.pth") + p.add_argument("--fps", type=int, default=2) + p.add_argument("--dpi", type=int, default=140) + p.add_argument("--repeat", action="store_true") + p.add_argument("--show-live", action="store_true") + p.add_argument("--no-video", action="store_true") + p.add_argument("--smoke", action="store_true") + return p + + +def main() -> None: + run(get_parser().parse_args()) + + +if __name__ == "__main__": + main() From 4fc0f91f8e3837e1fffa2f62be62707eb7b41f41 Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:18:56 -0700 Subject: [PATCH 07/11] Add gamma sweep runner --- scripts/run_gamma_sweep.sh | 26 ++++++++++++++++++++++++++ 1 file changed, 26 insertions(+) create mode 100644 scripts/run_gamma_sweep.sh diff --git a/scripts/run_gamma_sweep.sh b/scripts/run_gamma_sweep.sh new file mode 100644 index 0000000..45e5956 --- /dev/null +++ b/scripts/run_gamma_sweep.sh @@ -0,0 +1,26 @@ +#!/usr/bin/env bash +set -euo pipefail + +cd "$(dirname "$0")/.." +export PYTHONPATH="$PWD:${PYTHONPATH:-}" +export WANDB_MODE="${WANDB_MODE:-disabled}" +mkdir -p logs results/visualization/gamma_sweep + +DATASET="${DATASET:-sfm}" +DYNAMICS="${DYNAMICS:-doubleintegrator}" +NUM_EPISODES="${NUM_EPISODES:-20}" +GAMMA_COUNT="${GAMMA_COUNT:-21}" +SAFEMPPI_NUM_SAMPLES="${SAFEMPPI_NUM_SAMPLES:-2048}" +SAFEMPPI_HORIZON="${SAFEMPPI_HORIZON:-20}" +DEVICE="${DEVICE:-cuda}" +EXTRA_ARGS="${EXTRA_ARGS:-}" + +python -m cfm_mppi.visualization.live_gamma_compare \ + --dataset "$DATASET" \ + --dynamics "$DYNAMICS" \ + --num-episodes "$NUM_EPISODES" \ + --gamma-count "$GAMMA_COUNT" \ + --safemppi-num-samples "$SAFEMPPI_NUM_SAMPLES" \ + --safemppi-horizon "$SAFEMPPI_HORIZON" \ + --device "$DEVICE" \ + $EXTRA_ARGS 2>&1 | tee logs/gamma_sweep_${DATASET}_${DYNAMICS}.log From 1d432244a965ecb961faeaa1cb07cd544a3587be Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:19:14 -0700 Subject: [PATCH 08/11] Add overnight GPC training and gamma sweep runner --- scripts/train_gpc_then_gamma_sweep.sh | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) create mode 100644 scripts/train_gpc_then_gamma_sweep.sh diff --git a/scripts/train_gpc_then_gamma_sweep.sh b/scripts/train_gpc_then_gamma_sweep.sh new file mode 100644 index 0000000..41111f8 --- /dev/null +++ b/scripts/train_gpc_then_gamma_sweep.sh @@ -0,0 +1,24 @@ +#!/usr/bin/env bash +set -euo pipefail + +cd "$(dirname "$0")/.." +export PYTHONPATH="$PWD:${PYTHONPATH:-}" +export WANDB_MODE="${WANDB_MODE:-disabled}" +mkdir -p logs + +GPC_EPOCHS="${GPC_EPOCHS:-300}" +GPC_BATCH_SIZE="${GPC_BATCH_SIZE:-256}" +GPC_DEVICE="${GPC_DEVICE:-cuda}" +GPC_EXTRA_ARGS="${GPC_EXTRA_ARGS:-}" + +if [[ "${SKIP_GPC_TRAIN:-0}" != "1" ]]; then + [[ -f dataset/canonical/train.pt ]] || { echo "Missing dataset/canonical/train.pt. Run scripts/build_canonical_dataset.sh first." >&2; exit 1; } + [[ -f dataset/canonical/val.pt ]] || { echo "Missing dataset/canonical/val.pt. Run scripts/build_canonical_dataset.sh first." >&2; exit 1; } + python -m cfm_mppi.training.train_safe_cfm \ + --epochs "$GPC_EPOCHS" \ + --batch-size "$GPC_BATCH_SIZE" \ + --device "$GPC_DEVICE" \ + $GPC_EXTRA_ARGS 2>&1 | tee logs/train_gpc_overnight.log +fi + +DEVICE="${DEVICE:-$GPC_DEVICE}" scripts/run_gamma_sweep.sh From 5e00a7a22a1e77ec2c8db0d14b2037181b47e9cf Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:20:01 -0700 Subject: [PATCH 09/11] Add gamma grid check --- checks/gamma_grid_check.py | 9 +++++++++ 1 file changed, 9 insertions(+) create mode 100644 checks/gamma_grid_check.py diff --git a/checks/gamma_grid_check.py b/checks/gamma_grid_check.py new file mode 100644 index 0000000..8a1b56f --- /dev/null +++ b/checks/gamma_grid_check.py @@ -0,0 +1,9 @@ +from cfm_mppi.visualization.gamma_sweep_data import build_gamma_grid + + +def test_build_gamma_grid_includes_bounds_and_sorts(): + grid = build_gamma_grid([0.7, -0.2, 1.2, 0.3], count=5) + assert grid[0] == 0.0 + assert grid[-1] == 1.0 + assert grid == sorted(grid) + assert all(0.0 <= g <= 1.0 for g in grid) From f2b2155b5847490b9d4fc03d48fe56a949e2dc95 Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:21:22 -0700 Subject: [PATCH 10/11] Document gamma sweep workflow --- docs/GAMMA_SWEEP.md | 52 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 52 insertions(+) create mode 100644 docs/GAMMA_SWEEP.md diff --git a/docs/GAMMA_SWEEP.md b/docs/GAMMA_SWEEP.md new file mode 100644 index 0000000..80302bd --- /dev/null +++ b/docs/GAMMA_SWEEP.md @@ -0,0 +1,52 @@ +# Gamma Sweep Workflow + +This branch adds a three-panel animation comparing Mizuta CFM-MPPI with online safeMPPI over gamma in `[0, 1]`. + +Mizuta's pretrained checkpoint is expected at `output_dir/cfm_transformer/checkpoint.pth`. safeMPPI is evaluated online and does not require training. To train the safe contextual CFM/GPC model before the comparison, use the overnight script. + +Run the sweep: + +```bash +DATASET=sfm \ +DYNAMICS=doubleintegrator \ +NUM_EPISODES=20 \ +GAMMA_COUNT=21 \ +SAFEMPPI_NUM_SAMPLES=2048 \ +DEVICE=cuda \ +scripts/run_gamma_sweep.sh +``` + +Train safe contextual CFM/GPC first, then run the same sweep: + +```bash +GPC_EPOCHS=300 \ +GPC_BATCH_SIZE=256 \ +GPC_DEVICE=cuda \ +NUM_EPISODES=20 \ +GAMMA_COUNT=21 \ +SAFEMPPI_NUM_SAMPLES=2048 \ +scripts/train_gpc_then_gamma_sweep.sh +``` + +Direct module call: + +```bash +python -m cfm_mppi.visualization.live_gamma_compare \ + --dataset sfm \ + --dynamics doubleintegrator \ + --num-episodes 20 \ + --gamma-count 21 \ + --safemppi-num-samples 2048 \ + --device cuda +``` + +Outputs are written under `results/visualization/gamma_sweep////`: + +- `gamma_sweep_records.jsonl` +- `summary.csv` +- `summary.json` +- `summary.md` +- `live_gamma_sweep_last_frame.png` +- `live_gamma_sweep.mp4` or `live_gamma_sweep.gif` + +The panels are trajectory comparison, success/collision over gamma, and safety/performance/compute over gamma. From 46a40d863ea46356087e229e0217a60df4a31c09 Mon Sep 17 00:00:00 2001 From: DHLeexpress Date: Wed, 17 Jun 2026 14:21:31 -0700 Subject: [PATCH 11/11] Remove temporary probe file --- tmp_probe.txt | 1 - 1 file changed, 1 deletion(-) delete mode 100644 tmp_probe.txt diff --git a/tmp_probe.txt b/tmp_probe.txt deleted file mode 100644 index da0c4eb..0000000 --- a/tmp_probe.txt +++ /dev/null @@ -1 +0,0 @@ -probe