feat: integrate YOLOv10n model training pipeline with dataset and environment configuration
This commit is contained in:
parent
f98878f9c4
commit
e561a77fc0
13 changed files with 2147 additions and 2 deletions
7
.gitignore
vendored
7
.gitignore
vendored
|
|
@ -37,7 +37,12 @@ loko/static/
|
|||
deploy.sh
|
||||
generated_passwords.txt
|
||||
loko/generated_passwords.txt
|
||||
scratch/
|
||||
Caddyfile
|
||||
Caddyfile.prod
|
||||
.vscode/
|
||||
|
||||
# AI Training & ML Artifacts
|
||||
runs/
|
||||
*.pt
|
||||
loko/sign/ai/training/
|
||||
labels_signpanels/
|
||||
4
loko/sign/ai/__init__.py
Normal file
4
loko/sign/ai/__init__.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
from .detector import SignDetectionService
|
||||
from .catalog import SIGN_CATALOG, get_svg_url, match_sign_from_ocr, classify_sign_visual
|
||||
|
||||
__all__ = ["SignDetectionService", "SIGN_CATALOG", "get_svg_url", "match_sign_from_ocr", "classify_sign_visual"]
|
||||
619
loko/sign/ai/catalog.py
Normal file
619
loko/sign/ai/catalog.py
Normal file
|
|
@ -0,0 +1,619 @@
|
|||
"""
|
||||
Catalogue de référence des panneaux de signalisation routière (Code de la route belge / français / européen).
|
||||
Permet de matcher les détections d'objets et les lectures OCR avec les types de panneaux normalisés
|
||||
et leurs fichiers vectoriels SVG correspondants.
|
||||
"""
|
||||
import re
|
||||
from typing import Optional, Dict, Any, List
|
||||
|
||||
# Répertoire de base des SVGs statiques (les fichiers sur disque sont en majuscules, ex: F4A.svg, B1.svg)
|
||||
DEFAULT_SVG_BASE_PATH = "/static/assets/road_signs/2025/"
|
||||
|
||||
# Dictionnaire de référence des types de panneaux majeurs (clés en MAJUSCULES)
|
||||
SIGN_CATALOG = {
|
||||
# --- Priorité & Arrêt ---
|
||||
"B1": {
|
||||
"name_fr": "Cédez le passage",
|
||||
"name_nl": "Voorrang verlenen",
|
||||
"category": "priority",
|
||||
"shape": "triangle_down",
|
||||
"keywords": ["cedez", "passage", "voorrang", "verlenen"],
|
||||
},
|
||||
"B3": {
|
||||
"name_fr": "Route prioritaire",
|
||||
"name_nl": "Voorrangsweg",
|
||||
"category": "priority",
|
||||
"shape": "diamond",
|
||||
"keywords": ["priorite"],
|
||||
},
|
||||
"B5": {
|
||||
"name_fr": "Arrêt obligatoire (STOP)",
|
||||
"name_nl": "Verplichte stop (STOP)",
|
||||
"category": "priority",
|
||||
"shape": "octagon",
|
||||
"keywords": ["stop", "arret"],
|
||||
},
|
||||
"B7": {
|
||||
"name_fr": "Fin de route prioritaire",
|
||||
"name_nl": "Einde voorrangsweg",
|
||||
"category": "priority",
|
||||
"shape": "diamond",
|
||||
},
|
||||
"B9": {
|
||||
"name_fr": "Priorité de passage par rapport aux véhicules venant en sens inverse",
|
||||
"name_nl": "Voorrang ten opzichte van tegengesteld verkeer",
|
||||
"category": "priority",
|
||||
"shape": "square",
|
||||
},
|
||||
|
||||
# --- Interdiction & Limitation ---
|
||||
"C1": {
|
||||
"name_fr": "Sens interdit",
|
||||
"name_nl": "Verboden direction",
|
||||
"category": "prohibition",
|
||||
"shape": "circle",
|
||||
"keywords": ["sens interdit", "verboden"],
|
||||
},
|
||||
"C3": {
|
||||
"name_fr": "Accès interdit dans les deux sens à tout conducteur",
|
||||
"name_nl": "Verboden toegang in beide richtingen",
|
||||
"category": "prohibition",
|
||||
"shape": "circle",
|
||||
"keywords": ["acces interdit"],
|
||||
},
|
||||
"C11": {
|
||||
"name_fr": "Accès interdit aux cycles",
|
||||
"name_nl": "Verboden voor fietsers",
|
||||
"category": "prohibition",
|
||||
"shape": "circle",
|
||||
"keywords": ["velo", "cycle", "fiets"],
|
||||
},
|
||||
"C21": {
|
||||
"name_fr": "Accès interdit aux véhicules dont la masse dépasse la limite",
|
||||
"name_nl": "Verboden voor voertuigen met een hogere massa",
|
||||
"category": "prohibition",
|
||||
"shape": "circle",
|
||||
"has_numeric_value": True,
|
||||
"value_unit": "t",
|
||||
},
|
||||
"C43": {
|
||||
"name_fr": "Vitesse maximale autorisée",
|
||||
"name_nl": "Maximumsnelheid",
|
||||
"category": "prohibition",
|
||||
"shape": "circle",
|
||||
"has_numeric_value": True,
|
||||
"value_unit": "km/h",
|
||||
"keywords": ["30", "50", "70", "90", "120"],
|
||||
},
|
||||
|
||||
# --- Obligation ---
|
||||
"D1A": {
|
||||
"name_fr": "Direction obligatoire : tout droit",
|
||||
"name_nl": "Verplichte rijrichting: rechtdoor",
|
||||
"category": "obligation",
|
||||
"shape": "circle",
|
||||
},
|
||||
"D1B": {
|
||||
"name_fr": "Direction obligatoire : à droite",
|
||||
"name_nl": "Verplichte rijrichting: rechts",
|
||||
"category": "obligation",
|
||||
"shape": "circle",
|
||||
},
|
||||
"D1C": {
|
||||
"name_fr": "Direction obligatoire : à gauche",
|
||||
"name_nl": "Verplichte rijrichting: links",
|
||||
"category": "obligation",
|
||||
"shape": "circle",
|
||||
},
|
||||
"D7": {
|
||||
"name_fr": "Piste cyclable obligatoire",
|
||||
"name_nl": "Verplicht fietspad",
|
||||
"category": "obligation",
|
||||
"shape": "circle",
|
||||
},
|
||||
"D9": {
|
||||
"name_fr": "Partie de la voie publique réservée aux piétons et cyclistes",
|
||||
"name_nl": "Deel van de openbare weg voorbehouden voor voetgangers en fietsers",
|
||||
"category": "obligation",
|
||||
"shape": "circle",
|
||||
},
|
||||
|
||||
# --- Stationnement & Arrêt ---
|
||||
"E1": {
|
||||
"name_fr": "Stationnement interdit",
|
||||
"name_nl": "Parkeerverbod",
|
||||
"category": "parking",
|
||||
"shape": "circle",
|
||||
"keywords": ["stationnement interdit", "parkeerverbod"],
|
||||
},
|
||||
"E3": {
|
||||
"name_fr": "Arrêt et stationnement interdits",
|
||||
"name_nl": "Stilstand- en parkeerverbod",
|
||||
"category": "parking",
|
||||
"shape": "circle",
|
||||
"keywords": ["arret interdit"],
|
||||
},
|
||||
"E9A": {
|
||||
"name_fr": "Stationnement autorisé (Parking)",
|
||||
"name_nl": "Parkeren toegelaten (Parking)",
|
||||
"category": "parking",
|
||||
"shape": "square",
|
||||
"keywords": ["p", "parking"],
|
||||
},
|
||||
"E9B": {
|
||||
"name_fr": "Stationnement réservé aux personnes handicapées",
|
||||
"name_nl": "Parkeren voorbehouden voor personen met een handicap",
|
||||
"category": "parking",
|
||||
"shape": "square",
|
||||
"keywords": ["handicap", "pmr"],
|
||||
},
|
||||
|
||||
# --- Danger ---
|
||||
"A1A": {
|
||||
"name_fr": "Virage dangereux à gauche",
|
||||
"name_nl": "Gevaarlijke bocht naar links",
|
||||
"category": "danger",
|
||||
"shape": "triangle_up",
|
||||
},
|
||||
"A1B": {
|
||||
"name_fr": "Virage dangereux à droite",
|
||||
"name_nl": "Gevaarlijke bocht naar rechts",
|
||||
"category": "danger",
|
||||
"shape": "triangle_up",
|
||||
},
|
||||
"A15": {
|
||||
"name_fr": "Passage pour piétons (Danger)",
|
||||
"name_nl": "Voetgangersoversteekplaats (Gevaar)",
|
||||
"category": "danger",
|
||||
"shape": "triangle_up",
|
||||
},
|
||||
"A23": {
|
||||
"name_fr": "Endroit fréquenté par des enfants (École)",
|
||||
"name_nl": "Plaats waar veel kinderen komen (School)",
|
||||
"category": "danger",
|
||||
"shape": "triangle_up",
|
||||
"keywords": ["ecole", "enfants", "school"],
|
||||
},
|
||||
"A25": {
|
||||
"name_fr": "Ralentisseur de trafic / Dos d'âne",
|
||||
"name_nl": "Verkeersdrempel",
|
||||
"category": "danger",
|
||||
"shape": "triangle_up",
|
||||
"keywords": ["ralentisseur", "dos d'ane", "drempel"],
|
||||
},
|
||||
"A31": {
|
||||
"name_fr": "Travaux en cours",
|
||||
"name_nl": "Werken",
|
||||
"category": "danger",
|
||||
"shape": "triangle_up",
|
||||
"keywords": ["travaux", "werken"],
|
||||
},
|
||||
|
||||
# --- Zones & Indications ---
|
||||
"F4A": {
|
||||
"name_fr": "Début de Zone 30",
|
||||
"name_nl": "Begin van een Zone 30",
|
||||
"category": "zone",
|
||||
"shape": "rectangle",
|
||||
"keywords": ["zone 30", "zone30"],
|
||||
"has_numeric_value": True,
|
||||
"default_value": 30,
|
||||
},
|
||||
"F4B": {
|
||||
"name_fr": "Fin de Zone 30",
|
||||
"name_nl": "Einde van een Zone 30",
|
||||
"category": "zone",
|
||||
"shape": "rectangle",
|
||||
"keywords": ["fin zone 30", "einde zone 30"],
|
||||
},
|
||||
"F12A": {
|
||||
"name_fr": "Début d'une zone résidentielle / zone de rencontre",
|
||||
"name_nl": "Begin van een woonerf of erf",
|
||||
"category": "zone",
|
||||
"shape": "rectangle",
|
||||
"keywords": ["zone de rencontre", "woonerf"],
|
||||
},
|
||||
"F19": {
|
||||
"name_fr": "Voie à sens unique",
|
||||
"name_nl": "Eenrichtingsverkeer",
|
||||
"category": "indication",
|
||||
"shape": "square",
|
||||
},
|
||||
"F49": {
|
||||
"name_fr": "Passage pour piétons (Indication)",
|
||||
"name_nl": "Voetgangersoversteekplaats (Aanduiding)",
|
||||
"category": "indication",
|
||||
"shape": "square",
|
||||
},
|
||||
|
||||
# --- Panonceaux additionnels (Type M) ---
|
||||
"M1": {
|
||||
"name_fr": "Panonceau de distance",
|
||||
"name_nl": "Onderbord: afstand",
|
||||
"category": "panonceau",
|
||||
"shape": "rectangle",
|
||||
"has_numeric_value": True,
|
||||
"value_unit": "m",
|
||||
},
|
||||
"M2": {
|
||||
"name_fr": "Panonceau d'application ou d'exception (ex: Sauf riverains)",
|
||||
"name_nl": "Onderbord: uitzondering (bv: Uitgezonderd aangelanden)",
|
||||
"category": "panonceau",
|
||||
"shape": "rectangle",
|
||||
"keywords": ["sauf", "except", "uitgezonderd", "riverain", "aangelande", "velo", "fiets"],
|
||||
},
|
||||
"M3": {
|
||||
"name_fr": "Panonceau pour cyclistes / piétons",
|
||||
"name_nl": "Onderbord voor fietsers",
|
||||
"category": "panonceau",
|
||||
"shape": "rectangle",
|
||||
},
|
||||
"M4": {
|
||||
"name_fr": "Panonceau d'horaire ou de période",
|
||||
"name_nl": "Onderbord: tijdsperiode",
|
||||
"category": "panonceau",
|
||||
"shape": "rectangle",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_catalog_entry(code: str) -> Optional[Dict[str, Any]]:
|
||||
"""Recherche insensible à la casse dans le catalogue."""
|
||||
if not code:
|
||||
return None
|
||||
upper_code = code.strip().upper()
|
||||
return SIGN_CATALOG.get(upper_code)
|
||||
|
||||
|
||||
def get_svg_url(sign_code: str) -> str:
|
||||
"""
|
||||
Retourne l'URL du fichier SVG statique pour un code de panneau donné.
|
||||
Garantit l'utilisation des majuscules car les fichiers sous assets/road_signs/2025/ sont nommés en majuscules (ex: F4A.svg).
|
||||
"""
|
||||
if not sign_code:
|
||||
return ""
|
||||
code_upper = sign_code.strip().upper()
|
||||
return f"{DEFAULT_SVG_BASE_PATH}{code_upper}.svg"
|
||||
|
||||
|
||||
def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Analyse un texte OCR et tente d'associer un type de panneau normalisé.
|
||||
Exemples:
|
||||
- "STOP" -> B5
|
||||
- "ZONE 30" -> F4A (avec valeur 30)
|
||||
- "SAUF RIVERAINS" -> M2
|
||||
- "300 m" -> M1 (avec valeur 300)
|
||||
- "50" (seul dans un cercle) -> C43 (avec valeur 50)
|
||||
"""
|
||||
if not ocr_text:
|
||||
return None
|
||||
|
||||
cleaned = ocr_text.strip().upper()
|
||||
|
||||
# 1. STOP
|
||||
if "STOP" in cleaned:
|
||||
return {
|
||||
"code": "B5",
|
||||
"confidence": 0.96,
|
||||
"data": SIGN_CATALOG["B5"],
|
||||
"svg_url": get_svg_url("B5"),
|
||||
"matched_by": "text_stop",
|
||||
}
|
||||
|
||||
# 2. Zone 30 / Zone 20 / Fin de zone
|
||||
if "ZONE" in cleaned:
|
||||
if "FIN" in cleaned or "EINDE" in cleaned:
|
||||
return {
|
||||
"code": "F4B",
|
||||
"confidence": 0.92,
|
||||
"data": SIGN_CATALOG["F4B"],
|
||||
"svg_url": get_svg_url("F4B"),
|
||||
"matched_by": "text_end_zone",
|
||||
}
|
||||
# Détection vitesse de zone
|
||||
speed_match = re.search(r"\b(20|30|50)\b", cleaned)
|
||||
if speed_match:
|
||||
speed = int(speed_match.group(1))
|
||||
code = "F4A"
|
||||
return {
|
||||
"code": code,
|
||||
"confidence": 0.94,
|
||||
"data": SIGN_CATALOG["F4A"],
|
||||
"svg_url": get_svg_url(code),
|
||||
"value": speed,
|
||||
"matched_by": "text_zone_speed",
|
||||
}
|
||||
|
||||
# 3. Panonceaux d'exception ("Sauf ...", "Excepté ...", "Uitgezonderd ...")
|
||||
if re.search(r"\b(SAUF|EXCEPTE|EXCEPTE|UITGEZONDERD|RIVERAINS?|AANGELANDEN?|VELOS?|FIETS)\b", cleaned, re.IGNORECASE):
|
||||
return {
|
||||
"code": "M2",
|
||||
"confidence": 0.90,
|
||||
"data": SIGN_CATALOG["M2"],
|
||||
"svg_url": get_svg_url("M2"),
|
||||
"extracted_text": ocr_text.strip(),
|
||||
"matched_by": "text_exception",
|
||||
}
|
||||
|
||||
# 4. Panonceaux de distance ("300 m", "50m", "1.5 km")
|
||||
dist_match = re.search(r"\b(\d+(?:[.,]\d+)?)\s*(M|KM|METRES?|METERS?)\b", cleaned)
|
||||
if dist_match:
|
||||
val_str = dist_match.group(1).replace(",", ".")
|
||||
unit = dist_match.group(2).lower()
|
||||
val = float(val_str)
|
||||
if unit == "km":
|
||||
val *= 1000.0
|
||||
return {
|
||||
"code": "M1",
|
||||
"confidence": 0.88,
|
||||
"data": SIGN_CATALOG["M1"],
|
||||
"svg_url": get_svg_url("M1"),
|
||||
"value": val,
|
||||
"extracted_text": ocr_text.strip(),
|
||||
"matched_by": "text_distance",
|
||||
}
|
||||
|
||||
# 5. Limitation de vitesse pure ("30", "50", "70", "90", "110", "120")
|
||||
speed_alone_match = re.search(r"^\D*(\b(?:20|30|50|70|90|110|120)\b)\D*$", cleaned)
|
||||
if speed_alone_match:
|
||||
speed = int(speed_alone_match.group(1))
|
||||
return {
|
||||
"code": "C43",
|
||||
"confidence": 0.90,
|
||||
"data": SIGN_CATALOG["C43"],
|
||||
"svg_url": get_svg_url("C43"),
|
||||
"value": speed,
|
||||
"matched_by": "text_speed_limit",
|
||||
}
|
||||
|
||||
# 6. Tonnage ("3.5 t", "7.5t")
|
||||
ton_match = re.search(r"\b(\d+(?:[.,]\d+)?)\s*T\b", cleaned)
|
||||
if ton_match:
|
||||
val = float(ton_match.group(1).replace(",", "."))
|
||||
return {
|
||||
"code": "C21",
|
||||
"confidence": 0.88,
|
||||
"data": SIGN_CATALOG["C21"],
|
||||
"svg_url": get_svg_url("C21"),
|
||||
"value": val,
|
||||
"extracted_text": ocr_text.strip(),
|
||||
"matched_by": "text_tonnage",
|
||||
}
|
||||
|
||||
# 7. Parking PMR / Handicap
|
||||
if re.search(r"\b(HANDICAP|PMR|HANDICAPE)\b", cleaned):
|
||||
return {
|
||||
"code": "E9B",
|
||||
"confidence": 0.91,
|
||||
"data": SIGN_CATALOG["E9B"],
|
||||
"svg_url": get_svg_url("E9B"),
|
||||
"matched_by": "text_handicap",
|
||||
}
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def classify_sign_visual(crop_bgr: Any, ocr_text: str = "") -> Dict[str, Any]:
|
||||
"""
|
||||
Classifie un panneau par analyse de forme, couleur dominante (Bleu, Rouge, Jaune) et structure.
|
||||
Permet de discriminer les panneaux sans texte :
|
||||
- Rond Bleu -> D7 (Piste cyclable / vélo) ou D5 (Rond-point)
|
||||
- Triangle inversé -> B1 (Cédez le passage)
|
||||
- Triangle pointe en haut -> A15 (Passage piétons / Danger)
|
||||
- Octogone rouge -> B5 (STOP)
|
||||
- Losange jaune -> B3 (Route prioritaire)
|
||||
- Cercle bord rouge -> C3 (Accès interdit)
|
||||
"""
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
if not isinstance(crop_bgr, np.ndarray) or crop_bgr.size == 0:
|
||||
return {
|
||||
"code": "B1",
|
||||
"name_fr": "Panneau de signalisation",
|
||||
"name_nl": "Verkeersbord",
|
||||
"category": "priority",
|
||||
"svg_url": get_svg_url("B1"),
|
||||
"matched_by": "fallback_empty",
|
||||
"confidence": 0.50,
|
||||
}
|
||||
|
||||
h, w = crop_bgr.shape[:2]
|
||||
if h < 10 or w < 10:
|
||||
return {
|
||||
"code": "B1",
|
||||
"name_fr": "Panneau de signalisation",
|
||||
"name_nl": "Verkeersbord",
|
||||
"category": "priority",
|
||||
"svg_url": get_svg_url("B1"),
|
||||
"matched_by": "fallback_small",
|
||||
"confidence": 0.50,
|
||||
}
|
||||
|
||||
hsv = cv2.cvtColor(crop_bgr, cv2.COLOR_BGR2HSV)
|
||||
total_px = float(h * w)
|
||||
|
||||
# Masques couleur HSV
|
||||
blue_mask = cv2.inRange(hsv, np.array([95, 40, 30]), np.array([135, 255, 255]))
|
||||
red_mask1 = cv2.inRange(hsv, np.array([0, 50, 40]), np.array([12, 255, 255]))
|
||||
red_mask2 = cv2.inRange(hsv, np.array([160, 50, 40]), np.array([180, 255, 255]))
|
||||
red_mask = red_mask1 | red_mask2
|
||||
yellow_mask = cv2.inRange(hsv, np.array([15, 60, 60]), np.array([35, 255, 255]))
|
||||
|
||||
blue_ratio = np.count_nonzero(blue_mask) / total_px
|
||||
red_ratio = np.count_nonzero(red_mask) / total_px
|
||||
yellow_ratio = np.count_nonzero(yellow_mask) / total_px
|
||||
aspect_ratio = w / float(h)
|
||||
|
||||
# 1. PANNEAUX BLEUS (Famille D Obligation ou Famille F Indication)
|
||||
if blue_ratio > 0.12:
|
||||
if 0.75 <= aspect_ratio <= 1.35:
|
||||
# Extraction des éléments blancs centraux (pictogramme)
|
||||
gray = cv2.cvtColor(crop_bgr, cv2.COLOR_BGR2GRAY)
|
||||
center_crop = gray[int(h*0.18):int(h*0.82), int(w*0.18):int(w*0.82)]
|
||||
white_mask = center_crop > 165
|
||||
white_ratio = np.count_nonzero(white_mask) / float(center_crop.size)
|
||||
|
||||
if "P" in ocr_text.upper():
|
||||
code = "E9A"
|
||||
entry = SIGN_CATALOG.get("E9A", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Stationnement autorisé (Parking)"),
|
||||
"name_nl": entry.get("name_nl", "Parkeren toegelaten (Parking)"),
|
||||
"category": "parking",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_blue_parking",
|
||||
"confidence": 0.88,
|
||||
}
|
||||
|
||||
# Si rond bleu : D7 (Piste cyclable / vélo) vs D5 (Rond-point) vs D1A
|
||||
if "VELO" in ocr_text.upper() or "FIETS" in ocr_text.upper() or white_ratio > 0.03:
|
||||
code = "D7"
|
||||
entry = SIGN_CATALOG.get("D7", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Piste cyclable obligatoire"),
|
||||
"name_nl": entry.get("name_nl", "Verplicht fietspad"),
|
||||
"category": "obligation",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_blue_bike",
|
||||
"confidence": 0.90,
|
||||
}
|
||||
else:
|
||||
code = "D5"
|
||||
entry = SIGN_CATALOG.get("D5", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Sens giratoire obligatoire (Rond-point)"),
|
||||
"name_nl": entry.get("name_nl", "Verplicht rond punt"),
|
||||
"category": "obligation",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_blue_roundabout",
|
||||
"confidence": 0.86,
|
||||
}
|
||||
else:
|
||||
code = "F19"
|
||||
entry = SIGN_CATALOG.get("F19", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Voie à sens unique"),
|
||||
"name_nl": entry.get("name_nl", "Eenrichtingsverkeer"),
|
||||
"category": "indication",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_blue_rectangle",
|
||||
"confidence": 0.82,
|
||||
}
|
||||
|
||||
# 2. PANNEAUX ROUGE & BLANC
|
||||
if red_ratio > 0.06:
|
||||
top_half_red = np.count_nonzero(red_mask[:int(h*0.35), :])
|
||||
bottom_half_red = np.count_nonzero(red_mask[int(h*0.65):, :])
|
||||
|
||||
# Triangle inversé (B1 - Cédez le passage : large en haut, pointe en bas)
|
||||
if top_half_red > (bottom_half_red * 1.5) and aspect_ratio > 0.75:
|
||||
code = "B1"
|
||||
entry = SIGN_CATALOG.get("B1", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Cédez le passage"),
|
||||
"name_nl": entry.get("name_nl", "Voorrang verlenen"),
|
||||
"category": "priority",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_triangle_down_yield",
|
||||
"confidence": 0.92,
|
||||
}
|
||||
|
||||
# Triangle pointe en haut (Danger A... : base large en bas, pointe en haut)
|
||||
if bottom_half_red > (top_half_red * 1.4) and aspect_ratio > 0.75:
|
||||
code = "A15"
|
||||
entry = SIGN_CATALOG.get("A15", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Passage pour piétons (Danger)"),
|
||||
"name_nl": entry.get("name_nl", "Voetgangersoversteekplaats (Gevaar)"),
|
||||
"category": "danger",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_triangle_up_danger",
|
||||
"confidence": 0.88,
|
||||
}
|
||||
|
||||
# Analyse du centre pour discriminer C1 (Sens interdit), B5 (STOP) et C3 (Accès interdit)
|
||||
center_hsv = hsv[int(h * 0.28):int(h * 0.72), int(w * 0.15):int(w * 0.85)]
|
||||
c_h, c_w = center_hsv.shape[:2]
|
||||
c_sat = center_hsv[:, :, 1]
|
||||
c_val = center_hsv[:, :, 2]
|
||||
|
||||
# Le blanc = faible saturation chromatique (< 95) et luminosité relative (> 50)
|
||||
white_bar_mask = (c_sat < 95) & (c_val > 50)
|
||||
white_bar_ratio = np.count_nonzero(white_bar_mask) / float(max(1, c_h * c_w))
|
||||
|
||||
# Test C1 (Sens interdit : barre blanche horizontale continue au centre)
|
||||
if white_bar_ratio > 0.12 and "STOP" not in ocr_text.upper():
|
||||
row_white = np.sum(white_bar_mask, axis=1)
|
||||
active_rows = np.count_nonzero(row_white > (c_w * 0.28))
|
||||
if active_rows > 0 and active_rows < (c_h * 0.85):
|
||||
code = "C1"
|
||||
entry = SIGN_CATALOG.get("C1", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Sens interdit"),
|
||||
"name_nl": entry.get("name_nl", "Verboden richting"),
|
||||
"category": "prohibition",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_c1_horizontal_bar",
|
||||
"confidence": 0.95,
|
||||
}
|
||||
|
||||
# STOP (B5) : Mot STOP explicite ou forme octogonale rouge
|
||||
if "STOP" in ocr_text.upper():
|
||||
code = "B5"
|
||||
entry = SIGN_CATALOG.get("B5", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Arrêt obligatoire (STOP)"),
|
||||
"name_nl": entry.get("name_nl", "Verplichte stop (STOP)"),
|
||||
"category": "priority",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_stop_octagon",
|
||||
"confidence": 0.95,
|
||||
}
|
||||
|
||||
# Cercle d'interdiction (C3 / C43)
|
||||
code = "C3"
|
||||
entry = SIGN_CATALOG.get("C3", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Accès interdit dans les deux sens"),
|
||||
"name_nl": entry.get("name_nl", "Verboden toegang in beide richtingen"),
|
||||
"category": "prohibition",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_circle_prohibition",
|
||||
"confidence": 0.85,
|
||||
}
|
||||
|
||||
# 3. PANNEAUX JAUNES (Route prioritaire B3)
|
||||
if yellow_ratio > 0.08:
|
||||
code = "B3"
|
||||
entry = SIGN_CATALOG.get("B3", {})
|
||||
return {
|
||||
"code": code,
|
||||
"name_fr": entry.get("name_fr", "Route prioritaire"),
|
||||
"name_nl": entry.get("name_nl", "Voorrangsweg"),
|
||||
"category": "priority",
|
||||
"svg_url": get_svg_url(code),
|
||||
"matched_by": "visual_yellow_priority",
|
||||
"confidence": 0.89,
|
||||
}
|
||||
|
||||
# 4. Panneau générique par défaut
|
||||
return {
|
||||
"code": "B1",
|
||||
"name_fr": "Panneau de signalisation",
|
||||
"name_nl": "Verkeersbord",
|
||||
"category": "priority",
|
||||
"svg_url": get_svg_url("B1"),
|
||||
"matched_by": "visual_generic",
|
||||
"confidence": 0.60,
|
||||
}
|
||||
498
loko/sign/ai/detector.py
Normal file
498
loko/sign/ai/detector.py
Normal file
|
|
@ -0,0 +1,498 @@
|
|||
"""
|
||||
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",
|
||||
"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),
|
||||
}
|
||||
}
|
||||
BIN
loko/sign/ai/models/yolov10n.onnx
Normal file
BIN
loko/sign/ai/models/yolov10n.onnx
Normal file
Binary file not shown.
142
loko/sign/ai/train.py
Normal file
142
loko/sign/ai/train.py
Normal file
|
|
@ -0,0 +1,142 @@
|
|||
"""
|
||||
Pipeline de préparation de données, fine-tuning YOLO et export ONNX pour StreetUp / Loko.
|
||||
Permet d'entraîner le modèle de détection spécifique aux panneaux (sign) et panonceaux (sub_plate).
|
||||
"""
|
||||
import os
|
||||
import shutil
|
||||
import random
|
||||
from pathlib import Path
|
||||
import yaml
|
||||
|
||||
|
||||
def prepare_yolo_dataset(
|
||||
source_dir: str = "/home/kdt/Downloads/labels_signpanels",
|
||||
target_dir: str = "/home/kdt/StreetUp/Antigravity/loko/loko/sign/ai/training/dataset",
|
||||
train_ratio: float = 0.85,
|
||||
seed: int = 42,
|
||||
) -> Path:
|
||||
"""
|
||||
Scanne le dossier d'export X-AnyLabeling, nettoie et répartit les images/labels
|
||||
en sous-dossiers train/val, et génère le fichier data.yaml.
|
||||
"""
|
||||
src = Path(source_dir)
|
||||
dest = Path(target_dir)
|
||||
|
||||
if not src.exists():
|
||||
raise FileNotFoundError(f"Dossier source introuvable : {src}")
|
||||
|
||||
# Réinitialisation propre du dossier cible
|
||||
if dest.exists():
|
||||
shutil.rmtree(dest)
|
||||
|
||||
for split in ['train', 'val']:
|
||||
(dest / 'images' / split).mkdir(parents=True, exist_ok=True)
|
||||
(dest / 'labels' / split).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Récupération des paires (image, txt)
|
||||
valid_extensions = ('.png', '.jpg', '.jpeg', '.webp')
|
||||
image_files = [f for f in src.iterdir() if f.suffix.lower() in valid_extensions]
|
||||
|
||||
pairs = []
|
||||
for img_path in sorted(image_files):
|
||||
txt_path = img_path.with_suffix('.txt')
|
||||
if txt_path.exists():
|
||||
pairs.append((img_path, txt_path))
|
||||
else:
|
||||
print(f"⚠️ Label manquant pour l'image : {img_path.name}")
|
||||
|
||||
print(f"✓ {len(pairs)} paires (image + annotation) validées.")
|
||||
|
||||
# Mélange reproductible
|
||||
random.seed(seed)
|
||||
random.shuffle(pairs)
|
||||
|
||||
split_idx = max(1, int(len(pairs) * train_ratio))
|
||||
train_pairs = pairs[:split_idx]
|
||||
val_pairs = pairs[split_idx:]
|
||||
|
||||
print(f"📊 Répartition : {len(train_pairs)} en Entraînement (Train), {len(val_pairs)} en Validation (Val)")
|
||||
|
||||
def copy_split(pair_list, split_name):
|
||||
for idx, (img_p, txt_p) in enumerate(pair_list):
|
||||
safe_name = f"sign_{split_name}_{idx:04d}{img_p.suffix.lower()}"
|
||||
safe_txt_name = f"sign_{split_name}_{idx:04d}.txt"
|
||||
|
||||
shutil.copy2(img_p, dest / 'images' / split_name / safe_name)
|
||||
shutil.copy2(txt_p, dest / 'labels' / split_name / safe_txt_name)
|
||||
|
||||
copy_split(train_pairs, 'train')
|
||||
copy_split(val_pairs, 'val')
|
||||
|
||||
# Génération du data.yaml
|
||||
data_config = {
|
||||
'path': str(dest.resolve()),
|
||||
'train': 'images/train',
|
||||
'val': 'images/val',
|
||||
'names': {
|
||||
0: 'sign',
|
||||
1: 'sub_plate',
|
||||
}
|
||||
}
|
||||
|
||||
yaml_path = dest / 'data.yaml'
|
||||
with open(yaml_path, 'w', encoding='utf-8') as f:
|
||||
yaml.dump(data_config, f, default_flow_style=False, sort_keys=False)
|
||||
|
||||
print(f"✓ Fichier de configuration généré : {yaml_path}")
|
||||
return yaml_path
|
||||
|
||||
|
||||
def run_training(
|
||||
data_yaml_path: Path,
|
||||
epochs: int = 60,
|
||||
imgsz: int = 640,
|
||||
batch: int = 8,
|
||||
base_model: str = "yolo11n.pt",
|
||||
output_onnx_path: str = "/home/kdt/StreetUp/Antigravity/loko/loko/sign/ai/models/yolov10n.onnx",
|
||||
):
|
||||
"""
|
||||
Lance l'entraînement Ultralytics YOLO et exporte le modèle en ONNX optimisé.
|
||||
"""
|
||||
from ultralytics import YOLO
|
||||
|
||||
print(f"\n🚀 Démarrage du Fine-Tuning YOLO ({base_model}) sur vos {epochs} époques...")
|
||||
model = YOLO(base_model)
|
||||
|
||||
# Entraînement
|
||||
results = model.train(
|
||||
data=str(data_yaml_path),
|
||||
epochs=epochs,
|
||||
imgsz=imgsz,
|
||||
batch=batch,
|
||||
device="cpu",
|
||||
name="streetup_sign_detector",
|
||||
exist_ok=True,
|
||||
patience=20,
|
||||
save=True,
|
||||
)
|
||||
|
||||
# Récupération des meilleurs poids (.pt)
|
||||
best_weights = Path(model.trainer.save_dir) / "weights" / "best.pt"
|
||||
print(f"\n✓ Entraînement terminé ! Meilleurs poids enregistrés dans : {best_weights}")
|
||||
|
||||
# Export en ONNX
|
||||
print("📦 Exportation du modèle entraîné vers le format ONNX...")
|
||||
trained_model = YOLO(str(best_weights))
|
||||
exported_file = trained_model.export(
|
||||
format="onnx",
|
||||
opset=12,
|
||||
simplify=True,
|
||||
imgsz=imgsz,
|
||||
)
|
||||
|
||||
# Copie vers le répertoire de modèles Loko
|
||||
target_onnx = Path(output_onnx_path)
|
||||
target_onnx.parent.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copy2(exported_file, target_onnx)
|
||||
print(f"🎉 Modèle ONNX actif déployé avec succès dans : {target_onnx} ({target_onnx.stat().st_size / (1024*1024):.2f} Mo)")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
yaml_file = prepare_yolo_dataset()
|
||||
run_training(yaml_file, epochs=50)
|
||||
53
loko/sign/management/commands/setup_sign_ai.py
Normal file
53
loko/sign/management/commands/setup_sign_ai.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
"""
|
||||
Commande Django pour initialiser et vérifier l'environnement d'IA de signalisation :
|
||||
- Téléchargement du modèle YOLOv10-n ONNX
|
||||
- Initialisation des poids RapidOCR / PaddleOCR ONNX
|
||||
- Test de validation de l'inférence CPU
|
||||
"""
|
||||
import time
|
||||
import numpy as np
|
||||
import cv2
|
||||
from django.core.management.base import BaseCommand
|
||||
from sign.ai import SignDetectionService
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "Télécharge les modèles IA (YOLOv10-n ONNX & OCR) et valide l'inférence CPU."
|
||||
|
||||
def handle(self, *args, **options):
|
||||
self.stdout.write(self.style.NOTICE("=== Initialisation du module IA de Signalisation (StreetUp / Loko) ==="))
|
||||
|
||||
service = SignDetectionService.get_instance()
|
||||
self.stdout.write(f"1. Vérification du modèle YOLOv10-n dans : {service.yolo_model_path}")
|
||||
|
||||
start_dl = time.perf_counter()
|
||||
model_path = service._ensure_yolo_model()
|
||||
self.stdout.write(self.style.SUCCESS(f" ✓ Modèle YOLOv10-n présent ({model_path.stat().st_size / (1024*1024):.2f} Mo)"))
|
||||
|
||||
self.stdout.write("2. Chargement de la session ONNX Runtime...")
|
||||
session = service.get_yolo_session()
|
||||
self.stdout.write(self.style.SUCCESS(f" ✓ Session ONNX prête (Providers: {session.get_providers()})"))
|
||||
|
||||
self.stdout.write("3. Initialisation du moteur RapidOCR / PaddleOCR...")
|
||||
ocr = service.get_ocr_engine()
|
||||
self.stdout.write(self.style.SUCCESS(" ✓ Moteur OCR prêt."))
|
||||
|
||||
self.stdout.write("4. Exécution du test de validation d'inférence CPU...")
|
||||
# Image de test avec STOP et panonceau
|
||||
test_img = np.ones((600, 600, 3), dtype=np.uint8) * 240
|
||||
cv2.circle(test_img, (300, 200), 100, (0, 0, 200), -1)
|
||||
cv2.putText(test_img, "STOP", (240, 215), cv2.FONT_HERSHEY_SIMPLEX, 1.4, (255, 255, 255), 4)
|
||||
|
||||
res = service.analyze_image(test_img)
|
||||
|
||||
self.stdout.write(self.style.SUCCESS(f" ✓ Analyse réussie ! Panneaux détectés : {res['detected_count']}"))
|
||||
for p in res['panels']:
|
||||
self.stdout.write(f" - #{p['vertical_order']} [{p['code']}] {p['name_fr']} (OCR: '{p['ocr_text']}')")
|
||||
|
||||
perf = res['performance']
|
||||
self.stdout.write(self.style.NOTICE(f"5. Métriques de performance CPU :"))
|
||||
self.stdout.write(f" • Inférence YOLO : {perf['yolo_inference_ms']} ms")
|
||||
self.stdout.write(f" • Inférence OCR : {perf['ocr_inference_ms']} ms")
|
||||
self.stdout.write(f" • Temps total : {perf['total_processing_ms']} ms")
|
||||
|
||||
self.stdout.write(self.style.SUCCESS("=== Module IA opérationnel et prêt à l'emploi ==="))
|
||||
512
loko/sign/templates/sign/ai_demo.html
Normal file
512
loko/sign/templates/sign/ai_demo.html
Normal file
|
|
@ -0,0 +1,512 @@
|
|||
{% extends "base.html" %}
|
||||
{% load i18n static %}
|
||||
|
||||
{% block title %}{% translate "Banc de Test IA — Reconnaissance Signalisation & OCR" %}{% endblock %}
|
||||
|
||||
{% block head %}
|
||||
<style>
|
||||
.ai-hero-card {
|
||||
background: linear-gradient(135deg, #1e293b 0%, #0f172a 100%);
|
||||
border: 1px solid rgba(255, 255, 255, 0.1);
|
||||
border-radius: 1rem;
|
||||
color: #f8fafc;
|
||||
}
|
||||
.dropzone-box {
|
||||
border: 2px dashed #94a3b8;
|
||||
border-radius: 1rem;
|
||||
background: #f8fafc;
|
||||
transition: all 0.2s ease-in-out;
|
||||
cursor: pointer;
|
||||
}
|
||||
.dropzone-box:hover, .dropzone-box.dragover {
|
||||
border-color: #2563eb;
|
||||
background: #eff6ff;
|
||||
}
|
||||
.panel-card {
|
||||
border: 1px solid #e2e8f0;
|
||||
border-radius: 0.75rem;
|
||||
transition: transform 0.2s, box-shadow 0.2s;
|
||||
background: #ffffff;
|
||||
}
|
||||
.panel-card:hover {
|
||||
transform: translateY(-2px);
|
||||
box-shadow: 0 10px 25px -5px rgba(0, 0, 0, 0.1);
|
||||
}
|
||||
.svg-sign-preview {
|
||||
width: 84px;
|
||||
height: 84px;
|
||||
object-fit: contain;
|
||||
filter: drop-shadow(0 2px 4px rgba(0,0,0,0.15));
|
||||
}
|
||||
.perf-badge {
|
||||
font-family: monospace;
|
||||
font-size: 0.95rem;
|
||||
}
|
||||
.order-badge {
|
||||
width: 32px;
|
||||
height: 32px;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
border-radius: 50%;
|
||||
font-weight: 700;
|
||||
}
|
||||
</style>
|
||||
{% endblock %}
|
||||
|
||||
{% block content %}
|
||||
<div class="container py-4">
|
||||
|
||||
<!-- En-tête / Hero -->
|
||||
<div class="ai-hero-card p-4 mb-4 shadow-sm">
|
||||
<div class="d-flex flex-wrap align-items-center justify-content-between gap-3">
|
||||
<div>
|
||||
<div class="d-flex align-items-center gap-2 mb-2">
|
||||
<span class="badge bg-primary px-3 py-2 fs-6">
|
||||
<i class="bi bi-cpu me-1"></i> YOLOv10-n + PaddleOCR (CPU ONNX)
|
||||
</span>
|
||||
<span class="badge bg-success-subtle text-success border border-success-subtle px-2 py-1">
|
||||
<i class="bi bi-check-circle-fill me-1"></i> Phase 1 Opérationnelle
|
||||
</span>
|
||||
</div>
|
||||
<h2 class="fw-bold mb-1">
|
||||
<i class="bi bi-camera-fill text-warning me-2"></i>{% translate "Reconnaissance de Signalisation & OCR" %}
|
||||
</h2>
|
||||
<p class="text-slate-300 mb-0 opacity-75">
|
||||
{% translate "Analyse instantanée des photos de terrain : détection des panneaux, extraction du texte des panonceaux, ordonnancement vertical et visualisation SVG." %}
|
||||
</p>
|
||||
</div>
|
||||
<div>
|
||||
<a href="{% url 'sign:index' %}" class="btn btn-outline-light btn-sm">
|
||||
<i class="bi bi-arrow-left me-1"></i>{% translate "Retour à la Signalisation" %}
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Zone principale : Upload & Options -->
|
||||
<div class="row g-4 mb-4">
|
||||
<div class="col-lg-5">
|
||||
<div class="card border-0 shadow-sm rounded-3 h-100">
|
||||
<div class="card-body p-4 d-flex flex-column">
|
||||
<h5 class="fw-bold text-dark mb-3">
|
||||
<i class="bi bi-upload text-primary me-2"></i>{% translate "Source de l'image" %}
|
||||
</h5>
|
||||
|
||||
<!-- Dropzone -->
|
||||
<div id="dropzone" class="dropzone-box p-4 text-center mb-3">
|
||||
<input type="file" id="imageInput" accept="image/*" class="d-none" capture="environment">
|
||||
<i class="bi bi-cloud-arrow-up display-4 text-primary mb-2 d-block"></i>
|
||||
<h6 class="fw-bold mb-1">{% translate "Glissez une photo ou cliquez ici" %}</h6>
|
||||
<p class="text-muted small mb-2">{% translate "JPG, PNG, WebP (Prise de vue directe ou fichier)" %}</p>
|
||||
<div class="d-flex justify-content-center gap-2">
|
||||
<button type="button" class="btn btn-primary btn-sm px-3" onclick="document.getElementById('imageInput').click()">
|
||||
<i class="bi bi-folder2-open me-1"></i>{% translate "Parcourir" %}
|
||||
</button>
|
||||
<button type="button" class="btn btn-outline-secondary btn-sm px-3" id="btnCamera">
|
||||
<i class="bi bi-camera me-1"></i>{% translate "Caméra" %}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Réglage du seuil -->
|
||||
<div class="mb-3">
|
||||
<div class="d-flex justify-content-between align-items-center mb-1">
|
||||
<label for="confThreshold" class="form-label small fw-semibold text-muted mb-0">
|
||||
{% translate "Seuil de confiance minimum" %}
|
||||
</label>
|
||||
<span id="confValue" class="badge bg-light text-dark border">0.25</span>
|
||||
</div>
|
||||
<input type="range" class="form-range" id="confThreshold" min="0.10" max="0.90" step="0.05" value="0.25">
|
||||
</div>
|
||||
|
||||
<!-- Exemples rapides synthétiques -->
|
||||
<div class="mt-auto pt-3 border-top">
|
||||
<label class="form-label small fw-semibold text-muted mb-2">{% translate "Ou tester avec un exemple synthétique :" %}</label>
|
||||
<div class="d-flex flex-wrap gap-2">
|
||||
<button type="button" class="btn btn-outline-danger btn-sm" onclick="loadSyntheticExample('stop')">
|
||||
<i class="bi bi-octagon-fill me-1"></i>STOP + Sauf riverains
|
||||
</button>
|
||||
<button type="button" class="btn btn-outline-primary btn-sm" onclick="loadSyntheticExample('zone30')">
|
||||
<i class="bi bi-speedometer2 me-1"></i>Zone 30
|
||||
</button>
|
||||
<button type="button" class="btn btn-outline-warning btn-sm" onclick="loadSyntheticExample('distance')">
|
||||
<i class="bi bi-arrow-right me-1"></i>Panonceau 300 m
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Visualisation & Statut -->
|
||||
<div class="col-lg-7">
|
||||
<div class="card border-0 shadow-sm rounded-3 h-100">
|
||||
<div class="card-body p-4 d-flex flex-column justify-content-center align-items-center text-center" id="emptyStateContainer">
|
||||
<div class="py-5">
|
||||
<i class="bi bi-image text-muted display-1 opacity-25 d-block mb-3"></i>
|
||||
<h5 class="text-muted fw-bold">{% translate "Aucune image analysée pour le moment" %}</h5>
|
||||
<p class="text-muted small max-w-sm mb-0">
|
||||
{% translate "Importez une photo pour démarrer l'inférence YOLOv10-n et l'extraction OCR." %}
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Indicateur de chargement -->
|
||||
<div class="d-none text-center py-5 my-auto" id="loadingContainer">
|
||||
<div class="spinner-border text-primary mb-3" style="width: 3.5rem; height: 3.5rem;" role="status"></div>
|
||||
<h5 class="fw-bold text-dark mb-1">{% translate "Analyse en cours..." %}</h5>
|
||||
<p class="text-muted small mb-0">
|
||||
{% translate "Exécution de YOLOv10-n (Détection) + RapidOCR (Lecture) sur CPU ONNX" %}
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<!-- Résultat Image Annotée -->
|
||||
<div class="d-none w-100" id="imageResultContainer">
|
||||
<div class="d-flex justify-content-between align-items-center mb-2">
|
||||
<h6 class="fw-bold text-dark mb-0">
|
||||
<i class="bi bi-bounding-box text-primary me-2"></i>{% translate "Image avec détections & Bounding Boxes" %}
|
||||
</h6>
|
||||
<span id="detectedBadge" class="badge bg-success">0 panneau(x)</span>
|
||||
</div>
|
||||
<div class="position-relative text-center bg-dark rounded-3 overflow-hidden p-2">
|
||||
<img id="annotatedPreview" src="" alt="Détection" class="img-fluid rounded" style="max-height: 420px;">
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Section Résultats Détaillés & SVGs -->
|
||||
<div class="d-none" id="resultsSection">
|
||||
<!-- Métriques CPU -->
|
||||
<div class="row g-3 mb-4">
|
||||
<div class="col-md-4">
|
||||
<div class="card border-0 bg-primary bg-opacity-10 rounded-3 p-3 text-center">
|
||||
<span class="text-muted small fw-semibold">Inférence YOLOv10-n (CPU)</span>
|
||||
<h4 class="fw-bold text-primary mb-0" id="perfYolo">-- ms</h4>
|
||||
</div>
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<div class="card border-0 bg-info bg-opacity-10 rounded-3 p-3 text-center">
|
||||
<span class="text-muted small fw-semibold">Inférence OCR RapidOCR (CPU)</span>
|
||||
<h4 class="fw-bold text-info mb-0" id="perfOcr">-- ms</h4>
|
||||
</div>
|
||||
</div>
|
||||
<div class="col-md-4">
|
||||
<div class="card border-0 bg-success bg-opacity-10 rounded-3 p-3 text-center">
|
||||
<span class="text-muted small fw-semibold">Temps Total de Traitement</span>
|
||||
<h4 class="fw-bold text-success mb-0" id="perfTotal">-- ms</h4>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Liste des panneaux reconnus avec SVG -->
|
||||
<div class="card border-0 shadow-sm rounded-3 mb-4">
|
||||
<div class="card-header bg-white py-3 border-0">
|
||||
<h5 class="fw-bold text-dark mb-0">
|
||||
<i class="bi bi-list-check text-primary me-2"></i>{% translate "Panneaux Identifiés & Suggestions" %}
|
||||
</h5>
|
||||
</div>
|
||||
<div class="card-body p-4 pt-0">
|
||||
<div class="row g-3" id="panelsListContainer">
|
||||
<!-- Rempli dynamiquement en JS -->
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Payload JSON -->
|
||||
<div class="card border-0 shadow-sm rounded-3">
|
||||
<div class="card-header bg-white py-3 d-flex justify-content-between align-items-center">
|
||||
<span class="fw-semibold text-muted small">
|
||||
<i class="bi bi-code-square me-1"></i>{% translate "Réponse API JSON (POST /sign/api/detect/)" %}
|
||||
</span>
|
||||
<button class="btn btn-outline-secondary btn-sm" onclick="copyJsonPayload()">
|
||||
<i class="bi bi-clipboard me-1"></i>{% translate "Copier JSON" %}
|
||||
</button>
|
||||
</div>
|
||||
<div class="card-body p-0">
|
||||
<pre class="bg-dark text-light p-3 m-0 rounded-bottom small overflow-auto" style="max-height: 250px;" id="jsonPayload">{}</pre>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
<!-- Canvas invisible pour générer les exemples synthétiques -->
|
||||
<canvas id="synthCanvas" width="600" height="600" class="d-none"></canvas>
|
||||
|
||||
<script>
|
||||
const i18n = {
|
||||
invalidImage: "{{ _('Veuillez sélectionner un fichier image valide.')|escapejs }}",
|
||||
errorProcessing: "{{ _('Erreur lors du traitement')|escapejs }}",
|
||||
networkError: "{{ _('Erreur réseau lors de l\'analyse IA')|escapejs }}",
|
||||
noPanelFound: "{{ _('Aucun panneau spécifique identifié.')|escapejs }}",
|
||||
jsonCopied: "{{ _('Payload JSON copié dans le presse-papier !')|escapejs }}",
|
||||
};
|
||||
|
||||
let lastAnalysisResult = null;
|
||||
|
||||
// Gestion du slider de confiance
|
||||
const confSlider = document.getElementById('confThreshold');
|
||||
const confVal = document.getElementById('confValue');
|
||||
confSlider.addEventListener('input', () => {
|
||||
confVal.textContent = parseFloat(confSlider.value).toFixed(2);
|
||||
});
|
||||
|
||||
// Dropzone events
|
||||
const dropzone = document.getElementById('dropzone');
|
||||
const imageInput = document.getElementById('imageInput');
|
||||
|
||||
dropzone.addEventListener('dragover', (e) => {
|
||||
e.preventDefault();
|
||||
dropzone.classList.add('dragover');
|
||||
});
|
||||
dropzone.addEventListener('dragleave', () => dropzone.classList.remove('dragover'));
|
||||
dropzone.addEventListener('drop', (e) => {
|
||||
e.preventDefault();
|
||||
dropzone.classList.remove('dragover');
|
||||
if (e.dataTransfer.files.length > 0) {
|
||||
processFile(e.dataTransfer.files[0]);
|
||||
}
|
||||
});
|
||||
imageInput.addEventListener('change', () => {
|
||||
if (imageInput.files.length > 0) {
|
||||
processFile(imageInput.files[0]);
|
||||
}
|
||||
});
|
||||
|
||||
document.getElementById('btnCamera').addEventListener('click', () => {
|
||||
imageInput.click();
|
||||
});
|
||||
|
||||
function processFile(file) {
|
||||
if (!file.type.startsWith('image/')) {
|
||||
alert(i18n.invalidImage);
|
||||
return;
|
||||
}
|
||||
|
||||
const formData = new FormData();
|
||||
formData.append('image', file);
|
||||
formData.append('confidence_threshold', confSlider.value);
|
||||
formData.append('csrfmiddlewaretoken', '{{ csrf_token }}');
|
||||
|
||||
runAnalysis(formData);
|
||||
}
|
||||
|
||||
function runAnalysis(formData) {
|
||||
showLoading(true);
|
||||
|
||||
fetch("{% url 'sign:ai_demo' %}", {
|
||||
method: 'POST',
|
||||
body: formData,
|
||||
headers: {
|
||||
'X-Requested-With': 'XMLHttpRequest'
|
||||
}
|
||||
})
|
||||
.then(response => response.json())
|
||||
.then(data => {
|
||||
showLoading(false);
|
||||
if (data.status === 'success') {
|
||||
renderResults(data);
|
||||
} else {
|
||||
alert(data.message || i18n.errorProcessing);
|
||||
}
|
||||
})
|
||||
.catch(err => {
|
||||
showLoading(false);
|
||||
console.error(err);
|
||||
alert(i18n.networkError);
|
||||
});
|
||||
}
|
||||
|
||||
function showLoading(isLoading) {
|
||||
const emptyState = document.getElementById('emptyStateContainer');
|
||||
const loading = document.getElementById('loadingContainer');
|
||||
const imgResult = document.getElementById('imageResultContainer');
|
||||
|
||||
if (isLoading) {
|
||||
emptyState.classList.add('d-none');
|
||||
imgResult.classList.add('d-none');
|
||||
loading.classList.remove('d-none');
|
||||
} else {
|
||||
loading.classList.add('d-none');
|
||||
}
|
||||
}
|
||||
|
||||
function renderResults(data) {
|
||||
lastAnalysisResult = data;
|
||||
document.getElementById('imageResultContainer').classList.remove('d-none');
|
||||
document.getElementById('resultsSection').classList.remove('d-none');
|
||||
|
||||
// Image annotée
|
||||
document.getElementById('annotatedPreview').src = data.annotated_image;
|
||||
document.getElementById('detectedBadge').textContent = `${data.detected_count} panneau(x) détecté(s)`;
|
||||
|
||||
// Performances
|
||||
document.getElementById('perfYolo').textContent = `${data.performance.yolo_inference_ms} ms`;
|
||||
document.getElementById('perfOcr').textContent = `${data.performance.ocr_inference_ms} ms`;
|
||||
document.getElementById('perfTotal').textContent = `${data.performance.total_processing_ms} ms`;
|
||||
|
||||
// JSON
|
||||
document.getElementById('jsonPayload').textContent = JSON.stringify(data, null, 2);
|
||||
|
||||
// Cartes de panneaux
|
||||
const container = document.getElementById('panelsListContainer');
|
||||
container.innerHTML = '';
|
||||
|
||||
if (!data.panels || data.panels.length === 0) {
|
||||
container.innerHTML = '<div class="col-12 text-center text-muted py-3">' + i18n.noPanelFound + '</div>';
|
||||
return;
|
||||
}
|
||||
|
||||
data.panels.forEach(p => {
|
||||
const confPct = Math.round(p.confidence * 100);
|
||||
const confColor = confPct >= 75 ? 'bg-success' : (confPct >= 50 ? 'bg-warning' : 'bg-secondary');
|
||||
const svgSrc = p.svg_url || '';
|
||||
|
||||
const card = document.createElement('div');
|
||||
card.className = 'col-lg-6';
|
||||
card.innerHTML = `
|
||||
<div class="panel-card p-3 h-100 d-flex flex-column">
|
||||
<div class="d-flex align-items-start gap-3">
|
||||
<div class="text-center">
|
||||
<span class="order-badge bg-primary text-white mb-2 shadow-sm">
|
||||
#${p.vertical_order}
|
||||
</span>
|
||||
<div class="p-2 bg-light rounded-3 border d-flex align-items-center justify-content-center" style="width: 90px; height: 90px;">
|
||||
${svgSrc ? `<img src="${svgSrc}" alt="${p.code}" class="svg-sign-preview" onerror="this.onerror=null; this.parentElement.innerHTML='<i class=\\'bi bi-signpost-2 fs-1 text-secondary\\'></i>';">` : `<i class="bi bi-signpost-2 fs-1 text-secondary"></i>`}
|
||||
</div>
|
||||
<span class="badge bg-dark mt-2">${p.code || 'Inconnu'}</span>
|
||||
</div>
|
||||
<div class="flex-grow-1">
|
||||
<div class="d-flex justify-content-between align-items-start mb-1">
|
||||
<h6 class="fw-bold text-dark mb-0">${p.name_fr || p.name_nl || 'Panneau détecté'}</h6>
|
||||
<span class="badge ${confColor} px-2 py-1">${confPct}%</span>
|
||||
</div>
|
||||
<p class="text-muted small mb-2">${p.name_nl ? `<span class="fst-italic">${p.name_nl}</span>` : ''}</p>
|
||||
|
||||
<!-- Texte OCR / Panonceau -->
|
||||
${p.ocr_text ? `
|
||||
<div class="p-2 bg-light rounded-2 border border-secondary-subtle mb-2">
|
||||
<div class="d-flex align-items-center justify-content-between">
|
||||
<span class="small fw-semibold text-secondary">
|
||||
<i class="bi bi-fonts me-1"></i>Texte OCR extrait :
|
||||
</span>
|
||||
</div>
|
||||
<div class="fw-bold text-dark font-monospace mt-1">"${p.ocr_text}"</div>
|
||||
</div>
|
||||
` : ''}
|
||||
|
||||
<!-- Valeur numérique -->
|
||||
${p.signpanel_value !== null && p.signpanel_value !== undefined ? `
|
||||
<div class="small mb-1">
|
||||
<span class="text-muted fw-semibold">Valeur / Distance :</span>
|
||||
<span class="badge bg-info-subtle text-info-emphasis fw-bold">${p.signpanel_value}</span>
|
||||
</div>
|
||||
` : ''}
|
||||
|
||||
<div class="small text-muted mt-2">
|
||||
<i class="bi bi-arrows-move me-1"></i>Position verticale : <strong>#${p.vertical_order} sur le mât</strong>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
`;
|
||||
container.appendChild(card);
|
||||
});
|
||||
}
|
||||
|
||||
function copyJsonPayload() {
|
||||
if (!lastAnalysisResult) return;
|
||||
navigator.clipboard.writeText(JSON.stringify(lastAnalysisResult, null, 2))
|
||||
.then(() => alert(i18n.jsonCopied));
|
||||
}
|
||||
|
||||
// Génération d'exemples synthétiques pour tests immédiats sans photo réelle
|
||||
function loadSyntheticExample(type) {
|
||||
const canvas = document.getElementById('synthCanvas');
|
||||
const ctx = canvas.getContext('2d');
|
||||
ctx.fillStyle = '#f1f5f9';
|
||||
ctx.fillRect(0, 0, 600, 600);
|
||||
|
||||
if (type === 'stop') {
|
||||
// Cercle STOP
|
||||
ctx.fillStyle = '#dc2626';
|
||||
ctx.beginPath();
|
||||
ctx.arc(300, 200, 100, 0, Math.PI * 2);
|
||||
ctx.fill();
|
||||
ctx.fillStyle = '#ffffff';
|
||||
ctx.font = 'bold 44px Arial';
|
||||
ctx.textAlign = 'center';
|
||||
ctx.fillText('STOP', 300, 215);
|
||||
|
||||
// Panonceau Sauf riverains
|
||||
ctx.fillStyle = '#ffffff';
|
||||
ctx.fillRect(160, 370, 280, 90);
|
||||
ctx.strokeStyle = '#000000';
|
||||
ctx.lineWidth = 3;
|
||||
ctx.strokeRect(160, 370, 280, 90);
|
||||
ctx.fillStyle = '#000000';
|
||||
ctx.font = 'bold 24px Arial';
|
||||
ctx.fillText('Sauf riverains', 300, 425);
|
||||
} else if (type === 'zone30') {
|
||||
// Rectangle Zone 30
|
||||
ctx.fillStyle = '#ffffff';
|
||||
ctx.fillRect(150, 100, 300, 400);
|
||||
ctx.strokeStyle = '#000000';
|
||||
ctx.lineWidth = 4;
|
||||
ctx.strokeRect(150, 100, 300, 400);
|
||||
|
||||
ctx.fillStyle = '#000000';
|
||||
ctx.font = 'bold 36px Arial';
|
||||
ctx.textAlign = 'center';
|
||||
ctx.fillText('ZONE', 300, 170);
|
||||
|
||||
// Cercle rouge
|
||||
ctx.fillStyle = '#dc2626';
|
||||
ctx.beginPath();
|
||||
ctx.arc(300, 300, 90, 0, Math.PI * 2);
|
||||
ctx.fill();
|
||||
ctx.fillStyle = '#ffffff';
|
||||
ctx.beginPath();
|
||||
ctx.arc(300, 300, 72, 0, Math.PI * 2);
|
||||
ctx.fill();
|
||||
ctx.fillStyle = '#000000';
|
||||
ctx.font = 'bold 64px Arial';
|
||||
ctx.fillText('30', 300, 322);
|
||||
} else if (type === 'distance') {
|
||||
// Panneau Cédez le passage
|
||||
ctx.fillStyle = '#ffffff';
|
||||
ctx.beginPath();
|
||||
ctx.moveTo(150, 120);
|
||||
ctx.lineTo(450, 120);
|
||||
ctx.lineTo(300, 350);
|
||||
ctx.closePath();
|
||||
ctx.fill();
|
||||
ctx.strokeStyle = '#dc2626';
|
||||
ctx.lineWidth = 14;
|
||||
ctx.stroke();
|
||||
|
||||
// Panonceau 300 m
|
||||
ctx.fillStyle = '#ffffff';
|
||||
ctx.fillRect(180, 400, 240, 90);
|
||||
ctx.strokeStyle = '#000000';
|
||||
ctx.lineWidth = 3;
|
||||
ctx.strokeRect(180, 400, 240, 90);
|
||||
ctx.fillStyle = '#000000';
|
||||
ctx.font = 'bold 32px Arial';
|
||||
ctx.textAlign = 'center';
|
||||
ctx.fillText('300 m', 300, 455);
|
||||
}
|
||||
|
||||
const dataUrl = canvas.toDataURL('image/jpeg', 0.9);
|
||||
const formData = new FormData();
|
||||
formData.append('image_base64', dataUrl);
|
||||
formData.append('confidence_threshold', confSlider.value);
|
||||
formData.append('csrfmiddlewaretoken', '{{ csrf_token }}');
|
||||
|
||||
runAnalysis(formData);
|
||||
}
|
||||
</script>
|
||||
{% endblock %}
|
||||
|
|
@ -133,6 +133,11 @@
|
|||
<i class="bi bi-pencil me-1"></i>{% translate "Détail" %}
|
||||
</button>
|
||||
</li>
|
||||
<li class="nav-item ms-auto" role="presentation">
|
||||
<a href="{% url 'sign:ai_demo' %}" class="nav-link text-primary fw-semibold" title="{% translate 'Banc de test IA - Reconnaissance & OCR' %}">
|
||||
<i class="bi bi-camera-fill text-warning me-1"></i>{% translate "Banc IA" %}
|
||||
</a>
|
||||
</li>
|
||||
</ul>
|
||||
</div>
|
||||
|
||||
|
|
|
|||
190
loko/sign/tests_ai.py
Normal file
190
loko/sign/tests_ai.py
Normal file
|
|
@ -0,0 +1,190 @@
|
|||
"""
|
||||
Tests unitaires et d'intégration pour le module IA de reconnaissance de signalisation (sign.ai).
|
||||
"""
|
||||
import io
|
||||
import json
|
||||
import numpy as np
|
||||
import cv2
|
||||
from PIL import Image
|
||||
|
||||
from django.test import TestCase
|
||||
from django.urls import reverse
|
||||
from django.contrib.auth import get_user_model
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
|
||||
from rest_framework.test import APITestCase
|
||||
from rest_framework import status
|
||||
|
||||
from sign.ai import SignDetectionService, match_sign_from_ocr, get_svg_url
|
||||
from sign.models import SignPanelType
|
||||
|
||||
User = get_user_model()
|
||||
|
||||
|
||||
class SignCatalogMatcherTests(TestCase):
|
||||
"""Tests pour le matching de texte OCR avec le catalogue de panneaux."""
|
||||
|
||||
def test_match_stop_sign(self):
|
||||
res = match_sign_from_ocr("STOP")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "B5")
|
||||
self.assertEqual(res["data"]["category"], "priority")
|
||||
self.assertTrue(res["svg_url"].endswith("B5.svg"))
|
||||
|
||||
def test_match_zone_30(self):
|
||||
res = match_sign_from_ocr("ZONE 30")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "F4A")
|
||||
self.assertEqual(res["value"], 30)
|
||||
|
||||
def test_match_end_zone(self):
|
||||
res = match_sign_from_ocr("FIN DE ZONE")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "F4B")
|
||||
|
||||
def test_match_exception_panonceau(self):
|
||||
res = match_sign_from_ocr("Sauf riverains")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "M2")
|
||||
self.assertEqual(res["extracted_text"], "Sauf riverains")
|
||||
|
||||
def test_match_distance_panonceau(self):
|
||||
res = match_sign_from_ocr("300 m")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "M1")
|
||||
self.assertEqual(res["value"], 300.0)
|
||||
|
||||
def test_match_speed_limit(self):
|
||||
res = match_sign_from_ocr("50")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "C43")
|
||||
self.assertEqual(res["value"], 50)
|
||||
|
||||
def test_match_tonnage_sign(self):
|
||||
res = match_sign_from_ocr("3.5 t")
|
||||
self.assertIsNotNone(res)
|
||||
self.assertEqual(res["code"], "C21")
|
||||
self.assertEqual(res["value"], 3.5)
|
||||
|
||||
def test_visual_classification_blue_bike(self):
|
||||
from sign.ai import classify_sign_visual
|
||||
# Simuler un rond bleu
|
||||
img = np.ones((100, 100, 3), dtype=np.uint8) * 180
|
||||
cv2.circle(img, (50, 50), 45, (200, 50, 20), -1) # BGR Bleu
|
||||
cv2.circle(img, (50, 50), 15, (255, 255, 255), -1) # Blanc centre
|
||||
res = classify_sign_visual(img)
|
||||
self.assertEqual(res["code"], "D7")
|
||||
self.assertEqual(res["category"], "obligation")
|
||||
|
||||
def test_visual_classification_yellow_priority(self):
|
||||
from sign.ai import classify_sign_visual
|
||||
img = np.ones((100, 100, 3), dtype=np.uint8) * 180
|
||||
cv2.circle(img, (50, 50), 45, (0, 220, 220), -1) # BGR Jaune
|
||||
res = classify_sign_visual(img)
|
||||
self.assertEqual(res["code"], "B3")
|
||||
self.assertEqual(res["category"], "priority")
|
||||
|
||||
def test_visual_classification_sens_interdit_c1(self):
|
||||
from sign.ai import classify_sign_visual
|
||||
img = np.ones((100, 100, 3), dtype=np.uint8) * 180
|
||||
cv2.circle(img, (50, 50), 45, (30, 30, 220), -1) # Cercle rouge
|
||||
cv2.rectangle(img, (15, 42), (85, 58), (255, 255, 255), -1) # Barre blanche
|
||||
res = classify_sign_visual(img)
|
||||
self.assertEqual(res["code"], "C1")
|
||||
self.assertEqual(res["category"], "prohibition")
|
||||
|
||||
def test_get_svg_url(self):
|
||||
self.assertEqual(get_svg_url("B1"), "/static/assets/road_signs/2025/B1.svg")
|
||||
self.assertEqual(get_svg_url(""), "")
|
||||
|
||||
|
||||
class SignDetectionServiceTests(TestCase):
|
||||
"""Tests d'inférence du service SignDetectionService."""
|
||||
|
||||
def setUp(self):
|
||||
self.service = SignDetectionService.get_instance()
|
||||
|
||||
def _create_synthetic_image(self) -> np.ndarray:
|
||||
img = np.ones((500, 500, 3), dtype=np.uint8) * 240
|
||||
# Dessiner un panneau STOP
|
||||
cv2.circle(img, (250, 180), 90, (0, 0, 220), -1)
|
||||
cv2.putText(img, "STOP", (195, 195), cv2.FONT_HERSHEY_SIMPLEX, 1.3, (255, 255, 255), 3)
|
||||
# Panonceau sous le panneau
|
||||
cv2.rectangle(img, (150, 320), (350, 400), (255, 255, 255), -1)
|
||||
cv2.rectangle(img, (150, 320), (350, 400), (0, 0, 0), 2)
|
||||
cv2.putText(img, "Sauf velos", (165, 365), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 0, 0), 2)
|
||||
return img
|
||||
|
||||
def test_analyze_image_synthetic(self):
|
||||
img = self._create_synthetic_image()
|
||||
result = self.service.analyze_image(img)
|
||||
|
||||
self.assertEqual(result["status"], "success")
|
||||
self.assertGreaterEqual(result["detected_count"], 1)
|
||||
self.assertIn("performance", result)
|
||||
self.assertIn("yolo_inference_ms", result["performance"])
|
||||
self.assertIn("ocr_inference_ms", result["performance"])
|
||||
self.assertIn("annotated_image", result)
|
||||
self.assertTrue(result["annotated_image"].startswith("data:image/jpeg;base64,"))
|
||||
|
||||
# Vérifier l'ordonnancement vertical
|
||||
orders = [p["vertical_order"] for p in result["panels"]]
|
||||
self.assertEqual(orders, list(range(1, len(orders) + 1)))
|
||||
|
||||
|
||||
class SignAIApiAndViewsTests(APITestCase):
|
||||
"""Tests des endpoints API et de la vue de démo."""
|
||||
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(
|
||||
username="tester_ai",
|
||||
email="tester_ai@example.com",
|
||||
password="testpassword123"
|
||||
)
|
||||
self.client.force_login(self.user)
|
||||
|
||||
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
|
||||
img = np.ones((400, 400, 3), dtype=np.uint8) * 240
|
||||
cv2.circle(img, (200, 150), 80, (0, 0, 220), -1)
|
||||
cv2.putText(img, "STOP", (150, 165), cv2.FONT_HERSHEY_SIMPLEX, 1.1, (255, 255, 255), 3)
|
||||
|
||||
_, buf = cv2.imencode(".jpg", img)
|
||||
return SimpleUploadedFile("test_sign.jpg", buf.tobytes(), content_type="image/jpeg")
|
||||
|
||||
def test_ai_demo_view_get(self):
|
||||
url = reverse("sign:ai_demo")
|
||||
response = self.client.get(url)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertContains(response, "Reconnaissance de Signalisation & OCR")
|
||||
|
||||
def test_ai_demo_view_post_file(self):
|
||||
url = reverse("sign:ai_demo")
|
||||
file_obj = self._create_uploaded_image_file()
|
||||
response = self.client.post(
|
||||
url,
|
||||
{"image": file_obj, "confidence_threshold": "0.20"},
|
||||
HTTP_X_REQUESTED_WITH="XMLHttpRequest"
|
||||
)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
data = response.json()
|
||||
self.assertEqual(data["status"], "success")
|
||||
self.assertGreaterEqual(data["detected_count"], 1)
|
||||
|
||||
def test_api_detect_sign_post(self):
|
||||
url = reverse("sign:api_detect")
|
||||
file_obj = self._create_uploaded_image_file()
|
||||
response = self.client.post(
|
||||
url,
|
||||
{"image": file_obj, "confidence_threshold": 0.20},
|
||||
format="multipart"
|
||||
)
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
data = response.json()
|
||||
self.assertEqual(data["status"], "success")
|
||||
self.assertIn("panels", data)
|
||||
self.assertIn("performance", data)
|
||||
|
||||
def test_api_detect_sign_missing_image(self):
|
||||
url = reverse("sign:api_detect")
|
||||
response = self.client.post(url, {}, format="json")
|
||||
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
||||
|
|
@ -1,10 +1,13 @@
|
|||
from django.urls import path
|
||||
|
||||
from . import views
|
||||
from . import views_ai
|
||||
|
||||
app_name = "sign"
|
||||
urlpatterns = [
|
||||
path("", views.index, name="index"),
|
||||
path("ai/demo/", views_ai.SignAIDemoView.as_view(), name="ai_demo"),
|
||||
path("api/detect/", views_ai.DetectSignAPIView.as_view(), name="api_detect"),
|
||||
path("streets/", views.sign_streets_list, name="sign_streets_list"),
|
||||
path("streets/<int:street_id>/", views.sign_streets_detail, name="sign_streets_detail"),
|
||||
path("assets/<str:asset_model>/<int:asset_id>/", views.sign_assets_detail, name="sign_assets_detail"),
|
||||
|
|
|
|||
112
loko/sign/views_ai.py
Normal file
112
loko/sign/views_ai.py
Normal file
|
|
@ -0,0 +1,112 @@
|
|||
"""
|
||||
Vues et API pour la détection et l'OCR de panneaux de signalisation routière.
|
||||
"""
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, Dict
|
||||
|
||||
from django.views.generic import TemplateView
|
||||
from django.contrib.auth.mixins import LoginRequiredMixin
|
||||
from django.http import JsonResponse, HttpResponseBadRequest
|
||||
from django.utils.translation import gettext_lazy as _
|
||||
from django.views.decorators.csrf import csrf_exempt
|
||||
from django.utils.decorators import method_decorator
|
||||
|
||||
from rest_framework.views import APIView
|
||||
from rest_framework.response import Response
|
||||
from rest_framework import status
|
||||
from rest_framework.permissions import IsAuthenticated
|
||||
|
||||
from .ai import SignDetectionService, SIGN_CATALOG, get_svg_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SignAIDemoView(LoginRequiredMixin, TemplateView):
|
||||
"""
|
||||
Page de démonstration et banc de test interactif pour la reconnaissance de panneaux et panonceaux.
|
||||
Permet d'uploader ou capturer une photo et de visualiser instantanément les détections,
|
||||
les textes OCR, l'ordre vertical et les fichiers vectoriels SVG correspondants.
|
||||
"""
|
||||
template_name = "sign/ai_demo.html"
|
||||
|
||||
def get_context_data(self, **kwargs: Any) -> Dict[str, Any]:
|
||||
context = super().get_context_data(**kwargs)
|
||||
context["catalog_sample"] = list(SIGN_CATALOG.items())[:12]
|
||||
return context
|
||||
|
||||
def post(self, request, *args, **kwargs):
|
||||
image_file = request.FILES.get("image")
|
||||
image_base64 = request.POST.get("image_base64")
|
||||
confidence_threshold = float(request.POST.get("confidence_threshold", 0.25))
|
||||
|
||||
if not image_file and not image_base64:
|
||||
return JsonResponse(
|
||||
{"status": "error", "message": _("Aucune image fournie.")},
|
||||
status=400
|
||||
)
|
||||
|
||||
try:
|
||||
service = SignDetectionService.get_instance()
|
||||
if image_file:
|
||||
image_bytes = image_file.read()
|
||||
result = service.analyze_image(image_bytes, confidence_threshold=confidence_threshold)
|
||||
else:
|
||||
# Format data:image/...;base64,...
|
||||
if "," in image_base64:
|
||||
image_base64 = image_base64.split(",", 1)[1]
|
||||
import base64
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
result = service.analyze_image(image_bytes, confidence_threshold=confidence_threshold)
|
||||
|
||||
return JsonResponse(result)
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("Erreur lors de l'analyse IA de signalisation : %s", exc, exc_info=True)
|
||||
return JsonResponse(
|
||||
{"status": "error", "message": f"Erreur lors du traitement de l'image : {str(exc)}"},
|
||||
status=500
|
||||
)
|
||||
|
||||
|
||||
class DetectSignAPIView(APIView):
|
||||
"""
|
||||
Endpoint API REST pour l'analyse mobile/terrain d'une photo de signalisation.
|
||||
Accepte multipart/form-data ('image') ou JSON avec 'image_base64'.
|
||||
|
||||
Retourne la liste ordonnée des panneaux détectés, le texte OCR extrait,
|
||||
les valeurs numériques et le chemin du SVG officiel correspondant.
|
||||
"""
|
||||
permission_classes = [IsAuthenticated]
|
||||
|
||||
def post(self, request, *args, **kwargs):
|
||||
image_file = request.FILES.get("image")
|
||||
image_base64 = request.data.get("image_base64")
|
||||
confidence_threshold = float(request.data.get("confidence_threshold", 0.25))
|
||||
|
||||
if not image_file and not image_base64:
|
||||
return Response(
|
||||
{"status": "error", "message": "Paramètre 'image' ou 'image_base64' manquant."},
|
||||
status=status.HTTP_400_BAD_REQUEST
|
||||
)
|
||||
|
||||
try:
|
||||
service = SignDetectionService.get_instance()
|
||||
if image_file:
|
||||
image_bytes = image_file.read()
|
||||
else:
|
||||
if "," in image_base64:
|
||||
image_base64 = image_base64.split(",", 1)[1]
|
||||
import base64
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
|
||||
result = service.analyze_image(image_bytes, confidence_threshold=confidence_threshold)
|
||||
return Response(result, status=status.HTTP_200_OK)
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("Erreur API détection signalisation : %s", exc, exc_info=True)
|
||||
return Response(
|
||||
{"status": "error", "message": str(exc)},
|
||||
status=status.HTTP_500_INTERNAL_SERVER_ERROR
|
||||
)
|
||||
|
|
@ -42,3 +42,5 @@ ezdxf==1.4.4
|
|||
matplotlib==3.10.9
|
||||
boto3>=1.34.0
|
||||
markdown
|
||||
onnxruntime>=1.19.0
|
||||
rapidocr-onnxruntime>=1.2.0
|
||||
|
|
|
|||
Loading…
Reference in a new issue