Spaces:
Paused
Paused
Download satquery_engine/models/adapters/base.py from SM737/satquery-api: direct link, hf CLI and curl.
- Browser
- Download file 7.24 kB
-
https://huggingface.co/spaces/SM737/satquery-api/resolve/main/satquery_engine/models/adapters/base.py
- Command line
-
hf download hf://spaces/SM737/satquery-api/satquery_engine/models/adapters/base.py
-
curl -L -o base.py https://huggingface.co/spaces/SM737/satquery-api/resolve/main/satquery_engine/models/adapters/base.py
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 | |
| class AdaptedInput: | |
| tensor: np.ndarray | |
| selected_channels: tuple[int, ...] | |
| audit: dict[str, Any] | |
| 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) | |