Files
geointel/backend/app/services/yolo_adapter.py
T
Jens faeb58ef6d
GeoIntel release gates / Compile, test, contracts and builds (push) Successful in 1m49s
GeoIntel release gates / Python and npm vulnerability policy (push) Successful in 21s
GeoIntel release gates / Production AI image, SBOM and container scan (push) Successful in 5m39s
GeoIntel release gates / Deploy exact gated revision to Unraid (push) Failing after 58m43s
Initial public release
2026-08-31 21:56:53 +02:00

318 lines
11 KiB
Python

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)