mirror of
https://github.com/bec-project/ophyd_devices.git
synced 2026-09-06 21:00:56 +02:00
wip ad_roi_signal
This commit is contained in:
@@ -1,39 +1,90 @@
|
||||
"""Custom ROI processing Device for AreaDetector ROI/Stats plugin processing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
from functools import partial
|
||||
from time import time
|
||||
from typing import Any, Literal
|
||||
|
||||
import numpy as np
|
||||
from bec_lib.devicemanager import ScanInfo
|
||||
from bec_lib.logger import bec_logger
|
||||
from bec_lib.utils.rpc_utils import rgetattr
|
||||
from bec_server.scan_server.scans.scan_base import ScanInfo as ScanServerScanInfo
|
||||
from ophyd import Component as Cpt
|
||||
from ophyd import Device
|
||||
from ophyd import EpicsSignal, EpicsSignalRO, Kind, Signal
|
||||
from pydantic import ValidationError
|
||||
|
||||
from ophyd_devices.devices.areadetector.plugins import ROIPlugin_V35, StatsPlugin_V35
|
||||
from ophyd_devices import TransitionStatus
|
||||
from ophyd_devices.devices.areadetector.plugins import ROIPlugin_V35 as ROIPlugin
|
||||
from ophyd_devices.devices.areadetector.plugins import StatsPlugin_V35
|
||||
from ophyd_devices.utils.bec_roi_signals.roi_processing import (
|
||||
LITERAL_ROI_PROCESSING_CONFIG,
|
||||
ROIProcessing,
|
||||
)
|
||||
|
||||
|
||||
class StatsPlugin(StatsPlugin_V35):
|
||||
pass
|
||||
# plugin_type = None
|
||||
# codec = None
|
||||
# compressed_size = None
|
||||
"""Utility functions for the devices."""
|
||||
|
||||
|
||||
class ROIPlugin(ROIPlugin_V35):
|
||||
pass
|
||||
# plugin_type = None
|
||||
# codec = None
|
||||
# compressed_size = None
|
||||
def fetch_scan_info(scan_info: ScanInfo) -> ScanServerScanInfo:
|
||||
"""Fetch the scan parameters from the scan_info object and return them as a ScanServerScanInfo object."""
|
||||
info = scan_info.msg.info
|
||||
if isinstance(info["positions"], list):
|
||||
info["positions"] = np.array(info["positions"])
|
||||
info["num_monitored_readouts"] = scan_info.msg.num_monitored_readouts
|
||||
try:
|
||||
msg = ScanServerScanInfo.model_validate(info)
|
||||
except ValidationError: # This means we have an old scan_info object.
|
||||
info = deepcopy(info)
|
||||
# We need to convert a few parameters manually.
|
||||
info["scan_type"] = (
|
||||
"hardware_triggered" if info["scan_type"] == "fly" else "software_triggered"
|
||||
)
|
||||
msg = ScanServerScanInfo.model_validate(info)
|
||||
|
||||
return msg
|
||||
|
||||
|
||||
logger = bec_logger.logger
|
||||
|
||||
|
||||
class TSAcquireMode(IntEnum):
|
||||
"""Enum for TSAcquireMode values."""
|
||||
|
||||
FIXED_LENGTH = 0
|
||||
CIRCULAR_BUFFER = 1
|
||||
|
||||
|
||||
class TSReadMode(IntEnum):
|
||||
"""Enum for TSReadMode values. Recommend to use PASSIVE mode for most applications."""
|
||||
|
||||
PASSIVE = 0
|
||||
EVENT = 1
|
||||
IO_INTR = 2
|
||||
TEN_SECOND = 3
|
||||
FIVE_SECOND = 4
|
||||
TWO_SECOND = 5
|
||||
ONE_SECOND = 6
|
||||
HALF_SECOND = 7
|
||||
TWO_TENTHS_SECOND = 8
|
||||
ONE_TENTH_SECOND = 9
|
||||
|
||||
|
||||
class StatsPluginWithTSControl(StatsPlugin_V35):
|
||||
"""StatsPlugin with additional timestamp control signals."""
|
||||
|
||||
ts_acquire_status = Cpt(EpicsSignalRO, ":TS:TSAcquiring", kind=Kind.omitted, auto_monitor=True)
|
||||
ts_acquire = Cpt(EpicsSignal, "TS:TSAcquire", kind=Kind.omitted)
|
||||
ts_acquire_mode = Cpt(EpicsSignal, "TS:TSAcquireMode", kind=Kind.omitted)
|
||||
ts_read_mode = Cpt(EpicsSignal, "TS:TSRead.SCAN", kind=Kind.omitted)
|
||||
ts_read = Cpt(EpicsSignal, "TS:TSRead.PROC", kind=Kind.omitted)
|
||||
ts_current_index = Cpt(EpicsSignalRO, "TS:TSCurrentPoint", kind=Kind.omitted)
|
||||
ts_num_points = Cpt(EpicsSignalRO, "TS:TSNumPoints", kind=Kind.omitted, auto_monitor=True)
|
||||
|
||||
|
||||
@dataclass
|
||||
class StatsSubscription:
|
||||
"""Bookkeeping for callbacks subscribed to StatsPlugin output signals."""
|
||||
@@ -49,35 +100,35 @@ NDPLUGIN_STATS_CONFIG: LITERAL_ROI_PROCESSING_CONFIG = {
|
||||
"basic_statistics": {
|
||||
"enable_signal": "compute_statistics",
|
||||
"scalar_outputs": [
|
||||
"min",
|
||||
"min_value",
|
||||
"min_x",
|
||||
"min_y",
|
||||
"max",
|
||||
"max_value",
|
||||
"max_x",
|
||||
"max_y",
|
||||
"mean",
|
||||
"mean_value",
|
||||
"sigma",
|
||||
"total",
|
||||
"net",
|
||||
"sigma",
|
||||
],
|
||||
"waveform_outputs": [],
|
||||
"source_signals": {
|
||||
"min": "min_value",
|
||||
"min_x": "min_xy.x",
|
||||
"min_y": "min_xy.y",
|
||||
"max": "max_value",
|
||||
"max_x": "max_xy.x",
|
||||
"max_y": "max_xy.y",
|
||||
"mean": "mean_value",
|
||||
"total": "total",
|
||||
"net": "net",
|
||||
"sigma": "sigma_readout",
|
||||
"min_value": "ts_min_value",
|
||||
"min_x": "ts_min.x",
|
||||
"min_y": "ts_min.y",
|
||||
"max_value": "ts_max_value",
|
||||
"max_x": "ts_max.x",
|
||||
"max_y": "ts_max.y",
|
||||
"mean_value": "ts_mean_value",
|
||||
"sigma": "ts_sigma",
|
||||
"total": "ts_total",
|
||||
"net": "ts_net",
|
||||
"timestamp": "ts_timestamp",
|
||||
},
|
||||
},
|
||||
"centroid": {
|
||||
"enable_signal": "compute_centroid",
|
||||
"scalar_outputs": [
|
||||
"centroid_total",
|
||||
"centroid_x",
|
||||
"centroid_y",
|
||||
"sigma_x",
|
||||
@@ -92,74 +143,48 @@ NDPLUGIN_STATS_CONFIG: LITERAL_ROI_PROCESSING_CONFIG = {
|
||||
],
|
||||
"waveform_outputs": [],
|
||||
"source_signals": {
|
||||
"centroid_total": "centroid_total",
|
||||
"centroid_x": "centroid.x",
|
||||
"centroid_y": "centroid.y",
|
||||
"sigma_x": "sigma_x",
|
||||
"sigma_y": "sigma_y",
|
||||
"sigma_xy": "sigma_xy",
|
||||
"skew_x": "skew.x",
|
||||
"skew_y": "skew.y",
|
||||
"kurtosis_x": "kurtosis.x",
|
||||
"kurtosis_y": "kurtosis.y",
|
||||
"eccentricity": "eccentricity",
|
||||
"orientation": "orientation",
|
||||
},
|
||||
},
|
||||
"profiles": {
|
||||
"enable_signal": "compute_profiles",
|
||||
"scalar_outputs": ["profile_size_x", "profile_size_y", "cursor_x", "cursor_y"],
|
||||
"waveform_outputs": [
|
||||
"profile_average_x",
|
||||
"profile_average_y",
|
||||
"profile_threshold_x",
|
||||
"profile_threshold_y",
|
||||
"profile_centroid_x",
|
||||
"profile_centroid_y",
|
||||
"profile_cursor_x",
|
||||
"profile_cursor_y",
|
||||
],
|
||||
"source_signals": {
|
||||
"profile_size_x": "profile_size.x",
|
||||
"profile_size_y": "profile_size.y",
|
||||
"cursor_x": "cursor.x",
|
||||
"cursor_y": "cursor.y",
|
||||
"profile_average_x": "profile_average.x",
|
||||
"profile_average_y": "profile_average.y",
|
||||
"profile_threshold_x": "profile_threshold.x",
|
||||
"profile_threshold_y": "profile_threshold.y",
|
||||
"profile_centroid_x": "profile_centroid.x",
|
||||
"profile_centroid_y": "profile_centroid.y",
|
||||
"profile_cursor_x": "profile_cursor.x",
|
||||
"profile_cursor_y": "profile_cursor.y",
|
||||
},
|
||||
},
|
||||
"histogram": {
|
||||
"enable_signal": "compute_histogram",
|
||||
"scalar_outputs": ["hist_below", "hist_above", "hist_entropy"],
|
||||
"waveform_outputs": ["histogram", "histogram_x"],
|
||||
"source_signals": {
|
||||
"hist_below": "hist_below",
|
||||
"hist_above": "hist_above",
|
||||
"hist_entropy": "hist_entropy",
|
||||
"histogram": "histogram",
|
||||
"histogram_x": "histogram_x",
|
||||
"centroid_x": "ts_centroid.x",
|
||||
"centroid_y": "ts_centroid.y",
|
||||
"sigma_x": "ts_sigma_x.ts_sigma_x",
|
||||
"sigma_y": "ts_sigma_x.ts_sigma_y",
|
||||
"sigma_xy": "ts_sigma_xy",
|
||||
"skew_x": "ts_skew.x",
|
||||
"skew_y": "ts_skew.y",
|
||||
"kurtosis_x": "ts_kurtosis.x",
|
||||
"kurtosis_y": "ts_kurtosis.y",
|
||||
"eccentricity": "ts_eccentricity",
|
||||
"orientation": "ts_orientation",
|
||||
"timestamp": "ts_timestamp",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class AverageFramesForEachTrigger(Signal):
|
||||
"""Signal to compute the average number of frames per trigger."""
|
||||
|
||||
def put(self, value, **kwargs):
|
||||
if not isinstance(value, (bool, int, float)):
|
||||
raise ValueError("average_frames_per_trigger must be a boolean or numeric value.")
|
||||
value = bool(value)
|
||||
super().put(value, **kwargs)
|
||||
|
||||
|
||||
class ADROIProcessing(ROIProcessing):
|
||||
"""ROI processing signal for AD detector setups at PSI."""
|
||||
|
||||
roi1 = Cpt(ROIPlugin, "ROI1:", kind="normal")
|
||||
stats1 = Cpt(StatsPlugin, "Stats1:", kind="normal")
|
||||
stats1 = Cpt(StatsPluginWithTSControl, "Stats1:", kind="normal")
|
||||
average_frames_per_trigger = Cpt(
|
||||
AverageFramesForEachTrigger, "average_frames_per_trigger", kind=Kind.config, value=True
|
||||
)
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._stats_subscriptions: dict[str, StatsSubscription] = {}
|
||||
self._missing_stats_paths: set[str] = set()
|
||||
self._config_update_in_progress = False
|
||||
self.scan_server_scan_info: ScanServerScanInfo | None = None
|
||||
|
||||
def get_scalar_outputs(self) -> list[str]:
|
||||
scalar_outputs = []
|
||||
@@ -179,7 +204,7 @@ class ADROIProcessing(ROIProcessing):
|
||||
def wait_for_connection(self, *args, **kwargs):
|
||||
super().wait_for_connection(*args, **kwargs)
|
||||
# Subscribe stats plugin to ROI plugin, the self.root.cam assumes cam to be available on root object, so we trust roi1 to be properly configured.
|
||||
# self.roi1.nd_array_port.put(self.root.cam.port_name.get())
|
||||
self.roi1.nd_array_port.put(self.root.cam.port_name.get())
|
||||
self.stats1.nd_array_port.put(self.roi1.port_name.get())
|
||||
# ROIPlugin related callbacks
|
||||
self.selected_operations.subscribe(
|
||||
@@ -190,10 +215,6 @@ class ADROIProcessing(ROIProcessing):
|
||||
self._apply_roi_configuration()
|
||||
self._sync_stats_subscriptions()
|
||||
|
||||
def destroy(self):
|
||||
self._unsubscribe_all_stats()
|
||||
super().destroy()
|
||||
|
||||
def _on_config_update(self, value, **kwargs):
|
||||
self._config_update_in_progress = True
|
||||
try:
|
||||
@@ -216,6 +237,36 @@ class ADROIProcessing(ROIProcessing):
|
||||
self.roi1.size.x.put(int(self.width.get()))
|
||||
self.roi1.size.y.put(int(self.height.get()))
|
||||
|
||||
def _desired_stats_outputs(
|
||||
self,
|
||||
) -> dict[str, tuple[str, str, Literal["scalar", "waveform"], str]]:
|
||||
"""Determine the desired StatsPlugin outputs based on the selected operations."""
|
||||
desired = {}
|
||||
for operation in self.selected_operations.get():
|
||||
config = NDPLUGIN_STATS_CONFIG.get(operation)
|
||||
if not config:
|
||||
continue
|
||||
source_signals = config.get("source_signals", {})
|
||||
for result_name in config.get("scalar_outputs", []):
|
||||
dotted_name = source_signals.get(result_name)
|
||||
if dotted_name:
|
||||
key = f"{operation}.{result_name}"
|
||||
desired[key] = (operation, result_name, "scalar", dotted_name)
|
||||
for result_name in config.get("waveform_outputs", []):
|
||||
dotted_name = source_signals.get(result_name)
|
||||
if dotted_name:
|
||||
key = f"{operation}.{result_name}"
|
||||
desired[key] = (operation, result_name, "waveform", dotted_name)
|
||||
return desired
|
||||
|
||||
def _resolve_stats_signal(self, dotted_name: str) -> Any:
|
||||
"""Resolve a dotted name to an actual signal on the stats1 plugin."""
|
||||
signal = rgetattr(self.stats1, dotted_name, None)
|
||||
if signal is None:
|
||||
logger.warning(f"Could not resolve StatsPlugin signal: {dotted_name}")
|
||||
self._missing_stats_paths.add(dotted_name)
|
||||
return signal
|
||||
|
||||
def _sync_stats_subscriptions(self) -> None:
|
||||
"""Subscribe to selected StatsPlugin outputs and unsubscribe stale ones."""
|
||||
desired = self._desired_stats_outputs()
|
||||
@@ -253,36 +304,6 @@ class ADROIProcessing(ROIProcessing):
|
||||
|
||||
self._apply_stats_enable_signals()
|
||||
|
||||
def _desired_stats_outputs(
|
||||
self,
|
||||
) -> dict[str, tuple[str, str, Literal["scalar", "waveform"], str]]:
|
||||
if not self.active.get():
|
||||
return {}
|
||||
|
||||
selected_operations = set(self.selected_operations.get())
|
||||
desired = {}
|
||||
for operation, config in NDPLUGIN_STATS_CONFIG.items():
|
||||
if operation not in selected_operations:
|
||||
continue
|
||||
source_signals = config.get("source_signals", {})
|
||||
for result_name in config.get("scalar_outputs", []):
|
||||
signal_path = source_signals[result_name]
|
||||
desired[f"scalar:{operation}:{result_name}"] = (
|
||||
operation,
|
||||
result_name,
|
||||
"scalar",
|
||||
signal_path,
|
||||
)
|
||||
for result_name in config.get("waveform_outputs", []):
|
||||
signal_path = source_signals[result_name]
|
||||
desired[f"waveform:{operation}:{result_name}"] = (
|
||||
operation,
|
||||
result_name,
|
||||
"waveform",
|
||||
signal_path,
|
||||
)
|
||||
return desired
|
||||
|
||||
def _unsubscribe_all_stats(self) -> None:
|
||||
for subscription in self._stats_subscriptions.values():
|
||||
subscription.signal.unsubscribe(subscription.callback_id)
|
||||
@@ -294,26 +315,11 @@ class ADROIProcessing(ROIProcessing):
|
||||
enable_signal = config.get("enable_signal")
|
||||
if enable_signal is None:
|
||||
continue
|
||||
signal = self._resolve_stats_signal(enable_signal)
|
||||
signal = rgetattr(self.stats1, enable_signal, None) # Check if the attribute exists
|
||||
if signal is None:
|
||||
continue
|
||||
signal.put("Yes" if operation in selected_operations else "No")
|
||||
|
||||
def _resolve_stats_signal(self, signal_path: str):
|
||||
"""Resolve a dotted StatsPlugin attribute path to an ophyd signal."""
|
||||
try:
|
||||
signal = rgetattr(self.stats1, signal_path) # Check if the attribute exists
|
||||
except AttributeError:
|
||||
if signal_path not in self._missing_stats_paths:
|
||||
logger.warning(
|
||||
"StatsPlugin signal path %s is not available on %s.",
|
||||
signal_path,
|
||||
self.stats1.__class__.__name__,
|
||||
)
|
||||
self._missing_stats_paths.add(signal_path)
|
||||
return None
|
||||
return signal
|
||||
|
||||
def _on_stats_signal_update(
|
||||
self,
|
||||
*,
|
||||
@@ -325,16 +331,58 @@ class ADROIProcessing(ROIProcessing):
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""Publish a StatsPlugin update into the matching BEC result signal."""
|
||||
logger.info(f"StatsPlugin update for {operation}.{result_name} ({output_kind}): {value}")
|
||||
if not self._is_operation_active(operation):
|
||||
return
|
||||
if self.average_frames_per_trigger.get() is True:
|
||||
if isinstance(value, (list, np.ndarray)):
|
||||
value = value / len(value) # Average over the number of frames per trigger
|
||||
value = float(value) # Ensure the value is a float for list and np.ndarray types
|
||||
|
||||
signal = self.result_scalar if output_kind == "scalar" else self.result_waveform
|
||||
signal.put({result_name: {"value": value, "timestamp": timestamp or self._get_timestamp()}})
|
||||
signal.put({result_name: {"value": value, "timestamp": timestamp or time.time()}})
|
||||
|
||||
def _is_operation_active(self, operation: str) -> bool:
|
||||
return bool(self.active.get()) and operation in self.selected_operations.get()
|
||||
|
||||
################
|
||||
## Scan Hooks ##
|
||||
################
|
||||
|
||||
def on_connected(self):
|
||||
"""Hook called when the device is connected."""
|
||||
# TODO Consider using 'set' in future after wrapping EpicsSignal with proper set wrapper
|
||||
self.stats1.ts_acquire_mode.put(TSAcquireMode.FIXED_LENGTH.value)
|
||||
self.stats1.ts_read_mode.put(TSReadMode.PASSIVE.value)
|
||||
if self.stats1.ts_acquire_status.get() == 0:
|
||||
self.stats1.ts_acquire.put(0) # Start timestamp acquisition
|
||||
|
||||
def on_stage(self):
|
||||
"""Hook called when the device is staged."""
|
||||
self.root: PSIDeviceBase
|
||||
self.scan_server_scan_info = fetch_scan_info(self.root.scan_info)
|
||||
self.stats1.ts_num_points.put(self.scan_server_scan_info.frames_per_trigger)
|
||||
|
||||
def on_trigger(self) -> TransitionStatus:
|
||||
"""Hook called when the device is triggered."""
|
||||
# TODO add hook, needs to be triggered before the detector acquire starts if it is triggered manually.
|
||||
status = TransitionStatus(self.stats1.ts_acquire_status, transitions=[1, 0])
|
||||
self.stats1.ts_acquire.put(1) # Start timestamp acquisition
|
||||
return status
|
||||
|
||||
def on_stop(self):
|
||||
"""Hook called when the device is stopped."""
|
||||
self.stats1.ts_acquire.put(0) # Stop timestamp acquisition
|
||||
|
||||
def on_destroy(self):
|
||||
"""Hook called when the device is destroyed."""
|
||||
self.stats1.ts_acquire.put(0) # Stop timestamp acquisition
|
||||
self._unsubscribe_all_stats()
|
||||
|
||||
|
||||
#####################
|
||||
### Test Detector ###
|
||||
#####################
|
||||
|
||||
|
||||
import threading
|
||||
import traceback
|
||||
@@ -426,8 +474,21 @@ class MyDetector(PSIDeviceBase, ADBase):
|
||||
f"Error while polling array data for preview of {self.name}: {content}"
|
||||
)
|
||||
|
||||
def on_stage(self):
|
||||
"""Stage the detector and prepare for acquisition."""
|
||||
self.roi_processing.on_stage() # Stage the ROI processing
|
||||
|
||||
def on_trigger(self):
|
||||
"""Trigger the detector and wait for completion."""
|
||||
self.roi_processing.on_trigger() # Start timestamp acquisition
|
||||
|
||||
def on_stop(self):
|
||||
"""Stop the detector and timestamp acquisition."""
|
||||
self.roi_processing.on_stop() # Stop timestamp acquisition
|
||||
|
||||
def on_destroy(self):
|
||||
"""Clean up resources."""
|
||||
self.roi_processing.on_destroy() # Stop timestamp acquisition
|
||||
self._poll_thread_kill_event.set()
|
||||
if self._poll_thread.is_alive():
|
||||
self._poll_thread.join(timeout=0.5)
|
||||
|
||||
Reference in New Issue
Block a user