reach-down-vit-release / scripts /make_dataset.py
ckwolfe's picture
Publish complete fixed-wing ViT runs, dataset provenance, and evaluation artifacts
728cafc verified
Raw History Blame Contribute Delete
7.68 kB
#!/usr/bin/env python3
"""Deep reach iteration with domain randomization: terrains x states x platforms.
For every synthetic ground-truth terrain we sweep aircraft states, each with
a randomly drawn airframe (platform dynamics), landing-limit jitter, wind,
and DEM observation noise; solve the HJ reach-avoid problem (JAX, GPU) for
the *true* surface and dynamics; and emit one tile per (terrain, state):
inputs z_obs, slope_obs, rough_obs, hazard, margin_analytic (observed)
labels label_value (HJ), label_run (hazard-masked landability) (privileged)
meta platform + jittered limits + wind + state (conditioning scalars)
Safe landing is contingent on the platform: the same terrain yields different
labels for different airframes, and the network sees the airframe only
through its margin channel and conditioning scalars.
.venv/bin/python scripts/make_dataset.py --n-terrains 60 --states 4
"""
from __future__ import annotations
import argparse
import json
import math
import sys
from pathlib import Path
import numpy as np
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.paths import data_path
from reachdown.hj import HJConfig, hj_reach_margin
from reachdown.platforms import PLATFORMS, Platform, jittered_limits
from reachdown.synth import noisy_observation, sample_config
from reachdown.terrain import apply_hazards, roughness_m, slope_deg
WIND_LEVELS = (0.0, 3.0, 6.0) # m/s adversarial ball; discrete to bound JIT compiles
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--n-terrains", type=int, default=60)
ap.add_argument("--states", type=int, default=4, help="aircraft states per terrain")
ap.add_argument("--size", type=int, default=512, help="terrain edge, px")
ap.add_argument("--agl", type=float, nargs=2, default=(150.0, 550.0), metavar=("LO", "HI"))
ap.add_argument("--relief", type=float, nargs=2, default=(60.0, 220.0), metavar=("LO", "HI"),
help="generator-DR relief range (m), sampled per terrain")
ap.add_argument("--z-noise-max", type=float, default=1.0, help="obs noise sigma upper bound, m")
ap.add_argument("--safety-buffer", type=float, default=0.0,
help="extra spare altitude (m) required in the label; >0 => conservative")
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--out", type=Path, default=data_path("data", "tiles"))
ap.add_argument("--vehicle-config", type=Path, default=None,
help="JSON with {'vehicles': [platform configs]} to fine-tune a "
"new fleet; overrides the built-in PLATFORMS registry")
ap.add_argument("--no-jitter", action="store_true",
help="FW-2: disable landing-limit jitter (DR off)")
ap.add_argument("--wind-levels", type=float, nargs="+", default=None,
help="FW-2: override wind DR levels (e.g. 0 for no-wind)")
ap.add_argument("--fixed-generator", action="store_true",
help="FW-2: fixed terrain-generator statistics (no generator DR)")
args = ap.parse_args()
global WIND_LEVELS
if args.wind_levels is not None:
WIND_LEVELS = tuple(args.wind_levels)
if args.vehicle_config is not None:
cfgs = json.loads(args.vehicle_config.read_text())["vehicles"]
fleet = tuple(Platform.from_config(c) for c in cfgs)
print(f"fleet from {args.vehicle_config}: {[p.name for p in fleet]}")
else:
fleet = PLATFORMS
rng = np.random.default_rng(args.seed)
args.out.mkdir(parents=True, exist_ok=True)
n_total = args.n_terrains * args.states
for t in range(args.n_terrains):
# generator-statistics DR: each terrain draws its own world profile
cfg = (SynthConfig(size=args.size) if args.fixed_generator
else sample_config(rng, size=args.size, relief=tuple(args.relief)))
grid = synth_terrain(cfg, seed=args.seed + t)
hazard = synth_hazards(cfg, seed=args.seed + t)
extent = args.size * cfg.px
# surface-only quantities are shared across the per-state limit jitter
slope_clean = slope_deg(grid)
rough_clean = roughness_m(grid, fleet[0].limits.rough_window_m)
for s in range(args.states):
platform = fleet[rng.integers(len(fleet))]
limits = platform.limits if args.no_jitter else jittered_limits(platform.limits, rng)
wind = float(WIND_LEVELS[rng.integers(len(WIND_LEVELS))])
sigma_z = float(rng.uniform(0.0, args.z_noise_max))
# labels: clean surface, true dynamics/limits (privileged)
land = landability(grid, limits, slope=slope_clean, rough=rough_clean)
safe_run = apply_hazards(land.run, hazard)
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(*args.agl)),
heading=float(rng.uniform(-math.pi, math.pi)),
airspeed=platform.airspeed,
)
# inputs: what the aircraft would actually observe. Drawn BEFORE the
# HJ solve so a resumed run consumes the identical rng stream when it
# skips already-written tiles (hj_reach_margin draws nothing).
grid_obs = noisy_observation(grid, sigma_z, rng)
k = t * args.states + s
out_npz = args.out / f"tile_{k:05d}.npz"
if out_npz.exists():
print(f"[{k + 1}/{n_total}] terrain {t} exists, skip (resume)")
continue
margin_hj = hj_reach_margin(
grid, state, platform.polar, HJConfig(wind_max=wind)
)
# label is conservative by construction when --safety-buffer > 0:
# the safe set shrinks strictly inside the true HJ footprint
label_value = compose_value(safe_run, margin_hj, safety_buffer_m=args.safety_buffer)
np.savez_compressed(
out_npz,
z_obs=grid_obs.z,
slope_obs=slope_deg(grid_obs),
rough_obs=roughness_m(grid_obs, limits.rough_window_m),
hazard=hazard,
margin_analytic=reach_margin(grid_obs, state, platform.polar, wind_max=wind),
margin_hj=margin_hj.astype(np.float32),
label_value=label_value,
label_run=safe_run,
meta=json.dumps(
{
"px": cfg.px,
"terrain_seed": args.seed + t,
"relief_m": cfg.relief_m,
"spectral_beta": cfg.beta,
"platform": platform.name,
"slope_max_deg": limits.slope_max_deg,
"rough_max_m": limits.rough_max_m,
"run_length_m": limits.run_length_m,
"wind_max": wind,
"sigma_z": sigma_z,
"state": [state.x, state.y, state.z, state.heading],
}
),
)
print(f"[{k + 1}/{n_total}] terrain {t} {platform.name} wind={wind:.0f} "
f"sigma_z={sigma_z:.2f} safe&reachable={float((label_value > 0.5).mean()):.1%}")
print(f"wrote {n_total} tiles to {args.out}")
if __name__ == "__main__":
main()