feat: implement color-based sign verification and add ground truth management support

This commit is contained in:
kdeterme 2026-08-29 16:22:07 +02:00
parent a364b6161d
commit 577fb1232b
12 changed files with 2533 additions and 292 deletions

View file

@ -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",
]

View file

@ -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"])

View file

@ -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 = []
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)
# 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(file_path, size=224)
rgba = SyntheticSignAugmentor.load_png_to_numpy(fpath, size=224)
if rgba is not None:
templates_dict[c_idx] = rgba
if rgba is None:
# 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)
# Image canonique sur fond blanc
raw_bgr = rgba[:, :, :3]
raw_alpha = (rgba[:, :, 3] / 255.0)[:, :, np.newaxis]
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)
# 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 = []
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
canonical = (raw_bgr * raw_alpha + white_bg * (1.0 - raw_alpha)).astype(np.uint8)
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)
def to_tensor_norm(bgr_img: np.ndarray) -> np.ndarray:
rgb = cv2.cvtColor(bgr_img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
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),
}

View file

@ -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"]))

View file

@ -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

View file

@ -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

File diff suppressed because it is too large Load diff

View file

@ -32,12 +32,15 @@ from sign.ai import (
SyntheticSignAugmentor,
match_sign_from_ocr,
get_svg_url,
get_all_catalog_signs,
GroundTruthDatasetManager,
)
from sign.models import SignPanelType
User = get_user_model()
class SignCatalogMatcherTests(TestCase):
"""Tests pour le matching de texte OCR avec le catalogue de panneaux."""
@ -297,9 +300,11 @@ class SignAIApiAndViewsTests(APITestCase):
self.user = User.objects.create_user(
username="tester_ai",
email="tester_ai@example.com",
password="testpassword123"
password="testpassword123",
is_superuser=True
)
self.client.force_login(self.user)
self.client.force_authenticate(user=self.user)
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
img = np.ones((400, 400, 3), dtype=np.uint8) * 240
@ -404,3 +409,318 @@ class SignAIApiAndViewsTests(APITestCase):
inspections = SignPanelInspection.objects.filter(asset_object_id__in=[p1.id, p2.id])
self.assertEqual(inspections.count(), 2)
class GroundTruthAndCatalogAPITests(APITestCase):
"""Tests pour le gestionnaire de vérité terrain et les endpoints d'apprentissage actif."""
def setUp(self):
self.user = User.objects.create_user(username="ai_admin_tester", password="testpassword123", is_superuser=True)
self.non_admin_user = User.objects.create_user(username="standard_user", password="testpassword123", is_staff=False, is_superuser=False)
self.client.force_authenticate(user=self.user)
self.temp_gt_dir = tempfile.mkdtemp()
self.gt_manager = GroundTruthDatasetManager(base_dir=self.temp_gt_dir)
def _create_test_image(self) -> np.ndarray:
img = np.ones((200, 200, 3), dtype=np.uint8) * 200
# Dessiner un cercle bleu
cv2.circle(img, (100, 100), 60, (220, 50, 50), -1)
return img
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
img = Image.new("RGB", (200, 200), color=(100, 150, 200))
buf = io.BytesIO()
img.save(buf, format="JPEG")
buf.seek(0)
return SimpleUploadedFile("sample_ground_truth.jpg", buf.read(), content_type="image/jpeg")
def test_admin_permission_restrictions(self):
# 1. Utilisateur standard non-admin
self.client.force_login(self.non_admin_user)
self.client.force_authenticate(user=self.non_admin_user)
res_demo = self.client.get(reverse("sign:ai_demo"))
self.assertEqual(res_demo.status_code, status.HTTP_403_FORBIDDEN)
res_catalog = self.client.get(reverse("sign:api_catalog"))
self.assertEqual(res_catalog.status_code, status.HTTP_403_FORBIDDEN)
res_stats = self.client.get(reverse("sign:api_ground_truth_stats"))
self.assertEqual(res_stats.status_code, status.HTTP_403_FORBIDDEN)
# 2. Utilisateur administrateur
self.client.force_login(self.user)
self.client.force_authenticate(user=self.user)
res_demo_admin = self.client.get(reverse("sign:ai_demo"))
self.assertEqual(res_demo_admin.status_code, status.HTTP_200_OK)
res_catalog_admin = self.client.get(reverse("sign:api_catalog"))
self.assertEqual(res_catalog_admin.status_code, status.HTTP_200_OK)
def test_catalog_signs_function(self):
signs = get_all_catalog_signs()
self.assertGreater(len(signs), 10)
codes = [s["code"] for s in signs]
self.assertIn("E1", codes)
self.assertIn("B5", codes)
self.assertIn("C43", codes)
self.assertIn("ZE9A", codes)
def test_catalog_api_endpoint(self):
url = reverse("sign:api_catalog")
response = self.client.get(url, {"q": "ZE9A"})
self.assertEqual(response.status_code, status.HTTP_200_OK)
data = response.json()
self.assertGreater(data["match_count"], 0)
codes = [s["code"] for s in data["signs"]]
self.assertIn("ZE9A", codes)
def test_catalog_api_zone_filter(self):
url = reverse("sign:api_catalog")
response = self.client.get(url, {"category": "zone"})
self.assertEqual(response.status_code, status.HTTP_200_OK)
data = response.json()
self.assertGreater(data["match_count"], 0)
codes = [s["code"] for s in data["signs"]]
self.assertIn("ZE9A", codes)
def test_ground_truth_manager_save_and_stats(self):
img_np = self._create_test_image()
panels = [
{
"code": "E1",
"initial_code": "D1B",
"was_corrected": True,
"bbox": [20, 20, 180, 180],
"vertical_order": 1,
"ocr_text": "",
},
{
"code": "XD",
"initial_code": "XD",
"was_corrected": False,
"bbox": [80, 150, 120, 190],
"vertical_order": 2,
"ocr_text": "300 m",
}
]
res = self.gt_manager.save_ground_truth(
image_input=img_np,
panels=panels,
user_info={"username": "ai_admin_tester"},
notes="Test annotation E1"
)
self.assertEqual(res["status"], "success")
self.assertEqual(res["crops_saved"], 2)
stats = self.gt_manager.get_dataset_stats()
self.assertEqual(stats["total_images"], 1)
self.assertEqual(stats["total_crops"], 2)
self.assertEqual(stats["total_corrected_panels"], 1)
self.assertIn("E1", stats["classes_distribution"])
self.assertIn("XD", stats["classes_distribution"])
real_crops = self.gt_manager.get_real_crops_for_training()
self.assertIn("E1", real_crops)
self.assertEqual(len(real_crops["E1"]), 1)
def test_save_ground_truth_api(self):
url = reverse("sign:api_ground_truth")
file_obj = self._create_uploaded_image_file()
panels = [
{
"code": "E1",
"initial_code": "D1B",
"was_corrected": True,
"bbox": [10, 10, 190, 190],
"vertical_order": 1,
}
]
response = self.client.post(
url,
{
"image": file_obj,
"panels": json.dumps(panels),
"notes": "Correction terrain E1",
},
format="multipart"
)
self.assertEqual(response.status_code, status.HTTP_200_OK)
data = response.json()
self.assertEqual(data["status"], "success")
self.assertEqual(data["crops_saved"], 1)
self.assertIn("stats", data)
def test_ground_truth_stats_api(self):
url = reverse("sign:api_ground_truth_stats")
response = self.client.get(url)
self.assertEqual(response.status_code, status.HTTP_200_OK)
data = response.json()
self.assertIn("total_images", data)
self.assertIn("total_crops", data)
def test_specular_glare_and_synthetic_augmentation(self):
# Création d'une image RGBA synthétique
rgba = np.zeros((100, 100, 4), dtype=np.uint8)
rgba[:, :, 0] = 200 # Bleu
rgba[:, :, 2] = 20 # Rouge
rgba[:, :, 3] = 255 # Alpha
cv2.circle(rgba, (50, 50), 30, (0, 0, 220, 255), 8)
glare_applied = SyntheticSignAugmentor.apply_specular_glare(rgba[:, :, :3], rgba[:, :, 3])
self.assertEqual(glare_applied.shape, (100, 100, 3))
augmented = SyntheticSignAugmentor.augment_sign(rgba, size=224)
self.assertEqual(augmented.shape, (224, 224, 3))
self.assertEqual(augmented.dtype, np.uint8)
def test_hsv_color_profile_extraction_and_e1_filtering(self):
from sign.ai.catalog import extract_sign_color_profile, filter_and_rank_candidates_by_color
# 1. Image simulée E1 : Disque bleu avec bordure et diagonale rouge
e1_img = np.zeros((120, 120, 3), dtype=np.uint8)
# Fond extérieur noir/neutre
# Disque intérieur bleu
cv2.circle(e1_img, (60, 60), 50, (180, 50, 20), -1) # Bleu BGR
# Bordure et diagonale rouge
cv2.circle(e1_img, (60, 60), 50, (20, 20, 220), 8) # Rouge BGR
cv2.line(e1_img, (25, 95), (95, 25), (20, 20, 220), 8)
profile = extract_sign_color_profile(e1_img)
self.assertTrue(profile["has_red_and_blue"], f"Profile: {profile}")
self.assertGreater(profile["red_ratio"], 0.04)
self.assertGreater(profile["blue_ratio"], 0.05)
# Simulation de prédictions brutes erronées proposant D1B en 1ère position
raw_candidates = [
{"code": "D1B", "confidence": 0.65, "svg_url": "/static/D1B.svg"},
{"code": "E1", "confidence": 0.25, "svg_url": "/static/E1.svg"},
{"code": "E3", "confidence": 0.10, "svg_url": "/static/E3.svg"},
]
filtered = filter_and_rank_candidates_by_color(raw_candidates, e1_img)
self.assertEqual(filtered[0]["code"], "E1")
# D1B doit être éliminé en bas avec une confiance quasi nulle
d1b_item = next(c for c in filtered if c["code"] == "D1B")
self.assertLess(d1b_item["confidence"], 0.01)
def test_hsv_color_profile_pure_blue_filtering(self):
from sign.ai.catalog import extract_sign_color_profile, filter_and_rank_candidates_by_color
# Image simulée D1B : Disque bleu pur sans rouge avec flèche blanche
d1b_img = np.zeros((120, 120, 3), dtype=np.uint8)
cv2.circle(d1b_img, (60, 60), 50, (220, 80, 20), -1) # Bleu BGR
cv2.line(d1b_img, (60, 80), (60, 35), (255, 255, 255), 10) # Flèche blanche
profile = extract_sign_color_profile(d1b_img)
self.assertTrue(profile["is_pure_blue"], f"Profile: {profile}")
self.assertFalse(profile["has_red_and_blue"])
raw_candidates = [
{"code": "E1", "confidence": 0.60, "svg_url": "/static/E1.svg"},
{"code": "D1B", "confidence": 0.40, "svg_url": "/static/D1B.svg"},
]
filtered = filter_and_rank_candidates_by_color(raw_candidates, d1b_img)
self.assertEqual(filtered[0]["code"], "D1B")
e1_item = next(c for c in filtered if c["code"] == "E1")
self.assertLess(e1_item["confidence"], 0.01)
def test_no_red_f4a_exclusion_and_zone_parking(self):
from sign.ai.catalog import match_sign_from_ocr
# Panneau bleu/blanc sans rouge avec texte ZONE -> Doit être ZE9A et JAMAIS F4A
blue_white_img = np.ones((100, 100, 3), dtype=np.uint8) * 240
cv2.rectangle(blue_white_img, (20, 20), (80, 80), (200, 60, 20), -1) # Bleu
res = match_sign_from_ocr("ZONE Rappel Herhaling", crop_bgr=blue_white_img)
self.assertIsNotNone(res)
self.assertEqual(res["code"], "ZE9A")
self.assertNotEqual(res["code"], "F4A")
def test_gxc_arrow_distance_matching(self):
from sign.ai.catalog import match_sign_from_ocr
# Panonceau blanc avec distance courte "11m"
white_img = np.ones((60, 100, 3), dtype=np.uint8) * 250
res = match_sign_from_ocr("11m", crop_bgr=white_img)
self.assertIsNotNone(res)
self.assertEqual(res["code"], "GXC")
def test_type0_and_type0b_subplate_simplification(self):
from sign.ai.catalog import match_sign_from_ocr
# Panonceau textuel bleu
blue_img = np.ones((60, 100, 3), dtype=np.uint8) * 20
blue_img[:, :] = (180, 50, 20) # Bleu
res_blue = match_sign_from_ocr("Tarif specifique horaire centre", crop_bgr=blue_img)
self.assertIsNotNone(res_blue)
self.assertEqual(res_blue["code"], "TYPE0")
# Panonceau textuel blanc
white_img = np.ones((60, 100, 3), dtype=np.uint8) * 245
res_white = match_sign_from_ocr("Forfait 50e stationnement longue duree", crop_bgr=white_img)
self.assertIsNotNone(res_white)
self.assertEqual(res_white["code"], "TYPE0B")
def test_danger_inner_pictogram_bicycle_discrimination(self):
from sign.ai.catalog import discriminate_inner_pictogram
# Simuler un panneau triangulaire de danger avec un vélo (symbole horizontal et 2 roues)
tri_img = np.ones((140, 140, 3), dtype=np.uint8) * 240
# Bordure rouge
pts = np.array([[70, 10], [15, 125], [125, 125]], np.int32)
cv2.polylines(tri_img, [pts], isClosed=True, color=(20, 20, 220), thickness=8)
# Pictogramme vélo (deux roues noires distinctes en bas + cadre)
cv2.circle(tri_img, (50, 95), 8, (10, 10, 10), -1)
cv2.circle(tri_img, (90, 95), 8, (10, 10, 10), -1)
cv2.line(tri_img, (50, 95), (70, 75), (10, 10, 10), 3)
cv2.line(tri_img, (90, 95), (70, 75), (10, 10, 10), 3)
candidates = [
{"code": "A15", "confidence": 0.70}, # Faussement prédit A15 (piéton)
{"code": "A25", "confidence": 0.20}, # Vrai A25 (vélo)
{"code": "A14", "confidence": 0.10},
]
discriminated = discriminate_inner_pictogram(tri_img, candidates)
self.assertEqual(discriminated[0]["code"], "A25")
self.assertGreater(discriminated[0]["confidence"], 0.50)
def test_blue_vs_white_exception_subplate_discrimination(self):
from sign.ai.catalog import match_sign_from_ocr
# 1. Panonceau bleu avec texte blanc "EXCEPTE RIVERAINS" -> TYPE0 (JAMAIS M2)
blue_img = np.zeros((60, 120, 3), dtype=np.uint8)
blue_img[:, :] = (190, 60, 20) # Bleu pur
res_blue = match_sign_from_ocr("EXCEPTE RIVERAINS UITGEZ BEWONERS", crop_bgr=blue_img)
self.assertIsNotNone(res_blue)
self.assertEqual(res_blue["code"], "TYPE0")
self.assertNotEqual(res_blue["code"], "M2")
# 2. Panonceau blanc avec texte noir "EXCEPTE RIVERAINS" -> M2
white_img = np.ones((60, 120, 3), dtype=np.uint8) * 245
res_white = match_sign_from_ocr("EXCEPTE RIVERAINS UITGEZ BEWONERS", crop_bgr=white_img)
self.assertIsNotNone(res_white)
self.assertEqual(res_white["code"], "M2")
def test_blue_vs_white_distance_subplate_discrimination(self):
from sign.ai.catalog import match_sign_from_ocr
# 1. Panonceau bleu avec texte blanc "50 m" -> TYPEIA_50M (JAMAIS M1)
blue_img = np.zeros((50, 100, 3), dtype=np.uint8)
blue_img[:, :] = (200, 70, 20) # Bleu
res_blue = match_sign_from_ocr("50 m", crop_bgr=blue_img)
self.assertIsNotNone(res_blue)
self.assertEqual(res_blue["code"], "TYPEIA_50M")
self.assertNotEqual(res_blue["code"], "M1")
# 2. Panonceau blanc avec distance longue "500 m" -> M1
white_img = np.ones((50, 100, 3), dtype=np.uint8) * 240
res_white = match_sign_from_ocr("500 m", crop_bgr=white_img)
self.assertIsNotNone(res_white)
self.assertEqual(res_white["code"], "M1")

View file

@ -8,7 +8,12 @@ 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("api/catalog/", views_ai.SignCatalogSearchAPIView.as_view(), name="api_catalog"),
path("api/ground-truth/", views_ai.SaveGroundTruthAPIView.as_view(), name="api_ground_truth"),
path("api/ground-truth/stats/", views_ai.GroundTruthStatsAPIView.as_view(), name="api_ground_truth_stats"),
path("api/retrain/", views_ai.RetrainClassifierAPIView.as_view(), name="api_retrain"),
path("api/quick-create/", views_ai.QuickCreateSignWithAIView.as_view(), name="api_quick_create"),
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"),

View file

@ -7,7 +7,7 @@ import logging
from typing import Any, Dict
from django.views.generic import TemplateView
from django.contrib.auth.mixins import LoginRequiredMixin
from django.contrib.auth.mixins import LoginRequiredMixin, UserPassesTestMixin
from django.http import JsonResponse, HttpResponseBadRequest
from django.utils.translation import gettext_lazy as _
from django.views.decorators.csrf import csrf_exempt
@ -16,20 +16,55 @@ 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 rest_framework.permissions import IsAuthenticated, BasePermission
from .ai import SignDetectionService, SIGN_CATALOG, get_svg_url, SignClassifierEngine
from .ai import (
SignDetectionService,
SIGN_CATALOG,
get_svg_url,
get_all_catalog_signs,
SignClassifierEngine,
GroundTruthDatasetManager,
)
logger = logging.getLogger(__name__)
class SignAIDemoView(LoginRequiredMixin, TemplateView):
def is_admin_user(user) -> bool:
"""
Vérifie si l'utilisateur est administrateur (superuser, ou rôle 'admin').
"""
if not user or not user.is_authenticated:
return False
if user.is_superuser:
return True
try:
return user.config.roles.filter(name__iexact="admin").exists()
except Exception:
return False
class IsSignAIAdmin(BasePermission):
"""
Permission DRF accordée uniquement aux administrateurs (is_superuser, ou rôle admin).
"""
def has_permission(self, request, view):
return is_admin_user(request.user)
class SignAIDemoView(UserPassesTestMixin, TemplateView):
"""
Page de démonstration et banc de test interactif pour la reconnaissance de panneaux et panonceaux.
Réservé aux administrateurs.
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.
Intègre une interface de correction et d'apprentissage continu (Active Learning).
"""
template_name = "sign/ai_demo.html"
raise_exception = True
def test_func(self):
return is_admin_user(self.request.user)
def get_context_data(self, **kwargs: Any) -> Dict[str, Any]:
context = super().get_context_data(**kwargs)
@ -37,8 +72,14 @@ class SignAIDemoView(LoginRequiredMixin, TemplateView):
engine = SignClassifierEngine.get_instance()
context["classifier_is_trained"] = engine.is_trained()
context["classifier_meta"] = engine.get_metadata()
gt_manager = GroundTruthDatasetManager.get_instance()
context["ground_truth_stats"] = gt_manager.get_dataset_stats()
all_signs = get_all_catalog_signs()
context["catalog_total_count"] = len(all_signs)
return context
def post(self, request, *args, **kwargs):
image_file = request.FILES.get("image")
image_base64 = request.POST.get("image_base64")
@ -73,6 +114,7 @@ class SignAIDemoView(LoginRequiredMixin, TemplateView):
)
@method_decorator(csrf_exempt, name='dispatch')
class DetectSignAPIView(APIView):
"""
Endpoint API REST pour l'analyse mobile/terrain d'une photo de signalisation.
@ -288,3 +330,163 @@ class QuickCreateSignWithAIView(APIView):
}, status=status.HTTP_500_INTERNAL_SERVER_ERROR)
class SignCatalogSearchAPIView(APIView):
"""
Endpoint API pour la recherche, le filtrage et l'autocomplétion des panneaux
dans le catalogue officiel (500+ types).
Accepte les paramètres optionnels ?q=... (texte) et ?category=...
"""
permission_classes = [IsSignAIAdmin]
def get(self, request, *args, **kwargs):
query = request.GET.get("q", "").strip().lower()
category = request.GET.get("category", "").strip().lower()
all_signs = get_all_catalog_signs()
results = []
for s in all_signs:
code_lower = s.get("code", "").lower()
cat_lower = s.get("category", "").lower()
if category and category != "all":
# Permettre le matching souple (ex: ZE9A matche à la fois 'zone' et 'parking')
if category == "zone":
if not (cat_lower == "zone" or code_lower.startswith("z") or code_lower.startswith(("f4", "f103", "f105", "f12"))):
continue
elif category == "parking":
if not (cat_lower == "parking" or code_lower.startswith("e") or code_lower.startswith("ze") or code_lower.startswith("g")):
continue
elif cat_lower != category:
continue
if query:
fr_lower = s.get("name_fr", "").lower()
nl_lower = s.get("name_nl", "").lower()
if not (query in code_lower or query in fr_lower or query in nl_lower):
continue
results.append(s)
return Response({
"total_count": len(all_signs),
"match_count": len(results),
"signs": results[:500], # Large limite pour inclure l'ensemble des résultats
}, status=status.HTTP_200_OK)
@method_decorator(csrf_exempt, name='dispatch')
class SaveGroundTruthAPIView(APIView):
"""
Endpoint API pour enregistrer une image annotée/corrigée dans le dataset de vérité terrain (Ground Truth).
Découpe automatiquement chaque panneau selon sa boîte englobante et l'organise dans le dossier de classe.
"""
permission_classes = [IsSignAIAdmin]
def post(self, request, *args, **kwargs):
image_file = request.FILES.get("image")
image_base64 = request.data.get("image_base64")
panels_raw = request.data.get("panels")
notes = request.data.get("notes", "")
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
)
panels = []
if isinstance(panels_raw, str):
try:
panels = json.loads(panels_raw)
except Exception:
panels = []
elif isinstance(panels_raw, list):
panels = panels_raw
if not panels:
return Response(
{"status": "error", "message": "Aucune information de panneau fournie pour l'annotation."},
status=status.HTTP_400_BAD_REQUEST
)
try:
if image_file:
image_bytes = image_file.read()
else:
if "," in image_base64:
image_base64 = image_base64.split(",", 1)[1]
image_base64 = image_base64.replace(" ", "+").strip()
import base64
image_bytes = base64.b64decode(image_base64)
user_info = {}
if request.user and request.user.is_authenticated:
user_info = {
"id": request.user.id,
"username": request.user.username,
"email": getattr(request.user, "email", ""),
}
manager = GroundTruthDatasetManager.get_instance()
res = manager.save_ground_truth(
image_input=image_bytes,
panels=panels,
user_info=user_info,
notes=notes,
)
return Response(res, status=status.HTTP_200_OK)
except Exception as exc:
logger.error("Erreur lors de l'enregistrement de la vérité terrain : %s", exc, exc_info=True)
return Response(
{"status": "error", "message": str(exc)},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
@method_decorator(csrf_exempt, name='dispatch')
class GroundTruthStatsAPIView(APIView):
"""
Endpoint API pour consulter les statistiques du dataset d'apprentissage terrain.
"""
permission_classes = [IsSignAIAdmin]
def get(self, request, *args, **kwargs):
try:
manager = GroundTruthDatasetManager.get_instance()
stats = manager.get_dataset_stats()
return Response(stats, status=status.HTTP_200_OK)
except Exception as exc:
logger.error("Erreur consultation stats vérité terrain : %s", exc)
return Response({"status": "error", "message": str(exc)}, status=status.HTTP_500_INTERNAL_SERVER_ERROR)
@method_decorator(csrf_exempt, name='dispatch')
class RetrainClassifierAPIView(APIView):
"""
Endpoint API pour déclencher le ré-entraînement du classifieur ONNX avec intégration des données terrain.
"""
permission_classes = [IsSignAIAdmin]
def post(self, request, *args, **kwargs):
try:
engine = SignClassifierEngine.get_instance()
epochs = int(request.data.get("epochs", 12))
samples_per_class = int(request.data.get("samples_per_class", 50))
logger.info("Lancement du ré-entraînement classifieur IA (%d époques, %d samples/classe)...", epochs, samples_per_class)
result = engine.train_from_svgs(
epochs=epochs,
samples_per_class=samples_per_class,
include_ground_truth=True,
)
return Response(result, status=status.HTTP_200_OK)
except Exception as exc:
logger.error("Erreur lors du ré-entraînement du classifieur : %s", exc, exc_info=True)
return Response(
{"status": "error", "message": str(exc)},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)

View file

@ -1,3 +1,7 @@
# Les librairies IA en local (ultralytics, PyTorch, etc.) ont été supprimées.
# L'application utilise désormais un appel à un microservice externe pour la détection.
# Ce fichier reste vide ou pour d'éventuelles dépendances futures du client d'API.
# Dépendances pour l'entraînement et l'inférence du classifieur de panneaux
--extra-index-url https://download.pytorch.org/whl/cpu
torch>=2.2.0
torchvision>=0.17.0
onnxruntime>=1.17.0
rapidocr-onnxruntime>=1.3.0
opencv-python-headless>=4.8.0