File size: 7,675 Bytes
728cafc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
#!/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()