After story 4
This commit is contained in:
@@ -0,0 +1,98 @@
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
|
||||
|
||||
from batcher import DynamicBatcher
|
||||
from face_detector import Detection, FaceDetector
|
||||
from gpu_manager import GPUMemoryManager
|
||||
|
||||
|
||||
class Story04FaceDetectionTests(unittest.TestCase):
|
||||
def test_gpu_manager_respects_memory_thresholds(self):
|
||||
manager = GPUMemoryManager(
|
||||
max_memory_gb=18.0,
|
||||
reduce_threshold_gb=16.0,
|
||||
increase_threshold_gb=10.0,
|
||||
initial_batch_size=16,
|
||||
)
|
||||
|
||||
with patch.object(manager, "get_memory_stats", return_value=(17.0, 18.0)):
|
||||
manager.adjust_batch_size()
|
||||
self.assertLess(manager.current_batch_size, 16)
|
||||
|
||||
with patch.object(manager, "get_memory_stats", return_value=(8.0, 10.0)):
|
||||
manager.adjust_batch_size()
|
||||
self.assertGreater(manager.current_batch_size, manager.min_batch_size)
|
||||
|
||||
def test_dynamic_batcher_flushes_when_batch_full(self):
|
||||
processed_batches = []
|
||||
|
||||
def process_fn(batch):
|
||||
processed_batches.append(batch)
|
||||
return [len(batch)]
|
||||
|
||||
batcher = DynamicBatcher(
|
||||
process_fn=process_fn,
|
||||
max_batch_size=3,
|
||||
batch_timeout_ms=10000.0,
|
||||
)
|
||||
|
||||
for i in range(5):
|
||||
batcher.add(i)
|
||||
|
||||
self.assertEqual(len(processed_batches), 1)
|
||||
self.assertEqual(processed_batches[0], [0, 1, 2])
|
||||
self.assertEqual(batcher.queued_count, 2)
|
||||
|
||||
def test_face_detector_nms_removes_overlapping_boxes(self):
|
||||
detector = FaceDetector(
|
||||
engine_path="/nonexistent/model.trt",
|
||||
confidence_threshold=0.1,
|
||||
iou_threshold=0.45,
|
||||
max_faces_per_frame=10,
|
||||
)
|
||||
|
||||
duplicates = [
|
||||
Detection("frame.jpg", 10, 10, 50, 50, 0.9),
|
||||
Detection("frame.jpg", 12, 12, 48, 48, 0.8),
|
||||
Detection("frame.jpg", 100, 100, 150, 150, 0.75),
|
||||
]
|
||||
|
||||
kept = detector._nms(duplicates)
|
||||
self.assertEqual(len(kept), 2)
|
||||
self.assertAlmostEqual(kept[0].confidence, 0.9, places=5)
|
||||
|
||||
def test_face_detector_extracts_and_resizes_crops(self):
|
||||
detector = FaceDetector(
|
||||
engine_path="/nonexistent/model.trt",
|
||||
input_size=640,
|
||||
confidence_threshold=0.25,
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
from PIL import Image
|
||||
|
||||
frame_path = Path(tmpdir) / "frame.jpg"
|
||||
image = Image.new("RGB", (640, 480), color=(100, 150, 200))
|
||||
image.save(frame_path)
|
||||
|
||||
detections = [
|
||||
Detection(str(frame_path), 0, 0, 640, 480, 0.95),
|
||||
]
|
||||
cropped = detector.extract_crops(detections, output_dir=tmpdir, crop_size=(224, 224))
|
||||
|
||||
self.assertEqual(len(cropped), 1)
|
||||
self.assertTrue(Path(cropped[0].crop_path).exists())
|
||||
|
||||
with Image.open(cropped[0].crop_path) as crop:
|
||||
self.assertEqual(crop.size, (224, 224))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user