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) diff --git a/cfm_mppi/visualization/gamma_sweep_render.py b/cfm_mppi/visualization/gamma_sweep_render.py new file mode 100644 index 0000000..cf08278 --- /dev/null +++ b/cfm_mppi/visualization/gamma_sweep_render.py @@ -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 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" + ) 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() 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) 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. 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 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