feat: implement color-based sign verification and add ground truth management support
This commit is contained in:
parent
a364b6161d
commit
577fb1232b
12 changed files with 2533 additions and 292 deletions
|
|
@ -1,6 +1,7 @@
|
||||||
from .detector import SignDetectionService
|
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 .classifier import SignClassifierEngine, SyntheticSignAugmentor
|
||||||
|
from .ground_truth import GroundTruthDatasetManager
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"SignDetectionService",
|
"SignDetectionService",
|
||||||
|
|
@ -8,6 +9,10 @@ __all__ = [
|
||||||
"get_svg_url",
|
"get_svg_url",
|
||||||
"match_sign_from_ocr",
|
"match_sign_from_ocr",
|
||||||
"classify_sign_visual",
|
"classify_sign_visual",
|
||||||
|
"get_all_catalog_signs",
|
||||||
"SignClassifierEngine",
|
"SignClassifierEngine",
|
||||||
"SyntheticSignAugmentor",
|
"SyntheticSignAugmentor",
|
||||||
|
"GroundTruthDatasetManager",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,8 @@ et leurs fichiers vectoriels SVG correspondants.
|
||||||
"""
|
"""
|
||||||
import re
|
import re
|
||||||
from typing import Optional, Dict, Any, List
|
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)
|
# 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/"
|
DEFAULT_SVG_BASE_PATH = "/static/assets/road_signs/2025/"
|
||||||
|
|
@ -299,6 +301,60 @@ SIGN_CATALOG = {
|
||||||
"category": "panonceau",
|
"category": "panonceau",
|
||||||
"shape": "rectangle",
|
"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"
|
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é.
|
Tente d'associer un texte extrait par OCR à un type de panneau normalisé du catalogue.
|
||||||
Exemples:
|
Intègre une vérification colorimétrique stricte :
|
||||||
- "STOP" -> B5
|
- Si absence de rouge (red_ratio < 0.035), élimine strictement F4A (Zone 30), C43 (Limitation), C... et A...
|
||||||
- "50 km" ou "50" -> C43 (Limitation de vitesse 50 km/h)
|
- Détecte les panonceaux de distance / flèche montante GXC / XD
|
||||||
- "ZONE P" ou "ZONE ... EXCEPTE CARTE" -> ZE9A (Zone de stationnement)
|
- Simplifie les panonceaux textuels bruts : TYPE0 (fond bleu) ou TYPE0B (fond blanc)
|
||||||
- "ZONE 30" -> F4A (Zone 30)
|
|
||||||
- "SAUF RIVERAINS" -> M2
|
|
||||||
- "300 M" -> M1 (Distance)
|
|
||||||
"""
|
"""
|
||||||
if not ocr_text:
|
if not ocr_text:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
cleaned = ocr_text.strip().upper()
|
cleaned = re.sub(r"[^A-Za-z0-9\s/.,:-]", " ", ocr_text.upper())
|
||||||
cleaned_inline = re.sub(r"\s+", " ", cleaned)
|
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
|
# 1. STOP
|
||||||
if "STOP" in cleaned_inline:
|
if "STOP" in cleaned_inline:
|
||||||
return {
|
return {
|
||||||
"code": "B5",
|
"code": "B5",
|
||||||
"confidence": 0.96,
|
"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"),
|
"svg_url": get_svg_url("B5"),
|
||||||
"matched_by": "text_stop",
|
"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"),
|
"svg_url": get_svg_url("F4B"),
|
||||||
"matched_by": "text_end_zone",
|
"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 {
|
return {
|
||||||
"code": "ZE9A",
|
"code": code_zone,
|
||||||
"confidence": 0.95,
|
"confidence": 0.95,
|
||||||
"data": SIGN_CATALOG.get("ZE9A", {
|
"data": SIGN_CATALOG.get(code_zone, {
|
||||||
"name_fr": "Zone de stationnement réglementé",
|
"name_fr": "Zone de stationnement réglementé",
|
||||||
"name_nl": "Zone voor gereglementeerd parkeren",
|
"name_nl": "Zone voor gereglementeerd parkeren",
|
||||||
"category": "parking",
|
"category": "parking",
|
||||||
}),
|
}),
|
||||||
"svg_url": get_svg_url("ZE9A"),
|
"svg_url": get_svg_url(code_zone),
|
||||||
"extracted_text": ocr_text.strip(),
|
"extracted_text": ocr_text.strip(),
|
||||||
"matched_by": "text_zone_parking",
|
"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)
|
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))
|
speed = int(speed_match.group(1))
|
||||||
code = "F4A"
|
code = "F4A"
|
||||||
return {
|
return {
|
||||||
|
|
@ -404,6 +470,7 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]:
|
||||||
"value": speed,
|
"value": speed,
|
||||||
"matched_by": "text_zone_speed",
|
"matched_by": "text_zone_speed",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Zone piétonne
|
# Zone piétonne
|
||||||
if re.search(r"\b(PIETON|VOETGANGER|PIETONS|VOETGANGERS)\b", cleaned_inline):
|
if re.search(r"\b(PIETON|VOETGANGER|PIETONS|VOETGANGERS)\b", cleaned_inline):
|
||||||
return {
|
return {
|
||||||
|
|
@ -417,21 +484,21 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]:
|
||||||
"svg_url": get_svg_url("F103"),
|
"svg_url": get_svg_url("F103"),
|
||||||
"matched_by": "text_zone_pedestrian",
|
"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 {
|
return {
|
||||||
"code": "F4A",
|
"code": code_fallback,
|
||||||
"confidence": 0.88,
|
"confidence": 0.88,
|
||||||
"data": SIGN_CATALOG.get("F4A", {"name_fr": "Zone réglementée", "name_nl": "Gereglementeerde zone", "category": "indication"}),
|
"data": SIGN_CATALOG.get(code_fallback, {"name_fr": "Zone réglementée", "name_nl": "Gereglementeerde zone", "category": "indication"}),
|
||||||
"svg_url": get_svg_url("F4A"),
|
"svg_url": get_svg_url(code_fallback),
|
||||||
"extracted_text": ocr_text.strip(),
|
"extracted_text": ocr_text.strip(),
|
||||||
"matched_by": "text_zone",
|
"matched_by": "text_zone",
|
||||||
}
|
}
|
||||||
|
|
||||||
# 3. VITESSE MAXIMALE AUTORISÉE (C43 : "50", "50 km", "50 km/h", "30 km", "70 km/h", "90", "120")
|
# 3. VITESSE MAXIMALE AUTORISÉE (C43 : "50", "50 km/h", "30 km") -> STRICTEMENT CONDITIONNÉE À LA PRÉSENCE DE ROUGE
|
||||||
# Note : "50 km" ou "50 km/h" sur un panneau de limitation est une vitesse C43 et NON une distance M1 !
|
|
||||||
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)
|
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:
|
if speed_match and has_red:
|
||||||
# Exclure si le texte est explicitement une distance comme "50 m", "300 m", "1.5 km" (avec décimale ou mètres)
|
|
||||||
is_explicit_distance = bool(re.search(r"\b(\d+\s*M|\d+[,.]\d+\s*KM)\b", cleaned_inline))
|
is_explicit_distance = bool(re.search(r"\b(\d+\s*M|\d+[,.]\d+\s*KM)\b", cleaned_inline))
|
||||||
if not is_explicit_distance:
|
if not is_explicit_distance:
|
||||||
speed = int(speed_match.group(1))
|
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",
|
"matched_by": "text_speed_limit",
|
||||||
}
|
}
|
||||||
|
|
||||||
# 4. Panonceaux d'exception ("Sauf ...", "Excepté ...", "Uitgezonderd ...")
|
# 4. Panonceaux de distance ou flèche de zone (ex: "50m", "11m", "12 m", "300 m", "50 m")
|
||||||
if re.search(r"\b(SAUF|EXCEPTE|EXCEPTE|UITGEZONDERD|RIVERAINS?|AANGELANDEN?|VELOS?|FIETS)\b", cleaned_inline, re.IGNORECASE):
|
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)
|
||||||
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)
|
|
||||||
if dist_match:
|
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()
|
unit = (dist_match.group(2) or dist_match.group(4) or "m").lower()
|
||||||
val = float(val_str)
|
val = float(val_str)
|
||||||
if unit == "km":
|
if unit == "km":
|
||||||
val *= 1000.0
|
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 {
|
return {
|
||||||
"code": "M1",
|
"code": "M1",
|
||||||
"confidence": 0.88,
|
"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"),
|
"svg_url": get_svg_url("M1"),
|
||||||
"value": val,
|
"value": val,
|
||||||
"extracted_text": ocr_text.strip(),
|
"extracted_text": ocr_text.strip(),
|
||||||
"matched_by": "text_distance",
|
"matched_by": "text_distance_white_m1",
|
||||||
}
|
}
|
||||||
|
|
||||||
# 6. Tonnage ("3.5 t", "7.5t")
|
# 5. Stationnement Payant / Betalend
|
||||||
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
|
|
||||||
if re.search(r"\b(PAYANT|BETALEND|HORODATEUR|TICKET)\b", cleaned_inline):
|
if re.search(r"\b(PAYANT|BETALEND|HORODATEUR|TICKET)\b", cleaned_inline):
|
||||||
return {
|
return {
|
||||||
"code": "GVII_BETALEND",
|
"code": "GVII_BETALEND",
|
||||||
|
|
@ -528,7 +587,7 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]:
|
||||||
"matched_by": "text_parking_payant",
|
"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):
|
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 {
|
return {
|
||||||
"code": "GVIID_ELEKTRISCHE_WAGENS",
|
"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",
|
"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):
|
if re.search(r"\b(DISQUE|PARKEERSCHIJF)\b", cleaned_inline):
|
||||||
return {
|
return {
|
||||||
"code": "E9A_PARKEERSCHIJF",
|
"code": "E9A_PARKEERSCHIJF",
|
||||||
|
|
@ -558,9 +617,338 @@ def match_sign_from_ocr(ocr_text: str) -> Optional[Dict[str, Any]]:
|
||||||
"matched_by": "text_parking_disc",
|
"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
|
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]:
|
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.
|
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)
|
hsv = cv2.cvtColor(crop_bgr, cv2.COLOR_BGR2HSV)
|
||||||
total_px = float(h * w)
|
total_px = float(h * w)
|
||||||
|
|
||||||
# Masques couleur HSV
|
# Masques couleur HSV pour l'analyse spatiale de la forme
|
||||||
blue_mask = cv2.inRange(hsv, np.array([95, 40, 30]), np.array([135, 255, 255]))
|
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_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_mask2 = cv2.inRange(hsv, np.array([160, 50, 40]), np.array([180, 255, 255]))
|
||||||
red_mask = red_mask1 | red_mask2
|
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
|
profile = extract_sign_color_profile(crop_bgr)
|
||||||
red_ratio = np.count_nonzero(red_mask) / total_px
|
blue_ratio = profile["blue_ratio"]
|
||||||
yellow_ratio = np.count_nonzero(yellow_mask) / total_px
|
red_ratio = profile["red_ratio"]
|
||||||
|
yellow_ratio = profile["yellow_ratio"]
|
||||||
aspect_ratio = w / float(h)
|
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)
|
# 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
|
# Détection de texte ou lettre P / D dans l'OCR
|
||||||
ocr_clean = ocr_text.strip().upper()
|
ocr_clean = ocr_text.strip().upper()
|
||||||
if ocr_clean in ("P", "D", "🅿") or "PARKING" in ocr_clean or "PARKEREN" in ocr_clean:
|
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",
|
"matched_by": "visual_generic",
|
||||||
"confidence": 0.60,
|
"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"])
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -145,10 +145,63 @@ class SyntheticSignAugmentor:
|
||||||
|
|
||||||
return bg
|
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
|
@classmethod
|
||||||
def augment_sign(cls, rgba_sign: np.ndarray, size: int = 224) -> np.ndarray:
|
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.
|
sur le panneau RGBA et l'incruste sur un fond synthétique.
|
||||||
Retourne une image BGR 3 canaux de taille (size, size).
|
Retourne une image BGR 3 canaux de taille (size, size).
|
||||||
"""
|
"""
|
||||||
|
|
@ -157,13 +210,13 @@ class SyntheticSignAugmentor:
|
||||||
alpha = rgba_sign[:, :, 3].copy()
|
alpha = rgba_sign[:, :, 3].copy()
|
||||||
|
|
||||||
# 1. Déformation Perspective 3D (Angle de vue caméra smartphone / véhicule)
|
# 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
|
# Points sources
|
||||||
src_pts = np.float32([[0, 0], [w, 0], [w, h], [0, h]])
|
src_pts = np.float32([[0, 0], [w, 0], [w, h], [0, h]])
|
||||||
|
|
||||||
# Décalages de perspective aléatoires
|
# Décalages de perspective aléatoires
|
||||||
max_shift = 0.12
|
max_shift = 0.14
|
||||||
dx1 = random.uniform(-w * max_shift, w * max_shift)
|
dx1 = random.uniform(-w * max_shift, w * max_shift)
|
||||||
dy1 = random.uniform(-h * max_shift, h * max_shift)
|
dy1 = random.uniform(-h * max_shift, h * max_shift)
|
||||||
dx2 = random.uniform(-w * max_shift, w * 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_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)
|
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°)
|
# 2. Rotation légère (-12° à +12°)
|
||||||
rot_angle = random.uniform(-10, 10)
|
rot_angle = random.uniform(-12, 12)
|
||||||
M_rot = cv2.getRotationMatrix2D((size / 2.0, size / 2.0), rot_angle, 1.0)
|
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_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)
|
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]
|
alpha_norm = (warped_alpha.astype(np.float32) / 255.0)[:, :, np.newaxis]
|
||||||
bgr_float = warped_bgr.astype(np.float32)
|
bgr_float = warped_bgr.astype(np.float32)
|
||||||
|
|
||||||
grad_angle = random.uniform(0, 2 * math.pi)
|
grad_angle = random.uniform(0, 2 * math.pi)
|
||||||
gx, gy = math.cos(grad_angle), math.sin(grad_angle)
|
gx, gy = math.cos(grad_angle), math.sin(grad_angle)
|
||||||
y_coords, x_coords = np.mgrid[0:size, 0:size]
|
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 = 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.55, 1.45)[:, :, np.newaxis]
|
light_grad = np.clip(light_grad, 0.50, 1.50)[:, :, np.newaxis]
|
||||||
bgr_float = bgr_float * light_grad
|
bgr_float = bgr_float * light_grad
|
||||||
|
|
||||||
# Luminosité & Contraste globaux
|
# 5. Variations renforcées de Saturation et Luminosité HSV (peinture vieillie / plein soleil)
|
||||||
brightness = random.uniform(0.75, 1.25)
|
hsv = cv2.cvtColor(np.clip(bgr_float, 0, 255).astype(np.uint8), cv2.COLOR_BGR2HSV).astype(np.float32)
|
||||||
contrast = random.uniform(0.80, 1.25)
|
sat_factor = random.uniform(0.55, 1.45)
|
||||||
bgr_float = np.clip((bgr_float - 128.0) * contrast + 128.0 * brightness, 0, 255)
|
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
|
# 6. Contraste dynamique et correction Gamma
|
||||||
b_shift = random.uniform(0.92, 1.08)
|
contrast = random.uniform(0.70, 1.35)
|
||||||
g_shift = random.uniform(0.92, 1.08)
|
brightness_shift = random.uniform(-20, 25)
|
||||||
r_shift = random.uniform(0.92, 1.08)
|
bgr_float = np.clip((bgr_float - 128.0) * contrast + 128.0 + brightness_shift, 0, 255)
|
||||||
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)
|
|
||||||
|
|
||||||
# 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)
|
bg = cls.generate_random_background(size)
|
||||||
composite = (bgr_float * alpha_norm + bg * (1.0 - alpha_norm)).astype(np.uint8)
|
composite = (bgr_float * alpha_norm + bg * (1.0 - alpha_norm)).astype(np.uint8)
|
||||||
|
|
||||||
# 5. Flou optique & Bruit de capteur
|
# 9. Flou optique, flou de bougé & Bruit de capteur
|
||||||
if random.random() < 0.35:
|
if random.random() < 0.40:
|
||||||
ksize = random.choice([3, 5])
|
ksize = random.choice([3, 5])
|
||||||
composite = cv2.GaussianBlur(composite, (ksize, ksize), 0)
|
composite = cv2.GaussianBlur(composite, (ksize, ksize), 0)
|
||||||
|
|
||||||
if random.random() < 0.30:
|
if random.random() < 0.35:
|
||||||
noise = np.random.normal(0, random.uniform(2, 8), composite.shape).astype(np.int16)
|
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)
|
composite = np.clip(composite.astype(np.int16) + noise, 0, 255).astype(np.uint8)
|
||||||
|
|
||||||
return composite
|
return composite
|
||||||
|
|
@ -290,15 +355,17 @@ class SignClassifierEngine:
|
||||||
def train_from_svgs(
|
def train_from_svgs(
|
||||||
self,
|
self,
|
||||||
signs_dir: Union[str, Path] = DEFAULT_SIGNS_DIR,
|
signs_dir: Union[str, Path] = DEFAULT_SIGNS_DIR,
|
||||||
samples_per_class: int = 15,
|
samples_per_class: int = 50,
|
||||||
epochs: int = 10,
|
epochs: int = 10,
|
||||||
batch_size: int = 32,
|
batch_size: int = 32,
|
||||||
learning_rate: float = 0.001,
|
learning_rate: float = 0.001,
|
||||||
|
include_ground_truth: bool = True,
|
||||||
|
ground_truth_dir: Optional[Union[str, Path]] = None,
|
||||||
progress_callback: Optional[Any] = None,
|
progress_callback: Optional[Any] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Scanne le dossier des SVGs et PNGs, génère un dataset synthétique équilibré,
|
Scanne le dossier des SVGs et PNGs ainsi que les données de vérité terrain réelles,
|
||||||
entraîne MobileNetV3-Small et exporte vers ONNX.
|
génère un dataset enrichi, entraîne MobileNetV3-Small et exporte vers ONNX.
|
||||||
"""
|
"""
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
@ -308,62 +375,132 @@ class SignClassifierEngine:
|
||||||
|
|
||||||
start_time = time.perf_counter()
|
start_time = time.perf_counter()
|
||||||
template_files = self.discover_templates(signs_dir)
|
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)
|
num_classes = len(classes)
|
||||||
|
class_to_idx = {c: i for i, c in enumerate(classes)}
|
||||||
|
|
||||||
if num_classes < 2:
|
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)
|
logger.info("🔍 %d types de panneaux officiels découverts pour l'entraînement.", num_classes)
|
||||||
if progress_callback:
|
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
|
# 1. Chargement compact en mémoire des templates RGBA de base (~100 Mo max)
|
||||||
x_list = []
|
templates_dict: Dict[int, np.ndarray] = {}
|
||||||
y_list = []
|
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):
|
# 2. Chargement compact des vrais crops terrain annotés (~20 Mo max)
|
||||||
file_path = template_files[code]
|
real_crops_list: List[Tuple[np.ndarray, int]] = []
|
||||||
if file_path.suffix.lower() == '.svg':
|
for code, crop_paths in real_crops_map.items():
|
||||||
rgba = SyntheticSignAugmentor.render_svg_to_numpy(file_path, size=224)
|
if code in class_to_idx:
|
||||||
else:
|
c_idx = class_to_idx[code]
|
||||||
rgba = SyntheticSignAugmentor.load_png_to_numpy(file_path, size=224)
|
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:
|
canvas = np.ones((224, 224, 3), dtype=np.uint8) * 128
|
||||||
continue
|
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
|
# 3. Dataset PyTorch dynamique à génération à la volée (Ultra faible empreinte RAM < 150MB)
|
||||||
raw_bgr = rgba[:, :, :3]
|
class DynamicSignDataset(torch.utils.data.Dataset):
|
||||||
raw_alpha = (rgba[:, :, 3] / 255.0)[:, :, np.newaxis]
|
def __init__(self, tmpl_map, real_list, n_samples):
|
||||||
white_bg = np.ones((224, 224, 3), dtype=np.uint8) * 255
|
self.tmpl_map = tmpl_map
|
||||||
canonical = (raw_bgr * raw_alpha + white_bg * (1.0 - raw_alpha)).astype(np.uint8)
|
self.real_list = real_list
|
||||||
|
self.n_samples = n_samples
|
||||||
|
self.index_items = []
|
||||||
|
|
||||||
def to_tensor_norm(bgr_img: np.ndarray) -> np.ndarray:
|
for class_idx in tmpl_map:
|
||||||
rgb = cv2.cvtColor(bgr_img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
|
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)
|
mean = np.array([0.485, 0.456, 0.406], dtype=np.float32)
|
||||||
std = np.array([0.229, 0.224, 0.225], dtype=np.float32)
|
std = np.array([0.229, 0.224, 0.225], dtype=np.float32)
|
||||||
norm = (rgb - mean) / std
|
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))
|
dataset = DynamicSignDataset(templates_dict, real_crops_list, samples_per_class)
|
||||||
y_list.append(class_idx)
|
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
|
logger.info(
|
||||||
for _ in range(samples_per_class):
|
"📦 Dataset d'entraînement dynamique prêt : %d images (dont %d réelles/augmentées) pour %d classes.",
|
||||||
aug_bgr = SyntheticSignAugmentor.augment_sign(rgba, size=224)
|
total_samples, real_samples_count, num_classes
|
||||||
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)
|
|
||||||
if progress_callback:
|
if progress_callback:
|
||||||
progress_callback(f"Entraînement MobileNetV3 sur {total_samples} images ({epochs} époques)...")
|
progress_callback(
|
||||||
|
f"Entraînement MobileNetV3 sur {total_samples} images ({len(real_crops_list)} crops réels terrain, {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)
|
|
||||||
|
|
||||||
# 3. Modèle MobileNetV3-Small pré-entraîné
|
# 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"))
|
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()
|
model.train()
|
||||||
for epoch in range(1, epochs + 1):
|
for epoch in range(1, epochs + 1):
|
||||||
epoch_loss = 0.0
|
running_loss = 0.0
|
||||||
correct = 0
|
correct = 0
|
||||||
total = 0
|
total = 0
|
||||||
for batch_x, batch_y in loader:
|
for batch_x, batch_y in loader:
|
||||||
|
|
@ -392,24 +529,25 @@ class SignClassifierEngine:
|
||||||
loss.backward()
|
loss.backward()
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
|
|
||||||
epoch_loss += loss.item() * batch_x.size(0)
|
running_loss += loss.item() * batch_x.size(0)
|
||||||
_, predicted = outputs.max(1)
|
_, preds = torch.max(outputs, 1)
|
||||||
|
correct += torch.sum(preds == batch_y.data).item()
|
||||||
total += batch_y.size(0)
|
total += batch_y.size(0)
|
||||||
correct += predicted.eq(batch_y).sum().item()
|
|
||||||
|
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
acc = 100.0 * correct / max(1, total)
|
epoch_loss = running_loss / total
|
||||||
avg_loss = epoch_loss / max(1, total)
|
epoch_acc = (correct / total) * 100.0
|
||||||
logger.info("Époque %d/%d - Loss: %.4f - Précision: %.1f%%", epoch, epochs, avg_loss, acc)
|
logger.info("Epoch %d/%d - Loss: %.4f - Accuracy: %.1f%%", epoch, epochs, epoch_loss, epoch_acc)
|
||||||
if progress_callback:
|
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()
|
model.eval()
|
||||||
dummy_input = torch.randn(1, 3, 224, 224, device=device)
|
model.to("cpu")
|
||||||
self.onnx_path.parent.mkdir(parents=True, exist_ok=True)
|
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(
|
torch.onnx.export(
|
||||||
model,
|
model,
|
||||||
dummy_input,
|
dummy_input,
|
||||||
|
|
@ -422,18 +560,6 @@ class SignClassifierEngine:
|
||||||
dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}},
|
dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}},
|
||||||
dynamo=False
|
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
|
# 6. Sauvegarde des métadonnées
|
||||||
meta_data = {
|
meta_data = {
|
||||||
|
|
@ -442,7 +568,7 @@ class SignClassifierEngine:
|
||||||
"classes": classes,
|
"classes": classes,
|
||||||
"samples_per_class": samples_per_class,
|
"samples_per_class": samples_per_class,
|
||||||
"epochs": epochs,
|
"epochs": epochs,
|
||||||
"final_accuracy": round(acc, 2),
|
"final_accuracy": round(epoch_acc, 2),
|
||||||
"input_size": [224, 224],
|
"input_size": [224, 224],
|
||||||
"framework": "MobileNetV3-Small / ONNX",
|
"framework": "MobileNetV3-Small / ONNX",
|
||||||
"signs_dir": str(signs_dir)
|
"signs_dir": str(signs_dir)
|
||||||
|
|
@ -461,7 +587,7 @@ class SignClassifierEngine:
|
||||||
return {
|
return {
|
||||||
"status": "success",
|
"status": "success",
|
||||||
"num_classes": num_classes,
|
"num_classes": num_classes,
|
||||||
"accuracy": round(acc, 2),
|
"accuracy": round(epoch_acc, 2),
|
||||||
"training_time_seconds": round(total_time, 1),
|
"training_time_seconds": round(total_time, 1),
|
||||||
"onnx_path": str(self.onnx_path),
|
"onnx_path": str(self.onnx_path),
|
||||||
"size_mb": round(self.onnx_path.stat().st_size / (1024 * 1024), 2)
|
"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]:
|
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).
|
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:
|
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": []}
|
return {"status": "error", "code": None, "confidence": 0.0, "top_matches": []}
|
||||||
|
|
@ -524,24 +651,29 @@ class SignClassifierEngine:
|
||||||
exp_logits = np.exp(logits - np.max(logits))
|
exp_logits = np.exp(logits - np.max(logits))
|
||||||
probs = exp_logits / np.sum(exp_logits)
|
probs = exp_logits / np.sum(exp_logits)
|
||||||
|
|
||||||
# 4. Top-K
|
# 4. Top-K initial élargi
|
||||||
k = min(top_k, len(self._classes))
|
candidate_count = min(max(top_k * 4, 25), len(self._classes))
|
||||||
top_indices = np.argsort(probs)[::-1][:k]
|
top_indices = np.argsort(probs)[::-1][:candidate_count]
|
||||||
top_matches = []
|
raw_matches = []
|
||||||
for idx in top_indices:
|
for idx in top_indices:
|
||||||
code = self._classes[idx]
|
code = self._classes[idx]
|
||||||
conf = float(probs[idx])
|
conf = float(probs[idx])
|
||||||
top_matches.append({
|
raw_matches.append({
|
||||||
"code": code,
|
"code": code,
|
||||||
"confidence": round(conf, 4),
|
"confidence": round(conf, 4),
|
||||||
"svg_url": f"/static/assets/road_signs/2025/{code}.svg"
|
"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 {
|
return {
|
||||||
"status": "success",
|
"status": "success",
|
||||||
"code": best_match["code"],
|
"code": best_match["code"],
|
||||||
"confidence": best_match["confidence"],
|
"confidence": best_match["confidence"],
|
||||||
"svg_url": best_match["svg_url"],
|
"svg_url": best_match["svg_url"],
|
||||||
"top_matches": top_matches,
|
"top_matches": filtered_matches,
|
||||||
|
"color_profile": extract_sign_color_profile(crop_bgr),
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -364,8 +364,12 @@ class SignDetectionService:
|
||||||
)
|
)
|
||||||
total_classifier_ms += (time.perf_counter() - c_start) * 1000.0
|
total_classifier_ms += (time.perf_counter() - c_start) * 1000.0
|
||||||
|
|
||||||
# Tentative d'identification via l'OCR
|
# Extraction du profil de couleur
|
||||||
matched = match_sign_from_ocr(ocr_text)
|
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
|
code = None
|
||||||
name_fr = ""
|
name_fr = ""
|
||||||
|
|
@ -400,8 +404,8 @@ class SignDetectionService:
|
||||||
final_confidence = max(final_confidence, 0.95)
|
final_confidence = max(final_confidence, 0.95)
|
||||||
|
|
||||||
# 2. Vitesse maximale autorisée (C43 / C43_XX / ZC43) :
|
# 2. Vitesse maximale autorisée (C43 / C43_XX / ZC43) :
|
||||||
# Fusion : Si OCR extrait une vitesse OU que le classifieur neuronal a prédit C43
|
# 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):
|
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
|
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)
|
# 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:
|
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))
|
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) :
|
# 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"))):
|
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", "ZE9A_DISK", "F4A", "F4B", "F103", "ZC43"):
|
if matched and matched["code"] in ("ZE9A", "ZE9B", "ZE9A_DISK", "F4A", "F4B", "F103", "ZC43"):
|
||||||
code = matched["code"]
|
code = matched["code"]
|
||||||
name_fr = matched["data"]["name_fr"]
|
name_fr = matched["data"]["name_fr"]
|
||||||
name_nl = matched["data"]["name_nl"]
|
name_nl = matched["data"]["name_nl"]
|
||||||
|
|
@ -437,6 +441,10 @@ class SignDetectionService:
|
||||||
final_confidence = max(final_confidence, matched.get("confidence", 0.92))
|
final_confidence = max(final_confidence, matched.get("confidence", 0.92))
|
||||||
else:
|
else:
|
||||||
code = primary_nn_code
|
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, {})
|
entry = SIGN_CATALOG.get(code, {})
|
||||||
name_fr = entry.get("name_fr", f"Zone {code}")
|
name_fr = entry.get("name_fr", f"Zone {code}")
|
||||||
name_nl = entry.get("name_nl", f"Zone {code}")
|
name_nl = entry.get("name_nl", f"Zone {code}")
|
||||||
|
|
@ -445,9 +453,9 @@ class SignDetectionService:
|
||||||
matched_by = "ai_neural_classifier"
|
matched_by = "ai_neural_classifier"
|
||||||
final_confidence = max(final_confidence, float(classifier_pred.get("confidence", 0.85)))
|
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 (
|
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"
|
or det["class_name"] == "sub_plate"
|
||||||
):
|
):
|
||||||
code = matched["code"]
|
code = matched["code"]
|
||||||
|
|
@ -462,6 +470,14 @@ class SignDetectionService:
|
||||||
# 5. Réseau Neuronal MobileNetV3 (Classification visuelle fine des pictogrammes)
|
# 5. Réseau Neuronal MobileNetV3 (Classification visuelle fine des pictogrammes)
|
||||||
elif primary_nn_code and classifier_pred.get("confidence", 0.0) >= 0.02:
|
elif primary_nn_code and classifier_pred.get("confidence", 0.0) >= 0.02:
|
||||||
code = primary_nn_code
|
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)
|
svg_url = classifier_pred.get("svg_url") or get_svg_url(code)
|
||||||
matched_by = "ai_neural_classifier"
|
matched_by = "ai_neural_classifier"
|
||||||
final_confidence = float(classifier_pred.get("confidence", det["confidence"]))
|
final_confidence = float(classifier_pred.get("confidence", det["confidence"]))
|
||||||
|
|
|
||||||
255
loko/sign/ai/ground_truth.py
Normal file
255
loko/sign/ai/ground_truth.py
Normal 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
|
||||||
Binary file not shown.
|
|
@ -1,6 +1,6 @@
|
||||||
{
|
{
|
||||||
"created_at": "2026-08-22 18:42:48",
|
"created_at": "2026-08-29 15:30:09",
|
||||||
"num_classes": 502,
|
"num_classes": 503,
|
||||||
"classes": [
|
"classes": [
|
||||||
"A11",
|
"A11",
|
||||||
"A13",
|
"A13",
|
||||||
|
|
@ -429,6 +429,7 @@
|
||||||
"TYPEVF",
|
"TYPEVF",
|
||||||
"TYPEVG",
|
"TYPEVG",
|
||||||
"TYPEVI",
|
"TYPEVI",
|
||||||
|
"TYPEVII",
|
||||||
"TYPEVIIA_(+)2,5T",
|
"TYPEVIIA_(+)2,5T",
|
||||||
"TYPEVIIA_(+)2T",
|
"TYPEVIIA_(+)2T",
|
||||||
"TYPEVIIA_(+)3,5T",
|
"TYPEVIIA_(+)3,5T",
|
||||||
|
|
@ -505,9 +506,9 @@
|
||||||
"ZF111",
|
"ZF111",
|
||||||
"ZF113"
|
"ZF113"
|
||||||
],
|
],
|
||||||
"samples_per_class": 12,
|
"samples_per_class": 35,
|
||||||
"epochs": 10,
|
"epochs": 12,
|
||||||
"final_accuracy": 99.79,
|
"final_accuracy": 97.33,
|
||||||
"input_size": [
|
"input_size": [
|
||||||
224,
|
224,
|
||||||
224
|
224
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
|
|
@ -32,12 +32,15 @@ from sign.ai import (
|
||||||
SyntheticSignAugmentor,
|
SyntheticSignAugmentor,
|
||||||
match_sign_from_ocr,
|
match_sign_from_ocr,
|
||||||
get_svg_url,
|
get_svg_url,
|
||||||
|
get_all_catalog_signs,
|
||||||
|
GroundTruthDatasetManager,
|
||||||
)
|
)
|
||||||
from sign.models import SignPanelType
|
from sign.models import SignPanelType
|
||||||
|
|
||||||
User = get_user_model()
|
User = get_user_model()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class SignCatalogMatcherTests(TestCase):
|
class SignCatalogMatcherTests(TestCase):
|
||||||
"""Tests pour le matching de texte OCR avec le catalogue de panneaux."""
|
"""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(
|
self.user = User.objects.create_user(
|
||||||
username="tester_ai",
|
username="tester_ai",
|
||||||
email="tester_ai@example.com",
|
email="tester_ai@example.com",
|
||||||
password="testpassword123"
|
password="testpassword123",
|
||||||
|
is_superuser=True
|
||||||
)
|
)
|
||||||
self.client.force_login(self.user)
|
self.client.force_login(self.user)
|
||||||
|
self.client.force_authenticate(user=self.user)
|
||||||
|
|
||||||
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
|
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
|
||||||
img = np.ones((400, 400, 3), dtype=np.uint8) * 240
|
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])
|
inspections = SignPanelInspection.objects.filter(asset_object_id__in=[p1.id, p2.id])
|
||||||
self.assertEqual(inspections.count(), 2)
|
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")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,12 @@ urlpatterns = [
|
||||||
path("", views.index, name="index"),
|
path("", views.index, name="index"),
|
||||||
path("ai/demo/", views_ai.SignAIDemoView.as_view(), name="ai_demo"),
|
path("ai/demo/", views_ai.SignAIDemoView.as_view(), name="ai_demo"),
|
||||||
path("api/detect/", views_ai.DetectSignAPIView.as_view(), name="api_detect"),
|
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("api/quick-create/", views_ai.QuickCreateSignWithAIView.as_view(), name="api_quick_create"),
|
||||||
|
|
||||||
path("streets/", views.sign_streets_list, name="sign_streets_list"),
|
path("streets/", views.sign_streets_list, name="sign_streets_list"),
|
||||||
path("streets/<int:street_id>/", views.sign_streets_detail, name="sign_streets_detail"),
|
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"),
|
path("assets/<str:asset_model>/<int:asset_id>/", views.sign_assets_detail, name="sign_assets_detail"),
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,7 @@ import logging
|
||||||
from typing import Any, Dict
|
from typing import Any, Dict
|
||||||
|
|
||||||
from django.views.generic import TemplateView
|
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.http import JsonResponse, HttpResponseBadRequest
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.csrf import csrf_exempt
|
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.views import APIView
|
||||||
from rest_framework.response import Response
|
from rest_framework.response import Response
|
||||||
from rest_framework import status
|
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__)
|
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.
|
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,
|
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.
|
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"
|
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]:
|
def get_context_data(self, **kwargs: Any) -> Dict[str, Any]:
|
||||||
context = super().get_context_data(**kwargs)
|
context = super().get_context_data(**kwargs)
|
||||||
|
|
@ -37,8 +72,14 @@ class SignAIDemoView(LoginRequiredMixin, TemplateView):
|
||||||
engine = SignClassifierEngine.get_instance()
|
engine = SignClassifierEngine.get_instance()
|
||||||
context["classifier_is_trained"] = engine.is_trained()
|
context["classifier_is_trained"] = engine.is_trained()
|
||||||
context["classifier_meta"] = engine.get_metadata()
|
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
|
return context
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def post(self, request, *args, **kwargs):
|
def post(self, request, *args, **kwargs):
|
||||||
image_file = request.FILES.get("image")
|
image_file = request.FILES.get("image")
|
||||||
image_base64 = request.POST.get("image_base64")
|
image_base64 = request.POST.get("image_base64")
|
||||||
|
|
@ -73,6 +114,7 @@ class SignAIDemoView(LoginRequiredMixin, TemplateView):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@method_decorator(csrf_exempt, name='dispatch')
|
||||||
class DetectSignAPIView(APIView):
|
class DetectSignAPIView(APIView):
|
||||||
"""
|
"""
|
||||||
Endpoint API REST pour l'analyse mobile/terrain d'une photo de signalisation.
|
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)
|
}, 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
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,7 @@
|
||||||
# Les librairies IA en local (ultralytics, PyTorch, etc.) ont été supprimées.
|
# Dépendances pour l'entraînement et l'inférence du classifieur de panneaux
|
||||||
# L'application utilise désormais un appel à un microservice externe pour la détection.
|
--extra-index-url https://download.pytorch.org/whl/cpu
|
||||||
# Ce fichier reste vide ou pour d'éventuelles dépendances futures du client d'API.
|
torch>=2.2.0
|
||||||
|
torchvision>=0.17.0
|
||||||
|
onnxruntime>=1.17.0
|
||||||
|
rapidocr-onnxruntime>=1.3.0
|
||||||
|
opencv-python-headless>=4.8.0
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue