File size: 15,069 Bytes
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
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
"""The preprocessing *contract* β€” the single definition of "what the model eats".

This module is the answer to "how do I make inference input identical to training
input".  Training (`dwdata/dataset.py`) and inference (`infer/predict.py`) both
build their tensors here and nowhere else, and the resolved spec is serialised
into every checkpoint (`PreprocSpec.to_dict()`), so a checkpoint always carries
the recipe that produced it.

The contract, in order:

  1. **Read**  RGB uint8 (H, W, 3) + the ground sample distance in metres/pixel.
     For a GeoTIFF the GSD comes from the affine transform (converted to metres
     if the CRS is geographic); otherwise the caller supplies it.
  2. **Radiometric stretch** β€” per *scene* percentile stretch to full 8-bit range.
     Removes the sensor's exposure/white-balance from the input, which is what
     stops the network keying on absolute brightness.  Scene-level (not crop- or
     tile-level) so it is a single deterministic transform of the whole image.
  3. **Geometric canonicalisation** β€” resample so one pixel == `canonical_gsd_m`
     metres on the ground.  Heights are in metres and are NOT touched.
  4. **Tile** into `tile_size` windows (inference: overlapping + Hann-blended;
     training: one random crop, taken *before* the resample so it can never be
     padded).
  5. **Normalise** with the encoder's own mean/std (DINOv3-SAT does not use the
     ImageNet constants) -> float32 CHW.

Nothing here imports torch beyond the final tensor conversion, so the same code
runs in a DataLoader worker and in a CPU-only inference container.
"""

from __future__ import annotations

import json
from dataclasses import asdict, dataclass, field
from pathlib import Path

import numpy as np

# DINOv3 satellite checkpoints (SAT-493M) were pretrained with their own
# statistics, not ImageNet's.  v1/v2 used the ImageNet constants for both, which
# shifts every input by ~0.3 sigma before the encoder ever sees it.  We prefer the
# values the HF image processor reports and fall back to these.
DINOV3_SAT_MEAN = (0.430, 0.411, 0.296)
DINOV3_SAT_STD = (0.213, 0.156, 0.143)
IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)


@dataclass
class PreprocSpec:
    """Everything needed to turn an arbitrary image into a model input."""

    encoder_model_id: str = "facebook/dinov3-vitl16-pretrain-sat493m"
    mean: tuple = DINOV3_SAT_MEAN
    std: tuple = DINOV3_SAT_STD
    canonical_gsd_m: float = 0.5
    tile_size: int = 512
    patch: int = 16
    radiometric_stretch: bool = True
    stretch_lo_pct: float = 2.0
    stretch_hi_pct: float = 98.0
    # informational: what the height output means
    target: str = "nDSM_agl_metres"
    version: str = "v3"

    def to_dict(self) -> dict:
        return asdict(self)

    @classmethod
    def from_dict(cls, d: dict) -> "PreprocSpec":
        known = {f for f in cls.__dataclass_fields__}
        return cls(**{k: v for k, v in d.items() if k in known})

    @classmethod
    def from_config(cls, cfg, resolve_stats: bool = True) -> "PreprocSpec":
        mean, std = DINOV3_SAT_MEAN, DINOV3_SAT_STD
        if resolve_stats:
            mean, std = resolve_encoder_stats(
                cfg.encoder_model_id, getattr(cfg, "hf_token", "") or None
            )
        return cls(
            encoder_model_id=cfg.encoder_model_id,
            mean=tuple(mean), std=tuple(std),
            canonical_gsd_m=float(cfg.canonical_gsd_m),
            tile_size=int(cfg.tile_size),
            radiometric_stretch=bool(cfg.radiometric_stretch),
            stretch_lo_pct=float(cfg.stretch_lo_pct),
            stretch_hi_pct=float(cfg.stretch_hi_pct),
        )

    # -- step 5 -------------------------------------------------------
    def normalise(self, rgb_u8: np.ndarray) -> np.ndarray:
        """(H, W, 3) uint8 -> (3, H, W) float32, encoder-normalised."""
        m = np.asarray(self.mean, dtype=np.float32)
        s = np.asarray(self.std, dtype=np.float32)
        x = (rgb_u8.astype(np.float32) / 255.0 - m) / s
        return np.ascontiguousarray(np.transpose(x, (2, 0, 1)))

    def denormalise(self, chw: np.ndarray) -> np.ndarray:
        m = np.asarray(self.mean, dtype=np.float32).reshape(3, 1, 1)
        s = np.asarray(self.std, dtype=np.float32).reshape(3, 1, 1)
        x = np.clip(chw * s + m, 0, 1)
        return (np.transpose(x, (1, 2, 0)) * 255).astype(np.uint8)


def resolve_encoder_stats(model_id: str, token: str | None = None):
    """Ask the encoder's own HF image processor for its normalisation stats.

    Falls back to the published DINOv3-SAT constants (satellite checkpoints) or
    ImageNet (everything else) when the hub is unreachable.
    """
    try:
        from transformers import AutoImageProcessor

        proc = AutoImageProcessor.from_pretrained(model_id, token=token)
        mean = tuple(float(v) for v in proc.image_mean)
        std = tuple(float(v) for v in proc.image_std)
        if len(mean) == 3 and len(std) == 3 and all(s > 0 for s in std):
            print(f"[preproc] encoder stats from processor: mean={mean} std={std}")
            return mean, std
    except Exception as e:  # noqa: BLE001
        print(f"[preproc] image-processor lookup failed ({e}); using defaults")
    if "sat" in model_id.lower():
        return DINOV3_SAT_MEAN, DINOV3_SAT_STD
    return IMAGENET_MEAN, IMAGENET_STD


# ---------------------------------------------------------------------
# step 2 β€” radiometry
# ---------------------------------------------------------------------
def scene_stretch_bounds(
    rgb_u8: np.ndarray, lo_pct: float = 2.0, hi_pct: float = 98.0,
    max_samples: int = 4_000_000,
) -> tuple[np.ndarray, np.ndarray]:
    """Per-channel (lo, hi) byte values for a percentile stretch of one scene.

    For uint8 input the percentile comes from a 256-bin histogram, which is a
    single O(N) pass instead of `np.percentile`'s sort β€” ~15x faster on a 1 Mpx
    tile, and this runs once per *training crop*, so it was a real slice of the
    DataLoader budget.  Non-uint8 input falls back to `np.percentile`.
    """
    a = rgb_u8.reshape(-1, rgb_u8.shape[-1])
    if a.dtype == np.uint8:
        n, c = a.shape
        lo = np.empty(c, np.float32)
        hi = np.empty(c, np.float32)
        for ch in range(c):
            counts = np.bincount(a[:, ch], minlength=256)
            cum = np.cumsum(counts)
            lo[ch] = np.searchsorted(cum, lo_pct / 100.0 * n)
            hi[ch] = np.searchsorted(cum, hi_pct / 100.0 * n)
        return lo, np.maximum(hi, lo + 1.0)
    if a.shape[0] > max_samples:
        step = int(np.ceil(a.shape[0] / max_samples))
        a = a[::step]
    lo = np.percentile(a, lo_pct, axis=0).astype(np.float32)
    hi = np.percentile(a, hi_pct, axis=0).astype(np.float32)
    hi = np.maximum(hi, lo + 1.0)
    return lo, hi


def stretch_lut(lo: np.ndarray, hi: np.ndarray) -> np.ndarray:
    """(256, C) uint8 lookup table for `apply_stretch` β€” the stretch is pointwise."""
    v = np.arange(256, dtype=np.float32)[:, None]
    x = (v - np.asarray(lo, np.float32)) / (np.asarray(hi, np.float32) - lo)
    return (np.clip(x, 0.0, 1.0) * 255.0 + 0.5).astype(np.uint8)


def apply_stretch(rgb_u8: np.ndarray, lo: np.ndarray, hi: np.ndarray) -> np.ndarray:
    if rgb_u8.dtype == np.uint8 and rgb_u8.ndim == 3:
        return apply_stretch_lut(rgb_u8, stretch_lut(lo, hi))
    x = (rgb_u8.astype(np.float32) - lo) / (hi - lo)
    return (np.clip(x, 0.0, 1.0) * 255.0 + 0.5).astype(np.uint8)


def apply_stretch_lut(rgb_u8: np.ndarray, lut: np.ndarray) -> np.ndarray:
    """Apply a `stretch_lut` table β€” a gather over uint8, no float image at all."""
    out = np.empty_like(rgb_u8)
    for ch in range(rgb_u8.shape[-1]):
        np.take(lut[:, ch], rgb_u8[..., ch], out=out[..., ch])
    return out


def stretch_scene(rgb_u8: np.ndarray, spec: PreprocSpec) -> np.ndarray:
    if not spec.radiometric_stretch:
        return rgb_u8
    lo, hi = scene_stretch_bounds(rgb_u8, spec.stretch_lo_pct, spec.stretch_hi_pct)
    return apply_stretch(rgb_u8, lo, hi)


# ---------------------------------------------------------------------
# step 3 β€” geometry
# ---------------------------------------------------------------------
def resize(arr: np.ndarray, out_hw: tuple[int, int], order: str) -> np.ndarray:
    """order: 'bilinear' | 'nearest'.  Handles uint8 RGB, float32 and int labels."""
    from PIL import Image

    h, w = int(out_hw[0]), int(out_hw[1])
    if arr.shape[:2] == (h, w):
        return arr
    resample = Image.BILINEAR if order == "bilinear" else Image.NEAREST
    if arr.ndim == 3:
        return np.asarray(Image.fromarray(arr.astype(np.uint8)).resize((w, h), resample))
    if np.issubdtype(arr.dtype, np.integer):
        return np.asarray(
            Image.fromarray(arr.astype(np.int32), mode="I").resize((w, h), Image.NEAREST)
        ).astype(arr.dtype)
    return np.asarray(
        Image.fromarray(arr.astype(np.float32)).resize((w, h), resample), dtype=np.float32
    )


def gsd_to_shape(h: int, w: int, src_gsd_m: float, dst_gsd_m: float) -> tuple[int, int]:
    s = float(src_gsd_m) / float(dst_gsd_m)
    return max(1, int(round(h * s))), max(1, int(round(w * s)))


def to_canonical(rgb_u8: np.ndarray, src_gsd_m: float, spec: PreprocSpec) -> np.ndarray:
    out_hw = gsd_to_shape(rgb_u8.shape[0], rgb_u8.shape[1], src_gsd_m, spec.canonical_gsd_m)
    return resize(rgb_u8, out_hw, "bilinear")


def round_to_patch(n: int, patch: int, minimum: int) -> int:
    return max(minimum, int(round(n / patch)) * patch)


# ---------------------------------------------------------------------
# step 4 β€” tiling / blending (inference side)
# ---------------------------------------------------------------------
def hann2d(n: int) -> np.ndarray:
    w = np.hanning(n + 2)[1:-1]
    return np.clip(np.outer(w, w), 1e-3, None).astype(np.float32)


def tile_origins(extent: int, tile: int, overlap: float) -> list[int]:
    """Start offsets covering [0, extent) with `tile`-wide windows."""
    if extent <= tile:
        return [0]
    stride = max(1, int(round(tile * (1.0 - overlap))))
    xs = list(range(0, extent - tile + 1, stride))
    if xs[-1] != extent - tile:
        xs.append(extent - tile)
    return xs


# ---------------------------------------------------------------------
# step 1 β€” readers
# ---------------------------------------------------------------------
@dataclass
class SceneMeta:
    """What we know about an input image."""

    gsd_m: float
    gsd_source: str                 # "geotiff" | "user" | "assumed"
    georeferenced: bool = False
    transform: object | None = None  # affine.Affine, when georeferenced
    crs: object | None = None
    width: int = 0
    height: int = 0
    path: str = ""

    def summary(self) -> dict:
        return {
            "path": self.path, "width": self.width, "height": self.height,
            "gsd_m": self.gsd_m, "gsd_source": self.gsd_source,
            "georeferenced": self.georeferenced,
            "crs": str(self.crs) if self.crs is not None else None,
        }


def _metres_per_unit(crs, transform, width: int, height: int) -> float:
    """Convert one transform unit to metres.  Projected CRS -> already metres."""
    if crs is None:
        return 1.0
    try:
        if crs.is_geographic:
            # degrees -> metres at the scene's centre latitude
            import math

            lat = transform.f + transform.e * (height / 2.0)
            return 111_320.0 * max(0.15, math.cos(math.radians(lat)))
        units = (crs.linear_units or "metre").lower()
        if units.startswith(("met", "m")):
            return 1.0
        if units.startswith(("foot", "ft", "us survey")):
            return 0.3048006096012192 if "us" in units else 0.3048
    except Exception:  # noqa: BLE001
        pass
    return 1.0


TIF_SUFFIXES = {".tif", ".tiff", ".gtif", ".gtiff"}


def read_scene(
    path: str | Path,
    user_gsd_m: float = 0.0,
    assumed_gsd_m: float = 0.5,
    max_side: int = 0,
) -> tuple[np.ndarray, SceneMeta]:
    """Load any of PNG / JPG / (Geo)TIFF as (H, W, 3) uint8 + SceneMeta.

    GSD precedence: `user_gsd_m` > GeoTIFF transform > `assumed_gsd_m`.
    `max_side` (optional) downsamples enormous scenes before anything else; the
    reported GSD is scaled to match so metres stay correct.
    """
    path = Path(path)
    rgb, meta = None, None

    if path.suffix.lower() in TIF_SUFFIXES:
        try:
            import rasterio

            with rasterio.open(path) as src:
                idx = [1, 2, 3] if src.count >= 3 else [1] * 3
                rgb = src.read(indexes=idx).transpose(1, 2, 0)
                mpu = _metres_per_unit(src.crs, src.transform, src.width, src.height)
                native = float(abs(src.transform.a)) * mpu
                georef = src.crs is not None and abs(src.transform.a) > 0
                gsd = user_gsd_m if user_gsd_m > 0 else (native if georef else assumed_gsd_m)
                meta = SceneMeta(
                    gsd_m=float(gsd),
                    gsd_source="user" if user_gsd_m > 0 else ("geotiff" if georef else "assumed"),
                    georeferenced=bool(georef), transform=src.transform, crs=src.crs,
                    width=src.width, height=src.height, path=str(path),
                )
        except Exception as e:  # noqa: BLE001
            print(f"[preproc] rasterio read failed ({e}); falling back to PIL")

    if rgb is None:
        from PIL import Image

        Image.MAX_IMAGE_PIXELS = None
        im = Image.open(path).convert("RGB")
        rgb = np.asarray(im)
        meta = SceneMeta(
            gsd_m=float(user_gsd_m if user_gsd_m > 0 else assumed_gsd_m),
            gsd_source="user" if user_gsd_m > 0 else "assumed",
            georeferenced=False, width=im.width, height=im.height, path=str(path),
        )

    rgb = _to_uint8_rgb(rgb)
    meta.height, meta.width = rgb.shape[0], rgb.shape[1]

    if max_side and max(rgb.shape[:2]) > max_side:
        sc = max_side / max(rgb.shape[:2])
        rgb = resize(rgb, (int(rgb.shape[0] * sc), int(rgb.shape[1] * sc)), "bilinear")
        meta.gsd_m /= sc
        meta.height, meta.width = rgb.shape[0], rgb.shape[1]
        print(f"[preproc] downsampled scene to {rgb.shape[1]}x{rgb.shape[0]} "
              f"(gsd now {meta.gsd_m:.3f} m)")
    return rgb, meta


def _to_uint8_rgb(a: np.ndarray) -> np.ndarray:
    if a.ndim == 2:
        a = np.stack([a] * 3, -1)
    a = a[..., :3]
    if a.dtype == np.uint8:
        return np.ascontiguousarray(a)
    a = a.astype(np.float32)
    if a.max() <= 1.5:                       # float 0..1
        a = a * 255.0
    elif a.max() > 255.0:                    # uint16 / radiance
        lo, hi = scene_stretch_bounds(a.astype(np.float32), 1.0, 99.0)
        a = (a - lo) / (hi - lo) * 255.0
    return np.clip(a, 0, 255).astype(np.uint8)


def save_spec(spec: PreprocSpec, path: str | Path) -> None:
    Path(path).write_text(json.dumps(spec.to_dict(), indent=2))