from __future__ import annotations from contextlib import ExitStack, contextmanager from pathlib import Path import tempfile from typing import Any from collections.abc import Iterator from app.core.config import Settings from app.core.errors import AppError class YoloDetectionAdapter: def __init__(self, settings: Settings) -> None: self.settings = settings @staticmethod def dependencies_available() -> bool: try: import torch # noqa: F401 import ultralytics # noqa: F401 except Exception: return False return True def load_model(self, model_path: Path): if not model_path.exists() or not model_path.is_file(): raise AppError( code="DETECTION_MODEL_UNAVAILABLE", message="Configured YOLO model file does not exist", details={"model_path": str(model_path)}, status_code=503, ) if not self.dependencies_available(): raise AppError( code="DETECTION_DEPENDENCY_UNAVAILABLE", message="YOLO dependencies are not installed. Install backend optional extras with geointel-backend[ai].", status_code=503, ) self.validate_runtime() try: from ultralytics import YOLO except ImportError as exc: raise AppError( code="DETECTION_DEPENDENCY_UNAVAILABLE", message="YOLO dependencies are not importable. Install backend optional extras with geointel-backend[ai].", status_code=503, ) from exc try: return YOLO(str(model_path)) except Exception as exc: raise AppError( code="DETECTION_MODEL_LOAD_FAILED", message="Configured YOLO model could not be loaded", details={"model_path": str(model_path)}, status_code=503, ) from exc def validate_runtime(self) -> None: if not self.settings.yolo_require_cuda: return try: import torch except Exception as exc: raise AppError( code="DETECTION_ACCELERATOR_UNAVAILABLE", message="NVIDIA CUDA is required for configured YOLO inference, but PyTorch is not importable.", status_code=503, ) from exc if not torch.cuda.is_available(): raise AppError( code="DETECTION_ACCELERATOR_UNAVAILABLE", message="NVIDIA CUDA is required for configured YOLO inference, but no CUDA device is available.", details={"configured_device": self.settings.yolo_device}, status_code=503, ) if not str(self.settings.yolo_device).lower().startswith(("cuda", "0", "1", "2", "3")): raise AppError( code="DETECTION_ACCELERATOR_MISCONFIGURED", message="NVIDIA CUDA is required, but YOLO_DEVICE does not select a CUDA device.", details={"configured_device": self.settings.yolo_device}, status_code=503, ) def predict_tile(self, model, tile_path: Path, confidence_threshold: float) -> list[dict[str, Any]]: return self.predict_tiles(model, [tile_path], confidence_threshold)[0] def predict_tiles( self, model, tile_paths: list[Path], confidence_threshold: float, ) -> list[list[dict[str, Any]]]: """Run inference over several tiles per GPU call. One ``predict`` call per tile leaves an RTX-class card mostly idle on a run of a hundred tiles. Results are returned per tile, in the order the tiles were given, so the caller can still georeference each detection against its own tile transform. """ for tile_path in tile_paths: if not tile_path.exists() or not tile_path.is_file(): raise AppError( code="DETECTION_TILE_NOT_FOUND", message="Tile referenced by manifest does not exist", details={"tile_path": str(tile_path)}, status_code=422, ) batch_size = max(1, int(self.settings.yolo_batch_size or 1)) detections_per_tile: list[list[dict[str, Any]]] = [] for start in range(0, len(tile_paths), batch_size): batch = tile_paths[start : start + batch_size] with ExitStack() as stack: sources = [stack.enter_context(_prediction_source(path)) for path in batch] try: results = model.predict( source=sources, conf=float(confidence_threshold), imgsz=int(self.settings.yolo_image_size), device=self.settings.yolo_device, max_det=int(self.settings.yolo_max_detections), verbose=False, ) except AppError: raise except Exception as exc: raise AppError( code="DETECTION_INFERENCE_FAILED", message="Configured YOLO inference failed for a raster tile", details={"tile_path": str(batch[0]), "error": str(exc)}, status_code=503, ) from exc results = list(results) for offset in range(len(batch)): result = results[offset] if offset < len(results) else None detections_per_tile.append(_detections_from_result(result)) return detections_per_tile def _detections_from_result(result: Any) -> list[dict[str, Any]]: """Flatten one ultralytics result into the adapter's detection dicts.""" if result is None: return [] names = getattr(result, "names", {}) or {} boxes = getattr(result, "boxes", None) if boxes is None: return [] xyxy_values = _to_list(getattr(boxes, "xyxy", [])) confidence_values = _to_list(getattr(boxes, "conf", [])) class_values = _to_list(getattr(boxes, "cls", [])) detections: list[dict[str, Any]] = [] for index, bbox in enumerate(xyxy_values): class_id = int(class_values[index]) if index < len(class_values) else -1 detections.append( { "class_name": str(names.get(class_id, class_id)), "confidence": float(confidence_values[index]) if index < len(confidence_values) else 0.0, "bbox": [float(value) for value in bbox], "properties": {"class_id": class_id}, } ) return detections def _to_list(value: Any) -> list[Any]: if hasattr(value, "detach"): value = value.detach() if hasattr(value, "cpu"): value = value.cpu() if hasattr(value, "numpy"): value = value.numpy() if hasattr(value, "tolist"): return value.tolist() return list(value) def _rgb_band_indexes(dataset: Any) -> list[int]: """Pick the three bands that carry visible colour, in R, G, B order. Belgian orthophoto tiles are commonly 4-band RGB + near-infrared. Taking bands blindly would feed the detector an infrared channel as if it were colour, so an explicit colour interpretation wins when the raster has one. """ count = int(getattr(dataset, "count", 0) or 0) if count <= 0: raise ValueError("Raster tile has no bands") if count == 1: return [1, 1, 1] try: from rasterio.enums import ColorInterp interpretations = list(getattr(dataset, "colorinterp", ()) or ()) wanted = (ColorInterp.red, ColorInterp.green, ColorInterp.blue) if all(interpretation in interpretations for interpretation in wanted): return [interpretations.index(interpretation) + 1 for interpretation in wanted] except Exception: pass if count == 2: return [1, 1, 1] return [1, 2, 3] def _stretch_to_uint8(data: Any, valid: Any) -> Any: """Scale a (bands, H, W) array to 0-255 with one shared percentile stretch. ``uint8`` data is already display-ready and is passed through untouched; inventing a stretch for it would change pixel values the model was trained on. Anything wider (12-bit and 16-bit orthophotos, float reflectance) would otherwise be truncated to near-black by a plain dtype cast. The stretch bounds are computed over all bands together, not per band. A per-band stretch white-balances the tile and shifts every hue, while the detector learned on ordinary RGB orthophotos. """ import numpy as np if data.dtype == np.uint8: return data if valid is not None and valid.any(): sample = data[:, valid].reshape(-1) else: sample = data.reshape(-1) if sample.size == 0: return np.zeros(data.shape, dtype=np.uint8) low, high = (float(value) for value in np.percentile(sample.astype("float64"), (2.0, 98.0))) if not high > low: low, high = float(sample.min()), float(sample.max()) if not high > low: return np.full(data.shape, 0 if low == 0 else 255, dtype=np.uint8) scaled = (data.astype("float64") - low) * (255.0 / (high - low)) return np.clip(scaled, 0.0, 255.0).astype(np.uint8) def _read_tile_as_rgb(tile_path: Path) -> Any: """Read a raster tile into an (H, W, 3) uint8 array fit for inference.""" import numpy as np import rasterio with rasterio.open(tile_path) as dataset: indexes = _rgb_band_indexes(dataset) raw = dataset.read(indexes, masked=True) data = np.ma.getdata(raw) mask = np.ma.getmaskarray(raw) valid = ~mask.any(axis=0) rgb = np.moveaxis(_stretch_to_uint8(data, valid), 0, -1) # Nodata collars stay black instead of dragging the stretch toward zero. rgb = np.ascontiguousarray(rgb) rgb[~valid] = 0 return rgb @contextmanager def _prediction_source(tile_path: Path) -> Iterator[str]: """Yield a path to an 8-bit RGB rendering of ``tile_path`` for the model.""" temp_path: Path | None = None try: try: from PIL import Image except Exception: yield str(tile_path) return try: rgb = _read_tile_as_rgb(tile_path) except Exception: rgb = None if rgb is not None: with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as handle: temp_path = Path(handle.name) Image.fromarray(rgb).save(temp_path) yield str(temp_path) return # rasterio is unavailable or cannot read this file (a plain PNG/JPEG # fixture, for instance). Fall back to the previous PIL handling. try: with Image.open(tile_path) as image: if image.mode == "RGB" and len(image.getbands()) == 3: yield str(tile_path) return rgb_image = image.convert("RGB") with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as handle: temp_path = Path(handle.name) rgb_image.save(temp_path) yield str(temp_path) return except Exception: if temp_path is not None: raise yield str(tile_path) return finally: if temp_path is not None: temp_path.unlink(missing_ok=True)