From c1db7a8d8a547fef6a0c6915694e404ffa8ae031 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Tue, 28 Apr 2026 15:23:24 +0200 Subject: [PATCH] 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. --- src/aare/common/models.py | 24 +++++++++++------------ src/aare/daq/daq.py | 27 +++++++++----------------- src/aare/daq/mlbox.py | 23 ++++++++++------------ tests/unit/common/test_mlbox_model.py | 12 ++++++------ tests/unit/common/test_models_extra.py | 18 ++++++++--------- tests/unit/daq/test_mlbox.py | 26 ++++++++++++------------- 6 files changed, 59 insertions(+), 71 deletions(-) diff --git a/src/aare/common/models.py b/src/aare/common/models.py index 58cf5b83..52d23d94 100644 --- a/src/aare/common/models.py +++ b/src/aare/common/models.py @@ -431,12 +431,12 @@ class FluorescenceSpectrumOutputModel(BaseModel): class MLBoxType(Enum): - Loop_all = 0 - Pin = 1 - Crystal = 2 - Loop_face = 3 - Ice = 4 - Needle = 5 + LOOP_ALL = 0 + PIN = 1 + CRYSTAL = 2 + LOOP_FACE = 3 + ICE = 4 + NEEDLE = 5 class BoundingBoxModel(BaseModel): @@ -466,17 +466,17 @@ class MLOutputModel(BaseModel): @staticmethod def get_class_str(cls: MLBoxType) -> str: - if cls == MLBoxType.Loop_all: + if cls == MLBoxType.LOOP_ALL: return "Loop_all" - if cls == MLBoxType.Pin: + if cls == MLBoxType.PIN: return "Pin" - if cls == MLBoxType.Crystal: + if cls == MLBoxType.CRYSTAL: return "Crystal" - if cls == MLBoxType.Loop_face: + if cls == MLBoxType.LOOP_FACE: return "Loop_face" - if cls == MLBoxType.Ice: + if cls == MLBoxType.ICE: return "Ice" - if cls == MLBoxType.Needle: + if cls == MLBoxType.NEEDLE: return "Needle" return "Unknown" diff --git a/src/aare/daq/daq.py b/src/aare/daq/daq.py index 3aa962f3..00601291 100644 --- a/src/aare/daq/daq.py +++ b/src/aare/daq/daq.py @@ -43,7 +43,7 @@ from aare.common.models import ( SampleShortInfoList, AutofocusSettings, DAQStatusModel, BeamlineStatus, SessionStatus, SampleCameraSettings, ZoomModeEnum, SimpleScanParameters, MLBoxModel, FluorescenceSpectrumParameterModel, - FluorescenceSpectrumOutputModel, DAQOperation, LoopCenteringResult) + FluorescenceSpectrumOutputModel, DAQOperation, LoopCenteringResult, MLBoxType) from aare.common.automation_models import ( AutomationProgress, StepState, @@ -2362,7 +2362,7 @@ class AareDAQ: ) return top_left, n_y - + #TODO move wait_screenshot_sleep to save_screenshot_db!!! @log_timing(logger, "Loop center sequence") def __loop_center_sequence( self, @@ -2441,10 +2441,11 @@ class AareDAQ: boxes=boxes, bundle_image=bundle_image, ) - target = self._select_smargon_target( - calculated_target=target, - predicted_target=pred_target_point_smargon, - tolerance_um=500.0) + if cls in [MLBoxType.LOOP_ALL.value, MLBoxType.LOOP_FACE.value, MLBoxType.CRYSTAL.value]: + target = self._select_smargon_target( + calculated_target=target, + predicted_target=pred_target_point_smargon, + tolerance_um=500.0) logger.debug(f"calculated target: {target} compares to prediction: {pred_target_point_smargon}") except Exception as e: logger.error(f"Error getting ML box for angle {angle}") @@ -2460,10 +2461,10 @@ class AareDAQ: found_classes_count[int(c)] = found_classes_count.get(int(c), 0) + 1 logger.debug(f"classes found: {classes}") logger.debug(f"class found: {cls}") - if cls is not None and cls != 1: + if cls is not None and cls != MLBoxType.PIN.value and cls != MLBoxType.NEEDLE.value and cls != MLBoxType.ICE.value: targets_found_this_attempt += 1 - found_flag = True found_angle = angle + time_to_move_smargon = time.perf_counter() self.__devs.smargon_pos = target self.__devs.smargon_wait(60) @@ -2481,16 +2482,6 @@ class AareDAQ: ) self.save_screenshot_db(sample_id, f"{sample_id}_{angle}_{zoom_value:.0f}") - # if targets_found_this_attempt == 0: - # logger.error(f"No targets found in this attempt {attempt}") - # raise LoopCenteringFailed - # - # else: - # if targets_found_this_attempt >= len(base_angles): - # logger.debug( - # f"sucessfully found {targets_found_this_attempt} targets in attempt {attempt + 1} ") - # break - #if found_flag is not None and found_angle is not None: logger.debug(f"found a target at angle {found_angle} in attempt {attempt + 1}") base_angles = (0, 90) logger.debug(f"new base angles: {base_angles}") diff --git a/src/aare/daq/mlbox.py b/src/aare/daq/mlbox.py index 26417070..65e55bd2 100644 --- a/src/aare/daq/mlbox.py +++ b/src/aare/daq/mlbox.py @@ -1,14 +1,11 @@ from dataclasses import dataclass, field -from typing import Optional, Iterable, TYPE_CHECKING, Any +from typing import Optional, Iterable import cv2 import numpy as np import time -if TYPE_CHECKING: - from aarelcinfer_client.models import LatestPredictionModel -else: - LatestPredictionModel = Any +from aarelcinfer_client.models import LatestPredictionModel from aare.common.beamline import MXBeamline from aare.common.models import MLBoxModel, MLOutputModel, MLBoxType, BoundingBoxModel @@ -266,7 +263,7 @@ class MlBox: return {"left": left, "top": top, "overlap_x": overlap_x, "overlap_y": overlap_y} def box_relation(self, boxes: MLOutputModel, classes:list[str]|None = None): - pin: MLBoxModel | None = boxes.get_best_for_class(MLBoxType.Pin) + pin: MLBoxModel | None = boxes.get_best_for_class(MLBoxType.PIN) if not pin or not classes: return None @@ -305,7 +302,7 @@ class MlBox: def _filter_predictions(self, predictions: MLOutputModel, overlap_with_pin: Optional[float] = None, confidence_min: Optional[float] = None): - pin = predictions.get_best_for_class(MLBoxType.Pin) + pin = predictions.get_best_for_class(MLBoxType.PIN) keys_to_remove = [] for key, model in predictions.boxes.items(): if model is None or model.box is None: @@ -314,7 +311,7 @@ class MlBox: if confidence_min is not None and (model.conf or 0.0) < confidence_min: keys_to_remove.append(key) continue - if model.cls == MLBoxType.Pin: + if model.cls == MLBoxType.PIN: continue if overlap_with_pin is not None and pin: if self._check_overlap(model, pin) >= overlap_with_pin: @@ -410,7 +407,7 @@ class MlBox: Best MLBoxModel according to preferences """ if preferred_class is None: - order = (MLBoxType.Crystal, MLBoxType.Loop_face, MLBoxType.Loop_all, MLBoxType.Pin) + order = (MLBoxType.CRYSTAL, MLBoxType.LOOP_FACE, MLBoxType.LOOP_ALL, MLBoxType.PIN) else: if isinstance(preferred_class, MLBoxType): order = (preferred_class,) @@ -430,9 +427,9 @@ class MlBox: return None # Special handling: prefer loops over pin unless pin is significantly better - pin_box = best_boxes.get(MLBoxType.Pin) - loop_face_box = best_boxes.get(MLBoxType.Loop_face) - loop_all_box = best_boxes.get(MLBoxType.Loop_all) + pin_box = best_boxes.get(MLBoxType.PIN) + loop_face_box = best_boxes.get(MLBoxType.LOOP_FACE) + loop_all_box = best_boxes.get(MLBoxType.LOOP_ALL) best_loop = loop_face_box if loop_face_box else loop_all_box @@ -457,7 +454,7 @@ class MlBox: preferred_class: Optional[Iterable[int] | int | MLBoxType] = None) -> Optional[ MLBoxModel]: if preferred_class is None: - order = (MLBoxType.Crystal, MLBoxType.Loop_face, MLBoxType.Loop_all, MLBoxType.Pin) + order = (MLBoxType.CRYSTAL, MLBoxType.LOOP_FACE, MLBoxType.LOOP_ALL, MLBoxType.PIN) else: if isinstance(preferred_class, MLBoxType): order = (preferred_class,) diff --git a/tests/unit/common/test_mlbox_model.py b/tests/unit/common/test_mlbox_model.py index 3db87cc6..0ae92607 100644 --- a/tests/unit/common/test_mlbox_model.py +++ b/tests/unit/common/test_mlbox_model.py @@ -3,8 +3,8 @@ from aare.common.models import MLOutputModel, MLBoxType def test_add_box_generates_unique_keys(): model = MLOutputModel() - key1 = model.add_box(MLBoxType.Crystal, (1, 2, 3, 4), 0.8) - key2 = model.add_box(MLBoxType.Crystal, (5, 6, 7, 8), 0.9) + key1 = model.add_box(MLBoxType.CRYSTAL, (1, 2, 3, 4), 0.8) + key2 = model.add_box(MLBoxType.CRYSTAL, (5, 6, 7, 8), 0.9) assert key1 == "Crystal" assert key2 == "Crystal_2" @@ -12,14 +12,14 @@ def test_add_box_generates_unique_keys(): def test_get_best_for_class_returns_highest_confidence(): model = MLOutputModel() - model.add_box(MLBoxType.Pin, (1, 1, 2, 2), 0.3) - model.add_box(MLBoxType.Pin, (3, 3, 4, 4), 0.7) + model.add_box(MLBoxType.PIN, (1, 1, 2, 2), 0.3) + model.add_box(MLBoxType.PIN, (3, 3, 4, 4), 0.7) - best = model.get_best_for_class(MLBoxType.Pin) + best = model.get_best_for_class(MLBoxType.PIN) assert best is not None assert best.conf == 0.7 def test_get_best_for_class_returns_none_when_missing(): model = MLOutputModel() - assert model.get_best_for_class(MLBoxType.Crystal) is None \ No newline at end of file + assert model.get_best_for_class(MLBoxType.CRYSTAL) is None \ No newline at end of file diff --git a/tests/unit/common/test_models_extra.py b/tests/unit/common/test_models_extra.py index e2ace151..2c7c5032 100644 --- a/tests/unit/common/test_models_extra.py +++ b/tests/unit/common/test_models_extra.py @@ -59,7 +59,7 @@ def test_beam_mark_coeff_model_apply(): def test_ml_output_model_extra_methods(): model = MLOutputModel() - key = model.add_box(MLBoxType.Crystal, (1, 2, 3, 4), 0.8) + key = model.add_box(MLBoxType.CRYSTAL, (1, 2, 3, 4), 0.8) # Test get_box_model (line 499) box_model = model.get_box_model(key) @@ -74,22 +74,22 @@ def test_ml_output_model_extra_methods(): assert model.get_box_tuple_with_conf("NonExistent") is None # Test get_tuples_for_class (lines 523-527) - tuples = model.get_tuples_for_class(MLBoxType.Crystal) + tuples = model.get_tuples_for_class(MLBoxType.CRYSTAL) assert len(tuples) == 1 assert tuples[0] == (1, 2, 3, 4) # Test get_tuples_with_conf_for_class (lines 530-534) - tuples_conf = model.get_tuples_with_conf_for_class(MLBoxType.Crystal) + tuples_conf = model.get_tuples_with_conf_for_class(MLBoxType.CRYSTAL) assert len(tuples_conf) == 1 assert tuples_conf[0] == (1, 2, 3, 4, 0.8) # Test get_class_str (lines 470, 475-481) - assert MLOutputModel.get_class_str(MLBoxType.Loop_all) == "Loop_all" - assert MLOutputModel.get_class_str(MLBoxType.Pin) == "Pin" - assert MLOutputModel.get_class_str(MLBoxType.Crystal) == "Crystal" - assert MLOutputModel.get_class_str(MLBoxType.Loop_face) == "Loop_face" - assert MLOutputModel.get_class_str(MLBoxType.Ice) == "Ice" - assert MLOutputModel.get_class_str(MLBoxType.Needle) == "Needle" + assert MLOutputModel.get_class_str(MLBoxType.LOOP_ALL) == "Loop_all" + assert MLOutputModel.get_class_str(MLBoxType.PIN) == "Pin" + assert MLOutputModel.get_class_str(MLBoxType.CRYSTAL) == "Crystal" + assert MLOutputModel.get_class_str(MLBoxType.LOOP_FACE) == "Loop_face" + assert MLOutputModel.get_class_str(MLBoxType.ICE) == "Ice" + assert MLOutputModel.get_class_str(MLBoxType.NEEDLE) == "Needle" assert MLOutputModel.get_class_str(100) == "Unknown" def test_beamline_state_enum_display_name(): diff --git a/tests/unit/daq/test_mlbox.py b/tests/unit/daq/test_mlbox.py index 4af16fbf..4aef2039 100644 --- a/tests/unit/daq/test_mlbox.py +++ b/tests/unit/daq/test_mlbox.py @@ -37,21 +37,21 @@ def test_prediction_score(): out = MLOutputModel() assert MlBox._prediction_score(out) == (0, 0.0) - out.add_box(MLBoxType.Crystal, (0, 0, 10, 10), conf=0.9) + 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) + 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) + 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) + 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 @@ -84,8 +84,8 @@ def test_extract_bundle_meta(): 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) + 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, @@ -93,7 +93,7 @@ def test_preferred_class_prefers_loop_over_pin_when_margin_not_exceeded(): ) assert best is not None - assert best.cls == MLBoxType.Loop_face + assert best.cls == MLBoxType.LOOP_FACE def test_prediction_score_empty(): assert MlBox._prediction_score(None) == (0, 0.0) @@ -101,8 +101,8 @@ def test_prediction_score_empty(): 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) + 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) @@ -116,5 +116,5 @@ def test_best_by_class_keeps_highest_confidence_per_class(): 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 \ No newline at end of file + assert out.get_best_for_class(MLBoxType.PIN).conf == 0.8 + assert out.get_best_for_class(MLBoxType.CRYSTAL).conf == 0.6 \ No newline at end of file