Files
AareDAQ/tests/unit/daq/test_mlbox.py
T
appleb_m c1db7a8d8a
Build and Publish / test (push) Successful in 1m18s
Build and Publish / build (push) Successful in 14s
Build and Publish / Build and Deploy Docs (push) Successful in 35s
DAQ: refactored MLBOX and loop_Centering MLBOXTYPe. Loop_centering should no longer move if the only class found is pin, ice or needle. Refactored tests.
2026-04-28 15:23:24 +02:00

120 lines
4.0 KiB
Python

import pytest
import numpy as np
from unittest.mock import MagicMock, patch
from aare.daq.mlbox import MlBox
from aare.common.beamline import MXBeamline
from aare.common.models import MLOutputModel, MLBoxModel, MLBoxType
@pytest.fixture
def mock_wrapper():
return MagicMock()
@pytest.fixture
def mlbox(mock_wrapper):
return MlBox(MXBeamline.X10SA, wrapper=mock_wrapper)
def test_init_x10sa(mock_wrapper):
ml = MlBox(MXBeamline.X10SA, wrapper=mock_wrapper)
assert ml._MlBox__beamline == MXBeamline.X10SA
def test_predict_uses_injected_wrapper():
fake = MagicMock()
fake.get_latest_prediction_bundle.return_value = MagicMock()
ml = MlBox(MXBeamline.X10SA, wrapper=fake)
ml._fetch_prediction_bundle()
fake.get_latest_prediction_bundle.assert_called_once()
def test_init_unsupported():
with pytest.raises(NotImplementedError):
MlBox(MXBeamline.X06DA)
def test_prediction_score():
assert MlBox._prediction_score(None) == (0, 0.0)
out = MLOutputModel()
assert MlBox._prediction_score(out) == (0, 0.0)
out.add_box(MLBoxType.CRYSTAL, (0, 0, 10, 10), conf=0.9)
assert MlBox._prediction_score(out) == (1, 0.9)
def test_check_overlap():
box_a = MLBoxModel.from_tuple(MLBoxType.CRYSTAL, (0, 0, 10, 10), conf=1.0)
box_b = MLBoxModel.from_tuple(MLBoxType.CRYSTAL, (5, 5, 15, 15), conf=1.0)
assert MlBox._check_overlap(box_a, box_b) > 0
box_c = MLBoxModel.from_tuple(MLBoxType.CRYSTAL, (20, 20, 30, 30), conf=1.0)
assert MlBox._check_overlap(box_a, box_c) == 0
def test_filter_predictions(mlbox):
out = MLOutputModel()
out.add_box(MLBoxType.CRYSTAL, (0, 0, 10, 10), conf=0.4)
out.add_box(MLBoxType.CRYSTAL, (20, 20, 30, 30), conf=0.8)
mlbox._filter_predictions(out, confidence_min=0.5)
assert len(out.boxes) == 1
assert list(out.boxes.values())[0].conf == 0.8
@patch('cv2.imdecode')
def test_decode_bundle_image(mock_imdecode):
mock_imdecode.return_value = np.zeros((100, 100, 3))
img = MlBox._decode_bundle_image(b"fake_jpeg")
assert img is not None
assert img.shape == (100, 100, 3)
def test_extract_bundle_meta():
pred = MagicMock()
# Mocking what _extract_target_point expects:
# It tries to find 'target_point' in metadata or prediction itself.
# In src/aare/daq/mlbox.py:
# def _extract_target_point(prediction: LatestPredictionModel | None) -> tuple[float, float] | None:
# raw = getattr(prediction, "target_point", None)
# if raw is None:
# meta = getattr(prediction, "metadata", {})
# raw = meta.get("target_point")
pred.target_point = [100.0, 200.0]
pred.focus_score = 0.5
meta = MlBox._extract_bundle_meta(pred)
assert meta.target_point == (100.0, 200.0)
assert meta.focus == 0.5
def test_preferred_class_prefers_loop_over_pin_when_margin_not_exceeded():
boxes = MLOutputModel()
boxes.add_box(MLBoxType.PIN, (0, 0, 10, 10), 0.8)
boxes.add_box(MLBoxType.LOOP_FACE, (1, 1, 9, 9), 0.75)
best = MlBox.get_preferred_class_box_with_confidence_threshold(
boxes,
loop_preference_margin=0.1,
)
assert best is not None
assert best.cls == MLBoxType.LOOP_FACE
def test_prediction_score_empty():
assert MlBox._prediction_score(None) == (0, 0.0)
def test_prediction_score_uses_count_and_max_confidence():
model = MLOutputModel()
model.add_box(MLBoxType.PIN, (0, 0, 10, 10), 0.4)
model.add_box(MLBoxType.CRYSTAL, (1, 1, 9, 9), 0.8)
assert MlBox._prediction_score(model) == (2, 0.8)
def test_best_by_class_keeps_highest_confidence_per_class():
results = [
{"class": 1, "confidence": 0.4, "box": {"x1": 0, "y1": 0, "x2": 10, "y2": 10}},
{"class": 1, "confidence": 0.8, "box": {"x1": 1, "y1": 1, "x2": 11, "y2": 11}},
{"class": 2, "confidence": 0.6, "box": {"x1": 2, "y1": 2, "x2": 12, "y2": 12}},
]
out = MlBox._best_by_class(results)
assert out is not None
assert out.get_best_for_class(MLBoxType.PIN).conf == 0.8
assert out.get_best_for_class(MLBoxType.CRYSTAL).conf == 0.6