File size: 10,029 Bytes
16ce72c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
17aade5
 
16ce72c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
"""Single-image prediction, packaged for the Phase-0 3D viewer.

    # simplest form: everything defaults to the repo-root checkpoint + image
    python predict_image.py /teamspace/studios/this_studio/image.png \
        --ckpt /teamspace/studios/this_studio/best.pt

    # you know the ground sample distance -> the heights become metric
    python predict_image.py scene.png --ckpt best.pt --gsd 0.3

    # fastest look, no 8x dihedral averaging
    python predict_image.py scene.png --ckpt best.pt --no-tta

This is the v3 replacement for v2's `predict_image.py`, and it is a thin shell
around `infer/predict.py`'s machinery rather than a second implementation:

  * the preprocessing contract (encoder mean/std, canonical GSD, tile size,
    radiometric stretch) is read from the **checkpoint** via `PreprocSpec`, so
    this script cannot drift from training the way v2's did;
  * the image is *resampled to the canonical GSD and tiled*, never squashed to
    512x512 — v2's small-image path destroyed the metric scale outright;
  * inference goes through `infer.engine.predict_scene`, the same code path the
    final evaluation uses.

What it adds over `python -m infer.predict` is the **viewer contract**.
`DepthWizard/viewer/phase0_viewer.html` matches its file picker on substrings
("rgb" / "height") and decodes height with

    h_m = meta.height_min_m + (red/255) * (meta.height_max_m - meta.height_min_m)
    y_px = h_m / meta.gsd_m

so the export directory here is named to be picked up, and the 0-255 ramp is
stretched over a *robust* (percentile) height range instead of the raw min/max.
That matters: one hot pixel of LiDAR-noise-shaped prediction at 60 m would
otherwise compress a whole 15 m suburb into the bottom four grey levels.
"""

from __future__ import annotations

import argparse
import json
import sys
import time
from pathlib import Path

import numpy as np
import torch
from PIL import Image

_HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(_HERE))

from dwdata.preprocess import read_scene  # noqa: E402
from infer.engine import predict_scene  # noqa: E402
from infer.predict import load_model  # noqa: E402

Image.MAX_IMAGE_PIXELS = None
_VIEWER_HTML = _HERE.parents[1] / "viewer" / "phase0_viewer.html"


def robust_range(h: np.ndarray, lo_pct: float, hi_pct: float) -> tuple[float, float]:
    """Percentile height range used for the 8/16-bit display ramp.

    Clamped so the low end never rises above 0 m: ground *is* the reference
    surface of an nDSM and the viewer should draw it flat, not floating.
    """
    lo = float(np.nanpercentile(h, lo_pct)) if lo_pct > 0 else float(np.nanmin(h))
    hi = float(np.nanpercentile(h, hi_pct)) if hi_pct < 100 else float(np.nanmax(h))
    lo = min(lo, 0.0)
    return lo, max(hi, lo + 1e-3)


def write_viewer_export(out_dir: Path, stem: str, rgb_u8: np.ndarray,
                        height_m: np.ndarray, gsd_m: float, spec, meta,
                        clip_pct: tuple[float, float], elapsed_s: float) -> Path:
    out_dir.mkdir(parents=True, exist_ok=True)
    H, W = height_m.shape
    lo, hi = robust_range(height_m, *clip_pct)
    span = hi - lo

    # Height is encoded as 8-bit RGB, not a 16-bit greyscale PNG: R is the high
    # byte and G the low byte of a 16-bit fixed-point ramp.  The viewer reads
    # exactly one 8-bit channel back out of a 2D canvas (R), so it sees the plain
    # 0-255 ramp it expects, while R*256+G still recovers the full 16 bits for
    # any other consumer.  A true 16-bit PNG would leave the 16->8 downconversion
    # up to the browser's image decoder, which is not worth depending on.
    norm = np.clip((height_m - lo) / span, 0.0, 1.0)
    q16 = (norm * 65535.0 + 0.5).astype(np.uint16)
    enc = np.stack([(q16 >> 8).astype(np.uint8),
                    (q16 & 0xFF).astype(np.uint8),
                    np.zeros(q16.shape, np.uint8)], axis=-1)
    Image.fromarray(rgb_u8).save(out_dir / "rgb.png")
    Image.fromarray(enc, mode="RGB").save(out_dir / "height16.png")
    np.save(out_dir / "height_m.npy", height_m.astype(np.float32))

    payload = {
        # --- consumed by phase0_viewer.html -------------------------
        "stem": stem,
        "height_min_m": lo,
        "height_max_m": hi,
        "gsd_m": float(gsd_m),
        "size_px": [W, H],
        # --- provenance / diagnostics -------------------------------
        "product": "nDSM" if meta.gsd_source != "assumed" else "rDSM",
        "units": "metres_above_ground",
        "display_clip_pct": list(clip_pct),
        "height_true_min_m": float(np.nanmin(height_m)),
        "height_true_max_m": float(np.nanmax(height_m)),
        "height_mean_m": float(np.nanmean(height_m)),
        "height_median_m": float(np.nanmedian(height_m)),
        "frac_below_1m": float((height_m < 1.0).mean()),
        "vertical_quantum_m": span / 255.0,
        "scene": meta.summary(),
        "preproc": spec.to_dict(),
        "elapsed_s": round(elapsed_s, 1),
        "files": ["rgb.png", "height16.png", "height_m.npy"],
        "height_encode": "q = R*256+G (or just R/255 for 8-bit); "
                          "h_m = height_min_m + (q/65535)*(height_max_m-height_min_m)",
    }
    (out_dir / "meta.json").write_text(json.dumps(payload, indent=2))
    return out_dir, payload


def report(payload: dict, out_dir: Path) -> None:
    p = payload
    print(f"\nheight  {p['size_px'][0]}x{p['size_px'][1]}px @ {p['gsd_m']:.3f} m/px "
          f"({p['scene']['gsd_source']})")
    print(f"        true range   {p['height_true_min_m']:+.2f} .. {p['height_true_max_m']:+.2f} m")
    print(f"        display ramp {p['height_min_m']:+.2f} .. {p['height_max_m']:+.2f} m "
          f"({p['display_clip_pct'][0]}/{p['display_clip_pct'][1]} pct, "
          f"{p['vertical_quantum_m']:.3f} m per grey level)")
    print(f"        mean {p['height_mean_m']:.2f} m · median {p['height_median_m']:.2f} m · "
          f"{p['frac_below_1m'] * 100:.1f}% below 1 m")
    if p["frac_below_1m"] < 0.15:
        print("        WARNING: little flat ground — check --gsd; the model may be "
              "reading texture as terrain at this scale")
    if p["product"] == "rDSM":
        print("        NOTE: no scale metadata and no --gsd, so the GSD was assumed. "
              "Shapes are right, absolute metres are not.")
    print(f"\nsaved -> {out_dir}")
    print(f"  {', '.join(p['files'])}, meta.json")
    print("\nopen the viewer and pick rgb.png + height16.png + meta.json:")
    print(f"  python -m http.server -d {_VIEWER_HTML.parent} 8000")
    print(f"  then http://localhost:8000/{_VIEWER_HTML.name}")


def main() -> None:
    ap = argparse.ArgumentParser(
        description="v3 single-image prediction -> Phase-0 viewer export")
    ap.add_argument("image", nargs="?", default="/teamspace/studios/this_studio/image.png")
    ap.add_argument("--ckpt", default="/teamspace/studios/this_studio/best.pt")
    ap.add_argument("--out-dir", default=None,
                    help="default: <output_dir>/viewer_sample/<stem>")
    ap.add_argument("--gsd", type=float, default=0.0,
                    help="metres per pixel of the input; makes the output metric")
    ap.add_argument("--assumed-gsd", type=float, default=0.5,
                    help="fallback GSD when the image carries no scale "
                         "(default 0.5 = the model's canonical GSD, i.e. 1 px = 0.5 m)")
    ap.add_argument("--tta", dest="tta", action="store_true", default=True,
                    help="8x dihedral TTA (default on — nadir imagery has no up)")
    ap.add_argument("--no-tta", dest="tta", action="store_false")
    ap.add_argument("--tta-scales", default="1.0")
    ap.add_argument("--overlap", type=float, default=0.5,
                    help="tile overlap; 0.5 gives a smoother Hann blend than eval's 0.25")
    ap.add_argument("--batch-tiles", type=int, default=0,
                    help="tiles per forward pass; 0 = size from free GPU memory")
    ap.add_argument("--max-side", type=int, default=0, help="downsample huge scenes first")
    ap.add_argument("--clip-pct", default="0.5,99.5",
                    help="percentiles for the display ramp (use 0,100 for raw min/max)")
    ap.add_argument("--device", default="")
    ap.add_argument("--hf-token", default="")
    a = ap.parse_args()

    clip = tuple(float(v) for v in a.clip_pct.split(","))
    assert len(clip) == 2 and clip[0] < clip[1], "--clip-pct wants lo,hi"

    device = torch.device(a.device or ("cuda" if torch.cuda.is_available() else "cpu"))
    print(f"[predict] device: {device}")
    model, spec, cfg = load_model(a.ckpt, device, a.hf_token)

    rgb, meta = read_scene(a.image, user_gsd_m=a.gsd,
                           assumed_gsd_m=a.assumed_gsd or spec.canonical_gsd_m,
                           max_side=a.max_side)
    print(f"[predict] {a.image}: {meta.width}x{meta.height} @ {meta.gsd_m:.4f} m/px "
          f"({meta.gsd_source}) -> canonical {spec.canonical_gsd_m} m/px, "
          f"{spec.tile_size}px tiles, tta={a.tta}")
    if device.type == "cpu" and a.tta:
        print("[predict] CPU + TTA means 8 forward passes per tile; --no-tta for a quick look")

    scales = tuple(float(s) for s in a.tta_scales.split(",") if s.strip())
    amp_dt = (torch.bfloat16 if cfg.amp_dtype == "bf16" else torch.float16) \
        if device.type == "cuda" else None

    def prog(done, total):
        print(f"  tiles {done}/{total}", flush=True)

    t0 = time.time()
    height, _ = predict_scene(
        model, rgb, meta.gsd_m, spec, device, tta=a.tta, tta_scales=scales,
        amp_dtype=amp_dt, overlap=a.overlap, batch_tiles=a.batch_tiles, progress=prog,
    )
    elapsed = time.time() - t0

    stem = Path(a.image).stem
    out_dir = Path(a.out_dir) if a.out_dir else \
        Path(cfg.output_dir) / "viewer_sample" / stem
    out_dir, payload = write_viewer_export(
        out_dir, stem, rgb, height, meta.gsd_m, spec, meta, clip, elapsed)
    report(payload, out_dir)


if __name__ == "__main__":
    main()