Spaces:
Running on Zero
Running on Zero
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))
|