akashch1512's picture
First Push
16ce72c
Raw History Blame Contribute Delete
15.1 kB
"""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))