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
318 lines
11 KiB
Python
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)
|