diff --git a/loko/sign/ai/__init__.py b/loko/sign/ai/__init__.py index 1468bff..ca92b40 100644 --- a/loko/sign/ai/__init__.py +++ b/loko/sign/ai/__init__.py @@ -1,6 +1,7 @@ from .detector import SignDetectionService -from .catalog import SIGN_CATALOG, get_svg_url, match_sign_from_ocr, classify_sign_visual +from .catalog import SIGN_CATALOG, get_svg_url, match_sign_from_ocr, classify_sign_visual, get_all_catalog_signs from .classifier import SignClassifierEngine, SyntheticSignAugmentor +from .ground_truth import GroundTruthDatasetManager __all__ = [ "SignDetectionService", @@ -8,6 +9,10 @@ __all__ = [ "get_svg_url", "match_sign_from_ocr", "classify_sign_visual", + "get_all_catalog_signs", "SignClassifierEngine", "SyntheticSignAugmentor", + "GroundTruthDatasetManager", ] + + diff --git a/loko/sign/ai/catalog.py b/loko/sign/ai/catalog.py index 743732c..74e4812 100644 --- a/loko/sign/ai/catalog.py +++ b/loko/sign/ai/catalog.py @@ -5,6 +5,8 @@ et leurs fichiers vectoriels SVG correspondants. """ import re from typing import Optional, Dict, Any, List +import numpy as np +import cv2 # 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/" @@ -299,6 +301,60 @@ SIGN_CATALOG = { "category": "panonceau", "shape": "rectangle", }, + "TYPE0": { + "name_fr": "Panneau additionnel d'exception ou mention (fond bleu)", + "name_nl": "Blauw onderbord met witte tekst", + "category": "panonceau", + "shape": "rectangle", + }, + "TYPE0B": { + "name_fr": "Panneau additionnel d'exception ou mention (fond blanc)", + "name_nl": "Wit onderbord met zwarte tekst", + "category": "panonceau", + "shape": "rectangle", + }, + "TYPEIA_50M": { + "name_fr": "Panneau additionnel de distance (50 m - fond bleu)", + "name_nl": "Afstandsbord 50 m (blauwe achtergrond)", + "category": "panonceau", + "shape": "rectangle", + }, + "TYPEIA_200M": { + "name_fr": "Panneau additionnel de distance (200 m - fond bleu)", + "name_nl": "Afstandsbord 200 m (blauwe achtergrond)", + "category": "panonceau", + "shape": "rectangle", + }, + "TYPEIA_300M": { + "name_fr": "Panneau additionnel de distance (300 m - fond bleu)", + "name_nl": "Afstandsbord 300 m (blauwe achtergrond)", + "category": "panonceau", + "shape": "rectangle", + }, + "TYPEIA_GEN": { + "name_fr": "Panneau additionnel de distance (fond bleu)", + "name_nl": "Afstandsbord (blauwe achtergrond)", + "category": "panonceau", + "shape": "rectangle", + }, + "TYPEIB": { + "name_fr": "Panneau additionnel d'étendue avec flèches (fond bleu)", + "name_nl": "Uitgestrektheidsbord met pijlen (blauwe achtergrond)", + "category": "panonceau", + "shape": "rectangle", + }, + "GXC": { + "name_fr": "Début ou longueur de zone de stationnement (flèche montante)", + "name_nl": "Begin of lengte parkeerzone (opwaartse pijl)", + "category": "panonceau", + "shape": "rectangle", + }, + "XD": { + "name_fr": "Panonceau additionnel avec flèche", + "name_nl": "Onderbord met pijl", + "category": "panonceau", + "shape": "rectangle", + }, } @@ -340,29 +396,36 @@ def get_svg_url(sign_code: str) -> str: return f"{DEFAULT_SVG_BASE_PATH}{code_upper}.svg" -def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: +def match_sign_from_ocr(ocr_text: str, crop_bgr: Optional[np.ndarray] = None) -> Optional[Dict[str, Any]]: """ - Analyse un texte OCR et tente d'associer un type de panneau normalisé. - Exemples: - - "STOP" -> B5 - - "50 km" ou "50" -> C43 (Limitation de vitesse 50 km/h) - - "ZONE P" ou "ZONE ... EXCEPTE CARTE" -> ZE9A (Zone de stationnement) - - "ZONE 30" -> F4A (Zone 30) - - "SAUF RIVERAINS" -> M2 - - "300 M" -> M1 (Distance) + Tente d'associer un texte extrait par OCR à un type de panneau normalisé du catalogue. + Intègre une vérification colorimétrique stricte : + - Si absence de rouge (red_ratio < 0.035), élimine strictement F4A (Zone 30), C43 (Limitation), C... et A... + - Détecte les panonceaux de distance / flèche montante GXC / XD + - Simplifie les panonceaux textuels bruts : TYPE0 (fond bleu) ou TYPE0B (fond blanc) """ if not ocr_text: return None - cleaned = ocr_text.strip().upper() - cleaned_inline = re.sub(r"\s+", " ", cleaned) + cleaned = re.sub(r"[^A-Za-z0-9\s/.,:-]", " ", ocr_text.upper()) + cleaned_inline = re.sub(r"\s+", " ", cleaned).strip() + if not cleaned_inline: + return None + + # Extraction du profil de couleur si l'image est fournie + color_prof = extract_sign_color_profile(crop_bgr) if (crop_bgr is not None and isinstance(crop_bgr, np.ndarray) and crop_bgr.size > 0) else { + "red_ratio": 0.5, "blue_ratio": 0.5, "yellow_ratio": 0.0, "white_ratio": 0.5, + "has_red_and_blue": False, "is_pure_blue": False, "is_red_and_white": False, "is_yellow": False + } + has_red = bool(color_prof.get("red_ratio", 0.0) >= 0.035 or color_prof.get("has_red_and_blue")) + has_blue = bool(color_prof.get("blue_ratio", 0.0) >= 0.10) # 1. STOP if "STOP" in cleaned_inline: return { "code": "B5", "confidence": 0.96, - "data": SIGN_CATALOG["B5"], + "data": SIGN_CATALOG.get("B5", {"name_fr": "Arrêt obligatoire (STOP)", "name_nl": "Verplichte stop", "category": "priority"}), "svg_url": get_svg_url("B5"), "matched_by": "text_stop", } @@ -377,23 +440,26 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "svg_url": get_svg_url("F4B"), "matched_by": "text_end_zone", } - # Zone Stationnement / Parking ("ZONE P", "ZONE ... CARTE DE STATIONNEMENT", "ZONE ... PARKEERKAART", "ZONE ... DISQUE") - if re.search(r"\b(ZONE\s+P\b|PARKING|PARKEREN|STATIONNEMENT|PARKEER|DISQUE|PARKEERSCHIJF|CARTE|KAART)\b", cleaned_inline): + + # Zone Stationnement / Parking ("ZONE P", "ZONE ... CARTE", "ZONE ... DISQUE", ou ZONE SANS ROUGE avec BLEU) + if re.search(r"\b(ZONE\s+P\b|PARKING|PARKEREN|STATIONNEMENT|PARKEER|DISQUE|PARKEERSCHIJF|CARTE|KAART|RAPPEL|HERHALING)\b", cleaned_inline) or (has_blue and not has_red): + code_zone = "ZE9B" if re.search(r"\b(PMR|HANDICAP|GEHANDICAPT)\b", cleaned_inline) else "ZE9A" return { - "code": "ZE9A", + "code": code_zone, "confidence": 0.95, - "data": SIGN_CATALOG.get("ZE9A", { + "data": SIGN_CATALOG.get(code_zone, { "name_fr": "Zone de stationnement réglementé", "name_nl": "Zone voor gereglementeerd parkeren", "category": "parking", }), - "svg_url": get_svg_url("ZE9A"), + "svg_url": get_svg_url(code_zone), "extracted_text": ocr_text.strip(), "matched_by": "text_zone_parking", } - # Zone de vitesse ("ZONE 30", "ZONE 20", "ZONE 50") + + # Zone de vitesse ("ZONE 30", "ZONE 20", "ZONE 50") -> UNIQUEMENT SI DU ROUGE EST PRÉSENT speed_match = re.search(r"\b(20|30|50|70)\b", cleaned_inline) - if speed_match: + if speed_match and has_red: speed = int(speed_match.group(1)) code = "F4A" return { @@ -404,6 +470,7 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "value": speed, "matched_by": "text_zone_speed", } + # Zone piétonne if re.search(r"\b(PIETON|VOETGANGER|PIETONS|VOETGANGERS)\b", cleaned_inline): return { @@ -417,21 +484,21 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "svg_url": get_svg_url("F103"), "matched_by": "text_zone_pedestrian", } - # Zone générique + + # Zone générique : si pas de rouge, c'est une zone de stationnement ZE9A + code_fallback = "F4A" if has_red else "ZE9A" return { - "code": "F4A", + "code": code_fallback, "confidence": 0.88, - "data": SIGN_CATALOG.get("F4A", {"name_fr": "Zone réglementée", "name_nl": "Gereglementeerde zone", "category": "indication"}), - "svg_url": get_svg_url("F4A"), + "data": SIGN_CATALOG.get(code_fallback, {"name_fr": "Zone réglementée", "name_nl": "Gereglementeerde zone", "category": "indication"}), + "svg_url": get_svg_url(code_fallback), "extracted_text": ocr_text.strip(), "matched_by": "text_zone", } - # 3. VITESSE MAXIMALE AUTORISÉE (C43 : "50", "50 km", "50 km/h", "30 km", "70 km/h", "90", "120") - # Note : "50 km" ou "50 km/h" sur un panneau de limitation est une vitesse C43 et NON une distance M1 ! + # 3. VITESSE MAXIMALE AUTORISÉE (C43 : "50", "50 km/h", "30 km") -> STRICTEMENT CONDITIONNÉE À LA PRÉSENCE DE ROUGE speed_match = re.search(r"\b(10|20|30|40|50|60|70|80|90|100|110|120|130)\s*(?:KM(?:/H|/U)?|KPH)?\b", cleaned_inline) - if speed_match: - # Exclure si le texte est explicitement une distance comme "50 m", "300 m", "1.5 km" (avec décimale ou mètres) + if speed_match and has_red: is_explicit_distance = bool(re.search(r"\b(\d+\s*M|\d+[,.]\d+\s*KM)\b", cleaned_inline)) if not is_explicit_distance: speed = int(speed_match.group(1)) @@ -450,25 +517,51 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "matched_by": "text_speed_limit", } - # 4. Panonceaux d'exception ("Sauf ...", "Excepté ...", "Uitgezonderd ...") - if re.search(r"\b(SAUF|EXCEPTE|EXCEPTE|UITGEZONDERD|RIVERAINS?|AANGELANDEN?|VELOS?|FIETS)\b", cleaned_inline, re.IGNORECASE): - return { - "code": "M2", - "confidence": 0.90, - "data": SIGN_CATALOG.get("M2", {"name_fr": "Panonceau d'application ou d'exception", "name_nl": "Onderbord: uitzondering", "category": "panonceau"}), - "svg_url": get_svg_url("M2"), - "extracted_text": ocr_text.strip(), - "matched_by": "text_exception", - } - - # 5. Panonceaux de distance ("300 m", "50 m", "1.5 km", "2.0 km") - dist_match = re.search(r"\b(\d+(?:[.,]\d+)?)\s*(M|METRES?|METERS?)\b|\b(\d+[.,]\d+)\s*(KM)\b", cleaned_inline) + # 4. Panonceaux de distance ou flèche de zone (ex: "50m", "11m", "12 m", "300 m", "50 m") + dist_match = re.search(r"\b(\d+(?:[.,]\d+)?)\s*(M|METRES?|METERS?)\b|\b(\d+[.,]\d+)\s*(KM)\b|\b(\d+)\s*M\b", cleaned_inline) if dist_match: - val_str = (dist_match.group(1) or dist_match.group(3)).replace(",", ".") + val_str = (dist_match.group(1) or dist_match.group(3) or dist_match.group(5) or "0").replace(",", ".") unit = (dist_match.group(2) or dist_match.group(4) or "m").lower() val = float(val_str) if unit == "km": val *= 1000.0 + + int_val = int(val) + + # A) Si fond BLEU -> Panneau additionnel de distance bleu TYPEIA (TYPEIA_50M, TYPEIA_200M, TYPEIA_300M ou TYPEIA_GEN) + if has_blue and not has_red: + code_blue = f"TYPEIA_{int_val}M" if int_val in (50, 200, 300) else "TYPEIA_GEN" + return { + "code": code_blue, + "confidence": 0.96, + "data": SIGN_CATALOG.get(code_blue, { + "name_fr": f"Panneau additionnel de distance ({int_val} m - fond bleu)", + "name_nl": f"Afstandsbord {int_val} m (blauwe achtergrond)", + "category": "panonceau", + }), + "svg_url": get_svg_url(code_blue) or get_svg_url("TYPEIA_GEN") or get_svg_url("TYPE0"), + "value": val, + "extracted_text": ocr_text.strip(), + "matched_by": "text_distance_blue_typeia", + } + + # B) Si fond BLANC rectangulaire avec distance courte (ex: 11m, 12m, 25m) -> Flèche de zone GXC / XD + if val <= 100 and not has_red and not has_blue: + return { + "code": "GXC", + "confidence": 0.93, + "data": SIGN_CATALOG.get("GXC", { + "name_fr": f"Début / longueur de zone ({int_val} m)", + "name_nl": f"Begin / lengte van de zone ({int_val} m)", + "category": "panonceau", + }), + "svg_url": get_svg_url("GXC") or get_svg_url("XD"), + "value": val, + "extracted_text": ocr_text.strip(), + "matched_by": "text_distance_arrow_gxc", + } + + # C) Fond BLANC avec distance générique -> M1 (Panonceau blanc) return { "code": "M1", "confidence": 0.88, @@ -476,44 +569,10 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "svg_url": get_svg_url("M1"), "value": val, "extracted_text": ocr_text.strip(), - "matched_by": "text_distance", + "matched_by": "text_distance_white_m1", } - # 6. Tonnage ("3.5 t", "7.5t") - ton_match = re.search(r"\b(\d+(?:[.,]\d+)?)\s*T\b", cleaned_inline) - if ton_match: - val = float(ton_match.group(1).replace(",", ".")) - return { - "code": "C21", - "confidence": 0.88, - "data": SIGN_CATALOG.get("C21", {"name_fr": "Accès interdit aux véhicules dont la masse en charge dépasse le tonnage indiqué", "name_nl": "Verboden toegang voor voertuigen met een hogere massa dan aangeduid", "category": "prohibition"}), - "svg_url": get_svg_url("C21"), - "value": val, - "extracted_text": ocr_text.strip(), - "matched_by": "text_tonnage", - } - - # 7. Parking P ("P", "PARKING", "PARKEREN" ou lettre "D" isolée due à l'OCR sur le P) - if re.search(r"^\s*([PD])\s*$", cleaned_inline) or re.search(r"\b(PARKING|PARKEREN)\b", cleaned_inline): - return { - "code": "E9A", - "confidence": 0.96, - "data": SIGN_CATALOG.get("E9A", {"name_fr": "Stationnement autorisé (Parking)", "name_nl": "Parkeren toegelaten (Parking)", "category": "parking"}), - "svg_url": get_svg_url("E9A"), - "matched_by": "text_parking", - } - - # 8. Parking PMR / Handicap - if re.search(r"\b(HANDICAP|PMR|HANDICAPE|GEHANDICAPT)\b", cleaned_inline): - return { - "code": "E9B", - "confidence": 0.94, - "data": SIGN_CATALOG.get("E9B", {"name_fr": "Stationnement réservé aux personnes handicapées", "name_nl": "Parkeren voorbehouden voor personen met een handicap", "category": "parking"}), - "svg_url": get_svg_url("E9B"), - "matched_by": "text_handicap", - } - - # 9. Stationnement Payant / Betalend + # 5. Stationnement Payant / Betalend if re.search(r"\b(PAYANT|BETALEND|HORODATEUR|TICKET)\b", cleaned_inline): return { "code": "GVII_BETALEND", @@ -528,7 +587,7 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "matched_by": "text_parking_payant", } - # 10. Véhicules électriques en charge + # 6. Véhicules électriques en charge if re.search(r"\b(OPLADEND|OPLADEN|ELEKTRISCH|ELECTRIQUE|RECHARGE|CHARGE)\b", cleaned_inline) and re.search(r"\b(VEHICULE|VOERTUIG|WAGEN|AUTO)\b", cleaned_inline): return { "code": "GVIID_ELEKTRISCHE_WAGENS", @@ -543,7 +602,7 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "matched_by": "text_electric_vehicle", } - # 11. Disque de stationnement / Zone bleue + # 7. Disque de stationnement / Zone bleue if re.search(r"\b(DISQUE|PARKEERSCHIJF)\b", cleaned_inline): return { "code": "E9A_PARKEERSCHIJF", @@ -558,9 +617,338 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]: "matched_by": "text_parking_disc", } + # 8. Parking P ("P", "PARKING", "PARKEREN" ou lettre "D" isolée due à l'OCR) + if (re.search(r"^\s*([PD])\s*$", cleaned_inline) or re.search(r"\b(PARKING|PARKEREN)\b", cleaned_inline)) and (has_blue or not has_red): + return { + "code": "E9A", + "confidence": 0.96, + "data": SIGN_CATALOG.get("E9A", {"name_fr": "Stationnement autorisé (Parking)", "name_nl": "Parkeren toegelaten (Parking)", "category": "parking"}), + "svg_url": get_svg_url("E9A"), + "matched_by": "text_parking", + } + + # 9. Parking PMR / Handicap + if re.search(r"\b(HANDICAP|PMR|HANDICAPE|GEHANDICAPT)\b", cleaned_inline): + return { + "code": "E9B", + "confidence": 0.94, + "data": SIGN_CATALOG.get("E9B", {"name_fr": "Stationnement réservé aux personnes handicapées", "name_nl": "Parkeren voorbehouden voor personen met een handicap", "category": "parking"}), + "svg_url": get_svg_url("E9B"), + "matched_by": "text_handicap", + } + + # 10. Panonceaux d'application ou d'exception ("Sauf ...", "Excepté ...", "Uitgezonderd ...", "Riverains", "Bewoners") + if re.search(r"\b(SAUF|EXCEPTE|EXCEPTE|UITGEZONDERD|RIVERAINS?|BEWONERS?|AANGELANDEN?|VELOS?|FIETS)\b", cleaned_inline, re.IGNORECASE): + # A) Si fond BLEU -> TYPE0 (Panneau additionnel bleu avec texte blanc d'exception) + if has_blue and not has_red: + return { + "code": "TYPE0", + "confidence": 0.95, + "data": SIGN_CATALOG.get("TYPE0", { + "name_fr": "Panneau additionnel d'exception (fond bleu)", + "name_nl": "Blauw uitzonderingsbord (witte tekst)", + "category": "panonceau", + }), + "svg_url": get_svg_url("TYPE0"), + "extracted_text": ocr_text.strip(), + "matched_by": "text_exception_blue_type0", + } + + # B) Si fond BLANC -> M2 (Panonceau blanc d'exception avec pictogramme vélo) + return { + "code": "M2", + "confidence": 0.90, + "data": SIGN_CATALOG.get("M2", {"name_fr": "Panonceau d'application ou d'exception", "name_nl": "Onderbord: uitzondering", "category": "panonceau"}), + "svg_url": get_svg_url("M2"), + "extracted_text": ocr_text.strip(), + "matched_by": "text_exception_white_m2", + } + + # 11. Tonnage ("3.5 t", "7.5t") + ton_match = re.search(r"\b(\d+(?:[.,]\d+)?)\s*T\b", cleaned_inline) + if ton_match and has_red: + val = float(ton_match.group(1).replace(",", ".")) + return { + "code": "C21", + "confidence": 0.88, + "data": SIGN_CATALOG.get("C21", {"name_fr": "Accès interdit aux véhicules dont la masse en charge dépasse le tonnage indiqué", "name_nl": "Verboden toegang voor voertuigen met een hogere massa dan aangeduid", "category": "prohibition"}), + "svg_url": get_svg_url("C21"), + "value": val, + "extracted_text": ocr_text.strip(), + "matched_by": "text_tonnage", + } + + # 12. Panneaux additionnels génériques avec texte libre (simplification TYPE0 fond bleu / TYPE0B fond blanc) + if len(cleaned_inline.split()) >= 2: + if has_blue and not has_red: + return { + "code": "TYPE0", + "confidence": 0.90, + "data": SIGN_CATALOG.get("TYPE0", { + "name_fr": "Panneau additionnel bleu à texte blanc", + "name_nl": "Blauw onderbord met witte tekst", + "category": "panonceau", + }), + "svg_url": get_svg_url("TYPE0"), + "extracted_text": ocr_text.strip(), + "matched_by": "text_subplate_blue_type0", + } + elif not has_red: + return { + "code": "TYPE0B", + "confidence": 0.90, + "data": SIGN_CATALOG.get("TYPE0B", { + "name_fr": "Panneau additionnel blanc à texte noir", + "name_nl": "Wit onderbord met zwarte tekst", + "category": "panonceau", + }), + "svg_url": get_svg_url("TYPE0B"), + "extracted_text": ocr_text.strip(), + "matched_by": "text_subplate_white_type0b", + } + return None +def extract_sign_color_profile(crop_bgr: np.ndarray) -> Dict[str, Any]: + """ + Analyse l'histogramme HSV et la distribution spatiale des couleurs d'un panneau découpé. + Détecte avec précision les proportions de Rouge, Bleu, Jaune, Blanc. + """ + import cv2 + import numpy as np + + if not isinstance(crop_bgr, np.ndarray) or crop_bgr.size == 0 or crop_bgr.shape[0] < 5 or crop_bgr.shape[1] < 5: + return { + "red_ratio": 0.0, + "blue_ratio": 0.0, + "yellow_ratio": 0.0, + "white_ratio": 0.0, + "has_red_and_blue": False, + "is_pure_blue": False, + "is_red_and_white": False, + "is_yellow": False, + } + + h, w = crop_bgr.shape[:2] + # Cadrage central (80% au cœur du crop) pour éliminer le décor d'arrière-plan + margin_y = int(h * 0.10) + margin_x = int(w * 0.10) + center_bgr = crop_bgr[margin_y:max(margin_y + 1, h - margin_y), margin_x:max(margin_x + 1, w - margin_x)] + if center_bgr.size == 0: + center_bgr = crop_bgr + + hsv = cv2.cvtColor(center_bgr, cv2.COLOR_BGR2HSV) + total_pixels = float(max(1, center_bgr.shape[0] * center_bgr.shape[1])) + + # 1. Rouge (deux plages en HSV: 0-12 et 165-180 avec saturation et valeur suffisantes) + red_mask1 = cv2.inRange(hsv, np.array([0, 50, 45]), np.array([12, 255, 255])) + red_mask2 = cv2.inRange(hsv, np.array([165, 50, 45]), np.array([180, 255, 255])) + red_mask = red_mask1 | red_mask2 + red_ratio = np.count_nonzero(red_mask) / total_pixels + + # 2. Bleu (plage 90-138 avec saturation et valeur suffisantes) + blue_mask = cv2.inRange(hsv, np.array([90, 50, 40]), np.array([138, 255, 255])) + blue_ratio = np.count_nonzero(blue_mask) / total_pixels + + # 3. Jaune (plage 14-38, sat > 65, val > 70) + yellow_mask = cv2.inRange(hsv, np.array([14, 65, 70]), np.array([38, 255, 255])) + yellow_ratio = np.count_nonzero(yellow_mask) / total_pixels + + # 4. Blanc / Gris clair + white_mask = cv2.inRange(hsv, np.array([0, 0, 115]), np.array([180, 48, 255])) + white_ratio = np.count_nonzero(white_mask) / total_pixels + + # Signatures colorimétriques distinctives : + # A) ROUGE + BLEU (ex: E1, E2, E3, E4, ZE...) + has_red_and_blue = bool(red_ratio >= 0.045 and blue_ratio >= 0.070) + + # B) BLEU PUR sans rouge (ex: D1A, D1B, D3, D5, D7, F19...) + is_pure_blue = bool(blue_ratio >= 0.12 and red_ratio < 0.035) + + # C) ROUGE + BLANC sans bleu (ex: C1, C3, C43, B1, B5, A...) + is_red_and_white = bool(red_ratio >= 0.070 and blue_ratio < 0.040) + + # D) JAUNE prioritaire (ex: B3) + is_yellow = bool(yellow_ratio >= 0.080 and red_ratio < 0.040 and blue_ratio < 0.040) + + return { + "red_ratio": round(red_ratio, 3), + "blue_ratio": round(blue_ratio, 3), + "yellow_ratio": round(yellow_ratio, 3), + "white_ratio": round(white_ratio, 3), + "has_red_and_blue": has_red_and_blue, + "is_pure_blue": is_pure_blue, + "is_red_and_white": is_red_and_white, + "is_yellow": is_yellow, + } + + +def discriminate_inner_pictogram( + crop_bgr: np.ndarray, + candidates: List[Dict[str, Any]] +) -> List[Dict[str, Any]]: + """ + Pour les panneaux dont la forme extérieure est identique mais dont le pictogramme + intérieur est discriminant (notamment la famille Danger A... et Interdiction C...) : + Analyse la géométrie, l'aspect-ratio et la structure du pictogramme central noir + pour corriger les confusions (ex: A25 Vélo vs A15 Piéton, A14 Dos d'âne). + """ + if not candidates or crop_bgr is None or not isinstance(crop_bgr, np.ndarray) or crop_bgr.size == 0: + return candidates + + top_code = (candidates[0].get("code") or "").upper().strip() + is_danger_triangle = bool(top_code.startswith("A") or any((c.get("code") or "").startswith("A") for c in candidates[:3])) + + if not is_danger_triangle: + return candidates + + h, w = crop_bgr.shape[:2] + if h < 24 or w < 24: + return candidates + + # Région intérieure du pictogramme (environ 30% du haut à 85% du bas, 20% à 80% en largeur) + y1, y2 = int(0.32 * h), int(0.85 * h) + x1, x2 = int(0.18 * w), int(0.82 * w) + inner = crop_bgr[y1:y2, x1:x2] + if inner.size == 0: + return candidates + + gray = cv2.cvtColor(inner, cv2.COLOR_BGR2GRAY) + # Détection des pixels sombres du pictogramme intérieur (en ignorant les zones blanches/rouges) + dark_mask = (gray < 85).astype(np.uint8) * 255 + contours, _ = cv2.findContours(dark_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + + if not contours: + return candidates + + # Bounding box globale du symbole intérieur + all_pts = np.vstack([c for c in contours]) + bx, by, bw, bh = cv2.boundingRect(all_pts) + if bw < 5 or bh < 5: + return candidates + + sym_aspect_ratio = float(bw) / float(bh) + + # Analyse des composantes (détection des deux roues pour A25 bicyclette / cycliste) + has_bicycle_wheels = False + if len(contours) >= 2 and sym_aspect_ratio >= 1.15: + # Recherche de deux contours distincts alignés horizontalement dans la partie basse + boxes = [cv2.boundingRect(c) for c in contours if cv2.contourArea(c) >= 12] + if len(boxes) >= 2: + boxes_sorted = sorted(boxes, key=lambda b: b[0]) + left_box, right_box = boxes_sorted[0], boxes_sorted[-1] + y_diff = abs((left_box[1] + left_box[3] / 2) - (right_box[1] + right_box[3] / 2)) + if y_diff < bh * 0.35 and (right_box[0] - (left_box[0] + left_box[2])) > 0: + has_bicycle_wheels = True + + # Ajustement des multiplicateurs selon la morphologie du pictogramme + adjusted = [] + for cand in candidates: + code = (cand.get("code") or "").upper().strip() + conf = float(cand.get("confidence", 0.0)) + mult = 1.0 + + if sym_aspect_ratio >= 1.25 or has_bicycle_wheels: + # Symbole large / horizontal (ex: Vélo A25 / A21, Cassis A14) + if code in ("A25", "A21", "M12"): + mult = 3.5 if has_bicycle_wheels else 2.2 + elif code in ("A14", "A27", "A29"): + mult = 1.5 + elif code == "A15": # Piéton (silhouette verticale) fortement pénalisé si symbole large + mult = 0.25 + elif sym_aspect_ratio <= 0.95: + # Symbole vertical / allongé (ex: Piéton A15, Danger indéterminé A51) + if code == "A15": + mult = 2.5 + elif code in ("A25", "A21", "A14"): + mult = 0.35 + + adjusted.append({**cand, "confidence": conf * mult}) + + total_conf = sum(c["confidence"] for c in adjusted) + if total_conf > 0: + for c in adjusted: + c["confidence"] = round(c["confidence"] / total_conf, 4) + + adjusted.sort(key=lambda x: x["confidence"], reverse=True) + return adjusted + + +def filter_and_rank_candidates_by_color( + candidates: List[Dict[str, Any]], + crop_bgr: np.ndarray +) -> List[Dict[str, Any]]: + """ + Applique les règles physiques de compatibilité colorimétrique et géométrique sur les prédictions IA. + """ + if not candidates: + return [] + + profile = extract_sign_color_profile(crop_bgr) + red_ratio = profile.get("red_ratio", 0.0) + blue_ratio = profile.get("blue_ratio", 0.0) + has_red_blue = profile["has_red_and_blue"] + is_pure_blue = profile["is_pure_blue"] + is_red_white = profile["is_red_and_white"] + is_yellow = profile["is_yellow"] + + adjusted_candidates = [] + for cand in candidates: + code = (cand.get("code") or "").upper().strip() + conf = float(cand.get("confidence", 0.0)) + multiplier = 1.0 + + # RÈGLE ABSOLUE : S'il n'y a PAS de rouge (red_ratio < 0.035), élimination stricte des panneaux rouges ! + if red_ratio < 0.035: + if code in ("F4A", "C43", "C1", "C3", "B1", "B5") or code.startswith(("C43_", "A")): + multiplier = 0.0001 + elif code.startswith(("ZE9", "TYPE", "F", "D", "E9", "GXC", "XD")): + multiplier = 1.8 + + if has_red_blue: + if code.startswith("D") and not code.startswith("DISQUE"): + multiplier = 0.0001 + elif code.startswith(("E1", "E2", "E3", "E4", "E9", "ZE", "C")): + multiplier = 2.5 + elif is_pure_blue: + # Sur fond bleu pur : éliminer les panneaux rouges ET tous les panonceaux blancs M... et TYPE0B + if code.startswith(("E1", "E2", "E3", "C", "A", "B1", "B5", "F4A", "TYPE0B")) or (code.startswith("M") and not code.startswith("M12") and not code.startswith("MAX")): + multiplier = 0.0001 + elif code.startswith(("D", "F", "G", "TYPE", "E9")): + multiplier = 2.5 + elif is_red_white: + # Sur fond rouge/blanc ou blanc pur : éliminer les panneaux/panonceaux bleus + if code.startswith(("D", "E1", "E2", "E3", "E4", "TYPE0", "TYPEIA", "TYPEIB", "TYPEIC", "TYPEIV", "TYPEVI", "TYPEVII", "TYPEVIII", "TYPEX")): + multiplier = 0.0001 + elif code.startswith(("C", "A", "B", "Z", "M", "TYPE0B", "GXC", "XD")): + multiplier = 1.8 + elif is_yellow: + if code.startswith("B3"): + multiplier = 3.0 + elif code.startswith(("D", "E", "C", "A")): + multiplier = 0.01 + + new_conf = conf * multiplier + adjusted_candidates.append({ + **cand, + "confidence": new_conf, + "raw_confidence": conf, + "color_multiplier": multiplier, + }) + + total_conf = sum(c["confidence"] for c in adjusted_candidates) + if total_conf > 0: + for c in adjusted_candidates: + c["confidence"] = round(c["confidence"] / total_conf, 4) + + adjusted_candidates.sort(key=lambda x: x["confidence"], reverse=True) + + # Discrimination fine du pictogramme intérieur pour la famille A + final_candidates = discriminate_inner_pictogram(crop_bgr, adjusted_candidates) + return final_candidates + + 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. @@ -602,20 +990,35 @@ def classify_sign_visual(crop_bgr: Any, ocr_text: str = "") -> Dict[str, Any]: 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])) + # Masques couleur HSV pour l'analyse spatiale de la forme + blue_mask = cv2.inRange(hsv, np.array([90, 40, 30]), np.array([138, 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])) + yellow_mask = cv2.inRange(hsv, np.array([14, 60, 60]), np.array([38, 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 + profile = extract_sign_color_profile(crop_bgr) + blue_ratio = profile["blue_ratio"] + red_ratio = profile["red_ratio"] + yellow_ratio = profile["yellow_ratio"] aspect_ratio = w / float(h) + # 0. PANNEAUX ROUGE ET BLEU (Famille E1 / E3 : Stationnement interdit / Parquage interdit) + if profile["has_red_and_blue"]: + code = "E1" + entry = SIGN_CATALOG.get(code, {}) + return { + "code": code, + "name_fr": entry.get("name_fr", "Stationnement interdit (E1)"), + "name_nl": entry.get("name_nl", "Parkeerverbod (E1)"), + "category": "parking", + "svg_url": get_svg_url(code), + "matched_by": "visual_red_blue_parking_restriction", + "confidence": 0.94, + } + # 1. PANNEAUX BLEUS (Famille E9 Stationnement, D Obligation ou F Indication) - if blue_ratio > 0.10: + if blue_ratio > 0.10 and not profile["has_red_and_blue"]: # Détection de texte ou lettre P / D dans l'OCR ocr_clean = ocr_text.strip().upper() if ocr_clean in ("P", "D", "🅿") or "PARKING" in ocr_clean or "PARKEREN" in ocr_clean: @@ -808,3 +1211,103 @@ def classify_sign_visual(crop_bgr: Any, ocr_text: str = "") -> Dict[str, Any]: "matched_by": "visual_generic", "confidence": 0.60, } + + +def get_all_catalog_signs() -> List[Dict[str, Any]]: + """ + Retourne l'ensemble exhaustif des panneaux de signalisation répertoriés + (catalogue officiel + templates vectoriels SVG/PNG disponibles). + """ + from pathlib import Path + from django.conf import settings + + signs_dict: Dict[str, Dict[str, Any]] = {} + + # 1. Base depuis SIGN_CATALOG + for code, info in SIGN_CATALOG.items(): + code_upper = code.upper() + signs_dict[code_upper] = { + "code": code_upper, + "name_fr": info.get("name_fr", f"Panneau {code_upper}"), + "name_nl": info.get("name_nl", f"Verkeersbord {code_upper}"), + "category": info.get("category", "indication"), + "shape": info.get("shape", "other"), + "svg_url": get_svg_url(code_upper), + } + + # 2. Complément depuis les fichiers statiques de templates (500+ fichiers) + try: + base_dir = getattr(settings, "BASE_DIR", None) + if base_dir: + signs_dir = Path(base_dir) / "assets" / "static" / "assets" / "road_signs" / "2025" + if signs_dir.exists(): + for f in sorted(signs_dir.iterdir()): + if f.is_file() and f.suffix.lower() in ('.svg', '.png'): + code = f.stem.upper().strip() + if code not in signs_dict: + # Détermination automatique de la catégorie selon le préfixe + category = "indication" + if code.startswith("A"): + category = "danger" + elif code.startswith("B"): + category = "priority" + elif code.startswith("C"): + category = "prohibition" + elif code.startswith("D"): + category = "obligation" + elif code.startswith("E"): + category = "parking" + elif code.startswith("F"): + category = "indication" + elif code.startswith(("M", "TYPE", "X", "G")): + category = "panonceau" + elif code.startswith("Z"): + category = "zone" + elif code.startswith("S"): + category = "temporary" + + # Nom humanisé par défaut si non catalogué + name_fr = f"Panneau {code}" + name_nl = f"Verkeersbord {code}" + if category == "panonceau": + name_fr = f"Panonceau additionnel {code}" + name_nl = f"Onderbord {code}" + elif category == "zone": + name_fr = f"Panneau de zone {code}" + name_nl = f"Zonebord {code}" + + signs_dict[code] = { + "code": code, + "name_fr": name_fr, + "name_nl": name_nl, + "category": category, + "shape": "other", + "svg_url": get_svg_url(code), + } + except Exception: + pass + + # 3. Complément depuis la base de données Django si disponible + try: + from sign.models import SignPanelType + for pt in SignPanelType.objects.all(): + code = pt.code.upper().strip() + if code in signs_dict: + if pt.name_fr: + signs_dict[code]["name_fr"] = pt.name_fr + if pt.name_nl: + signs_dict[code]["name_nl"] = pt.name_nl + else: + signs_dict[code] = { + "code": code, + "name_fr": pt.name_fr or f"Panneau {code}", + "name_nl": pt.name_nl or f"Verkeersbord {code}", + "category": "indication", + "shape": "other", + "svg_url": get_svg_url(code), + } + except Exception: + pass + + return sorted(list(signs_dict.values()), key=lambda x: x["code"]) + diff --git a/loko/sign/ai/classifier.py b/loko/sign/ai/classifier.py index 8af3662..700dac1 100644 --- a/loko/sign/ai/classifier.py +++ b/loko/sign/ai/classifier.py @@ -145,10 +145,63 @@ class SyntheticSignAugmentor: return bg + @staticmethod + def apply_specular_glare(bgr_img: np.ndarray, alpha_mask: np.ndarray) -> np.ndarray: + """ + Simule des reflets métalliques / spéculaires du soleil ou des phares + sur le film rétro-réfléchissant du panneau métallique. + """ + if random.random() > 0.65: + return bgr_img + + h, w = bgr_img.shape[:2] + glare_mask = np.zeros((h, w), dtype=np.float32) + + glare_type = random.choice(["spot", "streak", "gradient"]) + if glare_type == "spot": + cx = random.randint(int(w * 0.2), int(w * 0.8)) + cy = random.randint(int(h * 0.2), int(h * 0.8)) + radius = random.randint(int(min(h, w) * 0.15), int(min(h, w) * 0.40)) + cv2.circle(glare_mask, (cx, cy), radius, 1.0, -1) + k = max(3, radius * 2 + 1) + if k % 2 == 0: + k += 1 + glare_mask = cv2.GaussianBlur(glare_mask, (k, k), 0) + elif glare_type == "streak": + angle = random.uniform(20, 70) + center = (random.randint(int(w * 0.3), int(w * 0.7)), random.randint(int(h * 0.3), int(h * 0.7))) + axes = (random.randint(int(w * 0.35), int(w * 0.75)), random.randint(int(h * 0.08), int(h * 0.20))) + cv2.ellipse(glare_mask, center, axes, angle, 0, 360, 1.0, -1) + glare_mask = cv2.GaussianBlur(glare_mask, (31, 31), 0) + else: + direction = random.choice(["top", "left", "diagonal"]) + if direction == "top": + for y in range(h): + glare_mask[y, :] = max(0.0, 1.0 - (y / float(max(1, int(h * 0.65))))) + elif direction == "left": + for x in range(w): + glare_mask[:, x] = max(0.0, 1.0 - (x / float(max(1, int(w * 0.65))))) + else: + for y in range(h): + for x in range(w): + glare_mask[y, x] = max(0.0, 1.0 - ((x + y) / float(max(1, w + h)) * 1.5)) + glare_mask = cv2.GaussianBlur(glare_mask, (25, 25), 0) + + # Restreindre le reflet uniquement à la surface du panneau (alpha) + alpha_norm = (alpha_mask.astype(np.float32) / 255.0) + glare_mask = glare_mask * alpha_norm + + glare_intensity = random.uniform(0.30, 0.75) + glare_3d = glare_mask[:, :, np.newaxis] * glare_intensity + + # Mélange vers blanc brillant avec léger impact de saturation + result = bgr_img.astype(np.float32) * (1.0 - glare_3d * 0.5) + 255.0 * glare_3d + return np.clip(result, 0, 255).astype(np.uint8) + @classmethod def augment_sign(cls, rgba_sign: np.ndarray, size: int = 224) -> np.ndarray: """ - Applique une suite de déformations physiques et colorimétriques réalistes + Applique une suite de déformations physiques, métalliques et colorimétriques réalistes sur le panneau RGBA et l'incruste sur un fond synthétique. Retourne une image BGR 3 canaux de taille (size, size). """ @@ -157,13 +210,13 @@ class SyntheticSignAugmentor: alpha = rgba_sign[:, :, 3].copy() # 1. Déformation Perspective 3D (Angle de vue caméra smartphone / véhicule) - scale = random.uniform(0.72, 0.96) + scale = random.uniform(0.70, 0.96) # Points sources src_pts = np.float32([[0, 0], [w, 0], [w, h], [0, h]]) # Décalages de perspective aléatoires - max_shift = 0.12 + max_shift = 0.14 dx1 = random.uniform(-w * max_shift, w * max_shift) dy1 = random.uniform(-h * max_shift, h * max_shift) dx2 = random.uniform(-w * max_shift, w * max_shift) @@ -183,48 +236,60 @@ class SyntheticSignAugmentor: warped_bgr = cv2.warpPerspective(bgr, M_persp, (size, size), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0)) warped_alpha = cv2.warpPerspective(alpha, M_persp, (size, size), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT, borderValue=0) - # 2. Rotation légère (-10° à +10°) - rot_angle = random.uniform(-10, 10) + # 2. Rotation légère (-12° à +12°) + rot_angle = random.uniform(-12, 12) M_rot = cv2.getRotationMatrix2D((size / 2.0, size / 2.0), rot_angle, 1.0) warped_bgr = cv2.warpAffine(warped_bgr, M_rot, (size, size), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0)) warped_alpha = cv2.warpAffine(warped_alpha, M_rot, (size, size), flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT, borderValue=0) - # 3. Éclairage & Ombrage réaliste (Gradient de soleil rasant ou ombre) + # 3. Reflets métalliques / spéculaires réalistes + warped_bgr = cls.apply_specular_glare(warped_bgr, warped_alpha) + + # 4. Éclairage directionnel & Ombrage (Gradient de soleil rasant ou ombre) alpha_norm = (warped_alpha.astype(np.float32) / 255.0)[:, :, np.newaxis] bgr_float = warped_bgr.astype(np.float32) grad_angle = random.uniform(0, 2 * math.pi) gx, gy = math.cos(grad_angle), math.sin(grad_angle) y_coords, x_coords = np.mgrid[0:size, 0:size] - light_grad = 1.0 + random.uniform(-0.35, 0.35) * (gx * (x_coords / float(size) - 0.5) + gy * (y_coords / float(size) - 0.5)) - light_grad = np.clip(light_grad, 0.55, 1.45)[:, :, np.newaxis] + light_grad = 1.0 + random.uniform(-0.40, 0.40) * (gx * (x_coords / float(size) - 0.5) + gy * (y_coords / float(size) - 0.5)) + light_grad = np.clip(light_grad, 0.50, 1.50)[:, :, np.newaxis] bgr_float = bgr_float * light_grad - # Luminosité & Contraste globaux - brightness = random.uniform(0.75, 1.25) - contrast = random.uniform(0.80, 1.25) - bgr_float = np.clip((bgr_float - 128.0) * contrast + 128.0 * brightness, 0, 255) + # 5. Variations renforcées de Saturation et Luminosité HSV (peinture vieillie / plein soleil) + hsv = cv2.cvtColor(np.clip(bgr_float, 0, 255).astype(np.uint8), cv2.COLOR_BGR2HSV).astype(np.float32) + sat_factor = random.uniform(0.55, 1.45) + val_factor = random.uniform(0.65, 1.35) + hsv[:, :, 1] = np.clip(hsv[:, :, 1] * sat_factor, 0, 255) + hsv[:, :, 2] = np.clip(hsv[:, :, 2] * val_factor, 0, 255) + bgr_float = cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2BGR).astype(np.float32) - # Teinte / vieillissement / saturation - b_shift = random.uniform(0.92, 1.08) - g_shift = random.uniform(0.92, 1.08) - r_shift = random.uniform(0.92, 1.08) - bgr_float[:, :, 0] *= b_shift - bgr_float[:, :, 1] *= g_shift - bgr_float[:, :, 2] *= r_shift - bgr_float = np.clip(bgr_float, 0, 255).astype(np.uint8) + # 6. Contraste dynamique et correction Gamma + contrast = random.uniform(0.70, 1.35) + brightness_shift = random.uniform(-20, 25) + bgr_float = np.clip((bgr_float - 128.0) * contrast + 128.0 + brightness_shift, 0, 255) - # 4. Composition sur fond d'environnement + gamma = random.uniform(0.75, 1.35) + inv_gamma = 1.0 / gamma + lut_table = np.array([((i / 255.0) ** inv_gamma) * 255 for i in np.arange(0, 256)]).astype(np.uint8) + bgr_float = cv2.LUT(bgr_float.astype(np.uint8), lut_table).astype(np.float32) + + # 7. Balance des couleurs / Température thermique (soleil couchant doré vs temps gris bleuté) + temp_shift = random.uniform(-0.08, 0.08) + bgr_float[:, :, 0] = np.clip(bgr_float[:, :, 0] * (1.0 - temp_shift), 0, 255) # Bleu + bgr_float[:, :, 2] = np.clip(bgr_float[:, :, 2] * (1.0 + temp_shift), 0, 255) # Rouge + + # 8. Composition sur fond d'environnement synthétique bg = cls.generate_random_background(size) composite = (bgr_float * alpha_norm + bg * (1.0 - alpha_norm)).astype(np.uint8) - # 5. Flou optique & Bruit de capteur - if random.random() < 0.35: + # 9. Flou optique, flou de bougé & Bruit de capteur + if random.random() < 0.40: ksize = random.choice([3, 5]) composite = cv2.GaussianBlur(composite, (ksize, ksize), 0) - if random.random() < 0.30: - noise = np.random.normal(0, random.uniform(2, 8), composite.shape).astype(np.int16) + if random.random() < 0.35: + noise = np.random.normal(0, random.uniform(2, 9), composite.shape).astype(np.int16) composite = np.clip(composite.astype(np.int16) + noise, 0, 255).astype(np.uint8) return composite @@ -290,15 +355,17 @@ class SignClassifierEngine: def train_from_svgs( self, signs_dir: Union[str, Path] = DEFAULT_SIGNS_DIR, - samples_per_class: int = 15, + samples_per_class: int = 50, epochs: int = 10, batch_size: int = 32, learning_rate: float = 0.001, + include_ground_truth: bool = True, + ground_truth_dir: Optional[Union[str, Path]] = None, progress_callback: Optional[Any] = None, ) -> Dict[str, Any]: """ - Scanne le dossier des SVGs et PNGs, génère un dataset synthétique équilibré, - entraîne MobileNetV3-Small et exporte vers ONNX. + Scanne le dossier des SVGs et PNGs ainsi que les données de vérité terrain réelles, + génère un dataset enrichi, entraîne MobileNetV3-Small et exporte vers ONNX. """ import torch import torch.nn as nn @@ -308,62 +375,132 @@ class SignClassifierEngine: start_time = time.perf_counter() template_files = self.discover_templates(signs_dir) - classes = sorted(list(template_files.keys())) + + # Récupération des vrais crops terrain annotés + real_crops_map: Dict[str, List[Path]] = {} + if include_ground_truth: + try: + from .ground_truth import GroundTruthDatasetManager + if ground_truth_dir: + real_crops_map = GroundTruthDatasetManager(base_dir=ground_truth_dir).get_real_crops_for_training() + elif Path(signs_dir).resolve() == Path(DEFAULT_SIGNS_DIR).resolve(): + real_crops_map = GroundTruthDatasetManager.get_instance().get_real_crops_for_training() + except Exception as e: + logger.warning("Impossible de charger les crops réels pour l'entraînement : %s", e) + + # Union des classes de templates et des classes terrain + all_class_keys = set(template_files.keys()) | set(real_crops_map.keys()) + classes = sorted(list(all_class_keys)) num_classes = len(classes) + class_to_idx = {c: i for i, c in enumerate(classes)} if num_classes < 2: - raise ValueError(f"Pas assez de templates trouvés ({num_classes}) pour entraîner le modèle.") + raise ValueError(f"Pas assez de classes trouvées ({num_classes}) pour entraîner le modèle.") logger.info("🔍 %d types de panneaux officiels découverts pour l'entraînement.", num_classes) if progress_callback: - progress_callback(f"Chargement de {num_classes} types de panneaux...") + progress_callback(f"Chargement et rendu des {num_classes} types de panneaux...") - # 1. Génération du dataset synthétique - x_list = [] - y_list = [] + # 1. Chargement compact en mémoire des templates RGBA de base (~100 Mo max) + templates_dict: Dict[int, np.ndarray] = {} + for code, fpath in template_files.items(): + if code in class_to_idx: + c_idx = class_to_idx[code] + if fpath.suffix.lower() == '.svg': + rgba = SyntheticSignAugmentor.render_svg_to_numpy(fpath, size=224) + else: + rgba = SyntheticSignAugmentor.load_png_to_numpy(fpath, size=224) + if rgba is not None: + templates_dict[c_idx] = rgba - for class_idx, code in enumerate(classes): - file_path = template_files[code] - if file_path.suffix.lower() == '.svg': - rgba = SyntheticSignAugmentor.render_svg_to_numpy(file_path, size=224) - else: - rgba = SyntheticSignAugmentor.load_png_to_numpy(file_path, size=224) + # 2. Chargement compact des vrais crops terrain annotés (~20 Mo max) + real_crops_list: List[Tuple[np.ndarray, int]] = [] + for code, crop_paths in real_crops_map.items(): + if code in class_to_idx: + c_idx = class_to_idx[code] + for cp in crop_paths: + try: + crop_bgr = cv2.imread(str(cp)) + if crop_bgr is None or crop_bgr.size == 0: + continue + h_c, w_c = crop_bgr.shape[:2] + scale = min(224.0 / h_c, 224.0 / w_c) + nw, nh = max(1, int(round(w_c * scale))), max(1, int(round(h_c * scale))) + resized = cv2.resize(crop_bgr, (nw, nh), interpolation=cv2.INTER_AREA) - if rgba is None: - continue + canvas = np.ones((224, 224, 3), dtype=np.uint8) * 128 + xo = (224 - nw) // 2 + yo = (224 - nh) // 2 + canvas[yo:yo + nh, xo:xo + nw] = resized + real_crops_list.append((canvas, c_idx)) + except Exception as crop_exc: + logger.debug("Erreur crop %s: %s", cp, crop_exc) - # Image canonique sur fond blanc - raw_bgr = rgba[:, :, :3] - raw_alpha = (rgba[:, :, 3] / 255.0)[:, :, np.newaxis] - white_bg = np.ones((224, 224, 3), dtype=np.uint8) * 255 - canonical = (raw_bgr * raw_alpha + white_bg * (1.0 - raw_alpha)).astype(np.uint8) + # 3. Dataset PyTorch dynamique à génération à la volée (Ultra faible empreinte RAM < 150MB) + class DynamicSignDataset(torch.utils.data.Dataset): + def __init__(self, tmpl_map, real_list, n_samples): + self.tmpl_map = tmpl_map + self.real_list = real_list + self.n_samples = n_samples + self.index_items = [] - def to_tensor_norm(bgr_img: np.ndarray) -> np.ndarray: - rgb = cv2.cvtColor(bgr_img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0 + for class_idx in tmpl_map: + self.index_items.append(("template_canon", class_idx)) + for _ in range(n_samples - 1): + self.index_items.append(("template_aug", class_idx)) + + for item_tuple in real_list: + self.index_items.append(("real_canon", item_tuple)) + for _ in range(3): + self.index_items.append(("real_aug", item_tuple)) + + def __len__(self): + return len(self.index_items) + + def __getitem__(self, idx): + itype, data = self.index_items[idx] + if itype == "template_canon": + c_idx = data + rgba = self.tmpl_map[c_idx] + raw_b = rgba[:, :, :3] + raw_a = (rgba[:, :, 3].astype(np.float32) / 255.0)[:, :, np.newaxis] + white_bg = np.ones((224, 224, 3), dtype=np.uint8) * 255 + img_bgr = (raw_b * raw_a + white_bg * (1.0 - raw_a)).astype(np.uint8) + lbl = c_idx + elif itype == "template_aug": + c_idx = data + rgba = self.tmpl_map[c_idx] + img_bgr = SyntheticSignAugmentor.augment_sign(rgba, size=224) + lbl = c_idx + elif itype == "real_canon": + img_bgr, lbl = data + else: # real_aug + base_bgr, lbl = data + rot_ang = random.uniform(-6, 6) + M = cv2.getRotationMatrix2D((112, 112), rot_ang, 1.0) + aug_real = cv2.warpAffine(base_bgr, M, (224, 224), borderMode=cv2.BORDER_REFLECT) + bright = random.uniform(0.85, 1.15) + img_bgr = np.clip(aug_real.astype(np.float32) * bright, 0, 255).astype(np.uint8) + + rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0 mean = np.array([0.485, 0.456, 0.406], dtype=np.float32) std = np.array([0.229, 0.224, 0.225], dtype=np.float32) norm = (rgb - mean) / std - return norm.transpose(2, 0, 1) + return torch.from_numpy(norm.transpose(2, 0, 1)), torch.tensor(lbl, dtype=torch.int64) - x_list.append(to_tensor_norm(canonical)) - y_list.append(class_idx) + dataset = DynamicSignDataset(templates_dict, real_crops_list, samples_per_class) + total_samples = len(dataset) + real_samples_count = len(real_crops_list) * 4 + loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, num_workers=0) - # Échantillons synthétiques avec variations réalistes - for _ in range(samples_per_class): - aug_bgr = SyntheticSignAugmentor.augment_sign(rgba, size=224) - x_list.append(to_tensor_norm(aug_bgr)) - y_list.append(class_idx) - - total_samples = len(x_list) - logger.info("📦 Dataset synthétique généré : %d images pour %d classes.", total_samples, num_classes) + logger.info( + "📦 Dataset d'entraînement dynamique prêt : %d images (dont %d réelles/augmentées) pour %d classes.", + total_samples, real_samples_count, num_classes + ) if progress_callback: - progress_callback(f"Entraînement MobileNetV3 sur {total_samples} images ({epochs} époques)...") - - # 2. Préparation des Tensors PyTorch - x_tensor = torch.tensor(np.array(x_list, dtype=np.float32)) - y_tensor = torch.tensor(np.array(y_list, dtype=np.int64)) - dataset = TensorDataset(x_tensor, y_tensor) - loader = DataLoader(dataset, batch_size=batch_size, shuffle=True) + progress_callback( + f"Entraînement MobileNetV3 sur {total_samples} images ({len(real_crops_list)} crops réels terrain, {epochs} époques)..." + ) # 3. Modèle MobileNetV3-Small pré-entraîné device = torch.device("cuda" if torch.cuda.is_available() else ("mps" if hasattr(torch.backends, "mps") and torch.backends.mps.is_available() else "cpu")) @@ -381,7 +518,7 @@ class SignClassifierEngine: model.train() for epoch in range(1, epochs + 1): - epoch_loss = 0.0 + running_loss = 0.0 correct = 0 total = 0 for batch_x, batch_y in loader: @@ -392,24 +529,25 @@ class SignClassifierEngine: loss.backward() optimizer.step() - epoch_loss += loss.item() * batch_x.size(0) - _, predicted = outputs.max(1) + running_loss += loss.item() * batch_x.size(0) + _, preds = torch.max(outputs, 1) + correct += torch.sum(preds == batch_y.data).item() total += batch_y.size(0) - correct += predicted.eq(batch_y).sum().item() scheduler.step() - acc = 100.0 * correct / max(1, total) - avg_loss = epoch_loss / max(1, total) - logger.info("Époque %d/%d - Loss: %.4f - Précision: %.1f%%", epoch, epochs, avg_loss, acc) + epoch_loss = running_loss / total + epoch_acc = (correct / total) * 100.0 + logger.info("Epoch %d/%d - Loss: %.4f - Accuracy: %.1f%%", epoch, epochs, epoch_loss, epoch_acc) if progress_callback: - progress_callback(f"Époque {epoch}/{epochs} : Précision {acc:.1f}% (Loss: {avg_loss:.4f})") + progress_callback(f"Époque {epoch}/{epochs} : Précision {epoch_acc:.1f}% (Loss {epoch_loss:.4f})") - # 5. Exportation vers ONNX + # 5. Export vers ONNX model.eval() - dummy_input = torch.randn(1, 3, 224, 224, device=device) - self.onnx_path.parent.mkdir(parents=True, exist_ok=True) + model.to("cpu") + dummy_input = torch.randn(1, 3, 224, 224, dtype=torch.float32) - try: + self.models_dir.mkdir(parents=True, exist_ok=True) + with torch.no_grad(): torch.onnx.export( model, dummy_input, @@ -422,18 +560,6 @@ class SignClassifierEngine: dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}}, dynamo=False ) - except TypeError: - torch.onnx.export( - model, - dummy_input, - str(self.onnx_path), - export_params=True, - opset_version=14, - do_constant_folding=True, - input_names=["input"], - output_names=["output"], - dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}} - ) # 6. Sauvegarde des métadonnées meta_data = { @@ -442,7 +568,7 @@ class SignClassifierEngine: "classes": classes, "samples_per_class": samples_per_class, "epochs": epochs, - "final_accuracy": round(acc, 2), + "final_accuracy": round(epoch_acc, 2), "input_size": [224, 224], "framework": "MobileNetV3-Small / ONNX", "signs_dir": str(signs_dir) @@ -461,7 +587,7 @@ class SignClassifierEngine: return { "status": "success", "num_classes": num_classes, - "accuracy": round(acc, 2), + "accuracy": round(epoch_acc, 2), "training_time_seconds": round(total_time, 1), "onnx_path": str(self.onnx_path), "size_mb": round(self.onnx_path.stat().st_size / (1024 * 1024), 2) @@ -488,7 +614,8 @@ class SignClassifierEngine: def predict(self, crop_bgr: np.ndarray, top_k: int = 5) -> Dict[str, Any]: """ Prédit le type officiel d'un panneau découpé (crop BGR). - Retourne le code officiel, la confiance et le top-K des alternatives. + Applique un filtrage physique par cohérence colorimétrique (histogramme HSV) + pour éliminer d'office les familles incompatibles (ex: D obligation quand rouge présent). """ if not isinstance(crop_bgr, np.ndarray) or crop_bgr.size == 0 or crop_bgr.shape[0] < 5 or crop_bgr.shape[1] < 5: return {"status": "error", "code": None, "confidence": 0.0, "top_matches": []} @@ -524,24 +651,29 @@ class SignClassifierEngine: exp_logits = np.exp(logits - np.max(logits)) probs = exp_logits / np.sum(exp_logits) - # 4. Top-K - k = min(top_k, len(self._classes)) - top_indices = np.argsort(probs)[::-1][:k] - top_matches = [] + # 4. Top-K initial élargi + candidate_count = min(max(top_k * 4, 25), len(self._classes)) + top_indices = np.argsort(probs)[::-1][:candidate_count] + raw_matches = [] for idx in top_indices: code = self._classes[idx] conf = float(probs[idx]) - top_matches.append({ + raw_matches.append({ "code": code, "confidence": round(conf, 4), "svg_url": f"/static/assets/road_signs/2025/{code}.svg" }) - best_match = top_matches[0] if top_matches else {"code": None, "confidence": 0.0, "svg_url": ""} + # 5. Filtrage & Réordonnancement par cohérence colorimétrique HSV + from .catalog import filter_and_rank_candidates_by_color, extract_sign_color_profile + filtered_matches = filter_and_rank_candidates_by_color(raw_matches, crop_bgr)[:top_k] + + best_match = filtered_matches[0] if filtered_matches else {"code": None, "confidence": 0.0, "svg_url": ""} return { "status": "success", "code": best_match["code"], "confidence": best_match["confidence"], "svg_url": best_match["svg_url"], - "top_matches": top_matches, + "top_matches": filtered_matches, + "color_profile": extract_sign_color_profile(crop_bgr), } diff --git a/loko/sign/ai/detector.py b/loko/sign/ai/detector.py index 9b26d2d..8bed8bf 100644 --- a/loko/sign/ai/detector.py +++ b/loko/sign/ai/detector.py @@ -364,8 +364,12 @@ class SignDetectionService: ) total_classifier_ms += (time.perf_counter() - c_start) * 1000.0 - # Tentative d'identification via l'OCR - matched = match_sign_from_ocr(ocr_text) + # Extraction du profil de couleur + from .catalog import extract_sign_color_profile + color_prof = extract_sign_color_profile(crop_bgr) + + # Tentative d'identification via l'OCR avec vérification colorimétrique + matched = match_sign_from_ocr(ocr_text, crop_bgr=crop_bgr) code = None name_fr = "" @@ -400,8 +404,8 @@ class SignDetectionService: final_confidence = max(final_confidence, 0.95) # 2. Vitesse maximale autorisée (C43 / C43_XX / ZC43) : - # Fusion : Si OCR extrait une vitesse OU que le classifieur neuronal a prédit C43 - elif (matched and matched["code"] == "C43") or (primary_nn_code and "C43" in primary_nn_code): + # Fusion : Si OCR extrait une vitesse OU que le classifieur neuronal a prédit C43 (avec ROUGE présent) + elif (matched and matched["code"] == "C43") or (primary_nn_code and "C43" in primary_nn_code and color_prof.get("red_ratio", 0) >= 0.035): val = matched.get("value") if matched else None # Si pas de valeur extraite de l'OCR, tenter d'extraire depuis le code neuronal (ex: C43_50 -> 50) if val is None and primary_nn_code: @@ -425,8 +429,8 @@ class SignDetectionService: final_confidence = max(final_confidence, 0.96 if nn_confirms_c43 else (matched.get("confidence", 0.90) if matched else 0.85)) # 3. Panneaux de Zone (Zone 30, Zone Parking ZE9A, Fin de zone) : - elif (matched and matched["code"] in ("ZE9A", "ZE9A_DISK", "F4A", "F4B", "F103", "ZC43")) or (primary_nn_code and primary_nn_code.startswith(("ZE9", "F4", "ZC"))): - if matched and matched["code"] in ("ZE9A", "ZE9A_DISK", "F4A", "F4B", "F103", "ZC43"): + elif (matched and matched["code"] in ("ZE9A", "ZE9B", "ZE9A_DISK", "F4A", "F4B", "F103", "ZC43")) or (primary_nn_code and primary_nn_code.startswith(("ZE9", "F4", "ZC"))): + if matched and matched["code"] in ("ZE9A", "ZE9B", "ZE9A_DISK", "F4A", "F4B", "F103", "ZC43"): code = matched["code"] name_fr = matched["data"]["name_fr"] name_nl = matched["data"]["name_nl"] @@ -437,6 +441,10 @@ class SignDetectionService: final_confidence = max(final_confidence, matched.get("confidence", 0.92)) else: code = primary_nn_code + # Si pas de rouge détecté et que c'est F4A, rectifier en ZE9A/ZE9B + if code.startswith("F4A") and color_prof.get("red_ratio", 0) < 0.035: + code = "ZE9B" if color_prof.get("blue_ratio", 0) > 0.08 else "ZE9A" + entry = SIGN_CATALOG.get(code, {}) name_fr = entry.get("name_fr", f"Zone {code}") name_nl = entry.get("name_nl", f"Zone {code}") @@ -445,9 +453,9 @@ class SignDetectionService: matched_by = "ai_neural_classifier" final_confidence = max(final_confidence, float(classifier_pred.get("confidence", 0.85))) - # 4. Matching OCR fort (STOP, Parking P, PMR, Payant, Recharge électrique, Tonnage, etc.) + # 4. Matching OCR fort (STOP, Parking P, PMR, Payant, Recharge électrique, Tonnage, Flèche distance, etc.) elif matched and ( - matched["code"] in ("B5", "E9A", "E9B", "GVII_BETALEND", "GVIID_ELEKTRISCHE_WAGENS", "E9A_PARKEERSCHIJF", "C21") + matched["code"] in ("B5", "E9A", "E9B", "GVII_BETALEND", "GVIID_ELEKTRISCHE_WAGENS", "E9A_PARKEERSCHIJF", "C21", "GXC", "TYPE0", "TYPE0B") or det["class_name"] == "sub_plate" ): code = matched["code"] @@ -462,6 +470,14 @@ class SignDetectionService: # 5. Réseau Neuronal MobileNetV3 (Classification visuelle fine des pictogrammes) elif primary_nn_code and classifier_pred.get("confidence", 0.0) >= 0.02: code = primary_nn_code + + # Simplification des panonceaux textuels complexes non-spécifiques + if det["class_name"] == "sub_plate" and code.startswith("TYPE") and code not in ("TYPEI", "TYPEII", "TYPEIII"): + if color_prof.get("blue_ratio", 0) >= 0.30 and color_prof.get("red_ratio", 0) < 0.035: + code = "TYPE0" + elif color_prof.get("red_ratio", 0) < 0.035: + code = "TYPE0B" + svg_url = classifier_pred.get("svg_url") or get_svg_url(code) matched_by = "ai_neural_classifier" final_confidence = float(classifier_pred.get("confidence", det["confidence"])) diff --git a/loko/sign/ai/ground_truth.py b/loko/sign/ai/ground_truth.py new file mode 100644 index 0000000..e49bdd7 --- /dev/null +++ b/loko/sign/ai/ground_truth.py @@ -0,0 +1,255 @@ +""" +Module de gestion du dataset de vérité terrain (Ground Truth / Active Learning) pour la signalisation routière. +Stocke les images de terrain annotées ou corrigées par les utilisateurs, +génère les découpes (crops) classées par code de panneau officiel (ex: E1, B5, XD) +et alimente le pipeline de ré-entraînement du classifieur neuronal. +""" +import io +import json +import logging +import os +import time +import uuid +from pathlib import Path +from typing import Any, Dict, 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 + +logger = logging.getLogger(__name__) + +GROUND_TRUTH_DIR = getattr( + settings, + "SIGN_AI_GROUND_TRUTH_DIR", + Path(__file__).resolve().parent / "training" / "dataset" / "ground_truth" +) + + +class GroundTruthDatasetManager: + """ + Gestionnaire pour la persistance et la réutilisation des annotations terrain (Ground Truth). + Structure des dossiers : + ground_truth/ + ├── images/ (photos complètes de terrain) + ├── crops/ + │ ├── E1/ (découpes de panneaux E1 réels) + │ ├── B5/ (découpes de panneaux B5 réels) + │ └── ... + └── annotations.json (index historique des annotations validées) + """ + _instance: Optional["GroundTruthDatasetManager"] = None + + def __init__(self, base_dir: Union[str, Path] = GROUND_TRUTH_DIR): + self.base_dir = Path(base_dir) + self.images_dir = self.base_dir / "images" + self.crops_dir = self.base_dir / "crops" + self.index_file = self.base_dir / "annotations.json" + self._ensure_directories() + + @classmethod + def get_instance(cls) -> "GroundTruthDatasetManager": + if cls._instance is None: + cls._instance = cls() + return cls._instance + + def _ensure_directories(self): + """Crée l'arborescence de dossiers si nécessaire.""" + self.images_dir.mkdir(parents=True, exist_ok=True) + self.crops_dir.mkdir(parents=True, exist_ok=True) + if not self.index_file.exists(): + with open(self.index_file, "w", encoding="utf-8") as f: + json.dump([], f, indent=2, ensure_ascii=False) + + def load_annotations(self) -> List[Dict[str, Any]]: + """Charge l'ensemble des annotations validées.""" + if not self.index_file.exists(): + return [] + try: + with open(self.index_file, "r", encoding="utf-8") as f: + return json.load(f) + except Exception as e: + logger.warning("Erreur lors de la lecture de %s: %s", self.index_file, e) + return [] + + def _save_annotations(self, items: List[Dict[str, Any]]): + """Sauvegarde atomique du fichier JSON d'annotations.""" + temp_file = self.index_file.with_suffix(".tmp") + with open(temp_file, "w", encoding="utf-8") as f: + json.dump(items, f, indent=2, ensure_ascii=False) + temp_file.replace(self.index_file) + + def save_ground_truth( + self, + image_input: Union[bytes, bytearray, io.BytesIO, np.ndarray, Image.Image], + panels: List[Dict[str, Any]], + user_info: Optional[Dict[str, Any]] = None, + notes: str = "" + ) -> Dict[str, Any]: + """ + Enregistre une photo de terrain et extrait les crops pour chaque panneau annoté/validé. + Retourne les métadonnées de l'enregistrement et les statistiques mises à jour. + """ + # 1. Normalisation de l'image en NumPy BGR + if isinstance(image_input, np.ndarray): + cv2_img = image_input + elif isinstance(image_input, Image.Image): + pil_img = ImageOps.exif_transpose(image_input) + cv2_img = cv2.cvtColor(np.array(pil_img), cv2.COLOR_RGB2BGR) + else: + if isinstance(image_input, io.BytesIO): + raw_bytes = image_input.getvalue() + else: + raw_bytes = bytes(image_input) + pil_img = Image.open(io.BytesIO(raw_bytes)) + 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) + + img_h, img_w = cv2_img.shape[:2] + timestamp_str = time.strftime("%Y%m%d_%H%M%S") + sample_uuid = str(uuid.uuid4())[:8] + image_filename = f"gt_{timestamp_str}_{sample_uuid}.jpg" + image_path = self.images_dir / image_filename + + # Sauvegarde de l'image haute qualité + cv2.imwrite(str(image_path), cv2_img, [int(cv2.IMWRITE_JPEG_QUALITY), 95]) + + saved_crops = [] + panels_metadata = [] + + for p_idx, panel in enumerate(panels): + raw_code = (panel.get("code") or "INCONNU").strip().upper() + bbox = panel.get("bbox") or [0, 0, img_w, img_h] + ocr_text = panel.get("ocr_text") or panel.get("signpanel_text") or "" + vertical_order = panel.get("vertical_order", p_idx + 1) + was_corrected = bool(panel.get("was_corrected", False)) + initial_code = panel.get("initial_code") or raw_code + + # Découpage du crop avec contrainte des limites de l'image + x1, y1, x2, y2 = [int(round(coord)) for coord in bbox] + x1 = max(0, min(img_w - 1, x1)) + y1 = max(0, min(img_h - 1, y1)) + x2 = max(x1 + 5, min(img_w, x2)) + y2 = max(y1 + 5, min(img_h, y2)) + + crop_bgr = cv2_img[y1:y2, x1:x2] + crop_filename = "" + crop_rel_path = "" + + if crop_bgr.size > 0 and raw_code: + # Dossier dédié pour la classe de panneau (ex: crops/E1/) + class_crop_dir = self.crops_dir / raw_code + class_crop_dir.mkdir(parents=True, exist_ok=True) + + crop_filename = f"{raw_code}_{timestamp_str}_{sample_uuid}_p{p_idx + 1}.jpg" + crop_path = class_crop_dir / crop_filename + cv2.imwrite(str(crop_path), crop_bgr, [int(cv2.IMWRITE_JPEG_QUALITY), 95]) + crop_rel_path = f"crops/{raw_code}/{crop_filename}" + saved_crops.append({ + "code": raw_code, + "file": crop_filename, + "path": crop_rel_path, + "width": crop_bgr.shape[1], + "height": crop_bgr.shape[0], + }) + + panels_metadata.append({ + "index": p_idx + 1, + "code": raw_code, + "initial_code": initial_code, + "was_corrected": was_corrected, + "bbox": [x1, y1, x2, y2], + "vertical_order": vertical_order, + "ocr_text": ocr_text, + "crop_rel_path": crop_rel_path, + }) + + # Enregistrement dans annotations.json + record = { + "id": sample_uuid, + "created_at": time.strftime("%Y-%m-%d %H:%M:%S"), + "image_filename": image_filename, + "image_dimensions": {"width": img_w, "height": img_h}, + "user": user_info or {}, + "notes": notes, + "panels_count": len(panels_metadata), + "panels": panels_metadata, + } + + all_records = self.load_annotations() + all_records.insert(0, record) + self._save_annotations(all_records) + + stats = self.get_dataset_stats() + logger.info( + "✓ Vérité terrain enregistrée : sample %s (%d panneaux, %d crops) - Total dataset: %d images", + sample_uuid, len(panels_metadata), len(saved_crops), stats["total_images"] + ) + + return { + "status": "success", + "sample_id": sample_uuid, + "image_filename": image_filename, + "crops_saved": len(saved_crops), + "stats": stats, + } + + def get_dataset_stats(self) -> Dict[str, Any]: + """Retourne les statistiques actuelles du dataset de vérité terrain.""" + records = self.load_annotations() + total_images = len(records) + total_crops = 0 + classes_count: Dict[str, int] = {} + corrected_count = 0 + + # Scan physique des crops pour garantir la concordance + if self.crops_dir.exists(): + for code_dir in self.crops_dir.iterdir(): + if code_dir.is_dir(): + count = len([f for f in code_dir.iterdir() if f.suffix.lower() in ('.jpg', '.jpeg', '.png', '.webp')]) + if count > 0: + classes_count[code_dir.name] = count + total_crops += count + + for r in records: + for p in r.get("panels", []): + if p.get("was_corrected"): + corrected_count += 1 + + top_classes = sorted(classes_count.items(), key=lambda x: x[1], reverse=True)[:10] + + return { + "total_images": total_images, + "total_crops": total_crops, + "distinct_classes": len(classes_count), + "total_corrected_panels": corrected_count, + "classes_distribution": classes_count, + "top_classes": [{"code": c, "count": n, "svg_url": get_svg_url(c)} for c, n in top_classes], + "last_updated": records[0]["created_at"] if records else None, + } + + def get_real_crops_for_training(self) -> Dict[str, List[Path]]: + """ + Retourne un dictionnaire {code_panneau: [chemins_fichiers_crops_reels]} + pour l'injection dans le pipeline de fine-tuning du classifieur. + """ + real_crops: Dict[str, List[Path]] = {} + if not self.crops_dir.exists(): + return real_crops + + for code_dir in self.crops_dir.iterdir(): + if code_dir.is_dir(): + code = code_dir.name.upper() + crops = [ + f for f in code_dir.iterdir() + if f.is_file() and f.suffix.lower() in ('.jpg', '.jpeg', '.png', '.webp') + ] + if crops: + real_crops[code] = sorted(crops) + + return real_crops diff --git a/loko/sign/ai/models/sign_classifier.onnx b/loko/sign/ai/models/sign_classifier.onnx index 8d97873..9d334a0 100644 Binary files a/loko/sign/ai/models/sign_classifier.onnx and b/loko/sign/ai/models/sign_classifier.onnx differ diff --git a/loko/sign/ai/models/sign_classifier_meta.json b/loko/sign/ai/models/sign_classifier_meta.json index 7200a2c..cf8ebbb 100644 --- a/loko/sign/ai/models/sign_classifier_meta.json +++ b/loko/sign/ai/models/sign_classifier_meta.json @@ -1,6 +1,6 @@ { - "created_at": "2026-08-22 18:42:48", - "num_classes": 502, + "created_at": "2026-08-29 15:30:09", + "num_classes": 503, "classes": [ "A11", "A13", @@ -429,6 +429,7 @@ "TYPEVF", "TYPEVG", "TYPEVI", + "TYPEVII", "TYPEVIIA_(+)2,5T", "TYPEVIIA_(+)2T", "TYPEVIIA_(+)3,5T", @@ -505,9 +506,9 @@ "ZF111", "ZF113" ], - "samples_per_class": 12, - "epochs": 10, - "final_accuracy": 99.79, + "samples_per_class": 35, + "epochs": 12, + "final_accuracy": 97.33, "input_size": [ 224, 224 diff --git a/loko/sign/templates/sign/ai_demo.html b/loko/sign/templates/sign/ai_demo.html index bf5e0d0..0e5ae1c 100644 --- a/loko/sign/templates/sign/ai_demo.html +++ b/loko/sign/templates/sign/ai_demo.html @@ -25,19 +25,28 @@ .panel-card { border: 1px solid #e2e8f0; border-radius: 0.75rem; - transition: transform 0.2s, box-shadow 0.2s; + transition: transform 0.2s, box-shadow 0.2s, border-color 0.2s; background: #ffffff; } .panel-card:hover { transform: translateY(-2px); box-shadow: 0 10px 25px -5px rgba(0, 0, 0, 0.1); } + .panel-card.is-corrected { + border-color: #f59e0b; + background: #fffbeb; + } .svg-sign-preview { width: 84px; height: 84px; object-fit: contain; filter: drop-shadow(0 2px 4px rgba(0,0,0,0.15)); } + .svg-sign-sm { + width: 48px; + height: 48px; + object-fit: contain; + } .perf-badge { font-family: monospace; font-size: 0.95rem; @@ -51,6 +60,40 @@ border-radius: 50%; font-weight: 700; } + .sign-select-grid { + display: grid; + grid-template-columns: repeat(auto-fill, minmax(135px, 1fr)); + gap: 0.75rem; + max-height: 380px; + overflow-y: auto; + } + .sign-select-item { + border: 2px solid #e2e8f0; + border-radius: 0.5rem; + padding: 0.5rem; + text-align: center; + background: #ffffff; + cursor: pointer; + transition: all 0.15s ease-in-out; + } + .sign-select-item:hover { + border-color: #3b82f6; + background: #eff6ff; + transform: scale(1.03); + } + .sign-select-item.selected { + border-color: #2563eb; + background: #dbeafe; + box-shadow: 0 0 0 2px rgba(37, 99, 235, 0.3); + } + .alt-chip { + cursor: pointer; + transition: all 0.15s ease-in-out; + } + .alt-chip:hover { + background-color: #e2e8f0 !important; + transform: translateY(-1px); + } {% endblock %} @@ -68,15 +111,21 @@ Phase 1 Opérationnelle + + {{ ground_truth_stats.total_images|default:"0" }} photo(s) terrain enregistrée(s) +
- {% translate "Analyse instantanée des photos de terrain : détection des panneaux, extraction du texte des panonceaux, ordonnancement vertical et visualisation SVG." %} + {% translate "Analyse instantanée des photos de terrain, correction interactive des panneaux et alimentation du dataset de vérité terrain pour l'apprentissage continu." %}
-+ {% translate "Vérifiez les détections. Cliquez sur 'Corriger ce panneau' pour rectifier le code officiel ou le panonceau si nécessaire." %} +
++ {% translate "Sauvegarde la photo et découpe automatiquement les panneaux validés pour les inclure dans les ré-entraînements du modèle." %} +
+