Instructions to use m1sc/reach-down-vit-release with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- DepthAnythingV2
How to use m1sc/reach-down-vit-release with DepthAnythingV2:
# Install from https://github.com/DepthAnything/Depth-Anything-V2 # Load the model and infer depth from an image import cv2 import torch from huggingface_hub import hf_hub_download from depth_anything_v2.dpt import DepthAnythingV2 # instantiate the model model = DepthAnythingV2(encoder="<ENCODER>", features=<NUMBER_OF_FEATURES>, out_channels=<OUT_CHANNELS>) # load the weights filepath = hf_hub_download(repo_id="m1sc/reach-down-vit-release", filename="depth_anything_v2_<ENCODER>.pth", repo_type="model") state_dict = torch.load(filepath, map_location="cpu") model.load_state_dict(state_dict) model.eval() raw_img = cv2.imread("your/image/path") depth = model.infer_image(raw_img) # HxW raw depth map in numpy - Notebooks
- Google Colab
- Kaggle
Download scripts/evaluate.py from m1sc/reach-down-vit-release: direct link, hf CLI and curl.
- Browser
- Download file 10 kB
-
https://huggingface.co/m1sc/reach-down-vit-release/resolve/main/scripts/evaluate.py
- Command line
-
hf download hf://m1sc/reach-down-vit-release/scripts/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/m1sc/reach-down-vit-release/resolve/main/scripts/evaluate.py
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() | |