from __future__ import annotations import importlib.util import hashlib import json import struct from pathlib import Path from typing import Any import yaml from satquery_engine.models.adapters import adapter_for from satquery_engine.models.manifest import HealthStatus, ModelManifest _CONTRACTS: dict[str, dict[str, Any]] = { "water_s2_surface": dict(adapter="GenericModelAdapter", required_files=("model.pth", "config.json"), required_packages=("torch", "segmentation_models_pytorch"), expected_resolution="Named six-band Sentinel-2, 8-30 m", normalization={"type":"scale", "factor":1.0/255.0}, input_dtype="float32", input_range="publisher-native Sentinel-2 DN"), "land_deepness": dict(adapter="GenericModelAdapter", required_files=("deeplabv3_landcover_4c.onnx",), required_packages=("onnxruntime",), expected_resolution="0.15-0.60 m aerial RGB", normalization={"type":"scale", "factor":1.0/255.0}, input_dtype="float32", input_range="RGB 0..255"), "building_satellite": dict(adapter="GenericModelAdapter", required_files=("model.safetensors", "config.json", "preprocessor_config.json"), required_packages=("torch", "transformers"), expected_resolution="VHR satellite RGB, <=1.5 m when known", normalization={"type":"checkpoint_native", "processor":"RfDetrImageProcessor"}, input_dtype="float32", input_range="RGB 0..255; production processor rescales and normalizes"), "satlas_aerial_swinb_si": dict(adapter="SatlasBuildingAdapter", required_files=("aerial_swinb_si.pth", "satlas_metadata.json"), required_packages=("torch", "satlaspretrain_models"), expected_resolution="0.5-2 m aerial RGB", normalization={"type":"scale", "factor":1.0/255.0}, input_dtype="float32", input_range="0..255 RGB"), "land_flair_hub": dict(adapter="FlairLandCoverAdapter", required_files=("FLAIR-HUB_LC-A_RGB_swinbase-upernet.safetensors", "configs_train/config_models.yaml", "configs_train/config_modalities.yaml", "configs_train/config_supervision.yaml"), required_packages=("torch", "segmentation_models_pytorch", "safetensors"), expected_resolution="FLAIR-HUB high-resolution aerial RGB", normalization={"type":"mean_std", "mean":[105.66,111.35,102.18], "std":[52.23,45.62,44.30]}, input_dtype="float32", input_range="0..255 RGB"), "water_shadow_rgb": dict(adapter="WaterShadowAdapter", required_files=("satquery_water_shadow_config.json", "satquery_water_shadow_state_dict.pt"), required_packages=("torch",), expected_resolution="RGB VHR", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="checkpoint-native RGB"), "landcover_primary": dict(adapter="FlairLandCoverAdapter", required_files=("best_checkpoint.pt", "class_map.json"), required_packages=("torch", "segmentation_models_pytorch"), expected_resolution="RGB VHR", normalization={"type":"scale", "factor":1.0/255.0}, input_dtype="float32", input_range="0..1 RGB"), "building_fallback": dict(adapter="DinoBuildingAdapter", required_files=("onnx/model.onnx", "model.ckpt", "README.md"), required_packages=("onnxruntime",), expected_resolution="VHR aerial or satellite RGB", normalization={"type":"mean_std", "mean":[0.4296737853,0.4001659668,0.3433337280], "std":[0.2056069389,0.1673855556,0.1598986423]}, input_dtype="float32", input_range="0..1 RGB"), "building_primary": dict(adapter="DinoBuildingAdapter", required_files=("onnx/model.onnx",), required_packages=("onnxruntime",), expected_resolution="0.2-0.75 m VHR RGB", normalization={"type":"mean_std","mean":[0.485,0.456,0.406],"std":[0.229,0.224,0.225]}, input_dtype="float32", input_range="0..1 reflectance/RGB"), "building_secondary": dict(adapter="DinoBuildingAdapter", required_files=("onnx/model_quantized.onnx",), required_packages=("onnxruntime",), expected_resolution="VHR RGB", normalization={"type":"mean_std","mean":[0.485,0.456,0.406],"std":[0.229,0.224,0.225]}, input_dtype="float32", input_range="0..1 RGB"), "water_finetuned": dict( adapter="SatlasWaterAdapter", required_files=("satquery_water_config.json", "satquery_water_state_dict.pt", "thresholds.json"), required_packages=("torch", "torchvision"), expected_resolution="Sentinel-2 10/20 m on a co-registered grid", normalization={ "type": "bundle_contract", "rgb": "clip(raw/3000,0,1)", "red_edge_nir_swir": "clip(raw/8160,0,1)", "source_scale": 10000, }, input_dtype="float32", input_range="Sentinel-2 digital numbers (nominally 0..10000)", ), "land_rgb": dict(adapter="FlairLandCoverAdapter", required_files=("FLAIR-INC_rgb_15cl_resnet34-deeplabv3_weights.pth",), required_packages=("torch","segmentation_models_pytorch"), expected_resolution="0.2-0.5 m aerial RGB", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="checkpoint-native RGB"), "land_s2": dict(adapter="BigEarthNetS2Adapter", required_files=("model.safetensors","config.json"), required_packages=("torch","transformers"), expected_resolution="Sentinel-2 10/20 m", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="Sentinel-2 reflectance"), "land_s1": dict(adapter="BigEarthNetS1Adapter", required_files=("model.safetensors","config.json"), required_packages=("torch","transformers"), expected_resolution="Sentinel-1 GRD", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="calibrated SAR"), "land_s1s2": dict(adapter="BigEarthNetFusionAdapter", required_files=("model.safetensors","config.json"), required_packages=("torch","transformers"), expected_resolution="co-registered Sentinel-1 and Sentinel-2", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="calibrated S1/S2"), "water_s2": dict(adapter="PrithviWaterAdapter", required_files=("Prithvi-EO-V2-300M-TL-Sen1Floods11.pt","config.yaml"), required_packages=("torch","terratorch","transformers"), expected_resolution="Sentinel-2 10/20 m", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="Sentinel-2 reflectance"), "earthdial": dict(adapter="EarthDialAdapter", required_files=("model.safetensors.index.json","config.json","tokenizer.model"), required_packages=("torch","transformers"), expected_resolution="rendered RGB or documented multispectral variant", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="display RGB only for RGB variant"), "remoteclip_rn50": dict(adapter="RemoteClipAdapter", required_files=("RemoteCLIP-RN50.pt",), required_packages=("torch","open_clip"), expected_resolution="rendered remote-sensing RGB", normalization={"type":"mean_std","mean":[0.48145466,0.4578275,0.40821073],"std":[0.26862954,0.26130258,0.27577711]}, input_dtype="float32", input_range="0..1 display RGB"), "remoteclip_vit_b_32": dict(adapter="RemoteClipAdapter", required_files=("RemoteCLIP-ViT-B-32.pt",), required_packages=("torch","open_clip"), expected_resolution="rendered remote-sensing RGB", normalization={"type":"mean_std","mean":[0.48145466,0.4578275,0.40821073],"std":[0.26862954,0.26130258,0.27577711]}, input_dtype="float32", input_range="0..1 display RGB"), "croma_base": dict(adapter="CromaAdapter", required_files=("CROMA_base.pt",), required_packages=("torch","croma"), expected_resolution="co-registered Sentinel-1/2", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="calibrated S1/S2"), "tinycd": dict(adapter="ChangeDetectionAdapter", required_files=("levir_best.pth",), required_packages=("torch",), expected_resolution="registered VHR RGB temporal pair", normalization={"type":"scale","factor":1.0}, input_dtype="float32", input_range="0..1 RGB"), "bit": dict(adapter="ChangeDetectionAdapter", required_files=("bit_r18_256x256_40k_levircd.pth",), required_packages=("torch","mmengine","mmseg"), expected_resolution="registered VHR RGB temporal pair", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="0..1 RGB"), "changerex": dict(adapter="ChangeDetectionAdapter", required_files=("ChangerEx_r18-512x512_40k_levircd.pth",), required_packages=("torch","mmengine","mmseg"), expected_resolution="registered VHR RGB temporal pair", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="0..1 RGB"), "ban": dict(adapter="ChangeDetectionAdapter", required_files=("ban_vit-l14-clip_mit-b0_512x512_40k_levircd.pth",), required_packages=("torch","mmengine","mmseg"), expected_resolution="registered VHR RGB temporal pair", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="0..1 RGB"), "changemamba": dict(adapter="ChangeDetectionAdapter", required_files=("changemamba*.pth",), required_packages=("torch","selective_scan"), expected_resolution="registered VHR RGB temporal pair", normalization={"type":"checkpoint_native"}, input_dtype="float32", input_range="0..1 RGB"), } class ModelHealthCheck: """File, adapter, dependency, and lightweight checkpoint validation.""" @staticmethod def _resolve_required(base: Path, pattern: str) -> Path | None: direct = base / pattern if direct.is_file(): return direct matches = [p for p in base.glob(pattern) if p.is_file()] return matches[0] if matches else None @staticmethod def _looks_like_lfs_pointer(path: Path) -> bool: if path.stat().st_size > 1024: return False return path.read_bytes()[:100].startswith(b"version https://git-lfs.github.com/spec/v1") @staticmethod def _check_safetensors(path: Path) -> None: with path.open("rb") as stream: header_len = struct.unpack(" path.stat().st_size - 8: raise ValueError("invalid safetensors header length") json.loads(stream.read(header_len)) def check(self, manifest: ModelManifest) -> ModelManifest: if not manifest.enabled: return manifest.with_health(HealthStatus.DISABLED, False, reason="disabled by manifest") if not manifest.local_path.exists(): return manifest.with_health(HealthStatus.MISSING, False, reason="model path does not exist") if not manifest.required_files: return manifest.with_health(HealthStatus.DEGRADED, False, reason="no validated checkpoint file contract is registered") missing: list[str] = [] files: list[Path] = [] for pattern in manifest.required_files: resolved = self._resolve_required(manifest.local_path, pattern) if manifest.local_path.is_dir() else manifest.local_path if resolved is None: missing.append(pattern) else: files.append(resolved) if missing: return manifest.with_health(HealthStatus.MISSING, False, missing_files=missing) try: adapter_for(manifest) for path in files: if not path.stat().st_size or self._looks_like_lfs_pointer(path): raise ValueError(f"{path.name} is empty or an unresolved Git LFS pointer") if path.suffix == ".safetensors": self._check_safetensors(path) elif path.suffix == ".onnx": import onnxruntime as ort # Avoid each startup graph allocating a full-machine thread # pool while model inference and API work share this host. options = ort.SessionOptions() options.intra_op_num_threads = 4 options.inter_op_num_threads = 1 session = ort.InferenceSession(str(path), sess_options=options, providers=["CPUExecutionProvider"]) if not session.get_inputs() or not session.get_outputs(): raise ValueError("ONNX graph has no inputs or outputs") if manifest.metadata.get("registry_key") == "building_fallback": # Health discovery must not execute an unused CPU model. # Actual inference validates finite outputs in buildings.py. expected = [1, 3, 256, 256] if session.get_inputs()[0].shape != expected or session.get_outputs()[0].shape != expected: raise ValueError("DINO fallback graph has an invalid schema") elif path.name.endswith("index.json"): payload = json.loads(path.read_text(encoding="utf-8")) shards = set(payload.get("weight_map", {}).values()) if not shards or any(not (path.parent / shard).is_file() for shard in shards): raise ValueError("checkpoint index references missing shards") elif path.suffix == ".json": payload = json.loads(path.read_text(encoding="utf-8")) if path.name == "satquery_buildings_config.json" and payload.get("architecture") != "SatlasBuildingNet": raise ValueError("building bundle architecture is incompatible") if path.name == "satquery_water_config.json" and payload.get("architecture") != "SatlasWaterNet": raise ValueError("water bundle architecture is incompatible") elif path.suffix in {".pt", ".pth", ".ckpt"}: with path.open("rb") as stream: if len(stream.read(16)) < 16: raise ValueError("checkpoint header is truncated") if path.name in {"satquery_buildings_state_dict.pt", "satquery_water_state_dict.pt"}: if importlib.util.find_spec("torch") is None: continue import torch state = torch.load(path, map_location="cpu", weights_only=True, mmap=True) if not isinstance(state, dict): raise ValueError("fine-tuned checkpoint is not a state dictionary") patch_key = "backbone.backbone.backbone.features.0.0.weight" head_key = "head.weight" if patch_key not in state or head_key not in state: raise ValueError("fine-tuned checkpoint is missing required model layers") expected_channels = 3 if "buildings" in path.name else 9 if tuple(state[patch_key].shape) != (128, expected_channels, 4, 4): raise ValueError("fine-tuned checkpoint input layer does not match its bundle contract") if tuple(state[head_key].shape[:2]) not in {(2, 64), (2, 96)}: raise ValueError("fine-tuned checkpoint output head does not match its bundle contract") if "water" in path.name: spectral_key = "spectral.net.0.weight" if spectral_key not in state or tuple(state[spectral_key].shape) != (32, 6, 3, 3): raise ValueError("water checkpoint spectral branch is incompatible") del state except Exception as exc: return manifest.with_health(HealthStatus.CORRUPT, False, reason=str(exc)) missing_packages = [name for name in manifest.required_packages if importlib.util.find_spec(name) is None] if missing_packages: return manifest.with_health( HealthStatus.DEGRADED, False, reason="runtime dependencies are missing", missing_packages=missing_packages, files_verified=[str(p) for p in files], preprocessing_verified=True, ) key = manifest.metadata.get("registry_key") if key == "building_satellite": from satquery_engine.services.buildings_rf import MODEL_SHA256 with (manifest.local_path / "model.safetensors").open("rb") as stream: digest = hashlib.file_digest(stream, "sha256").hexdigest() if digest != MODEL_SHA256: return manifest.with_health(HealthStatus.CORRUPT, False, reason="RF-DETR checkpoint identity mismatch") if key in {"land_s2", "land_s1", "land_s1s2"}: return manifest.with_health(HealthStatus.DEGRADED, False, reason="checkpoint structure verified; ConfigILM runtime and compatible sensor inference are not integrated") if key in {"satlas_aerial_swinb_si", "land_flair_hub"} or (key in {"building_primary", "water_finetuned"} and manifest.adapter in {"SatlasBuildingAdapter", "SatlasWaterAdapter"}): proof_path = manifest.local_path / "satquery_model_health.json" checkpoint = next((p for p in files if p.suffix in {".pth", ".pt", ".safetensors"}), None) if checkpoint is None or not proof_path.is_file(): return manifest.with_health(HealthStatus.DEGRADED, False, reason="smoke inference has not been verified") try: proof = json.loads(proof_path.read_text(encoding="utf-8")) with checkpoint.open("rb") as stream: digest = hashlib.file_digest(stream, "sha256").hexdigest() if proof.get("checkpoint_sha256") != digest or proof.get("inference_verified") is not True: raise ValueError("health proof does not match the checkpoint") if key == "land_flair_hub" and proof.get("output_shape") != [1, 19, 512, 512]: raise ValueError("FLAIR-HUB smoke output has the wrong class or spatial schema") if key == "satlas_aerial_swinb_si" and proof.get("checkpoint_id") != "Aerial_SwinB_SI": raise ValueError("Satlas checkpoint ID mismatch") if key == "building_primary" and proof.get("output_shape") != [1, 2, 512, 512]: raise ValueError("building task-head schema mismatch") if key == "water_finetuned" and proof.get("output_shape") != [1, 2, 512, 512]: raise ValueError("Sentinel-2 water task-head schema mismatch") except (OSError, ValueError, KeyError, json.JSONDecodeError) as exc: return manifest.with_health(HealthStatus.DEGRADED, False, reason=str(exc)) return manifest.with_health( HealthStatus.READY, True, files_verified=[str(p) for p in files], preprocessing_verified=True, smoke_test="graph_load" if any(p.suffix == ".onnx" for p in files) else "checkpoint_structure", ) _PROCESS_CACHE: dict[Path, dict[str, ModelManifest]] = {} class LocalModelRegistry: def __init__(self, model_root: Path, manifest_path: Path | None = None) -> None: self.model_root = model_root.resolve() self.manifest_path = manifest_path or self.model_root / "manifests" / "models.yaml" self._manifests: dict[str, ModelManifest] | None = None def load(self, refresh: bool = False) -> dict[str, ModelManifest]: if not refresh and self._manifests is not None: return self._manifests if not refresh and self.manifest_path in _PROCESS_CACHE: self._manifests = _PROCESS_CACHE[self.manifest_path] return self._manifests if not self.manifest_path.is_file(): self._manifests = {} return self._manifests payload = yaml.safe_load(self.manifest_path.read_text(encoding="utf-8")) or {} entries = payload.get("models", payload) manifests: dict[str, ModelManifest] = {} for key, raw in entries.items(): if not isinstance(raw, dict): continue contract = _CONTRACTS.get(key, {}) raw_path = Path(str(raw.get("path", key))) candidate_path = raw_path if raw_path.is_absolute() else self.model_root / raw_path if (candidate_path.is_dir() and (candidate_path / "satquery_buildings_config.json").is_file()) or "satquery_buildings_bundle" in str(raw_path): contract = dict( adapter="SatlasBuildingAdapter", required_files=("satquery_buildings_config.json", "satquery_buildings_state_dict.pt", "thresholds.json"), required_packages=("torch",), expected_resolution="0.2-0.75 m VHR RGB", normalization={"type": "scale", "factor": 1.0 / 255.0}, input_dtype="float32", input_range="0..1 RGB", ) if key != "water_shadow_rgb" and ( (candidate_path.is_dir() and (candidate_path / "satquery_water_config.json").is_file()) or "satquery_water_bundle" in str(raw_path) ): contract = dict(_CONTRACTS["water_finetuned"]) if (candidate_path.is_dir() and (candidate_path / "best_checkpoint.pt").is_file()) or "satquery_landcover_v1" in str(raw_path): contract = dict( adapter="FlairLandCoverAdapter", required_files=("best_checkpoint.pt", "class_map.json"), required_packages=("torch", "segmentation_models_pytorch"), expected_resolution="RGB VHR", normalization={"type": "scale", "factor": 1.0 / 255.0}, input_dtype="float32", input_range="0..1 RGB", ) if raw_path.is_absolute(): try: raw_path.resolve().relative_to(self.model_root) local_path = raw_path.resolve() except ValueError: local_path = self.model_root / key else: local_path = self.model_root / raw_path tasks = raw.get("task", ()) if isinstance(tasks, str): tasks = (tasks,) modalities = raw.get("modality", ()) if isinstance(modalities, str): modalities = (modalities,) bands = tuple(str(item) for item in raw.get("expected_channels", ())) input_size = tuple(int(item) for item in raw.get("input_size", ()) if isinstance(item, (int, float))) input_channels = len(bands) or (input_size[-1] if len(input_size) == 3 else 0) manifests[key] = ModelManifest( model_id=str(raw.get("model_id", key)), local_path=local_path, task=tuple(tasks), modality=tuple(modalities), expected_bands=bands, expected_band_order=bands, input_channels=input_channels, expected_resolution=str(contract.get("expected_resolution", raw.get("notes", {}).get("compatibility", "documented checkpoint domain"))), input_dtype=str(contract.get("input_dtype", "float32")), input_range=str(contract.get("input_range", "checkpoint-native")), normalization=dict(contract.get("normalization", {"type":"checkpoint_native"})), input_size=input_size, output_type=str(raw.get("output_type", "provider_output")), device="auto", precision="float32", adapter=str(contract.get("adapter", "GenericModelAdapter")), required_files=tuple(contract.get("required_files", ())), required_packages=tuple(contract.get("required_packages", ())), metadata={ "registry_key": key, "priority": raw.get("priority"), "display_name": raw.get("display_name", raw.get("model_id", key)), "base_model": raw.get("base_model"), "version": raw.get("version"), "training_dataset": raw.get("training_dataset"), "threshold_file": raw.get("threshold_file"), "policy": raw.get("policy", raw.get("priority")), }, ) checker = ModelHealthCheck() self._manifests = {key: checker.check(manifest) for key, manifest in manifests.items()} _PROCESS_CACHE[self.manifest_path] = self._manifests return self._manifests def get(self, key_or_id: str) -> ModelManifest | None: target = key_or_id.lower() for key, manifest in self.load().items(): if target in {key.lower(), manifest.model_id.lower(), manifest.model_id.rsplit("/", 1)[-1].lower()}: return manifest return None def dashboard(self) -> list[dict[str, Any]]: return [manifest.to_dict() for manifest in self.load().values()]