#!/usr/bin/env python3 """Episode-level evaluation: top-1 site safety and site regret. The number that matters for flight is not per-cell AUROC but whether the *selected* site is actually safe. For each episode (fresh terrain, fresh state, random platform, observation noise) we compare two site-selection policies against privileged ground truth: - **vit**: argmax of the fine-tuned ViT's value prediction - **analytic**: argmax of the onboard-computable baseline (oracle rules on the *observed* noisy heightfield x analytic glide margin) Ground truth = clean-surface landability x HJ reach margin. A pick is safe under the STRICT criterion when the 5x5 neighborhood minimum of the GT value at the chosen cell exceeds the safety threshold (the same criterion used by scripts/real_dem_eval.py and the closed-loop sim); the previous LENIENT point criterion (GT value at the single chosen cell) is also recorded for comparison. Regret is the GT value gap to the best cell. .venv/bin/python scripts/evaluate.py --episodes 48 """ from __future__ import annotations import argparse import math import sys from pathlib import Path import numpy as np import torch import torch.nn.functional as F from scipy import ndimage sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from reachdown import ( AircraftState, SynthConfig, compose_value, landability, reach_margin, synth_hazards, synth_terrain, ) from reachdown.data import normalize_inputs from reachdown.hj import HJConfig, hj_reach_margin from reachdown.model import LandingValueViT from reachdown.platforms import PLATFORMS from reachdown.synth import noisy_observation from reachdown.terrain import apply_hazards def main() -> None: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--checkpoint", type=Path, default=Path("runs/vit/best.pt")) ap.add_argument("--episodes", type=int, default=48) ap.add_argument("--seed-base", type=int, default=1000, help="disjoint from train seeds") ap.add_argument("--safe-threshold", type=float, default=0.5) ap.add_argument("--wind", type=float, default=0.0, help="adversarial wind bound (m/s): enters the GT oracle, the " "ViT wind channel, AND the wind-aware analytic baseline") ap.add_argument("--out-csv", type=Path, default=None, help="persist per-episode results to this CSV") args = ap.parse_args() def pick(vmap: np.ndarray) -> tuple[int, int]: """Robust argmax: score each cell by its neighborhood minimum so a lone over-confident pixel can't win.""" robust = ndimage.minimum_filter(vmap, size=5) return np.unravel_index(int(np.argmax(robust)), vmap.shape) # type: ignore[return-value] device = "cuda" if torch.cuda.is_available() else "cpu" ckpt = torch.load(args.checkpoint, map_location=device, weights_only=False) model = LandingValueViT( unfreeze_last_blocks=ckpt.get("unfreeze_blocks", 4), pretrained=ckpt.get("pretrained", True), ).to(device).eval() model.load_state_dict(ckpt["state_dict"]) size = ckpt["input_size"] zero_channels = ckpt.get("zero_channels", []) rng = np.random.default_rng(args.seed_base) cfg = SynthConfig() stats: dict[str, list] = { k: [] for k in ( "episode", "platform", "sigma_z", "strict_feasible", "vit_safe_strict", "vit_safe_lenient", "base_safe_strict", "base_safe_lenient", "vit_regret", "base_regret", "false_safe", ) } for ep in range(args.episodes): seed = args.seed_base + ep grid = synth_terrain(cfg, seed=seed) hazard = synth_hazards(cfg, seed=seed) platform = PLATFORMS[rng.integers(len(PLATFORMS))] sigma_z = float(rng.uniform(0.0, 1.0)) extent = cfg.size * cfg.px x0, y0 = rng.uniform(0.2 * extent, 0.8 * extent, size=2) state = AircraftState( x=float(x0), y=float(y0), z=grid.sample(float(x0), float(y0)) + float(rng.uniform(150.0, 550.0)), heading=float(rng.uniform(-math.pi, math.pi)), airspeed=platform.airspeed, ) # privileged ground truth land_gt = landability(grid, platform.limits) value_gt = compose_value( apply_hazards(land_gt.run, hazard), hj_reach_margin(grid, state, platform.polar, HJConfig(wind_max=args.wind)), ) best = float(value_gt.max()) if best < args.safe_threshold: continue # no feasible site exists; selection is untestable # strict site-safety criterion: the whole 5x5 neighborhood of the pick # must be GT-safe (matches real_dem_eval.py / closed_loop_sim.py) value_gt_min = ndimage.minimum_filter(value_gt, size=5) strict_feasible = float(value_gt_min.max()) >= args.safe_threshold # shared observation. The analytic baseline is wind-AWARE (first-order # headwind range penalty) so the comparison is fair, not a strawman; # the ViT's edge is its margin over a competent baseline. grid_obs = noisy_observation(grid, sigma_z, rng) land_obs = landability(grid_obs, platform.limits) margin_obs = reach_margin(grid_obs, state, platform.polar, wind_max=args.wind) # policy 1: ViT x = torch.from_numpy( normalize_inputs( grid_obs.z, land_obs.slope_deg, land_obs.rough_m, hazard, margin_obs, slope_max_deg=platform.limits.slope_max_deg, rough_max_m=platform.limits.rough_max_m, run_length_m=platform.limits.run_length_m, wind_max=args.wind, ) )[None].to(device) for ch in zero_channels: x[:, ch] = 0.0 x = F.interpolate(x, size=(size, size), mode="bilinear", align_corners=False) with torch.no_grad(), torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"): logits = model(x) pred = F.interpolate(logits.float(), size=grid.shape, mode="bilinear", align_corners=False)[0, 0].sigmoid().cpu().numpy() # policy 2: analytic baseline on the same observation value_base = compose_value(apply_hazards(land_obs.run, hazard), margin_obs) stats["episode"].append(ep) stats["platform"].append(platform.name) stats["sigma_z"].append(sigma_z) stats["strict_feasible"].append(float(strict_feasible)) for policy, vmap in (("vit", pred), ("base", value_base)): r, c = pick(vmap) gt = float(value_gt[r, c]) stats[f"{policy}_safe_lenient"].append(float(gt >= args.safe_threshold)) stats[f"{policy}_safe_strict"].append( float(value_gt_min[r, c] >= args.safe_threshold)) stats[f"{policy}_regret"].append(best - gt) # one-sided safety metric: does the ViT call cells safe that the HJ # oracle calls unsafe? (optimistic / false-safe rate — the number that # actually matters for a safety map; AUROC/MAE are symmetric and hide it) vit_safe_mask = pred >= args.safe_threshold oracle_unsafe = value_gt < args.safe_threshold stats["false_safe"].append( float((vit_safe_mask & oracle_unsafe).sum() / max(vit_safe_mask.sum(), 1)) ) print(f"[{ep + 1}/{args.episodes}] {platform.name} sigma_z={sigma_z:.2f} " f"vit_gt={stats['vit_regret'][-1]:.3f} base_gt={stats['base_regret'][-1]:.3f}") n = len(stats["vit_safe_strict"]) def cp_ci(k: int, m: int, alpha: float = 0.05) -> tuple[float, float]: """95% Clopper-Pearson (exact binomial) interval for k successes / m.""" from scipy import stats as sps lo = 0.0 if k == 0 else float(sps.beta.ppf(alpha / 2, k, m - k + 1)) hi = 1.0 if k == m else float(sps.beta.ppf(1 - alpha / 2, k + 1, m - k)) return lo, hi def rate(key: str, mask=None) -> str: a = np.array(stats[key], dtype=float) if mask is not None: a = a[mask] m = len(a) if m == 0: return "n/a (0 eps)" k = int(a.sum()) lo, hi = cp_ci(k, m) return f"{k / m:.1%} ({k}/{m}, 95% CI [{lo:.1%}, {hi:.1%}])" def ms(key: str) -> str: a = np.array(stats[key], dtype=float) return f"{a.mean():.4f} +/- {a.std():.4f}" print(f"\n{n} scoreable episodes (threshold {args.safe_threshold}, wind {args.wind} m/s):") print(f" top-1 site safety STRICT (5x5-min GT) : vit {rate('vit_safe_strict')} " f"analytic(wind-aware) {rate('base_safe_strict')}") print(f" top-1 site safety LENIENT (point GT) : vit {rate('vit_safe_lenient')} " f"analytic(wind-aware) {rate('base_safe_lenient')}") print(f" site regret : vit {ms('vit_regret')} analytic(wind-aware) {ms('base_regret')}") print(f" ViT false-safe rate (optimistic cells / ViT-safe cells): " f"{np.mean(stats['false_safe']):.1%} +/- {np.std(stats['false_safe']):.1%}") plats = np.array(stats["platform"]) print(" per-airframe strict site safety:") for p in sorted(set(plats)): mask = plats == p print(f" {p:12s} vit {rate('vit_safe_strict', mask)} " f"analytic {rate('base_safe_strict', mask)}") if args.out_csv is not None: import csv args.out_csv.parent.mkdir(parents=True, exist_ok=True) keys = [k for k in stats if k != "episode"] new = not args.out_csv.exists() with args.out_csv.open("a", newline="") as f: w = csv.writer(f) if new: w.writerow(["wind", "episode", *keys]) for i in range(n): w.writerow([args.wind, stats["episode"][i], *(stats[k][i] for k in keys)]) print(f" appended {n} per-episode rows to {args.out_csv}") if __name__ == "__main__": main()