Skip to content
149 changes: 149 additions & 0 deletions cfm_mppi/visualization/gamma_sweep_data.py
Original file line number Diff line number Diff line change
@@ -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)
79 changes: 79 additions & 0 deletions cfm_mppi/visualization/gamma_sweep_render.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
from __future__ import annotations

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
60 changes: 60 additions & 0 deletions cfm_mppi/visualization/gamma_sweep_summary.py
Original file line number Diff line number Diff line change
@@ -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"
)
98 changes: 98 additions & 0 deletions cfm_mppi/visualization/live_gamma_compare.py
Original file line number Diff line number Diff line change
@@ -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()
9 changes: 9 additions & 0 deletions checks/gamma_grid_check.py
Original file line number Diff line number Diff line change
@@ -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)
Loading