SM737's picture
Upload folder using huggingface_hub (part 4)
a358495 verified
Raw History Blame Contribute Delete
7.24 kB
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import numpy as np
from satquery_engine.models.manifest import ModelManifest
@dataclass(frozen=True)
class AdaptedInput:
tensor: np.ndarray
selected_channels: tuple[int, ...]
audit: dict[str, Any]
@dataclass(frozen=True)
class AdaptedWaterInput:
image_tensor: np.ndarray
spectral_tensor: np.ndarray
valid_mask: np.ndarray
selected_channels: tuple[int, ...]
audit: dict[str, Any]
class GenericModelAdapter:
"""Checkpoint-specific preprocessing contract; never guesses missing bands."""
def __init__(self, manifest: ModelManifest) -> None:
self.manifest = manifest
def preprocess(self, image: np.ndarray, selected_channels: tuple[int, ...]) -> AdaptedInput:
array = np.asarray(image)
if array.ndim != 3:
raise ValueError("Model input must be a channel-first 3-D scientific raster.")
if len(selected_channels) != self.manifest.input_channels:
raise ValueError(
f"{self.manifest.model_id} requires {self.manifest.input_channels} ordered channels; "
f"received {len(selected_channels)}."
)
zero_based = tuple(index - 1 for index in selected_channels)
if min(zero_based, default=0) < 0 or max(zero_based, default=0) >= array.shape[0]:
raise ValueError("Selected channel is outside the raster band range.")
tensor = array[list(zero_based)].astype(self.manifest.input_dtype, copy=False)
norm = self.manifest.normalization
if norm.get("type") == "mean_std":
mean = np.asarray(norm["mean"], dtype="float32")[:, None, None]
std = np.asarray(norm["std"], dtype="float32")[:, None, None]
tensor = (tensor - mean) / np.maximum(std, 1e-7)
elif norm.get("type") == "scale":
tensor = tensor * float(norm.get("factor", 1.0))
elif norm.get("type") not in {None, "identity", "checkpoint_native"}:
raise ValueError(f"Unsupported normalization contract: {norm.get('type')}")
return AdaptedInput(
tensor=tensor,
selected_channels=selected_channels,
audit={
"model_id": self.manifest.model_id,
"adapter": type(self).__name__,
"input_shape": list(array.shape),
"output_shape": list(tensor.shape),
"selected_channels": list(selected_channels),
"normalization": norm,
"dtype": str(tensor.dtype),
},
)
def postprocess(self, output: np.ndarray) -> np.ndarray:
result = np.asarray(output, dtype="float32")
if "probability" in self.manifest.output_type and (result.min(initial=0) < 0 or result.max(initial=0) > 1):
result = 1.0 / (1.0 + np.exp(-result))
return np.clip(result, 0.0, 1.0) if "probability" in self.manifest.output_type else result
# Public architectural name retained while concrete adapters stay explicit.
ModelAdapter = GenericModelAdapter
class DinoBuildingAdapter(GenericModelAdapter):
pass
class FlairLandCoverAdapter(GenericModelAdapter):
pass
class PrithviWaterAdapter(GenericModelAdapter):
pass
class BigEarthNetS2Adapter(GenericModelAdapter):
pass
class BigEarthNetS1Adapter(GenericModelAdapter):
pass
class BigEarthNetFusionAdapter(GenericModelAdapter):
pass
class EarthDialAdapter(GenericModelAdapter):
pass
class RemoteClipAdapter(GenericModelAdapter):
pass
class CromaAdapter(GenericModelAdapter):
pass
class ChangeDetectionAdapter(GenericModelAdapter):
pass
class SatlasBuildingAdapter(GenericModelAdapter):
def preprocess(self, image: np.ndarray, selected_channels: tuple[int, ...]) -> AdaptedInput:
array = np.asarray(image)
if array.ndim != 3 or len(selected_channels) != 3:
raise ValueError("Satlas building inference requires channel-first RGB imagery.")
zero_based = tuple(index - 1 for index in selected_channels)
tensor = array[list(zero_based)].astype("float32", copy=False)
if np.nanmax(tensor, initial=0.0) > 1.5:
tensor = tensor / 255.0
tensor = np.clip(tensor, 0.0, 1.0)
return AdaptedInput(
tensor=tensor,
selected_channels=selected_channels,
audit={
"model_id": self.manifest.model_id,
"adapter": type(self).__name__,
"input_shape": list(array.shape),
"output_shape": list(tensor.shape),
"selected_channels": list(selected_channels),
"normalization": "uint8/255 or identity for unit RGB",
"dtype": str(tensor.dtype),
},
)
class SatlasWaterAdapter(GenericModelAdapter):
def preprocess(self, image: np.ndarray, selected_channels: tuple[int, ...]) -> AdaptedWaterInput:
from satquery_engine.models.satlas_water_net import SOURCE_BANDS, prepare_water_inputs
array = np.asarray(image)
if array.ndim != 3 or len(selected_channels) != len(SOURCE_BANDS):
raise ValueError(
f"Satlas water inference requires {len(SOURCE_BANDS)} exact Sentinel-2 source bands."
)
zero_based = tuple(index - 1 for index in selected_channels)
if min(zero_based) < 0 or max(zero_based) >= array.shape[0]:
raise ValueError("Selected Sentinel-2 channel is outside the raster band range.")
source = {name: array[index] for name, index in zip(SOURCE_BANDS, zero_based)}
backbone, spectral, valid = prepare_water_inputs(source)
return AdaptedWaterInput(
image_tensor=backbone,
spectral_tensor=spectral,
valid_mask=valid,
selected_channels=selected_channels,
audit={
"model_id": self.manifest.model_id,
"adapter": type(self).__name__,
"input_shape": list(array.shape),
"image_tensor_shape": list(backbone.shape),
"spectral_tensor_shape": list(spectral.shape),
"selected_channels": list(selected_channels),
"normalization": self.manifest.normalization,
"dtype": str(backbone.dtype),
},
)
class WaterShadowAdapter(GenericModelAdapter):
pass
BigEarthNetAdapter = BigEarthNetS2Adapter
ChangeModelAdapter = ChangeDetectionAdapter
_ADAPTERS = {cls.__name__: cls for cls in (
GenericModelAdapter, DinoBuildingAdapter, SatlasBuildingAdapter, SatlasWaterAdapter, FlairLandCoverAdapter, PrithviWaterAdapter,
BigEarthNetS2Adapter, BigEarthNetS1Adapter, BigEarthNetFusionAdapter, BigEarthNetAdapter,
EarthDialAdapter, RemoteClipAdapter, CromaAdapter, ChangeDetectionAdapter, ChangeModelAdapter,
WaterShadowAdapter,
)}
def adapter_for(manifest: ModelManifest) -> GenericModelAdapter:
try:
adapter_type = _ADAPTERS[manifest.adapter]
except KeyError as exc:
raise ValueError(f"Unknown adapter {manifest.adapter!r} for {manifest.model_id}") from exc
return adapter_type(manifest)