loko/loko/sign/tests_ai.py

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)