Spaces:
Running on Zero
Running on Zero
Enhance GPU processing and batching capabilities; update README for clarity on test-time augmentation and batch sizes
17aade5 Download predict_image.py from akashch1512/SingleViewHeigthEstimation: direct link, hf CLI and curl.
- Browser
- Download file 10 kB
-
https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/predict_image.py
- Command line
-
hf download hf://spaces/akashch1512/SingleViewHeigthEstimation/predict_image.py
-
curl -L -o predict_image.py https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/predict_image.py
10 kB
| """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() | |