loko/loko/sign/ai/detector.py

499 lines
20 KiB
Python

"""
Service de détection et reconnaissance de panneaux de signalisation routière.
Combine YOLOv10-n (détection d'objets sans NMS sous ONNX Runtime)
et PaddleOCR / RapidOCR (lecture de texte de panonceaux sous ONNX Runtime).
"""
import os
import io
import time
import base64
import logging
import urllib.request
from pathlib import Path
from typing import Dict, Any, List, Optional, Tuple, Union
import cv2
import numpy as np
from PIL import Image, ImageOps
from django.conf import settings
from .catalog import SIGN_CATALOG, get_svg_url, match_sign_from_ocr
logger = logging.getLogger(__name__)
# URL officielle de téléchargement du modèle YOLOv10-n ONNX
YOLOV10N_ONNX_URL = "https://github.com/THU-MIG/yolov10/releases/download/v1.1/yolov10n.onnx"
# Noms des classes COCO pour référence
COCO_CLASSES = {
0: "person", 1: "bicycle", 2: "car", 3: "motorcycle", 5: "bus", 7: "truck",
9: "traffic light", 11: "stop sign", 12: "parking meter", 13: "bench",
}
def letterbox(
im: np.ndarray,
new_shape: Tuple[int, int] = (640, 640),
color: Tuple[int, int, int] = (114, 114, 114),
auto: bool = False,
scaleup: bool = True,
stride: int = 32
) -> Tuple[np.ndarray, float, Tuple[float, float]]:
"""Redimensionne et applique un padding (letterboxing) pour l'inférence YOLO."""
shape = im.shape[:2] # [hauteur, largeur]
if isinstance(new_shape, int):
new_shape = (new_shape, new_shape)
# Ratio d'échelle (nouveau / ancien)
r = min(new_shape[0] / shape[0], new_shape[1] / shape[1])
if not scaleup:
r = min(r, 1.0)
# Calcul du padding
new_unpad = int(round(shape[1] * r)), int(round(shape[0] * r))
dw, dh = new_shape[1] - new_unpad[0], new_shape[0] - new_unpad[1]
dw /= 2
dh /= 2
if shape[::-1] != new_unpad:
im = cv2.resize(im, new_unpad, interpolation=cv2.INTER_LINEAR)
top, bottom = int(round(dh - 0.1)), int(round(dh + 0.1))
left, right = int(round(dw - 0.1)), int(round(dw + 0.1))
im = cv2.copyMakeBorder(im, top, bottom, left, right, cv2.BORDER_CONSTANT, value=color)
return im, r, (dw, dh)
class SignDetectionService:
"""
Service Singleton pour l'inférence IA des panneaux de signalisation.
Initialise paresseusement les modèles ONNX Runtime pour économiser les ressources.
"""
_instance: Optional["SignDetectionService"] = None
def __init__(self):
self._yolo_session = None
self._ocr_engine = None
self.model_dir = getattr(
settings,
"SIGN_AI_MODEL_DIR",
Path(settings.BASE_DIR) / "sign" / "ai" / "models"
)
self.yolo_model_path = Path(self.model_dir) / "yolov10n.onnx"
@classmethod
def get_instance(cls) -> "SignDetectionService":
if cls._instance is None:
cls._instance = cls()
return cls._instance
def _ensure_yolo_model(self) -> Path:
"""Vérifie la présence du fichier modèle ONNX et le télécharge si besoin."""
os.makedirs(self.model_dir, exist_ok=True)
if not self.yolo_model_path.exists() or self.yolo_model_path.stat().st_size < 1000:
logger.info("Téléchargement du modèle YOLOv10-n ONNX depuis %s...", YOLOV10N_ONNX_URL)
req = urllib.request.Request(YOLOV10N_ONNX_URL, headers={"User-Agent": "Mozilla/5.0 Loko-AI"})
with urllib.request.urlopen(req, timeout=30) as resp, open(self.yolo_model_path, "wb") as f:
f.write(resp.read())
logger.info("Modèle YOLOv10-n téléchargé avec succès (%d octets)", self.yolo_model_path.stat().st_size)
return self.yolo_model_path
def get_yolo_session(self):
"""Retourne la session ONNX Runtime pour YOLOv10."""
if self._yolo_session is None:
import onnxruntime as ort
model_path = self._ensure_yolo_model()
sess_options = ort.SessionOptions()
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
sess_options.intra_op_num_threads = 4
self._yolo_session = ort.InferenceSession(
str(model_path),
sess_options=sess_options,
providers=["CPUExecutionProvider"]
)
logger.info("Session YOLOv10-n ONNX initialisée sur CPU.")
return self._yolo_session
def get_ocr_engine(self):
"""Retourne le moteur RapidOCR / PaddleOCR ONNX."""
if self._ocr_engine is None:
from rapidocr_onnxruntime import RapidOCR
self._ocr_engine = RapidOCR()
logger.info("Moteur RapidOCR initialisé avec succès.")
return self._ocr_engine
def load_image(self, image_input: Union[str, bytes, io.BytesIO, Image.Image]) -> Tuple[np.ndarray, Image.Image]:
"""Charge une image, applique la rotation EXIF et retourne (cv2_bgr, pil_image)."""
if isinstance(image_input, np.ndarray):
cv2_img = image_input
pil_img = Image.fromarray(cv2.cvtColor(image_input, cv2.COLOR_BGR2RGB))
return cv2_img, pil_img
elif isinstance(image_input, Image.Image):
pil_img = image_input
elif isinstance(image_input, (bytes, bytearray)):
pil_img = Image.open(io.BytesIO(image_input))
elif isinstance(image_input, io.BytesIO):
pil_img = Image.open(image_input)
elif isinstance(image_input, (str, Path)):
pil_img = Image.open(str(image_input))
else:
raise ValueError(f"Type d'image non supporté: {type(image_input)}")
# Correction automatique de l'orientation selon les tags EXIF de l'appareil photo
pil_img = ImageOps.exif_transpose(pil_img)
if pil_img.mode != "RGB":
pil_img = pil_img.convert("RGB")
cv2_img = cv2.cvtColor(np.array(pil_img), cv2.COLOR_RGB2BGR)
return cv2_img, pil_img
def _get_model_class_names(self, session) -> Dict[int, str]:
"""Extrait les noms de classes depuis les métadonnées ONNX ou utilise le catalogue par défaut."""
try:
import ast
meta = session.get_modelmeta().custom_metadata_map
if 'names' in meta:
names_raw = meta['names']
parsed = ast.literal_eval(names_raw)
return {int(k): str(v) for k, v in parsed.items()}
except Exception:
pass
return {0: "sign", 1: "sub_plate"}
def detect_yolo_boxes(
self,
cv2_img: np.ndarray,
confidence_threshold: float = 0.25
) -> Tuple[List[Dict[str, Any]], float]:
"""
Exécute le modèle ONNX (YOLOv10 end-to-end ou YOLOv8/v11) et retourne les détections filtrées.
"""
start_time = time.perf_counter()
session = self.get_yolo_session()
orig_h, orig_w = cv2_img.shape[:2]
letterbox_img, ratio, (dw, dh) = letterbox(cv2_img, (640, 640))
# Conversion BGR -> RGB et normalisation [0, 1]
rgb_img = cv2.cvtColor(letterbox_img, cv2.COLOR_BGR2RGB)
input_tensor = rgb_img.transpose((2, 0, 1)).astype(np.float32) / 255.0
input_tensor = np.expand_dims(input_tensor, axis=0) # (1, 3, 640, 640)
# Inférence ONNX
input_name = session.get_inputs()[0].name
outputs = session.run(None, {input_name: input_tensor})
raw = outputs[0]
names_dict = self._get_model_class_names(session)
results = []
if raw.ndim == 3 and raw.shape[2] == 6:
# Format YOLOv10 NMS-free : shape (1, 300, 6) -> [x1, y1, x2, y2, score, class_id]
for det in raw[0]:
x1, y1, x2, y2, score, cls_id = det
if score < confidence_threshold:
continue
x1 = (x1 - dw) / ratio
y1 = (y1 - dh) / ratio
x2 = (x2 - dw) / ratio
y2 = (y2 - dh) / ratio
x1 = max(0, min(orig_w - 1, int(round(x1))))
y1 = max(0, min(orig_h - 1, int(round(y1))))
x2 = max(0, min(orig_w - 1, int(round(x2))))
y2 = max(0, min(orig_h - 1, int(round(y2))))
if (x2 - x1) < 10 or (y2 - y1) < 10:
continue
cls_int = int(cls_id)
cls_name = names_dict.get(cls_int, COCO_CLASSES.get(cls_int, f"class_{cls_int}"))
results.append({
"bbox": [x1, y1, x2, y2],
"confidence": float(round(score, 3)),
"class_id": cls_int,
"class_name": cls_name,
})
elif raw.ndim == 3:
# Format YOLOv8 / YOLOv11 standard : shape (1, 4 + C, 8400)
pred = raw[0]
if pred.shape[0] < pred.shape[1]:
pred = pred.T # Transposition vers (8400, 4 + C)
boxes = pred[:, :4]
scores = pred[:, 4:]
class_ids = np.argmax(scores, axis=1)
confidences = np.max(scores, axis=1)
mask = confidences >= confidence_threshold
boxes = boxes[mask]
confidences = confidences[mask]
class_ids = class_ids[mask]
if len(boxes) > 0:
cx, cy, bw, bh = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]
x1 = ((cx - bw / 2.0) - dw) / ratio
y1 = ((cy - bh / 2.0) - dh) / ratio
w_box = bw / ratio
h_box = bh / ratio
boxes_for_nms = []
for i in range(len(x1)):
bx = max(0, min(orig_w - 1, int(round(x1[i]))))
by = max(0, min(orig_h - 1, int(round(y1[i]))))
bw_int = max(10, min(orig_w - bx, int(round(w_box[i]))))
bh_int = max(10, min(orig_h - by, int(round(h_box[i]))))
boxes_for_nms.append([bx, by, bw_int, bh_int])
indices = cv2.dnn.NMSBoxes(boxes_for_nms, confidences.tolist(), confidence_threshold, 0.45)
for idx in indices:
if isinstance(idx, (list, tuple, np.ndarray)):
idx = idx[0]
bx, by, bw_int, bh_int = boxes_for_nms[idx]
cls_int = int(class_ids[idx])
score = float(confidences[idx])
cls_name = names_dict.get(cls_int, f"class_{cls_int}")
results.append({
"bbox": [bx, by, bx + bw_int, by + bh_int],
"confidence": float(round(score, 3)),
"class_id": cls_int,
"class_name": cls_name,
})
elapsed_ms = (time.perf_counter() - start_time) * 1000.0
return results, elapsed_ms
def extract_ocr_from_crop(
self,
cv2_img: np.ndarray,
bbox: List[int]
) -> Tuple[str, List[Dict[str, Any]], float]:
"""Extrait le texte via RapidOCR sur la zone découpée."""
start_time = time.perf_counter()
x1, y1, x2, y2 = bbox
h, w = cv2_img.shape[:2]
# Marge de sécurité (padding 5%)
pad_x = int((x2 - x1) * 0.05)
pad_y = int((y2 - y1) * 0.05)
crop_x1 = max(0, x1 - pad_x)
crop_y1 = max(0, y1 - pad_y)
crop_x2 = min(w, x2 + pad_x)
crop_y2 = min(h, y2 + pad_y)
crop = cv2_img[crop_y1:crop_y2, crop_x1:crop_x2]
if crop.size == 0:
return "", [], 0.0
ocr_engine = self.get_ocr_engine()
ocr_result, _ = ocr_engine(crop)
lines = []
full_text_parts = []
if ocr_result:
for item in ocr_result:
# item: [box_points, text, score]
text = str(item[1]).strip() if len(item) > 1 else ""
try:
score = round(float(item[2]), 3) if len(item) > 2 else 1.0
except (ValueError, TypeError):
score = 1.0
if text:
full_text_parts.append(text)
lines.append({"text": text, "confidence": score})
full_text = " ".join(full_text_parts)
elapsed_ms = (time.perf_counter() - start_time) * 1000.0
return full_text, lines, elapsed_ms
def analyze_image(
self,
image_input: Union[str, bytes, io.BytesIO, Image.Image],
confidence_threshold: float = 0.25
) -> Dict[str, Any]:
"""
Pipeline complet d'analyse d'une image de signalisation.
Retourne les panneaux détectés, le texte OCR, l'ordre vertical, les SVGs associés
et l'image annotée en base64.
"""
total_start = time.perf_counter()
cv2_img, pil_img = self.load_image(image_input)
img_h, img_w = cv2_img.shape[:2]
# 1. Détection YOLO
yolo_boxes, yolo_ms = self.detect_yolo_boxes(cv2_img, confidence_threshold=confidence_threshold)
# Si aucun objet YOLO n'est détecté avec un modèle pré-entraîné COCO générique,
# ou si un seul panneau occupe toute l'image (photo cadrée de près),
# nous ajoutons l'image entière comme boîte candidate principale pour l'OCR et l'analyse.
candidate_boxes = []
if yolo_boxes:
candidate_boxes = yolo_boxes
else:
# Fallback de cadrage intelligent : pleine image + zone centrale
candidate_boxes.append({
"bbox": [0, 0, img_w, img_h],
"confidence": 0.50,
"class_id": -1,
"class_name": "panneau_principal",
})
# 2. Analyse OCR & Correspondance Catalogue pour chaque boîte
total_ocr_ms = 0.0
detected_panels = []
for idx, det in enumerate(candidate_boxes):
bbox = det["bbox"]
ocr_text, ocr_lines, ocr_ms = self.extract_ocr_from_crop(cv2_img, bbox)
total_ocr_ms += ocr_ms
# Tentative d'identification via l'OCR
matched = match_sign_from_ocr(ocr_text)
# Heuristique basée sur la classe COCO si disponible
code = None
name_fr = ""
name_nl = ""
category = "indication"
svg_url = ""
matched_by = "detection_generic"
val = None
if matched:
code = matched["code"]
name_fr = matched["data"]["name_fr"]
name_nl = matched["data"]["name_nl"]
category = matched["data"]["category"]
svg_url = matched["svg_url"]
matched_by = matched["matched_by"]
val = matched.get("value")
elif det["class_name"] == "stop sign":
code = "B5"
name_fr = "Arrêt obligatoire (STOP)"
name_nl = "Verplichte stop (STOP)"
category = "priority"
svg_url = get_svg_url("B5")
matched_by = "yolo_stop_sign"
elif det["class_name"] == "traffic light":
code = "SIGNALISATION_LUMINEUSE"
name_fr = "Feux de signalisation"
name_nl = "Verkeerslichten"
category = "trafficlights"
svg_url = "/static/assets/traffic_light_icon.svg"
matched_by = "yolo_traffic_light"
else:
# Classification visuelle basée sur la classe YOLO et l'image découpée (forme & couleur)
if det["class_name"] == "sub_plate":
code = "M2"
entry = SIGN_CATALOG.get("M2", {})
name_fr = entry.get("name_fr", "Panonceau additionnel")
name_nl = entry.get("name_nl", "Onderbord")
category = "panonceau"
svg_url = get_svg_url("M2")
matched_by = "yolo_sub_plate"
else:
# Découpage du panneau pour classification par forme et couleur
crop_bgr = cv2_img[bbox[1]:bbox[3], bbox[0]:bbox[2]]
from .catalog import classify_sign_visual
vis_res = classify_sign_visual(crop_bgr, ocr_text=ocr_text)
code = vis_res["code"]
name_fr = vis_res["name_fr"]
name_nl = vis_res["name_nl"]
category = vis_res["category"]
svg_url = vis_res["svg_url"]
matched_by = vis_res["matched_by"]
# Recherche en base de données pour associer le SignPanelType officiel si disponible
db_panel_type_id = None
try:
from sign.models import SignPanelType
db_type = SignPanelType.objects.filter(code__iexact=code).first()
if db_type:
db_panel_type_id = db_type.id
name_fr = db_type.name_fr or name_fr
name_nl = db_type.name_nl or name_nl
except Exception:
pass
detected_panels.append({
"id": idx + 1,
"bbox": bbox,
"confidence": det["confidence"],
"yolo_class": det["class_name"],
"code": code,
"name_fr": name_fr,
"name_nl": name_nl,
"category": category,
"svg_url": svg_url,
"matched_by": matched_by,
"ocr_text": ocr_text,
"ocr_lines": ocr_lines,
"signpanel_text": ocr_text if ocr_text else None,
"signpanel_value": val,
"signpanel_type_id": db_panel_type_id,
"ymin": bbox[1],
})
# 3. Ordonnancement vertical (Ordre de haut en bas sur le mât)
detected_panels.sort(key=lambda p: p["ymin"])
for order_idx, panel in enumerate(detected_panels, start=1):
panel["vertical_order"] = order_idx
# 4. Génération de l'image annotée
annotated_cv2 = cv2_img.copy()
for p in detected_panels:
x1, y1, x2, y2 = p["bbox"]
order = p["vertical_order"]
code = p["code"]
conf = int(p["confidence"] * 100)
# Couleur du rectangle (Vert pour haute confiance, Orange pour moyenne)
box_color = (46, 204, 113) if p["confidence"] >= 0.7 else (52, 152, 219)
cv2.rectangle(annotated_cv2, (x1, y1), (x2, y2), box_color, 3)
# Badge avec le numéro d'ordre et le code
label = f"#{order} {code} ({conf}%)"
if p["ocr_text"]:
label += f" - '{p['ocr_text'][:20]}'"
# Fond du texte
(label_w, label_h), baseline = cv2.getTextSize(label, cv2.FONT_HERSHEY_SIMPLEX, 0.6, 2)
cv2.rectangle(
annotated_cv2,
(x1, max(0, y1 - label_h - 10)),
(x1 + label_w + 10, y1),
box_color,
-1
)
cv2.putText(
annotated_cv2,
label,
(x1 + 5, max(label_h + 2, y1 - 5)),
cv2.FONT_HERSHEY_SIMPLEX,
0.6,
(255, 255, 255),
2,
cv2.LINE_AA
)
# Encodage de l'image annotée en base64 pour affichage immédiat
_, buffer = cv2.imencode(".jpg", annotated_cv2, [int(cv2.IMWRITE_JPEG_QUALITY), 85])
annotated_base64 = "data:image/jpeg;base64," + base64.b64encode(buffer).decode("utf-8")
total_elapsed_ms = (time.perf_counter() - total_start) * 1000.0
return {
"status": "success",
"success": True,
"image_dimensions": {"width": img_w, "height": img_h},
"detected_count": len(detected_panels),
"panels": detected_panels,
"annotated_image": annotated_base64,
"performance": {
"yolo_inference_ms": round(yolo_ms, 1),
"ocr_inference_ms": round(total_ocr_ms, 1),
"total_processing_ms": round(total_elapsed_ms, 1),
}
}