726 lines
29 KiB
Python
726 lines
29 KiB
Python
"""
|
|
Tests unitaires et d'intégration pour le module IA de reconnaissance de signalisation (sign.ai).
|
|
Couvre :
|
|
- Matcher OCR et catalogue
|
|
- SyntheticSignAugmentor et déformations réalistes
|
|
- SignClassifierEngine (auto-entraînement, export ONNX, inférence top-k)
|
|
- SignDetectionService (détection YOLO + OCR + Classifieur neuronal)
|
|
- Commande Django train_sign_classifier
|
|
- Endpoints API REST et vue de démonstration
|
|
"""
|
|
import io
|
|
import json
|
|
import os
|
|
import tempfile
|
|
from pathlib import Path
|
|
import numpy as np
|
|
import cv2
|
|
from PIL import Image
|
|
|
|
from django.test import TestCase
|
|
from django.urls import reverse
|
|
from django.contrib.auth import get_user_model
|
|
from django.core.files.uploadedfile import SimpleUploadedFile
|
|
from django.core.management import call_command
|
|
|
|
from rest_framework.test import APITestCase
|
|
from rest_framework import status
|
|
|
|
from sign.ai import (
|
|
SignDetectionService,
|
|
SignClassifierEngine,
|
|
SyntheticSignAugmentor,
|
|
match_sign_from_ocr,
|
|
get_svg_url,
|
|
get_all_catalog_signs,
|
|
GroundTruthDatasetManager,
|
|
)
|
|
from sign.models import SignPanelType
|
|
|
|
User = get_user_model()
|
|
|
|
|
|
|
|
class SignCatalogMatcherTests(TestCase):
|
|
"""Tests pour le matching de texte OCR avec le catalogue de panneaux."""
|
|
|
|
def test_match_stop_sign(self):
|
|
res = match_sign_from_ocr("STOP")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "B5")
|
|
self.assertEqual(res["data"]["category"], "priority")
|
|
self.assertTrue(res["svg_url"].endswith("B5.svg"))
|
|
|
|
def test_match_zone_30(self):
|
|
res = match_sign_from_ocr("ZONE 30")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "F4A")
|
|
self.assertEqual(res["value"], 30)
|
|
|
|
def test_match_end_zone(self):
|
|
res = match_sign_from_ocr("FIN DE ZONE")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "F4B")
|
|
|
|
def test_match_exception_panonceau(self):
|
|
res = match_sign_from_ocr("Sauf riverains")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "M2")
|
|
self.assertEqual(res["extracted_text"], "Sauf riverains")
|
|
|
|
def test_match_distance_panonceau(self):
|
|
res = match_sign_from_ocr("300 m")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "M1")
|
|
self.assertEqual(res["value"], 300.0)
|
|
|
|
def test_match_speed_limit(self):
|
|
res = match_sign_from_ocr("50")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "C43")
|
|
self.assertEqual(res["value"], 50)
|
|
|
|
def test_match_speed_limit_with_km_and_kmh(self):
|
|
res_km = match_sign_from_ocr("50 km")
|
|
self.assertIsNotNone(res_km)
|
|
self.assertEqual(res_km["code"], "C43")
|
|
self.assertEqual(res_km["value"], 50)
|
|
|
|
res_kmh = match_sign_from_ocr("70 km/h")
|
|
self.assertIsNotNone(res_kmh)
|
|
self.assertEqual(res_kmh["code"], "C43")
|
|
self.assertEqual(res_kmh["value"], 70)
|
|
|
|
def test_match_zone_parking_with_exemptions(self):
|
|
res = match_sign_from_ocr("ZONE P Excepte carte de stationnement Uitgezonderd parkeerkaart Rappel")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "ZE9A")
|
|
self.assertEqual(res["data"]["category"], "parking")
|
|
|
|
def test_match_tonnage_sign(self):
|
|
res = match_sign_from_ocr("3.5 t")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "C21")
|
|
self.assertEqual(res["value"], 3.5)
|
|
|
|
def test_match_parking_p_and_d(self):
|
|
res_p = match_sign_from_ocr("P")
|
|
self.assertIsNotNone(res_p)
|
|
self.assertEqual(res_p["code"], "E9A")
|
|
|
|
res_d = match_sign_from_ocr("D")
|
|
self.assertIsNotNone(res_d)
|
|
self.assertEqual(res_d["code"], "E9A")
|
|
|
|
def test_match_parking_payant(self):
|
|
res = match_sign_from_ocr("PAYANT BETALEND")
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "GVII_BETALEND")
|
|
self.assertEqual(res["data"]["category"], "parking")
|
|
|
|
def test_visual_classification_blue_parking_vertical(self):
|
|
from sign.ai import classify_sign_visual
|
|
# Simuler un rectangle bleu vertical (ex: 200 de haut x 130 de large)
|
|
img = np.zeros((200, 130, 3), dtype=np.uint8)
|
|
img[:, :] = (200, 50, 20) # BGR Bleu
|
|
res = classify_sign_visual(img, ocr_text="")
|
|
self.assertEqual(res["code"], "E9A")
|
|
self.assertEqual(res["category"], "parking")
|
|
|
|
def test_visual_classification_blue_bike(self):
|
|
from sign.ai import classify_sign_visual
|
|
img = np.ones((100, 100, 3), dtype=np.uint8) * 180
|
|
cv2.circle(img, (50, 50), 45, (200, 50, 20), -1)
|
|
cv2.circle(img, (50, 50), 15, (255, 255, 255), -1)
|
|
res = classify_sign_visual(img)
|
|
self.assertEqual(res["code"], "D7")
|
|
self.assertEqual(res["category"], "obligation")
|
|
|
|
def test_visual_classification_yellow_priority(self):
|
|
from sign.ai import classify_sign_visual
|
|
img = np.ones((100, 100, 3), dtype=np.uint8) * 180
|
|
cv2.circle(img, (50, 50), 45, (0, 220, 220), -1)
|
|
res = classify_sign_visual(img)
|
|
self.assertEqual(res["code"], "B3")
|
|
self.assertEqual(res["category"], "priority")
|
|
|
|
def test_visual_classification_sens_interdit_c1(self):
|
|
from sign.ai import classify_sign_visual
|
|
img = np.ones((100, 100, 3), dtype=np.uint8) * 180
|
|
cv2.circle(img, (50, 50), 45, (30, 30, 220), -1)
|
|
cv2.rectangle(img, (15, 42), (85, 58), (255, 255, 255), -1)
|
|
res = classify_sign_visual(img)
|
|
self.assertEqual(res["code"], "C1")
|
|
self.assertEqual(res["category"], "prohibition")
|
|
|
|
def test_get_svg_url(self):
|
|
self.assertEqual(get_svg_url("B1"), "/static/assets/road_signs/2025/B1.svg")
|
|
self.assertEqual(get_svg_url(""), "")
|
|
|
|
|
|
class SyntheticSignAugmentorTests(TestCase):
|
|
"""Tests pour le générateur de variations synthétiques."""
|
|
|
|
def test_letterbox_rgba(self):
|
|
# Image non carrée (100x200)
|
|
img = np.ones((100, 200, 4), dtype=np.uint8) * 255
|
|
letterboxed = SyntheticSignAugmentor.letterbox_rgba(img, target_size=224)
|
|
self.assertEqual(letterboxed.shape, (224, 224, 4))
|
|
|
|
def test_generate_random_background(self):
|
|
bg = SyntheticSignAugmentor.generate_random_background(size=224)
|
|
self.assertEqual(bg.shape, (224, 224, 3))
|
|
self.assertEqual(bg.dtype, np.uint8)
|
|
|
|
def test_augment_sign(self):
|
|
# Créer une image RGBA synthétique (disque rouge sur fond transparent)
|
|
rgba = np.zeros((224, 224, 4), dtype=np.uint8)
|
|
cv2.circle(rgba, (112, 112), 90, (0, 0, 220, 255), -1)
|
|
|
|
augmented = SyntheticSignAugmentor.augment_sign(rgba, size=224)
|
|
self.assertEqual(augmented.shape, (224, 224, 3))
|
|
self.assertEqual(augmented.dtype, np.uint8)
|
|
# Vérifier qu'il y a du contenu non noir
|
|
self.assertGreater(np.mean(augmented), 5.0)
|
|
|
|
|
|
class SignClassifierEngineTests(TestCase):
|
|
"""Tests du moteur d'auto-entraînement et d'inférence ONNX."""
|
|
|
|
def setUp(self):
|
|
self.temp_dir = tempfile.TemporaryDirectory()
|
|
self.models_dir = Path(self.temp_dir.name) / "models"
|
|
self.signs_dir = Path(self.temp_dir.name) / "signs"
|
|
self.signs_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
# Créer au moins 2 templates PNG de test
|
|
# 1. Panneau C1 (Sens interdit - rouge)
|
|
c1_img = np.zeros((200, 200, 4), dtype=np.uint8)
|
|
cv2.circle(c1_img, (100, 100), 90, (0, 0, 220, 255), -1)
|
|
cv2.rectangle(c1_img, (30, 85), (170, 115), (255, 255, 255, 255), -1)
|
|
cv2.imwrite(str(self.signs_dir / "C1.png"), c1_img)
|
|
|
|
# 2. Panneau D7 (Piste cyclable - bleu)
|
|
d7_img = np.zeros((200, 200, 4), dtype=np.uint8)
|
|
cv2.circle(d7_img, (100, 100), 90, (220, 100, 0, 255), -1)
|
|
cv2.circle(d7_img, (100, 100), 30, (255, 255, 255, 255), -1)
|
|
cv2.imwrite(str(self.signs_dir / "D7.png"), d7_img)
|
|
|
|
self.engine = SignClassifierEngine(models_dir=self.models_dir)
|
|
|
|
def tearDown(self):
|
|
self.temp_dir.cleanup()
|
|
|
|
def test_discover_templates(self):
|
|
templates = self.engine.discover_templates(signs_dir=self.signs_dir)
|
|
self.assertEqual(len(templates), 2)
|
|
self.assertIn("C1", templates)
|
|
self.assertIn("D7", templates)
|
|
|
|
def test_train_from_svgs_and_predict(self):
|
|
self.assertFalse(self.engine.is_trained())
|
|
|
|
# Entraîner un modèle miniature (2 époques, 4 samples par classe)
|
|
result = self.engine.train_from_svgs(
|
|
signs_dir=self.signs_dir,
|
|
samples_per_class=4,
|
|
epochs=2,
|
|
batch_size=8,
|
|
learning_rate=0.005,
|
|
)
|
|
|
|
self.assertEqual(result["status"], "success")
|
|
self.assertEqual(result["num_classes"], 2)
|
|
self.assertTrue(self.engine.is_trained())
|
|
self.assertTrue(Path(result["onnx_path"]).exists())
|
|
|
|
meta = self.engine.get_metadata()
|
|
self.assertEqual(meta["num_classes"], 2)
|
|
self.assertIn("C1", meta["classes"])
|
|
self.assertIn("D7", meta["classes"])
|
|
|
|
# Tester l'inférence sur un crop rouge (devrait être C1)
|
|
red_crop = np.ones((150, 150, 3), dtype=np.uint8) * 30
|
|
cv2.circle(red_crop, (75, 75), 60, (0, 0, 220), -1)
|
|
pred = self.engine.predict(red_crop, top_k=2)
|
|
|
|
self.assertEqual(pred["status"], "success")
|
|
self.assertIsNotNone(pred["code"])
|
|
self.assertGreaterEqual(len(pred["top_matches"]), 2)
|
|
self.assertGreaterEqual(pred["confidence"], 0.0)
|
|
|
|
def test_predict_invalid_crop(self):
|
|
empty_crop = np.array([])
|
|
pred = self.engine.predict(empty_crop)
|
|
self.assertEqual(pred["status"], "error")
|
|
|
|
|
|
class SignDetectionServiceTests(TestCase):
|
|
"""Tests d'inférence du service SignDetectionService."""
|
|
|
|
def setUp(self):
|
|
self.service = SignDetectionService.get_instance()
|
|
|
|
def _create_synthetic_image(self) -> np.ndarray:
|
|
img = np.ones((500, 500, 3), dtype=np.uint8) * 240
|
|
# Dessiner un panneau STOP
|
|
cv2.circle(img, (250, 180), 90, (0, 0, 220), -1)
|
|
cv2.putText(img, "STOP", (195, 195), cv2.FONT_HERSHEY_SIMPLEX, 1.3, (255, 255, 255), 3)
|
|
# Panonceau sous le panneau
|
|
cv2.rectangle(img, (150, 320), (350, 400), (255, 255, 255), -1)
|
|
cv2.rectangle(img, (150, 320), (350, 400), (0, 0, 0), 2)
|
|
cv2.putText(img, "Sauf velos", (165, 365), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 0, 0), 2)
|
|
return img
|
|
|
|
def test_analyze_image_synthetic(self):
|
|
img = self._create_synthetic_image()
|
|
result = self.service.analyze_image(img)
|
|
|
|
self.assertEqual(result["status"], "success")
|
|
self.assertGreaterEqual(result["detected_count"], 1)
|
|
self.assertIn("performance", result)
|
|
self.assertIn("yolo_inference_ms", result["performance"])
|
|
self.assertIn("ocr_inference_ms", result["performance"])
|
|
self.assertIn("classifier_inference_ms", result["performance"])
|
|
self.assertIn("annotated_image", result)
|
|
self.assertTrue(result["annotated_image"].startswith("data:image/jpeg;base64,"))
|
|
|
|
# Vérifier l'ordonnancement vertical et la présence des top_matches
|
|
orders = [p["vertical_order"] for p in result["panels"]]
|
|
self.assertEqual(orders, list(range(1, len(orders) + 1)))
|
|
for p in result["panels"]:
|
|
self.assertIn("top_matches", p)
|
|
self.assertIn("matched_by", p)
|
|
|
|
|
|
class SignAIApiAndViewsTests(APITestCase):
|
|
"""Tests des endpoints API et de la vue de démo."""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(
|
|
username="tester_ai",
|
|
email="tester_ai@example.com",
|
|
password="testpassword123",
|
|
is_superuser=True
|
|
)
|
|
self.client.force_login(self.user)
|
|
self.client.force_authenticate(user=self.user)
|
|
|
|
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
|
|
img = np.ones((400, 400, 3), dtype=np.uint8) * 240
|
|
cv2.circle(img, (200, 150), 80, (0, 0, 220), -1)
|
|
cv2.putText(img, "STOP", (150, 165), cv2.FONT_HERSHEY_SIMPLEX, 1.1, (255, 255, 255), 3)
|
|
|
|
_, buf = cv2.imencode(".jpg", img)
|
|
return SimpleUploadedFile("test_sign.jpg", buf.tobytes(), content_type="image/jpeg")
|
|
|
|
def test_ai_demo_view_get(self):
|
|
url = reverse("sign:ai_demo")
|
|
response = self.client.get(url)
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertContains(response, "Reconnaissance de Signalisation & OCR")
|
|
|
|
def test_ai_demo_view_post_file(self):
|
|
url = reverse("sign:ai_demo")
|
|
file_obj = self._create_uploaded_image_file()
|
|
response = self.client.post(
|
|
url,
|
|
{"image": file_obj, "confidence_threshold": "0.20"},
|
|
HTTP_X_REQUESTED_WITH="XMLHttpRequest"
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
data = response.json()
|
|
self.assertEqual(data["status"], "success")
|
|
self.assertGreaterEqual(data["detected_count"], 1)
|
|
|
|
def test_api_detect_sign_post(self):
|
|
url = reverse("sign:api_detect")
|
|
file_obj = self._create_uploaded_image_file()
|
|
response = self.client.post(
|
|
url,
|
|
{"image": file_obj, "confidence_threshold": 0.20},
|
|
format="multipart"
|
|
)
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertEqual(data["status"], "success")
|
|
self.assertIn("panels", data)
|
|
self.assertIn("performance", data)
|
|
|
|
def test_api_detect_sign_missing_image(self):
|
|
url = reverse("sign:api_detect")
|
|
response = self.client.post(url, {}, format="json")
|
|
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
|
|
|
def test_quick_create_sign_with_ai(self):
|
|
from assets.models import SignPanelType, SignStreet, SignPole, SignPanel
|
|
from inspections.models import SignPanelInspection
|
|
import json
|
|
|
|
SignStreet.objects.create(code="STR-AI-1", name_fr="Rue du Progrès")
|
|
SignPanelType.objects.get_or_create(code="C1", defaults={"name_fr": "Sens interdit"})
|
|
SignPanelType.objects.get_or_create(code="M2", defaults={"name_fr": "Panonceau vélo"})
|
|
|
|
url = reverse("sign:api_quick_create")
|
|
file_obj = self._create_uploaded_image_file()
|
|
|
|
panels_data = [
|
|
{
|
|
"signpanel_type_code": "C1",
|
|
"signpanel_text": "Sens unique",
|
|
"cleanliness": "clean",
|
|
"vertical_order": 1
|
|
},
|
|
{
|
|
"signpanel_type_code": "M2",
|
|
"signpanel_text": "Sauf cyclistes",
|
|
"cleanliness": "dirty",
|
|
"vertical_order": 2
|
|
}
|
|
]
|
|
|
|
response = self.client.post(
|
|
url,
|
|
{
|
|
"lat": "50.8503",
|
|
"lon": "4.3517",
|
|
"panels": json.dumps(panels_data),
|
|
"photos": file_obj
|
|
},
|
|
format="multipart"
|
|
)
|
|
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertTrue(data["success"])
|
|
self.assertEqual(data["panels_created"], 2)
|
|
|
|
pole = SignPole.objects.get(id=data["pole_id"])
|
|
self.assertIsNotNone(pole.geom)
|
|
self.assertEqual(pole.signpanels.count(), 2)
|
|
|
|
p1 = pole.signpanels.get(vertical_order=1)
|
|
self.assertEqual(p1.signpanel_type.code, "C1")
|
|
self.assertTrue(p1.is_compliant)
|
|
|
|
p2 = pole.signpanels.get(vertical_order=2)
|
|
self.assertEqual(p2.signpanel_type.code, "M2")
|
|
self.assertFalse(p2.is_compliant)
|
|
|
|
inspections = SignPanelInspection.objects.filter(asset_object_id__in=[p1.id, p2.id])
|
|
self.assertEqual(inspections.count(), 2)
|
|
|
|
|
|
class GroundTruthAndCatalogAPITests(APITestCase):
|
|
"""Tests pour le gestionnaire de vérité terrain et les endpoints d'apprentissage actif."""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="ai_admin_tester", password="testpassword123", is_superuser=True)
|
|
self.non_admin_user = User.objects.create_user(username="standard_user", password="testpassword123", is_staff=False, is_superuser=False)
|
|
self.client.force_authenticate(user=self.user)
|
|
self.temp_gt_dir = tempfile.mkdtemp()
|
|
self.gt_manager = GroundTruthDatasetManager(base_dir=self.temp_gt_dir)
|
|
|
|
def _create_test_image(self) -> np.ndarray:
|
|
img = np.ones((200, 200, 3), dtype=np.uint8) * 200
|
|
# Dessiner un cercle bleu
|
|
cv2.circle(img, (100, 100), 60, (220, 50, 50), -1)
|
|
return img
|
|
|
|
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
|
|
img = Image.new("RGB", (200, 200), color=(100, 150, 200))
|
|
buf = io.BytesIO()
|
|
img.save(buf, format="JPEG")
|
|
buf.seek(0)
|
|
return SimpleUploadedFile("sample_ground_truth.jpg", buf.read(), content_type="image/jpeg")
|
|
|
|
def test_admin_permission_restrictions(self):
|
|
# 1. Utilisateur standard non-admin
|
|
self.client.force_login(self.non_admin_user)
|
|
self.client.force_authenticate(user=self.non_admin_user)
|
|
res_demo = self.client.get(reverse("sign:ai_demo"))
|
|
self.assertEqual(res_demo.status_code, status.HTTP_403_FORBIDDEN)
|
|
|
|
res_catalog = self.client.get(reverse("sign:api_catalog"))
|
|
self.assertEqual(res_catalog.status_code, status.HTTP_403_FORBIDDEN)
|
|
|
|
res_stats = self.client.get(reverse("sign:api_ground_truth_stats"))
|
|
self.assertEqual(res_stats.status_code, status.HTTP_403_FORBIDDEN)
|
|
|
|
# 2. Utilisateur administrateur
|
|
self.client.force_login(self.user)
|
|
self.client.force_authenticate(user=self.user)
|
|
res_demo_admin = self.client.get(reverse("sign:ai_demo"))
|
|
self.assertEqual(res_demo_admin.status_code, status.HTTP_200_OK)
|
|
|
|
res_catalog_admin = self.client.get(reverse("sign:api_catalog"))
|
|
self.assertEqual(res_catalog_admin.status_code, status.HTTP_200_OK)
|
|
|
|
def test_catalog_signs_function(self):
|
|
signs = get_all_catalog_signs()
|
|
self.assertGreater(len(signs), 10)
|
|
codes = [s["code"] for s in signs]
|
|
self.assertIn("E1", codes)
|
|
self.assertIn("B5", codes)
|
|
self.assertIn("C43", codes)
|
|
self.assertIn("ZE9A", codes)
|
|
|
|
def test_catalog_api_endpoint(self):
|
|
url = reverse("sign:api_catalog")
|
|
response = self.client.get(url, {"q": "ZE9A"})
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertGreater(data["match_count"], 0)
|
|
codes = [s["code"] for s in data["signs"]]
|
|
self.assertIn("ZE9A", codes)
|
|
|
|
def test_catalog_api_zone_filter(self):
|
|
url = reverse("sign:api_catalog")
|
|
response = self.client.get(url, {"category": "zone"})
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertGreater(data["match_count"], 0)
|
|
codes = [s["code"] for s in data["signs"]]
|
|
self.assertIn("ZE9A", codes)
|
|
|
|
def test_ground_truth_manager_save_and_stats(self):
|
|
img_np = self._create_test_image()
|
|
panels = [
|
|
{
|
|
"code": "E1",
|
|
"initial_code": "D1B",
|
|
"was_corrected": True,
|
|
"bbox": [20, 20, 180, 180],
|
|
"vertical_order": 1,
|
|
"ocr_text": "",
|
|
},
|
|
{
|
|
"code": "XD",
|
|
"initial_code": "XD",
|
|
"was_corrected": False,
|
|
"bbox": [80, 150, 120, 190],
|
|
"vertical_order": 2,
|
|
"ocr_text": "300 m",
|
|
}
|
|
]
|
|
|
|
res = self.gt_manager.save_ground_truth(
|
|
image_input=img_np,
|
|
panels=panels,
|
|
user_info={"username": "ai_admin_tester"},
|
|
notes="Test annotation E1"
|
|
)
|
|
|
|
self.assertEqual(res["status"], "success")
|
|
self.assertEqual(res["crops_saved"], 2)
|
|
|
|
stats = self.gt_manager.get_dataset_stats()
|
|
self.assertEqual(stats["total_images"], 1)
|
|
self.assertEqual(stats["total_crops"], 2)
|
|
self.assertEqual(stats["total_corrected_panels"], 1)
|
|
self.assertIn("E1", stats["classes_distribution"])
|
|
self.assertIn("XD", stats["classes_distribution"])
|
|
|
|
real_crops = self.gt_manager.get_real_crops_for_training()
|
|
self.assertIn("E1", real_crops)
|
|
self.assertEqual(len(real_crops["E1"]), 1)
|
|
|
|
def test_save_ground_truth_api(self):
|
|
url = reverse("sign:api_ground_truth")
|
|
file_obj = self._create_uploaded_image_file()
|
|
|
|
panels = [
|
|
{
|
|
"code": "E1",
|
|
"initial_code": "D1B",
|
|
"was_corrected": True,
|
|
"bbox": [10, 10, 190, 190],
|
|
"vertical_order": 1,
|
|
}
|
|
]
|
|
|
|
response = self.client.post(
|
|
url,
|
|
{
|
|
"image": file_obj,
|
|
"panels": json.dumps(panels),
|
|
"notes": "Correction terrain E1",
|
|
},
|
|
format="multipart"
|
|
)
|
|
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertEqual(data["status"], "success")
|
|
self.assertEqual(data["crops_saved"], 1)
|
|
self.assertIn("stats", data)
|
|
|
|
def test_ground_truth_stats_api(self):
|
|
url = reverse("sign:api_ground_truth_stats")
|
|
response = self.client.get(url)
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertIn("total_images", data)
|
|
self.assertIn("total_crops", data)
|
|
|
|
def test_specular_glare_and_synthetic_augmentation(self):
|
|
# Création d'une image RGBA synthétique
|
|
rgba = np.zeros((100, 100, 4), dtype=np.uint8)
|
|
rgba[:, :, 0] = 200 # Bleu
|
|
rgba[:, :, 2] = 20 # Rouge
|
|
rgba[:, :, 3] = 255 # Alpha
|
|
cv2.circle(rgba, (50, 50), 30, (0, 0, 220, 255), 8)
|
|
|
|
glare_applied = SyntheticSignAugmentor.apply_specular_glare(rgba[:, :, :3], rgba[:, :, 3])
|
|
self.assertEqual(glare_applied.shape, (100, 100, 3))
|
|
|
|
augmented = SyntheticSignAugmentor.augment_sign(rgba, size=224)
|
|
self.assertEqual(augmented.shape, (224, 224, 3))
|
|
self.assertEqual(augmented.dtype, np.uint8)
|
|
|
|
def test_hsv_color_profile_extraction_and_e1_filtering(self):
|
|
from sign.ai.catalog import extract_sign_color_profile, filter_and_rank_candidates_by_color
|
|
|
|
# 1. Image simulée E1 : Disque bleu avec bordure et diagonale rouge
|
|
e1_img = np.zeros((120, 120, 3), dtype=np.uint8)
|
|
# Fond extérieur noir/neutre
|
|
# Disque intérieur bleu
|
|
cv2.circle(e1_img, (60, 60), 50, (180, 50, 20), -1) # Bleu BGR
|
|
# Bordure et diagonale rouge
|
|
cv2.circle(e1_img, (60, 60), 50, (20, 20, 220), 8) # Rouge BGR
|
|
cv2.line(e1_img, (25, 95), (95, 25), (20, 20, 220), 8)
|
|
|
|
profile = extract_sign_color_profile(e1_img)
|
|
self.assertTrue(profile["has_red_and_blue"], f"Profile: {profile}")
|
|
self.assertGreater(profile["red_ratio"], 0.04)
|
|
self.assertGreater(profile["blue_ratio"], 0.05)
|
|
|
|
# Simulation de prédictions brutes erronées proposant D1B en 1ère position
|
|
raw_candidates = [
|
|
{"code": "D1B", "confidence": 0.65, "svg_url": "/static/D1B.svg"},
|
|
{"code": "E1", "confidence": 0.25, "svg_url": "/static/E1.svg"},
|
|
{"code": "E3", "confidence": 0.10, "svg_url": "/static/E3.svg"},
|
|
]
|
|
|
|
filtered = filter_and_rank_candidates_by_color(raw_candidates, e1_img)
|
|
self.assertEqual(filtered[0]["code"], "E1")
|
|
# D1B doit être éliminé en bas avec une confiance quasi nulle
|
|
d1b_item = next(c for c in filtered if c["code"] == "D1B")
|
|
self.assertLess(d1b_item["confidence"], 0.01)
|
|
|
|
def test_hsv_color_profile_pure_blue_filtering(self):
|
|
from sign.ai.catalog import extract_sign_color_profile, filter_and_rank_candidates_by_color
|
|
|
|
# Image simulée D1B : Disque bleu pur sans rouge avec flèche blanche
|
|
d1b_img = np.zeros((120, 120, 3), dtype=np.uint8)
|
|
cv2.circle(d1b_img, (60, 60), 50, (220, 80, 20), -1) # Bleu BGR
|
|
cv2.line(d1b_img, (60, 80), (60, 35), (255, 255, 255), 10) # Flèche blanche
|
|
|
|
profile = extract_sign_color_profile(d1b_img)
|
|
self.assertTrue(profile["is_pure_blue"], f"Profile: {profile}")
|
|
self.assertFalse(profile["has_red_and_blue"])
|
|
|
|
raw_candidates = [
|
|
{"code": "E1", "confidence": 0.60, "svg_url": "/static/E1.svg"},
|
|
{"code": "D1B", "confidence": 0.40, "svg_url": "/static/D1B.svg"},
|
|
]
|
|
|
|
filtered = filter_and_rank_candidates_by_color(raw_candidates, d1b_img)
|
|
self.assertEqual(filtered[0]["code"], "D1B")
|
|
e1_item = next(c for c in filtered if c["code"] == "E1")
|
|
self.assertLess(e1_item["confidence"], 0.01)
|
|
|
|
def test_no_red_f4a_exclusion_and_zone_parking(self):
|
|
from sign.ai.catalog import match_sign_from_ocr
|
|
|
|
# Panneau bleu/blanc sans rouge avec texte ZONE -> Doit être ZE9A et JAMAIS F4A
|
|
blue_white_img = np.ones((100, 100, 3), dtype=np.uint8) * 240
|
|
cv2.rectangle(blue_white_img, (20, 20), (80, 80), (200, 60, 20), -1) # Bleu
|
|
res = match_sign_from_ocr("ZONE Rappel Herhaling", crop_bgr=blue_white_img)
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "ZE9A")
|
|
self.assertNotEqual(res["code"], "F4A")
|
|
|
|
def test_gxc_arrow_distance_matching(self):
|
|
from sign.ai.catalog import match_sign_from_ocr
|
|
|
|
# Panonceau blanc avec distance courte "11m"
|
|
white_img = np.ones((60, 100, 3), dtype=np.uint8) * 250
|
|
res = match_sign_from_ocr("11m", crop_bgr=white_img)
|
|
self.assertIsNotNone(res)
|
|
self.assertEqual(res["code"], "GXC")
|
|
|
|
def test_type0_and_type0b_subplate_simplification(self):
|
|
from sign.ai.catalog import match_sign_from_ocr
|
|
|
|
# Panonceau textuel bleu
|
|
blue_img = np.ones((60, 100, 3), dtype=np.uint8) * 20
|
|
blue_img[:, :] = (180, 50, 20) # Bleu
|
|
res_blue = match_sign_from_ocr("Tarif specifique horaire centre", crop_bgr=blue_img)
|
|
self.assertIsNotNone(res_blue)
|
|
self.assertEqual(res_blue["code"], "TYPE0")
|
|
|
|
# Panonceau textuel blanc
|
|
white_img = np.ones((60, 100, 3), dtype=np.uint8) * 245
|
|
res_white = match_sign_from_ocr("Forfait 50e stationnement longue duree", crop_bgr=white_img)
|
|
self.assertIsNotNone(res_white)
|
|
self.assertEqual(res_white["code"], "TYPE0B")
|
|
|
|
def test_danger_inner_pictogram_bicycle_discrimination(self):
|
|
from sign.ai.catalog import discriminate_inner_pictogram
|
|
|
|
# Simuler un panneau triangulaire de danger avec un vélo (symbole horizontal et 2 roues)
|
|
tri_img = np.ones((140, 140, 3), dtype=np.uint8) * 240
|
|
# Bordure rouge
|
|
pts = np.array([[70, 10], [15, 125], [125, 125]], np.int32)
|
|
cv2.polylines(tri_img, [pts], isClosed=True, color=(20, 20, 220), thickness=8)
|
|
# Pictogramme vélo (deux roues noires distinctes en bas + cadre)
|
|
cv2.circle(tri_img, (50, 95), 8, (10, 10, 10), -1)
|
|
cv2.circle(tri_img, (90, 95), 8, (10, 10, 10), -1)
|
|
cv2.line(tri_img, (50, 95), (70, 75), (10, 10, 10), 3)
|
|
cv2.line(tri_img, (90, 95), (70, 75), (10, 10, 10), 3)
|
|
|
|
candidates = [
|
|
{"code": "A15", "confidence": 0.70}, # Faussement prédit A15 (piéton)
|
|
{"code": "A25", "confidence": 0.20}, # Vrai A25 (vélo)
|
|
{"code": "A14", "confidence": 0.10},
|
|
]
|
|
|
|
discriminated = discriminate_inner_pictogram(tri_img, candidates)
|
|
self.assertEqual(discriminated[0]["code"], "A25")
|
|
self.assertGreater(discriminated[0]["confidence"], 0.50)
|
|
|
|
def test_blue_vs_white_exception_subplate_discrimination(self):
|
|
from sign.ai.catalog import match_sign_from_ocr
|
|
|
|
# 1. Panonceau bleu avec texte blanc "EXCEPTE RIVERAINS" -> TYPE0 (JAMAIS M2)
|
|
blue_img = np.zeros((60, 120, 3), dtype=np.uint8)
|
|
blue_img[:, :] = (190, 60, 20) # Bleu pur
|
|
res_blue = match_sign_from_ocr("EXCEPTE RIVERAINS UITGEZ BEWONERS", crop_bgr=blue_img)
|
|
self.assertIsNotNone(res_blue)
|
|
self.assertEqual(res_blue["code"], "TYPE0")
|
|
self.assertNotEqual(res_blue["code"], "M2")
|
|
|
|
# 2. Panonceau blanc avec texte noir "EXCEPTE RIVERAINS" -> M2
|
|
white_img = np.ones((60, 120, 3), dtype=np.uint8) * 245
|
|
res_white = match_sign_from_ocr("EXCEPTE RIVERAINS UITGEZ BEWONERS", crop_bgr=white_img)
|
|
self.assertIsNotNone(res_white)
|
|
self.assertEqual(res_white["code"], "M2")
|
|
|
|
def test_blue_vs_white_distance_subplate_discrimination(self):
|
|
from sign.ai.catalog import match_sign_from_ocr
|
|
|
|
# 1. Panonceau bleu avec texte blanc "50 m" -> TYPEIA_50M (JAMAIS M1)
|
|
blue_img = np.zeros((50, 100, 3), dtype=np.uint8)
|
|
blue_img[:, :] = (200, 70, 20) # Bleu
|
|
res_blue = match_sign_from_ocr("50 m", crop_bgr=blue_img)
|
|
self.assertIsNotNone(res_blue)
|
|
self.assertEqual(res_blue["code"], "TYPEIA_50M")
|
|
self.assertNotEqual(res_blue["code"], "M1")
|
|
|
|
# 2. Panonceau blanc avec distance longue "500 m" -> M1
|
|
white_img = np.ones((50, 100, 3), dtype=np.uint8) * 240
|
|
res_white = match_sign_from_ocr("500 m", crop_bgr=white_img)
|
|
self.assertIsNotNone(res_white)
|
|
self.assertEqual(res_white["code"], "M1")
|
|
|