ckwolfe's picture
Publish complete fixed-wing ViT runs, dataset provenance, and evaluation artifacts
728cafc verified
Raw History Blame Contribute Delete
10 kB
#!/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()