#!/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()