Files
camserver_sf/configuration/user_scripts/swissfel_spectral_processing.py
T
2026-08-13 17:01:22 +02:00

619 lines
27 KiB
Python

"""
Unified online spectral processing for SwissFEL PSSS/PMOS cameras.
Features:
- Dynamic Y-ROI via {camera}:SPC_ROI_YMIN / SPC_ROI_YMAX PVs
- Per-shot Gaussian fit + spectral moments (COM, RMS, skewness, IQR)
- Rolling averages (RAVG) for all scalar result PVs
- N-shot averaged spectrum with independent fit and moments, refit every
avg_fit_every_n_shots shots (config, default 10) rather than every shot -
it's the same fit cost as the main one and the running average barely
changes shot to shot; AVG-* PVs simply hold their last value in between
- Type-safe EPICS PV updates via _pv_safe()
- Robust fallbacks: axis PV -> x_axis argument, background arg or parameters key
- Optional gated photon-energy feedback: reuses the single-shot fit's centre
and FWHM above (no separate fit) behind internal/external validity checks;
held at last good value on a bad shot. Set "feedback": {"enabled": false}
in config to keep external_checks etc. around without paying the (small)
per-shot cost of evaluating them.
The feedback gate used to run its own separate fit (heavier smoothing) as
its own separate "stream" pipeline (swissfel_spectral_feedback.py) reading
SPECTRUM_X/Y back via bsread. Both of those were reverted/simplified after
measuring the actual cost with pv_update_timing.py: the separate stream
pipeline's bsread dispatch delivered updates in bursts rather than steadily
(this image pipeline runs directly off the camera frame rate, no such
dispatch layer), and the separate fit was a second full Gaussian fit per
shot, found to be a meaningful contributor to this pipeline dropping ~30%
of frames at the camera's ~100 Hz rate (processing_threads couldn't keep up).
"""
from logging import getLogger
from cam_server.pipeline.data_processing import functions
from cam_server.utils import create_thread_pvs, epics_lock
from collections import deque
import json
import numpy as np
import scipy.signal
import numba
import time
from threading import Thread
numba.set_num_threads(4)
_logger = getLogger(__name__)
# ---------------------------------------------------------------------------
# Shared state
# ---------------------------------------------------------------------------
global_roi = [0, 0] # [ymin, ymax] updated each frame from PVs
initialized = False
sent_pid = -1
buffer = deque(maxlen=5)
channel_pv_names = [] # polled each frame: [ymin, ymax, axis]
base_pv_names = [] # per-shot result PVs (ordered, for update thread)
all_pv_names = [] # base + RAVG + AVG (matches update thread tuple)
ARRAY_PVS = set() # PV names that carry waveform arrays
global_ravg_length = 100
ravg_buffers = {}
global_avg_length = 100
avg_buffer = None
avg_pv_names = []
# -- Gated feedback output: reuses the single-shot fit, held on bad shots
feedback_enabled = False
feedback_output_pv = None
feedback_valid_pv = None
feedback_internal_checks = [] # [{"metric": ..., "min": ..., "max": ..., "equals": ...}]
feedback_external_checks = [] # [{"pv": ..., "min": ..., "max": ..., "equals": ...}]
feedback_external_pv_names = []
# -- N-shot averaged-spectrum fit: only redone every avg_fit_every_n_shots
# (it's the same fit cost as the main one, and the running average barely
# changes shot to shot, so refitting every single shot is wasted work)
avg_fit_every_n_shots = 10
# ---------------------------------------------------------------------------
# Numba-accelerated spectrum extraction
# ---------------------------------------------------------------------------
@numba.njit(parallel=False)
def get_spectrum(image, background):
"""Row-wise sum of (image - background), returns 1-D float64 profile."""
y, x = image.shape
profile = np.zeros(x, dtype=np.float64)
for i in range(y):
for j in range(x):
profile[j] += image[i, j] - background[i, j]
return profile
# ---------------------------------------------------------------------------
# Background PV-update thread
# ---------------------------------------------------------------------------
def update_PVs(buf, *pv_names):
"""
Daemon thread: pops (value, ...) tuples from *buf* and writes each
value to the corresponding PV in *pv_names*.
"""
pvs = create_thread_pvs(list(pv_names))
while True:
time.sleep(0.1)
try:
rec = buf.popleft()
except IndexError:
continue
for pv, val in zip(pvs, rec):
if pv and pv.connected and val is not None:
try:
pv.put(val)
except Exception as e:
_logger.error(f"PV write error {pv.pvname}: {e}")
# ---------------------------------------------------------------------------
# Type-safe value coercion before EPICS put
# ---------------------------------------------------------------------------
def _pv_safe(val, pvname):
"""
Coerce *val* to the correct Python type for *pvname*:
- array PVs -> 1-D np.float64 array
- scalar PVs -> Python float
Raises TypeError / ValueError on mismatches to catch bugs early.
Returns None for None input (filtered by update_PVs).
"""
if val is None:
return None
if pvname in ARRAY_PVS:
if isinstance(val, np.ndarray):
if val.ndim == 1:
return val.astype(np.float64)
elif val.ndim == 0:
return np.array([val.item()], dtype=np.float64)
else:
raise ValueError(f"{pvname}: expected 1-D array, got shape {val.shape}")
elif isinstance(val, (list, tuple)):
return np.array(val, dtype=np.float64)
elif isinstance(val, (float, int, np.generic)):
return np.array([float(val)], dtype=np.float64)
else:
raise TypeError(f"{pvname}: expected array, got {type(val)}")
else:
if isinstance(val, np.ndarray):
if val.size == 1:
return float(val.item())
else:
raise ValueError(f"{pvname}: expected scalar, got array size {val.size}")
elif isinstance(val, (np.generic, float, int)):
return float(val)
else:
raise TypeError(f"{pvname}: expected scalar, got {type(val)}")
# ---------------------------------------------------------------------------
# Feedback validity check
# ---------------------------------------------------------------------------
def _check_passes(value, check):
"""
Evaluate one {"min":, "max":, "equals":} entry against *value*.
None never passes (disconnected PV / missing metric -> fail closed).
"""
if value is None:
return False
if "equals" in check:
return value == check["equals"]
if check.get("min") is not None and value < check["min"]:
return False
if check.get("max") is not None and value > check["max"]:
return False
return True
# ---------------------------------------------------------------------------
# One-time initialisation
# ---------------------------------------------------------------------------
def initialize(params):
"""
Build all PV name lists, initialise rolling-average and N-shot buffers,
and start the background update thread.
"""
global channel_pv_names, base_pv_names, all_pv_names, ARRAY_PVS
global global_ravg_length, global_avg_length, avg_buffer, avg_pv_names
global feedback_enabled, feedback_output_pv, feedback_valid_pv
global feedback_internal_checks, feedback_external_checks, feedback_external_pv_names
global avg_fit_every_n_shots
camera = params["camera_name"]
e_int = params["e_int_name"]
e_axis = params["e_axis_name"]
avg_fit_every_n_shots = params.get("avg_fit_every_n_shots", avg_fit_every_n_shots)
# -- Gated feedback (optional "feedback" block in config; set
# "enabled": false to keep the block - external_checks etc. - around
# without paying its per-shot cost) ----------------------------------
fb_cfg = params.get("feedback")
feedback_enabled = bool(fb_cfg) and fb_cfg.get("enabled", True)
if feedback_enabled:
feedback_output_pv = f"{camera}:{fb_cfg['output_pv']}"
feedback_valid_pv = f"{camera}:{fb_cfg['valid_pv']}"
feedback_internal_checks = fb_cfg.get("internal_checks", [])
feedback_external_checks = fb_cfg.get("external_checks", [])
feedback_external_pv_names = [c["pv"] for c in feedback_external_checks]
# -- PVs polled every frame (ROI + energy axis + feedback external checks)
ymin_pv = f"{camera}:SPC_ROI_YMIN"
ymax_pv = f"{camera}:SPC_ROI_YMAX"
channel_pv_names = [ymin_pv, ymax_pv, e_axis] + feedback_external_pv_names
# -- Per-shot result PVs (order defines the update-thread tuple) ---------
base_pv_names = [
e_int, # full spectrum waveform
f"{camera}:FIT-COM",
f"{camera}:FIT-FWHM",
f"{camera}:FIT-RMS",
f"{camera}:FIT-RES",
f"{camera}:FIT-SPECTRUM_Y", # fitted Gaussian waveform
f"{camera}:SPECT-COM",
f"{camera}:SPECT-RMS",
f"{camera}:SPECT-SKEW",
f"{camera}:SPECT-IQR",
f"{camera}:SPECT-RES",
]
if feedback_enabled:
base_pv_names = base_pv_names + [feedback_output_pv, feedback_valid_pv]
# -- Rolling-average PVs (all scalars except e_int, e_axis, params) -----
global_ravg_length = params.get("RAVG_length", global_ravg_length)
_ravg_exclude = {e_int, e_axis, f"{camera}:processing_parameters"}
if feedback_enabled:
_ravg_exclude |= {feedback_output_pv, feedback_valid_pv}
ravg_pv_names = [
f"{pv}-RAVG"
for pv in base_pv_names
if pv not in _ravg_exclude
]
# -- N-shot averaged spectrum PVs ----------------------------------------
global_avg_length = params.get("avg_nshots", global_avg_length)
avg_buffer = deque(maxlen=global_avg_length)
avg_pv_names = [
f"{camera}:AVG-FIT-COM",
f"{camera}:AVG-FIT-FWHM",
f"{camera}:AVG-FIT-RMS",
f"{camera}:AVG-FIT-RES",
f"{camera}:AVG-SPECT-COM",
f"{camera}:AVG-SPECT-RMS",
f"{camera}:AVG-SPECT-SKEW",
f"{camera}:AVG-SPECT-IQR",
f"{camera}:AVG-SPECT-RES",
f"{camera}:AVG-SPECTRUM_Y",
f"{camera}:AVG-FIT-SPECTRUM_Y",
]
# -- Full ordered list for the update thread -----------------------------
all_pv_names = base_pv_names + ravg_pv_names + avg_pv_names
# -- Array PV registry (used by _pv_safe) --------------------------------
ARRAY_PVS = {
e_int,
f"{camera}:FIT-SPECTRUM_Y",
f"{camera}:FIT-SPECTRUM_Y-RAVG",
f"{camera}:AVG-SPECTRUM_Y",
f"{camera}:AVG-FIT-SPECTRUM_Y",
}
# -- Launch background thread --------------------------------------------
t = Thread(target=update_PVs, args=(buffer, *all_pv_names), daemon=True)
t.start()
_logger.info(f"Spectral processing initialised for {camera} "
f"(RAVG={global_ravg_length}, AVG={global_avg_length}, "
f"feedback={'on:' + feedback_output_pv if feedback_enabled else 'off'})")
# ---------------------------------------------------------------------------
# Savitzky-Golay helper (guards against short spectra)
# ---------------------------------------------------------------------------
def _smooth(spectrum, polyorder=3, preferred_window=51):
n = len(spectrum)
max_odd = max(polyorder + 2, (n // 2) * 2 - 1) # largest valid odd window
win = min(preferred_window, max_odd)
if win > polyorder:
try:
return scipy.signal.savgol_filter(spectrum, window_length=win,
polyorder=polyorder, mode="interp")
except Exception:
pass
return spectrum.copy()
# ---------------------------------------------------------------------------
# Main processing entry point
# ---------------------------------------------------------------------------
def process_image(image, pulse_id, timestamp, x_axis, y_axis, parameters,
bsdata=None, background=None):
"""
Per-pulse entry point called by cam_server.
Steps
-----
1. First-call initialisation.
2. Read dynamic Y-ROI from EPICS PVs (SPC_ROI_YMIN / SPC_ROI_YMAX).
3. Resolve energy axis (PV preferred, x_axis argument as fallback).
4. Subtract pixel background, apply Y-ROI, extract spectrum.
5. Smooth, Gaussian fit, spectral moments (COM, RMS, skewness, IQR).
6. Compute rolling averages (RAVG) for all scalar PVs.
7. Compute N-shot averaged spectrum with independent fit and moments.
8. Queue EPICS update; return full result dict.
"""
global initialized, sent_pid, global_roi
global ravg_buffers, avg_buffer
try:
# ------------------------------------------------------------------ #
# 1. One-time init #
# ------------------------------------------------------------------ #
if not initialized:
initialize(parameters)
initialized = True
camera = parameters["camera_name"]
# ------------------------------------------------------------------ #
# 2. Dynamic Y-ROI from PVs #
# SATOP31-PMOS132-2D:SPC_ROI_YMIN / SPC_ROI_YMAX #
# ------------------------------------------------------------------ #
ymin_pv, ymax_pv, axis_pv, *ext_pv_objs = create_thread_pvs(channel_pv_names)
if ymin_pv and ymin_pv.connected:
try:
global_roi[0] = int(ymin_pv.value)
except Exception:
pass
if ymax_pv and ymax_pv.connected:
try:
global_roi[1] = int(ymax_pv.value)
except Exception:
pass
# ------------------------------------------------------------------ #
# 3. Energy axis: PV first, x_axis argument as fallback #
# ------------------------------------------------------------------ #
axis = None
if axis_pv and axis_pv.connected and axis_pv.value is not None:
try:
axis = np.asarray(axis_pv.value, dtype=np.float64)
except Exception:
axis = None
if axis is None and x_axis is not None:
try:
axis = np.asarray(x_axis, dtype=np.float64)
_logger.debug("Energy axis PV unavailable - using x_axis argument")
except Exception:
axis = None
if axis is None:
_logger.warning("No valid energy axis available; skipping pulse")
return {}
# ------------------------------------------------------------------ #
# 4. Image pre-processing and ROI crop #
# ------------------------------------------------------------------ #
proc_img = image.astype(np.float32) - np.float32(parameters.get("pixel_bkg", 0.0))
nrows, ncols = proc_img.shape
# Trim axis to image width
axis = axis[:ncols]
if len(axis) < ncols:
_logger.warning(f"Energy axis shorter than image width ({len(axis)} < {ncols})")
return {}
# Resolve background image (parameters key wins over function argument)
bg_param = parameters.pop("background_data", None)
if isinstance(bg_param, np.ndarray) and bg_param.shape == proc_img.shape:
bg_img = bg_param.astype(np.float32)
elif isinstance(background, np.ndarray) and background.shape == proc_img.shape:
bg_img = background.astype(np.float32)
else:
bg_img = None
# Apply Y-ROI (clamp to image bounds)
ymin, ymax = global_roi
if (isinstance(ymin, (int, np.integer)) and
isinstance(ymax, (int, np.integer)) and
0 <= ymin and ymax > ymin):
ymin = max(0, min(ymin, nrows - 1))
ymax = max(ymin + 1, min(ymax, nrows))
img_roi = proc_img[ymin:ymax, :]
bg_roi = bg_img[ymin:ymax, :] if bg_img is not None else None
_logger.debug(f"Applying Y-ROI [{ymin}:{ymax}] (of {nrows} rows)")
else:
img_roi = proc_img
bg_roi = bg_img
if ymin != 0 or ymax != 0:
_logger.warning(f"Y-ROI [{ymin}:{ymax}] invalid; using full image")
# Spectrum extraction
if bg_roi is not None:
try:
spectrum = get_spectrum(img_roi, bg_roi)
except Exception as e:
_logger.warning(f"get_spectrum failed, falling back to sum: {e}")
spectrum = np.sum(img_roi - bg_roi, axis=0).astype(np.float64)
else:
spectrum = np.sum(img_roi, axis=0).astype(np.float64)
n_roi_rows = img_roi.shape[0]
# Full-image spectrum (all rows, no ROI) for diagnostics
if bg_img is not None:
spectrum_all = np.sum(proc_img - bg_img, axis=0).astype(np.float64)
else:
spectrum_all = np.sum(proc_img, axis=0).astype(np.float64)
# ------------------------------------------------------------------ #
# 5. Per-shot: smooth -> Gaussian fit -> spectral moments #
# ------------------------------------------------------------------ #
smoothed = _smooth(spectrum)
s_min, s_max = float(smoothed.min()), float(smoothed.max())
amplitude = s_max - s_min
skip = amplitude <= n_roi_rows * 1.5
fit_ok = not skip
try:
offset, amp_fit, center, sigma = functions.gauss_fit_psss(
smoothed[::2], axis[::2],
offset=s_min, amplitude=amplitude, skip=skip, maxfev=10
)
center = float(center)
sigma = float(abs(sigma)) # sigma must be positive
except Exception:
# Fall back to weighted moments
total = float(np.sum(smoothed))
center = float(np.sum(axis * smoothed) / total) if total else np.nan
sigma = float(np.sqrt(np.sum((axis - center)**2 * smoothed) / total)) if total else np.nan
offset, amp_fit = s_min, amplitude
fit_spectrum = offset + amp_fit * np.exp(-((axis - center)**2) / (2 * sigma**2))
# Normalised moments of smoothed spectrum
sm_total = np.sum(smoothed)
sm_norm = smoothed / sm_total if sm_total else smoothed
spect_com = float(np.sum(axis * sm_norm))
spect_std = float(np.sqrt(np.sum((axis - spect_com)**2 * sm_norm)))
spect_skew = (float(np.sum((axis - spect_com)**3 * sm_norm) / spect_std**3)
if spect_std else np.nan)
cum = np.cumsum(sm_norm)
e25 = float(np.interp(0.25, cum, axis))
e75 = float(np.interp(0.75, cum, axis))
spect_iqr = e75 - e25
# ------------------------------------------------------------------ #
# 5a. Gated feedback: reuses the single-shot fit above (center, #
# sigma/FWHM) - no separate fit - + validity checks, held at #
# last good value on a bad shot #
# ------------------------------------------------------------------ #
if feedback_enabled:
# A fit centre outside the measured energy axis is an
# extrapolation, never a real photon energy - always reject it.
in_range = bool(np.isfinite(center) and len(axis)
and axis[0] <= center <= axis[-1])
fb_metrics = {
"_FIT_OK": fit_ok,
"_CENTER": center,
"_SIGMA": sigma,
"_FWHM": 2.355 * sigma,
"_AMPLITUDE": amplitude,
}
fb_valid = in_range
for chk in feedback_internal_checks:
if not _check_passes(fb_metrics.get(chk["metric"]), chk):
fb_valid = False
break
if fb_valid and feedback_external_checks:
for pv_obj, chk in zip(ext_pv_objs, feedback_external_checks):
pv_val = pv_obj.value if (pv_obj and pv_obj.connected) else None
if not _check_passes(pv_val, chk):
fb_valid = False
break
# ------------------------------------------------------------------ #
# 6. Rolling averages (RAVG) for all scalar base PVs #
# ------------------------------------------------------------------ #
_ravg_exclude = {parameters["e_int_name"], parameters["e_axis_name"],
f"{camera}:processing_parameters"}
if feedback_enabled:
_ravg_exclude |= {feedback_output_pv, feedback_valid_pv}
per_shot_results = {
parameters["e_int_name"]: spectrum,
f"{camera}:FIT-COM": center,
f"{camera}:FIT-FWHM": 2.355 * sigma,
f"{camera}:FIT-RMS": sigma,
f"{camera}:FIT-RES": (2.355 * sigma / center * 1000) if center else np.nan,
f"{camera}:FIT-SPECTRUM_Y": fit_spectrum,
f"{camera}:SPECT-COM": spect_com,
f"{camera}:SPECT-RMS": spect_std,
f"{camera}:SPECT-SKEW": spect_skew,
f"{camera}:SPECT-IQR": spect_iqr,
f"{camera}:SPECT-RES": (spect_iqr / spect_com * 1000) if spect_com else np.nan,
}
if feedback_enabled:
per_shot_results[feedback_valid_pv] = 1.0 if fb_valid else 0.0
per_shot_results[feedback_output_pv] = center if fb_valid else None
ravg_results = {}
for pv in base_pv_names:
if pv in _ravg_exclude:
continue
buf = ravg_buffers.setdefault(pv, deque(maxlen=global_ravg_length))
val = per_shot_results.get(pv)
if pv in ARRAY_PVS:
# Array PV (e.g. FIT-SPECTRUM_Y): element-wise rolling average
if isinstance(val, np.ndarray) and val.ndim == 1:
buf.append(val.astype(np.float64))
if buf:
ravg_results[f"{pv}-RAVG"] = np.mean(np.stack(buf), axis=0)
else:
# Scalar PV: simple rolling mean
if val is not None and np.isscalar(val) and np.isfinite(val):
buf.append(val)
if buf:
ravg_results[f"{pv}-RAVG"] = float(np.mean(buf))
# ------------------------------------------------------------------ #
# 7. N-shot averaged spectrum: fit + moments, only every #
# avg_fit_every_n_shots shots (same fit cost as the main one, #
# and the running average barely changes shot to shot - the #
# AVG-* PVs simply hold their last value on skipped shots) #
# ------------------------------------------------------------------ #
avg_buffer.append(spectrum)
if pulse_id % avg_fit_every_n_shots == 0:
avg_spectrum = np.mean(np.stack(avg_buffer), axis=0)
sm_avg = _smooth(avg_spectrum)
a_min, a_max = float(sm_avg.min()), float(sm_avg.max())
amp_a = a_max - a_min
skip_a = amp_a <= n_roi_rows * 1.5
try:
offs_a, amp_fit_a, center_a, sigma_a = functions.gauss_fit_psss(
sm_avg[::2], axis[::2],
offset=a_min, amplitude=amp_a, skip=skip_a, maxfev=10
)
center_a = float(center_a)
sigma_a = float(abs(sigma_a))
except Exception:
total_a = float(np.sum(sm_avg))
center_a = float(np.sum(axis * sm_avg) / total_a) if total_a else np.nan
sigma_a = float(np.sqrt(np.sum((axis - center_a)**2 * sm_avg) / total_a)) if total_a else np.nan
offs_a, amp_fit_a = a_min, amp_a
fit_avg_spectrum = np.abs(
offs_a + amp_fit_a * np.exp(-((axis - center_a)**2) / (2 * sigma_a**2))
)
sn_total = np.sum(sm_avg)
sn_norm = sm_avg / sn_total if sn_total else sm_avg
com_a = float(np.sum(axis * sn_norm))
std_a = float(np.sqrt(np.sum((axis - com_a)**2 * sn_norm)))
skew_a = (float(np.sum((axis - com_a)**3 * sn_norm) / std_a**3)
if std_a else np.nan)
cum_a = np.cumsum(sn_norm)
iqr_a = float(np.interp(0.75, cum_a, axis) - np.interp(0.25, cum_a, axis))
res_a = (iqr_a / com_a * 1000) if com_a else np.nan
avg_results = {
f"{camera}:AVG-FIT-COM": center_a,
f"{camera}:AVG-FIT-FWHM": 2.355 * sigma_a,
f"{camera}:AVG-FIT-RMS": sigma_a,
f"{camera}:AVG-FIT-RES": (2.355 * sigma_a / center_a * 1000) if center_a else np.nan,
f"{camera}:AVG-SPECT-COM": com_a,
f"{camera}:AVG-SPECT-RMS": std_a,
f"{camera}:AVG-SPECT-SKEW": skew_a,
f"{camera}:AVG-SPECT-IQR": iqr_a,
f"{camera}:AVG-SPECT-RES": res_a,
f"{camera}:AVG-SPECTRUM_Y": avg_spectrum,
f"{camera}:AVG-FIT-SPECTRUM_Y": fit_avg_spectrum,
}
else:
avg_results = {}
# ------------------------------------------------------------------ #
# 8. Merge, queue EPICS update, return #
# ------------------------------------------------------------------ #
# Extras not pushed to EPICS but useful for downstream panels
extra = {
parameters["e_axis_name"]: axis,
f"{camera}:SPECTRUM_Y_SUM": float(np.sum(spectrum)),
f"{camera}:SPECTRUM_Y_ALL": spectrum_all,
f"{camera}:processing_parameters": json.dumps({"roi": global_roi}),
}
full_results = {**per_shot_results, **ravg_results, **avg_results, **extra}
if epics_lock.acquire(False):
try:
if pulse_id > sent_pid:
sent_pid = pulse_id
entry = tuple(
_pv_safe(full_results.get(pv), pv)
for pv in all_pv_names
)
buffer.append(entry)
finally:
epics_lock.release()
return full_results
except Exception as ex:
_logger.exception(f"process_image error: {ex}")
return {}