ewoksxes 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (40) hide show
  1. ewoksxes/__init__.py +0 -0
  2. ewoksxes/tasks/__init__.py +19 -0
  3. ewoksxes/tasks/calibrate_energy.py +69 -0
  4. ewoksxes/tasks/combine_spectra.py +117 -0
  5. ewoksxes/tasks/compute_flat_and_mask.py +176 -0
  6. ewoksxes/tasks/compute_roi.py +50 -0
  7. ewoksxes/tasks/fit_polynomial_2d.py +68 -0
  8. ewoksxes/tasks/flat_field_correction.py +102 -0
  9. ewoksxes/tasks/load_raw_data.py +66 -0
  10. ewoksxes/tasks/save_flat_and_mask.py +71 -0
  11. ewoksxes/tasks/save_spectrum.py +113 -0
  12. ewoksxes/tasks/utils.py +33 -0
  13. ewoksxes/tests/__init__.py +0 -0
  14. ewoksxes/tests/conftest.py +52 -0
  15. ewoksxes/tests/data/flat.edf +0 -0
  16. ewoksxes/tests/data/mask.npy +0 -0
  17. ewoksxes/tests/data/von_hamos_0000.h5 +0 -0
  18. ewoksxes/tests/test_calibrate_energy.py +137 -0
  19. ewoksxes/tests/test_combine_spectra.py +134 -0
  20. ewoksxes/tests/test_compute_flat_and_mask.py +55 -0
  21. ewoksxes/tests/test_compute_roi.py +90 -0
  22. ewoksxes/tests/test_fit_polynomial_2d.py +31 -0
  23. ewoksxes/tests/test_flat_field_correction.py +84 -0
  24. ewoksxes/tests/test_flat_field_workflow.py +107 -0
  25. ewoksxes/tests/test_load_raw_data.py +37 -0
  26. ewoksxes/tests/test_save_flat_and_mask.py +55 -0
  27. ewoksxes/tests/test_save_spectrum.py +141 -0
  28. ewoksxes/tests/test_utils.py +43 -0
  29. ewoksxes/tests/test_xes_calibration_workflow.py +115 -0
  30. ewoksxes/tests/test_xes_processing_workflow.py +99 -0
  31. ewoksxes/workflows/__init__.py +0 -0
  32. ewoksxes/workflows/xes_calibration.json +55 -0
  33. ewoksxes/workflows/xes_flat_field.json +63 -0
  34. ewoksxes/workflows/xes_processing.json +75 -0
  35. ewoksxes-0.0.1.dist-info/METADATA +116 -0
  36. ewoksxes-0.0.1.dist-info/RECORD +40 -0
  37. ewoksxes-0.0.1.dist-info/WHEEL +5 -0
  38. ewoksxes-0.0.1.dist-info/entry_points.txt +2 -0
  39. ewoksxes-0.0.1.dist-info/licenses/LICENSE.md +20 -0
  40. ewoksxes-0.0.1.dist-info/top_level.txt +1 -0
ewoksxes/__init__.py ADDED
File without changes
@@ -0,0 +1,19 @@
1
+ from .calibrate_energy import CalibrateEnergy
2
+ from .combine_spectra import CombineSpectra
3
+ from .compute_flat_and_mask import ComputeFlatAndMask
4
+ from .compute_roi import ComputeROI
5
+ from .fit_polynomial_2d import FitPolynomial2D
6
+ from .flat_field_correction import FlatFieldCorrection
7
+ from .load_raw_data import LoadRawDataAverage
8
+ from .save_flat_and_mask import SaveFlatAndMask
9
+
10
+ __all__ = [
11
+ "CalibrateEnergy",
12
+ "CombineSpectra",
13
+ "ComputeFlatAndMask",
14
+ "ComputeROI",
15
+ "FitPolynomial2D",
16
+ "FlatFieldCorrection",
17
+ "LoadRawDataAverage",
18
+ "SaveFlatAndMask",
19
+ ]
@@ -0,0 +1,69 @@
1
+ import logging
2
+
3
+ import numpy as np
4
+ from ewokscore import Task
5
+
6
+ logger = logging.getLogger(__name__)
7
+
8
+
9
+ class CalibrateEnergy(
10
+ Task,
11
+ input_names=["spectra", "kb_px", "vtc_px", "e_kb", "e_vtc"],
12
+ output_names=["energies", "spectra", "slope", "intercept"],
13
+ ):
14
+ """
15
+ Two-point linear calibration from pixel -> energy for each ROI.
16
+
17
+ Inputs
18
+ ------
19
+ spectra : list[np.ndarray]
20
+ One 1D spectrum per ROI.
21
+ kb_px : list[float] | np.ndarray
22
+ Pixel index of the Kβ reference for each ROI.
23
+ vtc_px : list[float] | np.ndarray
24
+ Pixel index of the VTC reference for each ROI.
25
+ e_kb : float
26
+ Known energy of the Kβ line (eV).
27
+ e_vtc : float
28
+ Known energy of the VTC line (eV).
29
+
30
+ Outputs
31
+ -------
32
+ energies : list[np.ndarray]
33
+ Energy axis per ROI (same length as its spectrum).
34
+ spectra : list[np.ndarray]
35
+ Pass-through of the input spectra (pipeline convenience).
36
+ slope : np.ndarray
37
+ ΔE / Δpixel per ROI.
38
+ intercept : np.ndarray
39
+ Intercept per ROI so that E = slope * x + intercept.
40
+ """
41
+
42
+ def run(self):
43
+ spectra = list(self.inputs.spectra)
44
+ kb = np.asarray(self.inputs.kb_px, dtype=float).ravel()
45
+ vtc = np.asarray(self.inputs.vtc_px, dtype=float).ravel()
46
+ e_kb = float(self.inputs.e_kb)
47
+ e_vtc = float(self.inputs.e_vtc)
48
+
49
+ if len(spectra) != kb.size or kb.size != vtc.size:
50
+ raise ValueError(
51
+ "Lengths must match: len(spectra) == len(kb_px) == len(vtc_px)"
52
+ )
53
+
54
+ dv = vtc - kb
55
+ if np.any(dv == 0.0):
56
+ raise ValueError("kb_px and vtc_px must differ for every ROI")
57
+
58
+ slope = (e_vtc - e_kb) / dv # shape (n_rois,)
59
+ intercept = e_kb - slope * kb # shape (n_rois,)
60
+
61
+ energies = []
62
+ for s, m, b in zip(spectra, slope, intercept):
63
+ x = np.arange(s.shape[0], dtype=float)
64
+ energies.append(m * x + b)
65
+
66
+ self.outputs.energies = energies
67
+ self.outputs.spectra = spectra
68
+ self.outputs.slope = slope
69
+ self.outputs.intercept = intercept
@@ -0,0 +1,117 @@
1
+ import logging
2
+
3
+ import numpy as np
4
+ from ewokscore import Task
5
+ from scipy.interpolate import interp1d
6
+
7
+ logger = logging.getLogger(__name__)
8
+
9
+
10
+ class CombineSpectra(
11
+ Task,
12
+ input_names=["energies", "spectra", "energy_range", "n_points"],
13
+ optional_input_names=["normalize", "norm_range"],
14
+ output_names=["energy", "summed_spectrum"],
15
+ ):
16
+ """
17
+ Optionally normalize spectra, then interpolate each onto a common energy axis
18
+ and sum them.
19
+
20
+ Inputs
21
+ ------
22
+ - energies: list of 1D energy arrays (or lists), one per spectrum
23
+ - spectra: list of 1D spectra (same length as energies)
24
+ - energy_range: (emin, emax) -> the range of the FINAL output energy axis
25
+ - n_points: int >= 2 -> number of points on the FINAL output energy axis
26
+
27
+ Optional
28
+ --------
29
+ - normalize: bool (default False). When True, each spectrum is normalized by
30
+ its integral over 'norm_range' before interpolation/summing.
31
+ - norm_range: (emin, emax) normalization window (required when normalize=True)
32
+
33
+ Outputs
34
+ -------
35
+ - energy: np.ndarray of length n_points, spanning [emin, emax]
36
+ - summed_spectrum: np.ndarray of same length, sum of all interpolated spectra
37
+ """
38
+
39
+ def run(self):
40
+ raw_energies = self.inputs.energies
41
+ raw_spectra = self.inputs.spectra
42
+ energy_range = self.inputs.energy_range
43
+ n_points = self.inputs.n_points
44
+
45
+ # Validate and USE energy_range as the FINAL axis range
46
+ try:
47
+ emin, emax = energy_range
48
+ except (TypeError, ValueError):
49
+ raise ValueError(f"energy_range must be (emin, emax); got {energy_range}")
50
+ if emin >= emax:
51
+ raise ValueError(f"Invalid energy_range: {energy_range}")
52
+ if not isinstance(n_points, int) or n_points < 2:
53
+ raise ValueError(f"n_points must be integer >= 2, got {n_points}")
54
+
55
+ # Convert to arrays
56
+ energies_list = [np.asarray(e, dtype=float) for e in raw_energies]
57
+ spectra = [np.asarray(s, dtype=float) for s in raw_spectra]
58
+
59
+ if len(energies_list) != len(spectra):
60
+ raise ValueError("energies_list and spectra must have the same length")
61
+
62
+ # Optional normalization
63
+ if bool(self.get_input_value("normalize", False)):
64
+ norm_range = self.get_input_value("norm_range", None)
65
+ if norm_range is None or len(norm_range) != 2:
66
+ raise ValueError(
67
+ "When normalize=True, norm_range=(emin, emax) must be "
68
+ f"provided; got {norm_range}"
69
+ )
70
+ nmin, nmax = float(norm_range[0]), float(norm_range[1])
71
+
72
+ for idx, (energies, spectrum) in enumerate(zip(energies_list, spectra)):
73
+ mask = (energies >= nmin) & (energies <= nmax)
74
+ if not np.any(mask):
75
+ logger.warning(
76
+ f"Spectrum {idx}: no points in norm_range {nmin}-{nmax}; "
77
+ "leaving spectrum unchanged"
78
+ )
79
+ continue
80
+ norm_factor = np.trapezoid(spectrum[mask], energies[mask])
81
+ if norm_factor == 0 or not np.isfinite(norm_factor):
82
+ logger.warning(
83
+ f"Spectrum {idx}: invalid normalization factor "
84
+ f"{norm_factor}; leaving spectrum unchanged"
85
+ )
86
+ continue
87
+ spectra[idx] = spectrum / norm_factor
88
+ logger.info("CombineSpectra: normalization completed.")
89
+
90
+ # Build FINAL energy axis USING energy_range
91
+ energy = np.linspace(emin, emax, n_points, dtype=float)
92
+ summed = np.zeros_like(energy)
93
+
94
+ # Interpolate & sum
95
+ for idx, (energies, spectrum) in enumerate(zip(energies_list, spectra)):
96
+ if energies.shape[0] != spectrum.shape[0]:
97
+ common = min(energies.shape[0], spectrum.shape[0])
98
+ logger.warning(
99
+ f"Spectrum {idx} length {spectrum.shape[0]} and energies "
100
+ f"{energies.shape[0]} mismatch; trimming to first {common} "
101
+ "points for interpolation."
102
+ )
103
+ energies_to_use = energies[:common]
104
+ spectrum_to_use = spectrum[:common]
105
+ else:
106
+ energies_to_use = energies
107
+ spectrum_to_use = spectrum
108
+
109
+ interp = interp1d(
110
+ energies_to_use, spectrum_to_use, bounds_error=False, fill_value=0.0
111
+ )
112
+ summed += interp(energy)
113
+ logger.info(f"CombineSpectra: added spectrum {idx} to sum.")
114
+
115
+ self.outputs.energy = energy
116
+ self.outputs.summed_spectrum = summed
117
+ logger.info("CombineSpectra completed.")
@@ -0,0 +1,176 @@
1
+ import logging
2
+
3
+ import numpy as np
4
+ from ewokscore import Task
5
+
6
+ from .utils import poly2d_eval
7
+
8
+ logger = logging.getLogger(__name__)
9
+
10
+
11
+ def _expand_columns(mask: np.ndarray, expand: int) -> np.ndarray:
12
+ """Binary dilate a column mask by ±expand columns (no extra deps)."""
13
+ if expand <= 0:
14
+ return mask
15
+ out = mask.copy()
16
+ for k in range(1, expand + 1):
17
+ out[:, k:] |= mask[:, :-k]
18
+ out[:, :-k] |= mask[:, k:]
19
+ return out
20
+
21
+
22
+ def _expand_rows(mask: np.ndarray, expand: int) -> np.ndarray:
23
+ """Binary dilate a row mask by ±expand rows (no extra deps)."""
24
+ if expand <= 0:
25
+ return mask
26
+ out = mask.copy()
27
+ for k in range(1, expand + 1):
28
+ out[k:, :] |= mask[:-k, :]
29
+ out[:-k, :] |= mask[k:, :]
30
+ return out
31
+
32
+
33
+ class ComputeFlatAndMask(
34
+ Task,
35
+ input_names=["image", "coeffs"],
36
+ optional_input_names=[
37
+ # base validity thresholds
38
+ "min_intensity", # discard pixels with raw <= this (default: 0.0)
39
+ "min_fitted", # discard where fitted surface <= this (default: 1e-6)
40
+ # gap detection via residual ratio
41
+ "ratio_threshold", # ratio=image/fitted below this => bad (default: 0.2)
42
+ # fraction of bad pixels in a column to flag gap (default: 0.8)
43
+ "column_fraction_threshold",
44
+ "row_fraction_threshold", # same for rows (default: 0.95)
45
+ "detect_columns", # default True
46
+ "detect_rows", # default False (usually gaps are vertical)
47
+ "gap_expand", # dilate gap bands by this many pixels (default: 1)
48
+ # flat output
49
+ "invert", # output flat = fitted/image if True (default True)
50
+ "clip_flat", # optional clip for flat magnitudes, e.g. (0.25, 4.0)
51
+ ],
52
+ output_names=["flat", "mask"],
53
+ ):
54
+ """
55
+ Compute flat-field and detect gap bands *automatically* from the 2D fit residuals.
56
+
57
+ Pipeline:
58
+ 1) Evaluate fitted surface S = poly2d(x,y; coeffs).
59
+ 2) Build a base invalid mask:
60
+ - image <= min_intensity OR ~finite OR S <= min_fitted
61
+ 3) Residual ratio R = image / S on valid base pixels; elsewhere treat as 0.
62
+ 4) Candidate bad pixels = base_invalid OR (R <= ratio_threshold).
63
+ 5) Gap detection (bands):
64
+ - per-column bad fraction; mark columns >= column_fraction_threshold as gaps.
65
+ - optional per-row bad fraction for horizontal bands.
66
+ - optional dilation by 'gap_expand' px.
67
+ 6) Final valid mask = NOT(gap_bands) AND NOT(base_invalid)
68
+ 7) Flat:
69
+ - if invert: flat = S / image on valid pixels (else flat = image / S)
70
+ - fill invalid pixels with 1.0 (so correction leaves them unchanged;
71
+ mask will zero them later).
72
+ - optional clipping.
73
+
74
+ Outputs:
75
+ - flat (float32, finite)
76
+ - mask (float32, 1 valid / 0 invalid). Gaps are 0 in the mask.
77
+ """
78
+
79
+ def run(self):
80
+ img = np.asarray(self.inputs.image, dtype=np.float64)
81
+ coeffs = np.asarray(self.inputs.coeffs, dtype=np.float64)
82
+
83
+ # --- Parameters & defaults
84
+ min_intensity = float(self.get_input_value("min_intensity", 0.0))
85
+ min_fitted = float(self.get_input_value("min_fitted", 1e-6))
86
+
87
+ ratio_thr = float(self.get_input_value("ratio_threshold", 0.2))
88
+ col_fr_thr = float(self.get_input_value("column_fraction_threshold", 0.8))
89
+ row_fr_thr = float(self.get_input_value("row_fraction_threshold", 0.95))
90
+ detect_columns = bool(self.get_input_value("detect_columns", True))
91
+ detect_rows = bool(self.get_input_value("detect_rows", False))
92
+ gap_expand = int(self.get_input_value("gap_expand", 1))
93
+
94
+ invert = bool(self.get_input_value("invert", True))
95
+ clip_flat = self.get_input_value("clip_flat", None) # e.g., (0.25, 4.0) or None
96
+
97
+ H, W = img.shape
98
+ y = np.arange(H, dtype=np.float64)
99
+ x = np.arange(W, dtype=np.float64)
100
+ X, Y = np.meshgrid(x, y)
101
+
102
+ # --- Evaluate fitted surface
103
+ fitted = poly2d_eval((X, Y), coeffs)
104
+
105
+ # --- Base invalid mask
106
+ base_invalid = (
107
+ (img <= min_intensity)
108
+ | ~np.isfinite(img)
109
+ | (fitted <= min_fitted)
110
+ | ~np.isfinite(fitted)
111
+ )
112
+
113
+ # --- Residual ratio
114
+ ratio = np.zeros_like(img, dtype=np.float64)
115
+ good = ~base_invalid
116
+ ratio[good] = img[good] / np.maximum(fitted[good], min_fitted)
117
+ candidate_bad = base_invalid | (ratio <= ratio_thr)
118
+
119
+ # --- Gap detection (bands)
120
+ gap_bands = np.zeros_like(candidate_bad, dtype=bool)
121
+
122
+ if detect_columns:
123
+ col_bad_frac = candidate_bad.mean(axis=0) # per-column fraction bad
124
+ gap_cols = col_bad_frac >= col_fr_thr
125
+ if gap_cols.any():
126
+ gap_bands[:, gap_cols] = True
127
+ logger.info(
128
+ "Detected %d gap columns (thr=%.2f)",
129
+ int(gap_cols.sum()),
130
+ col_fr_thr,
131
+ )
132
+
133
+ if detect_rows:
134
+ row_bad_frac = candidate_bad.mean(axis=1) # per-row fraction bad
135
+ gap_rows = row_bad_frac >= row_fr_thr
136
+ if gap_rows.any():
137
+ gap_bands[gap_rows, :] = True
138
+ logger.info(
139
+ "Detected %d gap rows (thr=%.2f)",
140
+ int(gap_rows.sum()),
141
+ row_fr_thr,
142
+ )
143
+
144
+ # --- Expand/dilate gap bands to cover edges
145
+ if detect_columns and gap_expand > 0:
146
+ gap_bands = _expand_columns(gap_bands, gap_expand)
147
+ if detect_rows and gap_expand > 0:
148
+ gap_bands = _expand_rows(gap_bands, gap_expand)
149
+
150
+ # --- Final valid mask
151
+ valid = (~gap_bands) & (~base_invalid)
152
+
153
+ # --- Flat computation
154
+ # invert=True => flat = fitted / image (typical: multiply raw by flat)
155
+ # invert=False => flat = image / fitted
156
+ eps = 1e-12
157
+ flat = np.ones_like(img, dtype=np.float64)
158
+ if invert:
159
+ flat[valid] = fitted[valid] / np.maximum(img[valid], eps)
160
+ else:
161
+ flat[valid] = img[valid] / np.maximum(fitted[valid], eps)
162
+
163
+ if clip_flat is not None:
164
+ lo, hi = clip_flat
165
+ flat = np.clip(flat, lo, hi)
166
+
167
+ # Fill invalids with 1.0 (so downstream (image*flat*mask) → masked zeros)
168
+ flat[~valid] = 1.0
169
+
170
+ self.outputs.flat = flat.astype(np.float32)
171
+ self.outputs.mask = valid.astype(np.float32)
172
+ logger.info(
173
+ "ComputeFlatAndMask: valid=%.2f%%, gaps=%.2f%%",
174
+ 100.0 * valid.mean(),
175
+ 100.0 * gap_bands.mean(),
176
+ )
@@ -0,0 +1,50 @@
1
+ import logging
2
+
3
+ from ewokscore import Task
4
+
5
+ logger = logging.getLogger(__name__)
6
+
7
+
8
+ class ComputeROI(
9
+ Task,
10
+ input_names=["image_corrected", "roi_list"],
11
+ output_names=["spectra"],
12
+ ):
13
+ """
14
+ Extract mean 1D spectra for each ROI.
15
+
16
+ roi_list can be:
17
+ - [(y1, y2), (y1, y2), ...] for full-width horizontal stripes
18
+ - [(y1, y2, x1, x2), ...] for rectangular ROIs
19
+
20
+ For each ROI:
21
+ roi = image_corrected[y1:y2, x1:x2]
22
+ spectra = mean over axis=0 (vertical collapse)
23
+
24
+ The result is a list of 1D spectra.
25
+ """
26
+
27
+ def run(self):
28
+ image = self.inputs.image_corrected
29
+ result_spectra = []
30
+
31
+ for roi in self.inputs.roi_list:
32
+ if len(roi) == 2:
33
+ # full-width stripe
34
+ y1, y2 = roi
35
+ x1, x2 = 0, image.shape[1]
36
+ elif len(roi) == 4:
37
+ # rectangle with x-limits
38
+ y1, y2, x1, x2 = roi
39
+ else:
40
+ raise ValueError(
41
+ f"ROI {roi} must have 2 values (y1,y2) or 4 values (y1,y2,x1,x2)"
42
+ )
43
+
44
+ logger.info(f"Computing spectra for ROI: y={y1}:{y2}, x={x1}:{x2}")
45
+ roi_data = image[y1:y2, x1:x2]
46
+ spectrum = roi_data.mean(axis=0)
47
+ result_spectra.append(spectrum)
48
+
49
+ self.outputs.spectra = result_spectra
50
+ logger.info(f"Extracted {len(result_spectra)} spectra")
@@ -0,0 +1,68 @@
1
+ import logging
2
+
3
+ import numpy as np
4
+ from ewokscore import Task
5
+ from scipy.optimize import curve_fit
6
+
7
+ from .utils import poly2d_eval
8
+
9
+ logger = logging.getLogger(__name__)
10
+
11
+
12
+ class FitPolynomial2D(
13
+ Task,
14
+ input_names=["image"],
15
+ optional_input_names=["degree", "mask_zeros"],
16
+ output_names=["coeffs"],
17
+ ):
18
+ """
19
+ Fit a 2D cubic polynomial (16 coefficients) to a flatfield-like image.
20
+
21
+ - Intended for XES von Hamos detectors (gapped columns).
22
+ - Zeros (dead/gap pixels) can be excluded from the fit.
23
+
24
+ Inputs
25
+ ------
26
+ image : 2D float array
27
+ degree : int (default=3) # only cubic supported
28
+ mask_zeros : bool (default=True)
29
+
30
+ Output
31
+ ------
32
+ coeffs : (16,) float64
33
+ """
34
+
35
+ def run(self):
36
+ img = np.asarray(self.inputs.image, dtype=np.float64)
37
+ degree = int(self.get_input_value("degree", 3))
38
+ mask_zeros = bool(self.get_input_value("mask_zeros", True))
39
+
40
+ if degree != 3:
41
+ raise NotImplementedError(
42
+ "Only cubic (degree=3) 2D polynomial fitting is supported"
43
+ )
44
+
45
+ H, W = img.shape
46
+ y = np.arange(H, dtype=np.float64)
47
+ x = np.arange(W, dtype=np.float64)
48
+ X, Y = np.meshgrid(x, y)
49
+
50
+ # Exclude obvious dead/gap pixels from the fit
51
+ mask = (img == 0.0) if mask_zeros else np.zeros_like(img, dtype=bool)
52
+
53
+ xdata = np.vstack((X.ravel()[~mask.ravel()], Y.ravel()[~mask.ravel()]))
54
+ zdata = img.ravel()[~mask.ravel()]
55
+
56
+ # Initial guess (all ones works fine for this basis)
57
+ p0 = np.ones(16, dtype=np.float64)
58
+
59
+ def model(xy, *params):
60
+ return poly2d_eval(xy, np.asarray(params))
61
+
62
+ logger.info(
63
+ "Fitting 2D cubic polynomial to flatfield (zeros masked: %s)...", mask_zeros
64
+ )
65
+ coeffs, _ = curve_fit(model, xdata, zdata, p0=p0, maxfev=200000)
66
+
67
+ self.outputs.coeffs = coeffs.astype(np.float64)
68
+ logger.info("FitPolynomial2D completed. Coeff count: %d", coeffs.size)
@@ -0,0 +1,102 @@
1
+ import logging
2
+ import os
3
+
4
+ import fabio
5
+ import numpy as np
6
+ from ewokscore import Task
7
+
8
+ logger = logging.getLogger(__name__)
9
+
10
+
11
+ class FlatFieldCorrection(
12
+ Task,
13
+ input_names=["image", "flat_field_path"],
14
+ optional_input_names=["mask_path", "i0"],
15
+ output_names=["image_corrected"],
16
+ ):
17
+ """
18
+ Apply flat-field and (optionally) mask corrections to the raw image,
19
+ then normalize by I0.
20
+
21
+ Supported file formats:
22
+ - Flat-field: .edf, .tif, .tiff, .npy
23
+ - Mask: .edf, .tif, .tiff, .msk, .npy
24
+
25
+ If no mask is provided, a mask of ones is used.
26
+
27
+ Correction:
28
+ corrected = ((image * flat) * mask) / i0
29
+
30
+ i0 is a scalar (incident flux) used to normalize the result.
31
+ If omitted, defaults to 1.
32
+ """
33
+
34
+ SUPPORTED_EXTENSIONS_FLAT = (".edf", ".tif", ".tiff", ".npy")
35
+ SUPPORTED_EXTENSIONS_MASK = (".edf", ".tif", ".tiff", ".msk", ".npy")
36
+
37
+ def _load_image_file(self, path: str, label: str, allowed_exts) -> np.ndarray:
38
+ """Helper to load image files with validation."""
39
+ if not os.path.exists(path):
40
+ raise FileNotFoundError(f"{label} file does not exist: {path}")
41
+
42
+ ext = os.path.splitext(path)[1].lower()
43
+ if ext not in allowed_exts:
44
+ raise ValueError(
45
+ f"{label} file format not supported ({path}). "
46
+ f"Supported extensions: {allowed_exts}"
47
+ )
48
+
49
+ logger.info(f"Loading {label} file: {path}")
50
+ if ext == ".npy":
51
+ data = np.load(path).astype(np.float32)
52
+ else:
53
+ data = fabio.open(path).data.astype(np.float32)
54
+
55
+ logger.info(f"{label} shape: {data.shape}")
56
+ return data
57
+
58
+ def run(self):
59
+ raw_image = self.inputs.image
60
+ logger.info(
61
+ f"Starting flat-field correction, raw image shape: {raw_image.shape}"
62
+ )
63
+
64
+ # Load flat-field
65
+ flat = self._load_image_file(
66
+ self.inputs.flat_field_path, "Flat-field", self.SUPPORTED_EXTENSIONS_FLAT
67
+ )
68
+
69
+ # Load mask if provided
70
+ mask_path = self.get_input_value("mask_path", None)
71
+ if mask_path:
72
+ mask = self._load_image_file(
73
+ mask_path, "Mask", self.SUPPORTED_EXTENSIONS_MASK
74
+ )
75
+ else:
76
+ logger.info("No mask file provided; using an all-ones mask.")
77
+ mask = np.ones_like(raw_image, dtype=np.float32)
78
+
79
+ # Validate shapes
80
+ if flat.shape != raw_image.shape:
81
+ raise ValueError(
82
+ f"Flat-field shape {flat.shape} does not match "
83
+ f"raw image shape {raw_image.shape}"
84
+ )
85
+ if mask.shape != raw_image.shape:
86
+ raise ValueError(
87
+ f"Mask shape {mask.shape} does not match "
88
+ f"raw image shape {raw_image.shape}"
89
+ )
90
+
91
+ # Apply correction
92
+ corrected = (raw_image * flat) * mask
93
+
94
+ # I0 normalization
95
+ i0_value = self.get_input_value("i0", 1.0)
96
+ if i0_value != 1.0:
97
+ logger.info(f"Normalizing by I0 = {i0_value}")
98
+ corrected = corrected / i0_value
99
+
100
+ self.outputs.image_corrected = corrected
101
+ logger.info("Flat-field correction completed successfully.")
102
+ logger.info(f"Corrected image shape: {corrected.shape}")
@@ -0,0 +1,66 @@
1
+ import logging
2
+ import os
3
+
4
+ import h5py
5
+ import numpy as np
6
+ from ewokscore import Task
7
+
8
+ logger = logging.getLogger(__name__)
9
+
10
+
11
+ class LoadRawDataAverage(
12
+ Task,
13
+ input_names=["bliss_scan_data_url"],
14
+ output_names=["image"],
15
+ ):
16
+ """
17
+ Load raw detector data from a Bliss scan data URL.
18
+
19
+ The input `bliss_scan_data_url` must be of the form:
20
+
21
+ /full/path/to/sample_dataset.h5::scan.1/measurement/detector
22
+
23
+ Where:
24
+ - `/full/path/to/sample_dataset.h5` is the HDF5 file.
25
+ - `scan.1/measurement/detector` is the internal HDF5 dataset path.
26
+
27
+ The task loads the dataset. If the dataset is 3D (multiple frames),
28
+ it averages along the first dimension to produce a 2D image.
29
+ """
30
+
31
+ def run(self):
32
+ # Split the URL into file path and dataset path
33
+ url = self.inputs.bliss_scan_data_url
34
+ if "::" not in url:
35
+ raise ValueError(
36
+ f"Invalid bliss_scan_data_url format: {url}. "
37
+ "Expected file.h5::dataset/path"
38
+ )
39
+
40
+ file_path, dataset_path = url.split("::", 1)
41
+
42
+ logger.info(f"Opening HDF5 file: {file_path}")
43
+ logger.info(f"Reading dataset: {dataset_path}")
44
+
45
+ if not os.path.exists(file_path):
46
+ raise FileNotFoundError(f"HDF5 file does not exist: {file_path}")
47
+
48
+ with h5py.File(file_path, "r") as f:
49
+ if dataset_path not in f:
50
+ raise KeyError(f"Dataset '{dataset_path}' not found in {file_path}")
51
+ data = f[dataset_path][:]
52
+
53
+ logger.info(f"Raw data shape: {data.shape}")
54
+
55
+ if data.ndim == 3:
56
+ logger.info("Averaging along the first axis of the 3D dataset")
57
+ image = data.mean(axis=0)
58
+ elif data.ndim == 2:
59
+ image = data
60
+ else:
61
+ raise ValueError(
62
+ f"Unexpected data shape {data.shape}. Expected 2D or 3D array."
63
+ )
64
+
65
+ self.outputs.image = image.astype(np.float32)
66
+ logger.info(f"Final image shape: {self.outputs.image.shape}")