feat: introduce switching gridscan analysis modes

with raster service refactor
This commit is contained in:
2026-09-08 13:28:03 +02:00
committed by perl_d
parent c12fb34057
commit fcc5aea1a0
4 changed files with 209 additions and 236 deletions
+1 -3
View File
@@ -7,12 +7,10 @@ requires-python = ">=3.11"
dependencies = [
"uv",
"gunicorn",
# >=0.7: DataCollectionParameters.transmission is a 0-1 fraction, which
# the scan panels rely on (older releases held an int percentage).
"aarecommon>=0.7.3",
"pydantic>=2.11",
"numpy",
"jfjoch_client>=1.0.0rc165",
"jfjoch_client>=1.0.0rc166",
"pyJWT",
"pyzmq",
"opencv-python-headless",
+10
View File
@@ -10,6 +10,7 @@ import numpy as np
import redis
import redis_lock
from aarecommon.config.beamline import cfg_get
from aarecommon.config.config_models import GridscanAnalysisMode
from aarecommon.config.logger import setup_logger
from aarecommon.errors.exception_handler import BeamlineBusyException
from aarecommon.math.coordinate import AerotechCoordinate, Coordinate
@@ -656,6 +657,15 @@ class BeamlineConfig:
self._client.set(f"{self._bl}:beam_size_x", data.x)
self._client.set(f"{self._bl}:beam_size_y", data.y)
@property
def gridscan_analysis_mode(self) -> GridscanAnalysisMode:
data = self._client.get(f"{self._bl}:gridscan_analysis_mode")
return GridscanAnalysisMode(data) if data is not None else GridscanAnalysisMode.FindXtal
@gridscan_analysis_mode.setter
def gridscan_analysis_mode(self, val: GridscanAnalysisMode) -> GridscanAnalysisMode:
return self._client.set(f"{self._bl}:gridscan_analysis_mode", str(val))
def _get_settings(self) -> BeamlineSettingsModel:
tmp = self._client.get(f"{self._bl}:settings")
if tmp is None:
+182 -233
View File
@@ -1,8 +1,10 @@
import copy
import re
import time
from math import ceil, floor
from aarecommon.config.beamline import cfg_get
from aarecommon.config.config_models import GridscanAnalysisMode
from aarecommon.config.logger_events import (
geom_log_context,
log_ml_bundle_meta,
@@ -13,22 +15,29 @@ from aarecommon.config.logger_events import (
from aarecommon.errors.exception_handler import RasterScanException
from aarecommon.math.coordinate import AerotechCoordinate, Coordinate, SmargonCoordinate
from aarecommon.math.find_xtal import (
center_of_grid,
compute_crystal_score_array,
get_best_b_factor,
get_best_res,
get_xtal_size,
raster_highest_score,
)
from aarecommon.math.jfjoch_gridscan_union import GridScanResult, Thresholds, analyse
from aarecommon.math.raster_grid import grid_to_image_id
from aarecommon.models.models import BeamlineStateEnum, MLBoxType
from aarecommon.models.raster_grid import (
CenterOfMassModel,
CompletedRasterGrid,
CompletedRasterGridElem,
RasterGridRequest,
)
from aareDB import SampleEventType
from bec_lib.redis_connector.streams import dataclass
from jfjoch_client import ScanResult
from jfjoch_client.exceptions import NotFoundException
from jfjoch_client.models.grid_scan import GridScan
from aare.daq.daq import DAQStatusModel
from aare.daq.mlbox import MLBoxPredictionResult
from aare.daq.operations.common.ml_bounding_box import (
MLRasterPlan,
@@ -39,10 +48,92 @@ from aare.daq.operations.raster.models import RasterContext
from aare.devices.area_detector import AutoEnum
@dataclass
class _AnalysisResults:
center_smargon_coord: SmargonCoordinate
com: CenterOfMassModel
v2_results: GridScanResult | None = None
class RasterService:
def __init__(self, *, context: RasterContext, logger):
self.ctx = context
self.logger = logger
self.ingestor = self.ctx.services.ingestion or self.ctx.deps.aare
self.send_sample_event = (
self.ctx.services.events.send
if self.ctx.services.events is not None
else self.ctx.deps.aare.send_sample_event
)
def _get_anaylsis_results(
self,
mode: GridscanAnalysisMode,
scan_result: ScanResult,
request: RasterGridRequest,
*,
daq: DAQStatusModel | None = None,
thresholds: Thresholds | None = None,
z_um: float | None = None,
):
if mode == GridscanAnalysisMode.FindXtal:
com = raster_highest_score(scan_result.images)
smargon_coord = self._smargon_offset_from_scan_result(scan_result, request, com)
v2_results = None
else:
step_x_um = (request.grid_size_mm.x / request.n_x) * 1000
step_y_um = (request.grid_size_mm.y / request.n_y) * 1000
grid_scan = GridScan(
n_fast=request.n_x, step_x_um=step_x_um, step_y_um=step_y_um, snake=True
)
v2_results = analyse(scan_result, grid_scan, daq=daq, thresholds=thresholds, z_um=z_um)
self.logger.info(
f"experimental gridscan analaysis results: \n {v2_results.model_dump_json(indent=2)}"
)
if v2_results.centre is not None:
com = CenterOfMassModel(
n_x=v2_results.centre.nx,
n_y=v2_results.centre.ny,
max_image=v2_results.centre.image_number,
)
smargon_coord = SmargonCoordinate(
sh_mm=request.smargon_top_left.sh_mm + v2_results.centre.offset_mm(),
phi_deg=request.smargon_top_left.phi_deg,
chi_deg=request.smargon_top_left.chi_deg,
)
else:
com = center_of_grid(request.n_x, request.n_y)
smargon_coord = self._smargon_offset_from_scan_result(scan_result, request, com)
return _AnalysisResults(smargon_coord, com, v2_results)
def _smargon_offset_from_scan_result(
self, scan_result: ScanResult, request: RasterGridRequest, com: CenterOfMassModel
):
target_coor = com.get_com_mm(request)
target_coor_offset = self.ctx.sample_geometry.smargon_nudge(target_coor)
self.logger.info(
"Calculated raster centre offset",
extra=merge_log_context(
sample_log_context(self.ctx.sample),
raster_request_log_context(request),
{
"centre_offset_x_mm": target_coor_offset.x,
"centre_offset_y_mm": target_coor_offset.y,
"centre_offset_z_mm": target_coor_offset.z,
"grid_half_width_x_mm": x,
"grid_half_height_y_mm": y,
"top_left_x_mm": getattr(request.smargon_top_left.sh_mm, "x", None),
"top_left_y_mm": getattr(request.smargon_top_left.sh_mm, "y", None),
"top_left_z_mm": getattr(request.smargon_top_left.sh_mm, "z", None),
},
),
)
self.logger.info(f"moving Smargon to grid centre offset {target_coor_offset}")
return SmargonCoordinate(
sh_mm=request.smargon_top_left.sh_mm + target_coor_offset,
phi_deg=request.smargon_top_left.phi_deg,
chi_deg=request.smargon_top_left.chi_deg,
)
@staticmethod
def _scan_result_image_count(scan_result) -> int | None:
@@ -165,61 +256,6 @@ class RasterService:
x1, y1, x2, y2 = box
return x1 <= margin or y1 <= margin or x2 >= width - margin or y2 >= height - margin
def _zoom_to_fit_box(self, plan: MLRasterPlan) -> tuple | None:
"""Pick the box that drives zoom-to-fit: loop_all if detected and not
clipped at the frame edge, otherwise loop_face."""
if plan.loop_all_box is not None and not self._box_touches_frame_edge(
plan.loop_all_box, plan.image_width, plan.image_height
):
return plan.loop_all_box
return plan.loop_face_box
def _apply_zoom_to_fit(self, plan: MLRasterPlan) -> None:
"""Zoom so the chosen loop box covers at most 1/3 of the frame."""
box = self._zoom_to_fit_box(plan)
if box is None or not plan.image_width or not plan.image_height:
return
x1, y1, x2, y2 = box
box_px = max(x2 - x1, y2 - y1)
if box_px <= 0:
return
screen_px = min(plan.image_width, plan.image_height)
current_zoom = self.ctx.deps.devs.zoom
pixel_in_mm_now = self.ctx.deps.cfg.pixel_to_mm(current_zoom)
box_mm = box_px * pixel_in_mm_now
# Target mm-per-pixel so the box spans exactly screen/3.
target_pixel_in_mm = 3.0 * box_mm / screen_px
try:
target_zoom = self.ctx.deps.cfg.zoom_for_pixel_to_mm(target_pixel_in_mm)
except (ValueError, ZeroDivisionError) as e:
self.logger.warning(f"Skipping auto-center zoom-to-fit: {e}")
return
zoom_min = float(cfg_get("daq.hardware.zoom_min", 1.0))
zoom_max = float(cfg_get("daq.hardware.zoom_max", 1000.0))
target_zoom = max(zoom_min, min(zoom_max, target_zoom))
self.logger.info(
"Auto-center zoom-to-fit",
extra=merge_log_context(
sample_log_context(self.ctx.sample),
{
"box_px": box_px,
"screen_px": screen_px,
"current_zoom": current_zoom,
"target_zoom": target_zoom,
"used_loop_all": box is plan.loop_all_box,
},
),
)
self.ctx.deps.devs.samcam_auto(AutoEnum.AUTO)
self.ctx.deps.devs.set_zoom(target_zoom, wait=True)
time.sleep(0.2)
self.ctx.deps.devs.samcam_auto(AutoEnum.ONCE)
def auto_center_line_scan_top_left(
self,
*,
@@ -386,17 +422,72 @@ class RasterService:
return top_left, n_y
def execute(
self, request: RasterGridRequest, wait_for_screenshot: float | None = None
) -> CompletedRasterGridElem:
def _do_gridscan_move_on_aerotech(self, request: RasterGridRequest, expected_time: float):
self.ctx.deps.devs.set_dtz(request.dtz, wait=True)
self.ctx.deps.devs.aerotech.grid_scan(
grid_elem_size_y_um=request.grid_size_mm.y * 1000,
grid_elem_size_x_um=request.grid_size_mm.x * 1000,
grid_elem_count_x=request.n_x,
grid_elem_count_y=request.n_y,
time_sec=request.exp_time_s,
run_async=True,
)
self.ctx.deps.devs.aerotech.wait_till_done(
timeout=int(round(expected_time + expected_time * 0.1 + 60, 0))
)
if isinstance(self.ctx.deps.cfg.abr_meas_pos, Coordinate):
coord = self.ctx.deps.cfg.abr_meas_pos
else:
coord = self.ctx.deps.cfg.abr_meas_pos.at_mm
self.ctx.deps.devs.aerotech_pos = AerotechCoordinate(
at_mm=coord, omega_deg=self.ctx.deps.devs.aerotech_omega
)
self.ctx.deps.devs.aerotech.wait_till_done(timeout=60)
def _check_setup(self, request: RasterGridRequest, wait_for_screenshot_s: float | None):
total_time = request.exp_time_s * request.n_x * request.n_y + request.n_y * 0.3
if wait_for_screenshot is None:
wait_for_screenshot = self.ctx.settings.wait_for_screenshot_s
if wait_for_screenshot_s is None:
wait_for_screenshot_s = self.ctx.settings.wait_for_screenshot_s
if total_time > 1200:
raise RasterScanException(f"Raster scan is too long {total_time}s > 20min")
return total_time, wait_for_screenshot_s
def _initial_gridscan_ingestion(
self,
request: RasterGridRequest,
scan_result: ScanResult,
wait_for_screenshot_s: float,
sample_id: int,
):
self.logger.debug("moving to XtalSnapshot to take a screenshot of the sample")
if self.ctx.services.state is None:
raise RuntimeError("RasterService requires services.state")
self.ctx.services.state.set_state(BeamlineStateEnum.XtalSnapshot)
time.sleep(wait_for_screenshot_s)
self.ctx.services.screenshots.save_to_db(
sample_id, f"{sample_id}_post_raster_{int(request.omega_deg)}deg"
)
self.ingestor.ingest_gridscan(
sample=self.ctx.sample,
raster_result=scan_result,
raster_request=request,
geom=self.ctx.sample_geometry,
com=None,
beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom),
)
def _additional_analysis_data_ingestion(self, results: GridScanResult): ...
def execute(
self, request: RasterGridRequest, wait_for_screenshot: float | None = None
) -> CompletedRasterGridElem:
total_time, wait_for_screenshot = self._check_setup(request, wait_for_screenshot)
status = self.ctx.status
smargon_top_left = request.smargon_top_left
@@ -421,43 +512,11 @@ class RasterService:
if self.ctx.sample is not None and self.ctx.sample.db_id is not None:
self.ctx.deps.aare.create_gridscan_run(self.ctx.sample, request, status)
if not self.ctx.deps.cfg.simulated_detector:
self.ctx.deps.jfjoch.wait_till_running(timeout=60.0)
else:
self.logger.info("Simulated detector mode enabled; faking jfjoch intilalisation.")
self.ctx.deps.jfjoch.wait_till_running(timeout=60.0)
self.logger.debug(
f"Starting grid scan with {request.n_x}x{request.n_y} points, exp time {request.exp_time_s}s"
)
if not self.ctx.deps.cfg.simulated_detector:
# The detector move (sa2dc) is non-blocking; ensure it has reached
# the requested distance before triggering the scan, so data is
# collected at the dtz already reported to JFJoch/DB.
self.ctx.deps.devs.set_dtz(request.dtz, wait=True)
self.ctx.deps.devs.aerotech.grid_scan(
grid_elem_size_y_um=request.grid_size_mm.y * 1000,
grid_elem_size_x_um=request.grid_size_mm.x * 1000,
grid_elem_count_x=request.n_x,
grid_elem_count_y=request.n_y,
time_sec=request.exp_time_s,
run_async=True,
)
self.ctx.deps.devs.aerotech.wait_till_done(
timeout=int(round(total_time + total_time * 0.1 + 60, 0))
)
if isinstance(self.ctx.deps.cfg.abr_meas_pos, Coordinate):
coord = self.ctx.deps.cfg.abr_meas_pos
else:
coord = self.ctx.deps.cfg.abr_meas_pos.at_mm
self.ctx.deps.devs.aerotech_pos = AerotechCoordinate(
at_mm=coord, omega_deg=self.ctx.deps.devs.aerotech_omega
)
self.ctx.deps.devs.aerotech.wait_till_done(timeout=60)
x = None
y = None
self._do_gridscan_move_on_aerotech(request, total_time)
scan_result = self.ctx.deps.jfjoch.wait_till_done(60)
if scan_result is None:
@@ -474,64 +533,10 @@ class RasterService:
self.logger.info(
f"Raster scan results from JFJoch: file {scan_result.file_prefix} with images {scan_result.images}"
)
com = raster_highest_score(scan_result.images)
if com is None:
self.logger.info("Calcualted COM is None -> using centre image")
if request.n_x == 1:
x = request.grid_size_mm.x / 2.0
else:
x = ((request.n_x - 1) * request.grid_size_mm.x) / 2.0
y = ((request.n_y - 1) * request.grid_size_mm.y) / 2.0
target_coor_offset = self.ctx.sample_geometry.smargon_nudge(Coordinate(x=x, y=y))
self.logger.info(
"Calculated raster centre offset",
extra=merge_log_context(
sample_log_context(self.ctx.sample),
raster_request_log_context(request),
{
"centre_offset_x_mm": target_coor_offset.x,
"centre_offset_y_mm": target_coor_offset.y,
"centre_offset_z_mm": target_coor_offset.z,
"grid_half_width_x_mm": x,
"grid_half_height_y_mm": y,
"top_left_x_mm": getattr(request.smargon_top_left.sh_mm, "x", None),
"top_left_y_mm": getattr(request.smargon_top_left.sh_mm, "y", None),
"top_left_z_mm": getattr(request.smargon_top_left.sh_mm, "z", None),
},
),
)
else:
self.logger.info("Calcualted COM is not None, proceeding")
target_coor = com.get_com_mm(request)
target_coor_offset = self.ctx.sample_geometry.smargon_nudge(target_coor)
self.logger.info(
"Calculated raster centre offset",
extra=merge_log_context(
sample_log_context(self.ctx.sample),
raster_request_log_context(request),
{
"centre_offset_x_mm": target_coor_offset.x,
"centre_offset_y_mm": target_coor_offset.y,
"centre_offset_z_mm": target_coor_offset.z,
"grid_half_width_x_mm": x,
"grid_half_height_y_mm": y,
"top_left_x_mm": getattr(request.smargon_top_left.sh_mm, "x", None),
"top_left_y_mm": getattr(request.smargon_top_left.sh_mm, "y", None),
"top_left_z_mm": getattr(request.smargon_top_left.sh_mm, "z", None),
},
),
)
self.logger.info(f"moving Smargon to grid centre offset {target_coor_offset}")
target_smargon = SmargonCoordinate(
sh_mm=request.smargon_top_left.sh_mm + target_coor_offset,
phi_deg=request.smargon_top_left.phi_deg,
chi_deg=request.smargon_top_left.chi_deg,
analysis_results = self._get_anaylsis_results(
self.ctx.deps.cfg.gridscan_analysis_mode, scan_result, request, daq=self.ctx.status
)
self.ctx.deps.devs.smargon_pos = target_smargon
self.ctx.deps.devs.smargon_pos = analysis_results.center_smargon_coord
self.ctx.deps.devs.smargon_wait(timeout=180)
self.logger.info(
@@ -555,68 +560,30 @@ class RasterService:
else None
)
if sample_id:
self.logger.debug("moving to XtalSnapshot to take a screenshot of the sample")
if self.ctx.services.state is None:
raise RuntimeError("RasterService requires services.state")
self.ctx.services.state.set_state(BeamlineStateEnum.XtalSnapshot)
self._initial_gridscan_ingestion(
request, scan_result, wait_for_screenshot, sample_id
)
if analysis_results.v2_results is not None:
self._additional_analysis_data_ingestion(analysis_results.v2_results)
if wait_for_screenshot and wait_for_screenshot > 0:
time.sleep(wait_for_screenshot)
self.ctx.services.screenshots.save_to_db(
sample_id, f"{sample_id}_post_raster_{int(request.omega_deg)}deg"
diffraction_image_id = analysis_results.com.max_image
diffraction_image_filename = (
f"{sample_id}_best_diffraction_from_raster_image_{diffraction_image_id}"
)
if self.ctx.services.ingestion is not None:
self.ctx.services.ingestion.ingest_gridscan(
sample=self.ctx.sample,
raster_result=scan_result,
raster_request=request,
geom=self.ctx.sample_geometry,
com=None,
beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom),
)
else:
self.ctx.deps.aare.ingest_gridscan(
sample=self.ctx.sample,
raster_result=scan_result,
raster_request=request,
geom=self.ctx.sample_geometry,
com=None,
beam_mark_pxl=self.ctx.deps.cfg.get_beam_mark(self.ctx.deps.devs.zoom),
)
if com is not None and com.max_image is not None:
diffraction_image_filename = f"{self.ctx.sample.db_id}_best_diffraction_from_raster_image_{com.max_image}"
diffraction_image_id = com.max_image
else:
diffraction_image_filename = (
f"{self.ctx.sample.db_id}_diffraction_image_near_grid_scan"
)
diffraction_image_id = self._grid_image_id_from_centre_offset(
x_mm=x, y_mm=y, request=request
)
self._upload_raster_diffraction_preview(
sample_id=self.ctx.sample.db_id,
sample_id=sample_id,
filename=diffraction_image_filename,
image_id=diffraction_image_id,
scan_result=scan_result,
request=request,
)
if self.ctx.services.events is not None:
self.ctx.services.events.send(
self.ctx.sample.db_id,
SampleEventType.RASTERED,
comment=f"Raster completed at {request.omega_deg:.1f} deg",
)
else:
self.ctx.deps.aare.send_sample_event(
self.ctx.sample.db_id,
event_type=SampleEventType.RASTERED,
comment=f"Raster completed at {request.omega_deg:.1f} deg",
)
self.send_sample_event(
sample_id,
SampleEventType.RASTERED,
comment=f"Raster completed at {request.omega_deg:.1f} deg",
)
self.logger.info(
"Raster finished",
@@ -671,18 +638,9 @@ class RasterService:
),
)
if self.ctx.services.events is not None:
self.ctx.services.events.send(
sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
else:
self.ctx.deps.aare.send_sample_event(
sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
self.send_sample_event(
sample.db_id, SampleEventType.RASTERING, comment=f"Raster at {geom.omega_deg:.1f} deg"
)
plan = self.ml_raster_plan(sample.db_id, f"ml_{geom.omega_deg:.2f}deg")
@@ -773,18 +731,12 @@ class RasterService:
res1 = self.execute(grid)
grid.omega_deg += 90
if self.ctx.services.events is not None:
self.ctx.services.events.send(
self.ctx.sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
else:
self.ctx.deps.aare.send_sample_event(
self.ctx.sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
self.send_sample_event(
self.ctx.sample.db_id,
SampleEventType.RASTERING,
comment=f"Raster at {geom.omega_deg:.1f} deg",
)
self.ctx.deps.devs.aerotech_omega = grid.omega_deg
grid.n_x = 1
@@ -813,12 +765,9 @@ class RasterService:
)
status = self.ctx.status
if not self.ctx.deps.cfg.simulated_detector:
self.logger.info(f"initialise detector for raster at {grid.omega_deg}")
self.ctx.deps.jfjoch.measure_raster(grid, status)
self.logger.info("detector initialised")
else:
self.logger.info("Simulated detector mode enabled; using fake raster result.")
self.logger.info(f"initialise detector for raster at {grid.omega_deg}")
self.ctx.deps.jfjoch.measure_raster(grid, status)
self.logger.info("detector initialised")
if self.ctx.services.state is None:
raise RuntimeError("RasterService requires services.state")
+16
View File
@@ -846,6 +846,22 @@ async def local_contact_restart_device(device: str, token: str = Depends(oauth2_
return result
@app.get("/local_contact/get_gridscan_mode")
async def get_gridscan_mode(token: str = Depends(oauth2_scheme)) -> GridscanAnalysisMode:
"""
Toggle the TELL blower. Staff only.
Args:
token: OAuth2 access token.
Returns:
Dictionary with status and message.
"""
data = auth.parse_token(token)
auth.check_jwt_staff_only(data)
return cfg.get_gridscan_analysis_mode()
@app.post("/local_contact/resync/detector_metadata")
async def local_contact_resync_detector_metadata(token: str = Depends(oauth2_scheme)) -> dict:
"""