406 lines
15 KiB
Python
406 lines
15 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,
|
|
)
|
|
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"
|
|
)
|
|
self.client.force_login(self.user)
|
|
|
|
def _create_uploaded_image_file(self) -> SimpleUploadedFile:
|
|
img = np.ones((400, 400, 3), dtype=np.uint8) * 240
|
|
cv2.circle(img, (200, 150), 80, (0, 0, 220), -1)
|
|
cv2.putText(img, "STOP", (150, 165), cv2.FONT_HERSHEY_SIMPLEX, 1.1, (255, 255, 255), 3)
|
|
|
|
_, buf = cv2.imencode(".jpg", img)
|
|
return SimpleUploadedFile("test_sign.jpg", buf.tobytes(), content_type="image/jpeg")
|
|
|
|
def test_ai_demo_view_get(self):
|
|
url = reverse("sign:ai_demo")
|
|
response = self.client.get(url)
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertContains(response, "Reconnaissance de Signalisation & OCR")
|
|
|
|
def test_ai_demo_view_post_file(self):
|
|
url = reverse("sign:ai_demo")
|
|
file_obj = self._create_uploaded_image_file()
|
|
response = self.client.post(
|
|
url,
|
|
{"image": file_obj, "confidence_threshold": "0.20"},
|
|
HTTP_X_REQUESTED_WITH="XMLHttpRequest"
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
data = response.json()
|
|
self.assertEqual(data["status"], "success")
|
|
self.assertGreaterEqual(data["detected_count"], 1)
|
|
|
|
def test_api_detect_sign_post(self):
|
|
url = reverse("sign:api_detect")
|
|
file_obj = self._create_uploaded_image_file()
|
|
response = self.client.post(
|
|
url,
|
|
{"image": file_obj, "confidence_threshold": 0.20},
|
|
format="multipart"
|
|
)
|
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
|
data = response.json()
|
|
self.assertEqual(data["status"], "success")
|
|
self.assertIn("panels", data)
|
|
self.assertIn("performance", data)
|
|
|
|
def test_api_detect_sign_missing_image(self):
|
|
url = reverse("sign:api_detect")
|
|
response = self.client.post(url, {}, format="json")
|
|
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
|
|
|
|
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)
|