diff --git a/docs/polymer_overlap_reconciliation.md b/docs/polymer_overlap_reconciliation.md new file mode 100644 index 000000000..b5021df39 --- /dev/null +++ b/docs/polymer_overlap_reconciliation.md @@ -0,0 +1,34 @@ +# Polymer / `dev` reconciliation ledger + +Audit base: `origin/dev` at `dca541308cc6fa500b677cade5365969ab0db552`. +Polymer source was inspected from the preserved `polymers` branch and the exact +pre-maintenance backup branch. The curated branch is `paper/polymers`. + +| Polymer path / symbol | Current `dev` counterpart | Relationship and behavior/API differences | Callers / tests | Resolution | Risk | +|---|---|---|---|---|---| +| `core/ml/cnn2d.py:CNN2d`, `MultiChannelCNN2d` | `core/ml/cnn.py:CNN2d` | Similar U-Net purpose, but not checkpoint-equivalent. Defaults differ for skip connections and dropout placement; the paper network also has configurable convolution kernel size and per-output activations. The archived training runner confirms the paper checkpoint used dropout `0.0` (the historical inference constructor's `0.2` default is incompatible). These differences change parameter layout and numerical output. | Polymer notebooks, archived training runner, `BraggPeaksPolymer`; new resolver/numerical tests. | Retain the exact architecture privately in `diffraction/polymer_models.py`, pin dropout `0.0`, and continue using `dev` CNN for all non-polymer callers. | High if consolidated; paper weights cannot safely load into the `dev` class. | +| Polymer `core/utils/augment_dp.py` and backup | `core/utils/augment_dp.py:DPAugmentor` | Same purpose. `dev` has the maintained RNG/device-aware implementation. Polymer variants include historical experiments and a large backup. The immutable release archive preserves the exact training augmentor separately. | Generator/training archive; existing core augmentor tests. | Reuse `dev`; do not transplant either polymer copy. | Low for inference; retraining reproducibility depends on the private archive. | +| Polymer `core/io/file_readers.py:read_4dstem` | `core/io/file_readers.py:read_4dstem` | Conflicting extensions: polymer reshapes selected 3D frame stacks and optionally transposes scan axes; `dev` has maintained 4D loading, metadata overrides, and hot-pixel filtering. Scan-axis interpretation is acquisition-specific. | Tutorials load a canonical 4D scan; core reader tests exercise `dev`. | Reuse `dev`; defer the 3D-stack feature to an independent reader PR with acquisition fixtures. | Medium: silently choosing a scan axis can transpose real-space coordinates. | +| Polymer `origin_finding.py` and `polar4dstem.find_origin` | `diffractive_imaging/origin_models.py:CenterOfMassOriginModel` | Complementary algorithms. `dev` estimates center of mass for ptychography; polymer angular-uniformity search minimizes polar angular variation and returns row/column pixel origins. | Polymer polar workflow and focused origin/polar tests; ptychography callers use COM model. | Retain angular origin finding in `diffraction/polar_transform.py`; leave COM API unchanged. | Medium: algorithms are not interchangeable for masked/anisotropic patterns. | +| `diffraction/peak_detection.py` | No general Bragg-peak detector on `dev` | Polymer code provides strict local maxima, quadratic subpixel refinement, peak pairing, and central-beam selection in row/column convention. | `BraggPeaksPolymer`; numerical coordinate regression. | Retain, initially scoped to diffraction. | Medium: coordinate order must remain `(row, col)` internally and `(x, y)` only at plotting boundaries. | +| `diffraction/polar_transform.py` | `diffractive_imaging/complex_probe.py` polar coordinate helpers and `_torch_polar` in direct ptychography | Complementary. `dev` helpers build frequency grids or convert tensor components; polymer code resamples whole 4D scans about per-pattern origins, optionally corrects ellipses, and defines explicit angular folding. | `BraggPeaksPolymer`, origin workflow, numerical orientation regression. | Retain the Torch-native scan transform plus minimal `Polar4dstem`; do not expose the experimental `polar.py` / `polar_new.py` duplicates. | High: angle direction, row/column origin order, and whether Friedel partners are sampled or summed affect flowline orientation and intensity. | +| Polymer `core/utils/utils.py:parse_reciprocal_units` | `core/utils/utils.py:electron_wavelength_angstrom`; calibrated dataset metadata | Complementary. Polymer parsing converts reciprocal nm to reciprocal Å; angular sampling requires voltage-dependent wavelength conversion. The historical parser accepted ambiguous substrings. | `BraggPeaksPolymer.pixels_to_inv_A`; unit regression. | Keep strict parser in `diffraction/polymer_utils.py`, reuse `dev` electron wavelength function, warn on unknown units, and preserve the documented 300 kV compatibility default. | High: a factor-of-ten unit error changes every reported radial position. | +| Polymer `sample_average_from_image` and broad probe-fit helpers | `dev` generic array utilities / filtering | Only local polar-neighborhood averaging is required. Probe circle/ellipse fitting and broad utility additions are unrelated. | Polymer peak intensity extraction. | Retain the local averaging helper only; do not expand generic utilities. | Low; angular axis wraps while radial axis does not. | +| Polymer `Dataset`, `Dataset4dstem`, `polar4dstem` changes | Current `Dataset*` classes | Most polymer dataset edits are experimental or debug-only (including a constructor print). `Polar4dstem` is a genuinely distinct `(scan_y, scan_x, phi, r)` calibrated container. | Polar transform; existing dataset tests. | Reuse current datasets unchanged and add only `Polar4dstem`. | Medium: polar axes are `(phi, r)`, unlike Cartesian `(qy, qx)`. | +| `BraggPeaksPolymer` normalization and BatchNorm adaptation | No `dev` equivalent | New capability. It caches scan-level percentiles, supports ROI masks, either adapts BatchNorm running statistics for deterministic eval or retains train-batch behavior, and performs masked inference. | Paper tutorial and forthcoming GPU comparison. | Retain. Named model specs carry normalization metadata; legacy caller-supplied normalization functions remain supported. | High: changing normalization or BN mode changes peak counts. | +| `BraggPeaksPolymer` visualization, count maps, peak figure export, flowlines | General `core.visualization.show_2d` | Complementary. Polymer methods create domain-specific overlays and flowline orientation/color composites; generic display remains useful underneath. | Paper figure workflow. | Retain domain-specific methods and reuse `show_2d`. | Medium: orientation convention and cyclic color mapping are scientific outputs. | +| `polar.py`, `polar_new.py`, `polymer_analytical_functions.py`, Kirkland table | Torch `polar_transform.py` or no paper inference caller | Multiple experimental/analytical paths, not required by the selected model inference/tutorial path. | Historical notebooks only. | Defer; preserve in checkpoint and private archive, omit from PR. | Low for the paper workflow; revisit as separate physics APIs. | +| Polymer generator/training launchers, notebooks, results, workspace files, grain clustering | No relevant `dev` API | Out of scope or private/research state. | Private archive and independent WIP backup. | Exclude from public polymer PR. | None to inference; disclosure/repository hygiene risk if included. | + +## Final disposition summary + +- Reused from `dev`: maintained CNNs for non-polymer callers, augmentation, file readers, + generic datasets/utilities, electron wavelength conversion, and visualization. +- Retained separately: exact checkpoint architecture, `BraggPeaksPolymer`, peak detection, + angular origin/polar transforms, `Polar4dstem`, strict reciprocal-unit conversion, and + polar neighborhood averaging. +- Deferred: acquisition-specific 3D-stack reshaping and all experimental analytical polar + implementations. +- Required follow-up before converting the draft PR to ready: compare pinned-model output + tensors, peak coordinates/counts, reciprocal radii, and flowline orientation against the + archived paper environment on the experimental scan. diff --git a/src/quantem/__init__.py b/src/quantem/__init__.py index ba70f629f..db9d38c33 100644 --- a/src/quantem/__init__.py +++ b/src/quantem/__init__.py @@ -10,5 +10,6 @@ from quantem import imaging as imaging from quantem import diffractive_imaging as diffractive_imaging +from quantem import diffraction as diffraction __version__ = version("quantem") diff --git a/src/quantem/core/datastructures/__init__.py b/src/quantem/core/datastructures/__init__.py index dfb5b47ac..c149f811b 100644 --- a/src/quantem/core/datastructures/__init__.py +++ b/src/quantem/core/datastructures/__init__.py @@ -5,3 +5,4 @@ from quantem.core.datastructures.dataset4d import Dataset4d as Dataset4d from quantem.core.datastructures.dataset3d import Dataset3d as Dataset3d from quantem.core.datastructures.dataset2d import Dataset2d as Dataset2d +from quantem.core.datastructures.polar4dstem import Polar4dstem as Polar4dstem diff --git a/src/quantem/core/datastructures/polar4dstem.py b/src/quantem/core/datastructures/polar4dstem.py new file mode 100644 index 000000000..e832e26d3 --- /dev/null +++ b/src/quantem/core/datastructures/polar4dstem.py @@ -0,0 +1,408 @@ +from typing import TYPE_CHECKING, Any + +import numpy as np +from numpy.typing import NDArray +from scipy.ndimage import map_coordinates + +if TYPE_CHECKING: + from .dataset4dstem import Dataset4dstem + +from quantem.core.datastructures.dataset4d import Dataset4d + + +class Polar4dstem(Dataset4d): + """4D-STEM dataset in polar coordinates (scan_y, scan_x, phi, r).""" + + def __init__( + self, + array: NDArray | Any, + name: str, + origin: NDArray | tuple | list | float | int, + sampling: NDArray | tuple | list | float | int, + units: list[str] | tuple | list, + signal_units: str = "arb. units", + metadata: dict | None = None, + _token: object | None = None, + ): + if metadata is None: + metadata = {} + mdata_keys_polar = [ + "polar_radial_min", + "polar_radial_max", + "polar_radial_step", + "polar_num_annular_bins", + "polar_two_fold_rotation_symmetry", + "polar_origin_row", + "polar_origin_col", + "polar_ellipse_params", + ] + for k in mdata_keys_polar: + if k not in metadata: + metadata[k] = None + super().__init__( + array=array, + name=name, + origin=origin, + sampling=sampling, + units=units, + signal_units=signal_units, + metadata=metadata, + _token=_token, + ) + + @classmethod + def from_array( + cls, + array: NDArray | Any, + name: str | None = None, + origin: NDArray | tuple | list | float | int | None = None, + sampling: NDArray | tuple | list | float | int | None = None, + units: list[str] | tuple | list | None = None, + signal_units: str = "arb. units", + metadata: dict | None = None, + ) -> "Polar4dstem": + array = np.asarray(array) + if array.ndim != 4: + raise ValueError("Polar4dstem.from_array expects a 4D array.") + if origin is None: + origin = np.zeros(4, dtype=float) + if sampling is None: + sampling = np.ones(4, dtype=float) + if units is None: + units = ["pixels", "pixels", "deg", "pixels"] + if metadata is None: + metadata = {} + return cls( + array=array, + name=name if name is not None else "Polar 4D-STEM dataset", + origin=origin, + sampling=sampling, + units=units, + signal_units=signal_units, + metadata=metadata, + _token=cls._token, + ) + + @property + def n_phi(self) -> int: + return int(self.array.shape[2]) + + @property + def n_r(self) -> int: + return int(self.array.shape[3]) + + +def _precompute_polar_coords( + ny: int, + nx: int, + origin_row: float, + origin_col: float, + ellipse_params: tuple[float, float, float] | None, + num_annular_bins: int, + radial_min: float, + radial_max: float | None, + radial_step: float, + two_fold_rotation_symmetry: bool, +) -> tuple[NDArray, NDArray, NDArray, float]: + origin_row = float(origin_row) + origin_col = float(origin_col) + if radial_step <= 0: + raise ValueError("radial_step must be > 0.") + if num_annular_bins < 1: + raise ValueError("num_annular_bins must be >= 1.") + if radial_max is None: + r_row_pos = origin_row + r_row_neg = (ny - 1) - origin_row + r_col_pos = origin_col + r_col_neg = (nx - 1) - origin_col + radial_max_eff = float(min(r_row_pos, r_row_neg, r_col_pos, r_col_neg)) + else: + radial_max_eff = float(radial_max) + if radial_max_eff <= radial_min: + radial_max_eff = radial_min + radial_step + radial_bins = np.arange(radial_min, radial_max_eff, radial_step, dtype=np.float64) + if radial_bins.size == 0: + radial_bins = np.array([radial_min], dtype=np.float64) + if two_fold_rotation_symmetry: + phi_range = np.pi + else: + phi_range = 2.0 * np.pi + phi_bins = np.linspace(0.0, phi_range, num_annular_bins, endpoint=False, dtype=np.float64) + phi_grid, r_grid = np.meshgrid(phi_bins, radial_bins, indexing="ij") + if ellipse_params is None: + x = r_grid * np.cos(phi_grid) + y = r_grid * np.sin(phi_grid) + else: + if len(ellipse_params) != 3: + raise ValueError("ellipse_params must be (a, b, theta_deg).") + a, b, theta_deg = ellipse_params + theta = np.deg2rad(theta_deg) + alpha = phi_grid - theta + u = (a / b) * r_grid * np.cos(alpha) + v_prime = r_grid * np.sin(alpha) + cos_t = np.cos(theta) + sin_t = np.sin(theta) + x = u * cos_t - v_prime * sin_t + y = u * sin_t + v_prime * cos_t + coords_y = y + origin_row + coords_x = x + origin_col + coords = np.stack((coords_y, coords_x), axis=0) + return coords, phi_bins, radial_bins, radial_max_eff + + +def find_origin( + data, + *, + ellipse_params=None, + num_annular_bins=180, + radial_min=0.0, + radial_max=None, + radial_step=1.0, + two_fold_rotation_symmetry=False, +): + """ + Placeholder for future automatic diffraction center finding method. + """ + if len(data.array.shape) == 2: + ny, nx = data.array.shape + scan_y, scan_x = 1, 1 + elif len(data.array.shape) == 4: + scan_y, scan_x, ny, nx = data.array.shape + else: + raise ValueError("find_origin only supports 2D or 4D-STEM datasets for now.") + + origin_array = np.zeros((scan_y, scan_x, 2), dtype=float) + + max_steps = 1000 # prevent infinite loops + + # start with center of image for now + estimated_origin_row = (ny - 1) / 2.0 + estimated_origin_col = (nx - 1) / 2.0 + + for y_pos in range(scan_y): + for x_pos in range(scan_x): + print(f"Finding origin for scan pos ({y_pos}, {x_pos})") + + coords_cache = {} + + polar = data.polar_transform( + origin_array=[estimated_origin_row, estimated_origin_col], + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + two_fold_rotation_symmetry=two_fold_rotation_symmetry, + scan_pos=(y_pos, x_pos), + ) + + min_r = int(np.floor(0.1 * polar.shape[1])) + max_r = int(np.ceil(0.9 * polar.shape[1])) + std_est_origin = polar[:, min_r:max_r].std(axis=0) + std_est_origin_sum = std_est_origin.sum() + + origin_row = int(round(estimated_origin_row)) + origin_col = int(round(estimated_origin_col)) + coords_cache[(origin_row, origin_col)] = std_est_origin_sum + + if y_pos == 0 and x_pos == 0: + print(f"Initial std sum at estimated origin: {std_est_origin_sum}") + + converged = False + best = std_est_origin_sum + steps = 0 + while not converged and steps < max_steps: + steps += 1 + moved = False + + neighbors = [ + (origin_row + dr, origin_col + dc) + for dr in (-1, 0, 1) + for dc in (-1, 0, 1) + if not (dr == 0 and dc == 0) + ] + neighbors = [(r, c) for (r, c) in neighbors if 0 <= r < ny and 0 <= c < nx] + + for origin_r, origin_c in neighbors: + if (origin_r, origin_c) not in coords_cache: + polar = data.polar_transform( + origin_array=[origin_r, origin_c], + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + two_fold_rotation_symmetry=two_fold_rotation_symmetry, + scan_pos=(y_pos, x_pos), + ) + std_test = polar[:, min_r:max_r].std(axis=0) + coords_cache[(origin_r, origin_c)] = std_test.sum() + + if coords_cache[(origin_r, origin_c)] < best: + origin_row = origin_r + origin_col = origin_c + best = coords_cache[(origin_r, origin_c)] + moved = True + print(f"Moved to ({origin_row}, {origin_col}) with std sum {best}") + + if not moved: + converged = True + + if y_pos == 0 and x_pos == 0: + print(f"Final std sum at found origin ({origin_row}, {origin_col}): {best}") + origin_array[y_pos, x_pos, 0] = origin_row + origin_array[y_pos, x_pos, 1] = origin_col + + return origin_array + + +def dataset4dstem_polar_transform( + self: "Dataset4dstem", + origin_array: NDArray | None = None, + ellipse_params: tuple[float, float, float] | None = None, + num_annular_bins: int = 180, + radial_min: float = 0.0, + radial_max: float | None = None, + radial_step: float = 1.0, + two_fold_rotation_symmetry: bool = False, + name: str | None = None, + signal_units: str | None = None, + scan_pos: tuple[int, int] | None = None, +) -> Polar4dstem: + if self.array.ndim != 4: + raise ValueError("polar_transform requires a 4D-STEM dataset (ndim=4).") + scan_y, scan_x, ny, nx = self.array.shape + + # Standardize origin_array input + origin_array = np.asarray(origin_array) if origin_array is not None else None + if origin_array is None: + center = np.array([(ny - 1) / 2.0, (nx - 1) / 2.0], dtype=float) + origins = np.broadcast_to(center, (scan_y, scan_x, 2)).copy() + elif origin_array.shape == (2,): + origins = np.empty((scan_y, scan_x, 2), dtype=float) + origins[...] = origin_array + elif origin_array.shape == (scan_y, scan_x, 2): + origins = origin_array + else: + raise ValueError( + "origin_array must have shape None, (2,) or (scan_y, scan_x, 2)." + f" Got {origin_array.shape}." + ) + + # If scan_pos is provided, compute polar transform only for that position + if scan_pos is not None: + iy, ix = scan_pos + dp = self.array[iy, ix] # (ny, nx) view + r0 = float(origins[iy, ix, 0]) + c0 = float(origins[iy, ix, 1]) + + coords, phi_bins, radial_bins, radial_max_eff = _precompute_polar_coords( + ny=ny, + nx=nx, + origin_row=r0, + origin_col=c0, + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + two_fold_rotation_symmetry=two_fold_rotation_symmetry, + ) + polar2d = map_coordinates(dp, coords, order=1, mode="constant", cval=0.0) # (phi, r) + return polar2d + + # Otherwise, compute polar transform for all scan positions + # Determine one overall radial_max if not provided + if radial_max is None: + r_row_pos = origins[:, :, 0] + r_row_neg = (ny - 1) - origins[:, :, 0] + r_col_pos = origins[:, :, 1] + r_col_neg = (nx - 1) - origins[:, :, 1] + radial_max_eff_array = np.minimum.reduce([r_row_pos, r_row_neg, r_col_pos, r_col_neg]) + radial_max = float(max(radial_max_eff_array.min(), radial_min + radial_step)) + + # Precompute polar coords only once, using the origin from the first probe position + origin_row_f = float(origins[0, 0, 0]) + origin_col_f = float(origins[0, 0, 1]) + coords, phi_bins, radial_bins, radial_max_eff = _precompute_polar_coords( + ny=ny, + nx=nx, + origin_row=origin_row_f, + origin_col=origin_col_f, + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + two_fold_rotation_symmetry=two_fold_rotation_symmetry, + ) + n_phi = phi_bins.size + n_r = radial_bins.size + result_dtype = np.result_type(self.array.dtype, np.float32) + out = np.empty((scan_y, scan_x, n_phi, n_r), dtype=result_dtype) + + for iy in range(scan_y): + for ix in range(scan_x): + dp = self.array[iy, ix] + r0 = float(origins[iy, ix, 0]) + c0 = float(origins[iy, ix, 1]) + + coords, _, _, radial_max_eff = _precompute_polar_coords( + ny=ny, + nx=nx, + origin_row=r0, + origin_col=c0, + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + two_fold_rotation_symmetry=two_fold_rotation_symmetry, + ) + out[iy, ix] = map_coordinates( + dp, + coords, + order=1, + mode="constant", + cval=0.0, + ) + + phi_range = np.pi if two_fold_rotation_symmetry else 2.0 * np.pi + phi_step_deg = (phi_range / float(n_phi)) * (180.0 / np.pi) + sampling = np.zeros(4, dtype=float) + origin = np.zeros(4, dtype=float) + sampling[0:2] = np.asarray(self.sampling)[0:2] + sampling[2] = phi_step_deg + sampling[3] = float(np.asarray(self.sampling)[-1]) * radial_step + origin[0:2] = np.asarray(self.origin)[0:2] + origin[2] = 0.0 + origin[3] = radial_min * float(np.asarray(self.sampling)[-1]) + units = [ + self.units[0], + self.units[1], + "deg", + self.units[-1], + ] + metadata = dict(self.metadata) + metadata.update( + { + "polar_radial_min": float(radial_min), + "polar_radial_max": float(radial_max_eff), + "polar_radial_step": float(radial_step), + "polar_num_annular_bins": int(n_phi), + "polar_two_fold_rotation_symmetry": bool(two_fold_rotation_symmetry), + "polar_origin_row": float(origins[0, 0, 0]), + "polar_origin_col": float(origins[0, 0, 1]), + "polar_ellipse_params": tuple(ellipse_params) if ellipse_params is not None else None, + } + ) + return Polar4dstem( + array=out, + name=name if name is not None else f"{self.name}_polar", + origin=origin, + sampling=sampling, + units=units, + signal_units=signal_units if signal_units is not None else self.signal_units, + metadata=metadata, + _token=Polar4dstem._token, + ) diff --git a/src/quantem/diffraction/__init__.py b/src/quantem/diffraction/__init__.py index e69de29bb..6fae8ef7d 100644 --- a/src/quantem/diffraction/__init__.py +++ b/src/quantem/diffraction/__init__.py @@ -0,0 +1,62 @@ +"""Diffraction analysis interfaces.""" + +from quantem.diffraction.bragg_peaks import BraggPeaksPolymer, ScanMaskEditor +from quantem.diffraction.polymer_models import ( + PAPER_MODEL_ID, + PAPER_MODEL_VERSION, + PolymerModelError, + PolymerModelResolution, + resolve_polymer_model, +) +from quantem.diffraction.polymer_ice import ( + IceDetectionResult, + IceFlaggerDebug, + IceFlaggerParams, + apply_ice_mask_to_vector, + collect_peak_widths, + compute_global_intensity_threshold, + detect_ice, + flag_ice_peaks_in_dataset, + flag_ice_peaks_in_pattern, + measure_peak_widths, + sharpness_mask, + plot_q_intensity_density, +) +from quantem.diffraction.polymer_normalization import ( + GlobalPercentileNormalization, + GlobalPercentileStrategy, + LegacyNormalizationAdapter, + NormalizationStrategy, + PerImageMinMaxPercentileNormalization, + PerImageMinMaxPercentileStrategy, + resolve_normalization_strategy, +) + +__all__ = [ + "BraggPeaksPolymer", + "ScanMaskEditor", + "GlobalPercentileNormalization", + "GlobalPercentileStrategy", + "IceDetectionResult", + "IceFlaggerDebug", + "IceFlaggerParams", + "measure_peak_widths", + "sharpness_mask", + "LegacyNormalizationAdapter", + "NormalizationStrategy", + "PAPER_MODEL_ID", + "PAPER_MODEL_VERSION", + "PolymerModelError", + "PolymerModelResolution", + "PerImageMinMaxPercentileNormalization", + "PerImageMinMaxPercentileStrategy", + "apply_ice_mask_to_vector", + "compute_global_intensity_threshold", + "collect_peak_widths", + "detect_ice", + "flag_ice_peaks_in_dataset", + "flag_ice_peaks_in_pattern", + "plot_q_intensity_density", + "resolve_polymer_model", + "resolve_normalization_strategy", +] diff --git a/src/quantem/diffraction/bragg_peaks.py b/src/quantem/diffraction/bragg_peaks.py new file mode 100644 index 000000000..8562f9dad --- /dev/null +++ b/src/quantem/diffraction/bragg_peaks.py @@ -0,0 +1,8206 @@ +# from collections.abc import Sequence +import warnings +import tempfile +from typing import Tuple + +import matplotlib.pyplot as plt +import numpy as np +from numpy.typing import ArrayLike +from scipy.ndimage import gaussian_filter, map_coordinates, label +from tqdm import tqdm +import torch +from quantem.core.datastructures.dataset3d import Dataset3d +from quantem.core.datastructures.dataset4dstem import Dataset4dstem +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.polymer_models import ( + PAPER_MODEL_ID, + PAPER_MODEL_VERSION, + MultiChannelCNN2d, + build_polymer_model, + resolve_polymer_model, +) +from quantem.diffraction.polymer_normalization import ( + LegacyNormalizationAdapter, + NormalizationStrategy, + resolve_normalization_strategy, +) +from quantem.core.datastructures import Vector +from quantem.core.visualization import show_2d +from quantem.diffraction.polar_transform import ( + find_origin as find_origin_angular_uniformity, + polar_transform as karen_polar_transform, + polar_transform_peaks as karen_polar_transform_peaks, +) +from quantem.diffraction.peak_detection import detect_blobs, find_central_beam_from_peaks +from quantem.diffraction.orientation_correlation import ( + calculate_orientation_correlation as _calculate_orientation_correlation, +) +from quantem.core.utils.utils import electron_wavelength_angstrom +from quantem.diffraction.polymer_utils import parse_reciprocal_units, sample_average_from_image +from emdfile import tqdmnd +from scipy.ndimage import gaussian_filter1d +from scipy.signal import find_peaks, peak_widths +import ipywidgets as widgets +from ipywidgets import IntSlider, Button, HBox, VBox, interactive_output +from IPython.display import clear_output, display +from pathlib import Path +from mpl_toolkits.axes_grid1.inset_locator import inset_axes +from matplotlib.patches import Ellipse, Rectangle +from matplotlib.colors import BoundaryNorm, hsv_to_rgb, rgb_to_hsv + +# Default 1/e decay length for the orientation-correlation slope fit, as a +# fraction of the fitted lobe's distance span. Chosen on real scan data +# (pg3T2 07-03-2024/61): it holds the fitted intercept within ~1.7 degrees of +# the measured boundary at zero separation across all six ring pairs, versus +# up to 7 degrees for an unweighted fit, while keeping enough effective points +# for a stable slope. See plot_orientation_correlation(slope_weight_scale=...). +SLOPE_WEIGHT_FRACTION = 0.10 + + +def _apply_zoom_crop(data, zoom_factor, center=None): + """Crop data to center region based on zoom factor.""" + if zoom_factor == 1.0: + return data, (0, data.shape[0], 0, data.shape[1]) + + h, w = data.shape + new_h, new_w = int(h / zoom_factor), int(w / zoom_factor) + new_h = max(1, min(h, new_h)) + new_w = max(1, min(w, new_w)) + + if center is None: + center_y, center_x = (h - 1) / 2, (w - 1) / 2 + else: + center_y, center_x = center + + top = int(round(center_y - (new_h - 1) / 2)) + left = int(round(center_x - (new_w - 1) / 2)) + top = min(max(top, 0), h - new_h) + left = min(max(left, 0), w - new_w) + + return data[top:top+new_h, left:left+new_w], (top, top+new_h, left, left+new_w) + + +def _mean_intensity_map(dataset_cartesian, scan_shape): + Ry, Rx = scan_shape + return np.array( + [ + [np.mean(dataset_cartesian[i, j].array) for j in range(Rx)] + for i in range(Ry) + ] + ) + + +def _resolve_intensity_map( + dataset_cartesian, + intensity_map, + scan_shape, + *, + validate=True, + announce_upsample=False, +): + Ry, Rx = scan_shape + if intensity_map is None: + return _mean_intensity_map(dataset_cartesian, scan_shape), 1 + + map_shape = intensity_map.shape[:2] + upsample_factor = map_shape[0] // Ry + if validate: + if upsample_factor != map_shape[1] // Rx: + raise ValueError("Inconsistent upsample factors") + if map_shape[0] % Ry != 0 or map_shape[1] % Rx != 0: + raise ValueError( + f"intensity_map shape {intensity_map.shape} not integer multiple of ({Ry}, {Rx})" + ) + if announce_upsample: + print(f"Auto-detected upsample_factor: {upsample_factor}") + return intensity_map, upsample_factor + + +def _intensity_display_limits(intensity_map): + is_rgb_map = intensity_map.ndim == 3 and intensity_map.shape[2] in (3, 4) + if is_rgb_map: + return is_rgb_map, None, None + finite = np.isfinite(intensity_map) + if not np.any(finite): + return is_rgb_map, 0.0, 1.0 + vmin, vmax = np.quantile(intensity_map[finite], [0.01, 0.99]) + return is_rgb_map, vmin, vmax + + +def _normalized_dp( + dataset_cartesian, + ry_data, + rx_data, + *, + norm_upper_quantile=None, + norm_power=1.0, + copy_data=True, +): + dp_data = dataset_cartesian[ry_data, rx_data].array + if copy_data: + dp_data = dp_data.copy() + if norm_upper_quantile is not None: + dp_data = np.clip(dp_data, 0, np.quantile(dp_data, norm_upper_quantile)) + if norm_power != 1.0: + m = np.nanmax(dp_data) + if np.isfinite(m) and m > 0: + dp_data = (dp_data / m) ** norm_power * m + return dp_data + + +def _display_center(image_centers, ry_data, rx_data, image_shape): + center_y, center_x = image_shape[0] / 2, image_shape[1] / 2 + if image_centers is not None: + stored_center = image_centers[:, ry_data, rx_data] + if np.all(np.isfinite(stored_center)) and not np.allclose(stored_center, 0): + center_y, center_x = stored_center + return center_y, center_x + + +def _has_peak_positions(peaks_x, peaks_y): + return ( + peaks_x is not None + and peaks_y is not None + and len(peaks_x) > 0 + and len(peaks_y) > 0 + ) + + +def _vector_field_flat(vector, field): + """Return one current-Vector field as a one-dimensional NumPy array.""" + return vector.select_fields(field).flatten()[:, 0] + + +def _vector_field_cell(vector, field, row, col): + """Return one current-Vector field from a scan cell as a 1D array.""" + return vector.select_fields(field)[row, col].array[:, 0] + + +def _central_peak_index(peaks_x, peaks_y, peaks_r_invA, center, max_dist=None): + """Index of the detected central-beam peak, or ``None``. + + Defined as the detected peak nearest the calibrated beam ``center`` (from + ``image_centers`` / ``find_central_beams_4d``), but only when it lies within + ``max_dist`` pixels of it. The filled central-beam marker itself is always drawn at + ``center``; this index only flags which detected peak, if any, to drop from the + open-circle set so a ring is not drawn on top of the beam. + + ``peaks_r_invA`` is unused (kept for call-site compatibility): selecting the beam by + smallest polar radius made the marker jump to an off-center low-q Bragg peak when the + beam itself was not detected as a peak. + """ + if not _has_peak_positions(peaks_x, peaks_y): + return None + center_y, center_x = center + distances = np.sqrt( + (np.asarray(peaks_x) - center_x) ** 2 + (np.asarray(peaks_y) - center_y) ** 2 + ) + idx = int(np.argmin(distances)) + if max_dist is not None and distances[idx] > max_dist: + return None + return idx + + +def _central_beam_max_dist(image_shape): + """Pixel radius within which a detected peak counts as the central beam. + + Small enough that finite-q Bragg peaks are never mistaken for the beam, generous + enough to absorb a few-pixel disagreement between center-finding and peak detection. + """ + return max(4.0, 0.03 * min(image_shape[0], image_shape[1])) + + +def _zoom_peak_overlay( + dp_data, + peaks_x, + peaks_y, + peaks_r_invA, + peak_ints, + central_idx, + zoom, + fallback_center, +): + if zoom == 1: + return dp_data, peaks_x, peaks_y, peaks_r_invA, peak_ints, central_idx, fallback_center + + dp_data, ranges = _apply_zoom_crop(dp_data, zoom, center=fallback_center) + top, bot, left, right = ranges + display_center = (fallback_center[0] - top, fallback_center[1] - left) + + if _has_peak_positions(peaks_x, peaks_y): + mask = (top <= peaks_y) & (peaks_y < bot) & (left <= peaks_x) & (peaks_x < right) + kept_indices = np.flatnonzero(mask) + if central_idx is not None: + central_matches = np.flatnonzero(kept_indices == central_idx) + central_idx = int(central_matches[0]) if len(central_matches) else None + peaks_y = peaks_y[mask] - top + peaks_x = peaks_x[mask] - left + peaks_r_invA = peaks_r_invA[mask] if peaks_r_invA is not None else None + peak_ints = peak_ints[mask] if peak_ints is not None else None + + return dp_data, peaks_x, peaks_y, peaks_r_invA, peak_ints, central_idx, display_center + + +def _polar_peak_bins( + polar_r, + polar_theta, + max_radius_invA, + num_radial_bins, + num_annular_bins, + two_fold_symmetry, +): + r_bins = polar_r / max_radius_invA * num_radial_bins + theta_period = np.pi if two_fold_symmetry else 2 * np.pi + theta_bins = polar_theta / theta_period * num_annular_bins + return r_bins, theta_bins + + +def _plot_bragg_peaks_on_ax( + ax, + peaks_x, + peaks_y, + peaks_r_invA, + peak_intensities, + central_idx, + *, + radial_range=None, + show_all_peaks=False, + selected_peak_color="red", + other_peak_color="gray", + central_beam_color="red", + peak_intensity_mode="size", + peak_size_range=(30, 300), + peak_cmap="hot", + peak_vmin=None, + peak_vmax=None, + crosshair_width_peaks=2, + crosshair_scaling_peaks=1, + crosshair_scaling_central_beam=1, + peak_marker="o", + peak_marker_facecolors="none", + peak_marker_size=None, + peak_alpha=0.8, + central_alpha=0.95, + central_linewidth=2, + add_colorbar=False, + center=None, + show_center=True, + show_central_beam=True, +): + # show_central_beam=False fully suppresses the central-beam marker (both the + # provided-center dot and the detected-central-peak dot); the central peak is + # still excluded from the open-circle set via non_central below. + plot_detected_center = (center is None or not show_center) and show_central_beam + if center is not None and show_center and show_central_beam: + center_y, center_x = center + ax.scatter( + center_x, + center_y, + s=120 * crosshair_scaling_central_beam, + alpha=central_alpha, + linewidths=central_linewidth, + edgecolors="k", + facecolors=central_beam_color, + marker="o", + zorder=10, + ) + + if peaks_r_invA is None or len(peaks_r_invA) == 0: + return + if not _has_peak_positions(peaks_x, peaks_y): + return + + central_style = dict( + edgecolors="k", + facecolors=central_beam_color, + marker="o", + zorder=10, + ) + + if radial_range is not None: + mask = (peaks_r_invA >= radial_range[0]) & (peaks_r_invA < radial_range[1]) + if show_all_peaks and np.any(~mask): + out_indices = np.where(~mask)[0] + if plot_detected_center and central_idx is not None and central_idx in out_indices: + ax.scatter( + peaks_x[central_idx], + peaks_y[central_idx], + s=30, + alpha=0.95, + linewidths=2, + **central_style, + ) + other_out_mask = out_indices != central_idx + if np.any(other_out_mask): + ax.scatter( + peaks_x[out_indices[other_out_mask]], + peaks_y[out_indices[other_out_mask]], + c=other_peak_color, + s=30, + alpha=0.5, + marker="x", + linewidths=1.5, + ) + if not np.any(mask): + return + in_range_indices = np.where(mask)[0] + if central_idx is not None and central_idx in in_range_indices: + central_idx = np.where(in_range_indices == central_idx)[0][0] + else: + central_idx = None + peaks_x, peaks_y = peaks_x[mask], peaks_y[mask] + peak_intensities = peak_intensities[mask] if peak_intensities is not None else None + + if central_idx is not None: + if plot_detected_center: + ax.scatter( + peaks_x[central_idx], + peaks_y[central_idx], + s=120 * crosshair_scaling_central_beam, + alpha=central_alpha, + linewidths=central_linewidth, + **central_style, + ) + non_central = np.ones(len(peaks_x), dtype=bool) + non_central[central_idx] = False + else: + non_central = np.ones(len(peaks_x), dtype=bool) + + if not np.any(non_central): + return + + if peak_intensities is not None and peak_intensity_mode is not None: + int_subset = peak_intensities[non_central] + int_min = peak_vmin if peak_vmin is not None else np.min(int_subset) + int_max = peak_vmax if peak_vmax is not None else np.max(int_subset) + norm_int = ( + (int_subset - int_min) / (int_max - int_min) + if int_max > int_min + else np.ones_like(int_subset) + ) + if peak_intensity_mode == "color": + colors, sizes = plt.cm.get_cmap(peak_cmap)(norm_int), 100 + elif peak_intensity_mode == "size": + colors = selected_peak_color + sizes = peak_size_range[0] + norm_int * (peak_size_range[1] - peak_size_range[0]) + elif peak_intensity_mode == "both": + colors = plt.cm.get_cmap(peak_cmap)(norm_int) + sizes = peak_size_range[0] + norm_int * (peak_size_range[1] - peak_size_range[0]) + else: + colors, sizes = selected_peak_color, 100 + if peak_marker_size is not None: + sizes = peak_marker_size + + scatter_kwargs = dict( + s=sizes * crosshair_scaling_peaks, + alpha=peak_alpha, + marker=peak_marker, + facecolors=peak_marker_facecolors, + linewidths=crosshair_width_peaks, + zorder=5, + ) + if peak_marker_facecolors == "none": + ax.scatter( + peaks_x[non_central], + peaks_y[non_central], + edgecolors=colors, + **scatter_kwargs, + ) + else: + ax.scatter( + peaks_x[non_central], + peaks_y[non_central], + c=colors, + **scatter_kwargs, + ) + + if peak_intensity_mode in ["color", "both"]: + sm = plt.cm.ScalarMappable( + cmap=peak_cmap, norm=plt.Normalize(vmin=int_min, vmax=int_max) + ) + sm.set_array([]) + if add_colorbar: + plt.colorbar(sm, ax=ax, pad=0.02, fraction=0.046).set_label( + "Peak Intensity", fontsize=8 + ) + else: + scatter_kwargs = dict( + s=100, + alpha=peak_alpha, + marker=peak_marker, + facecolors=peak_marker_facecolors, + linewidths=2, + zorder=5, + ) + if peak_marker_facecolors == "none": + ax.scatter( + peaks_x[non_central], + peaks_y[non_central], + edgecolors=selected_peak_color, + **scatter_kwargs, + ) + else: + ax.scatter( + peaks_x[non_central], + peaks_y[non_central], + c=selected_peak_color, + **scatter_kwargs, + ) + + +def _draw_peaks_data_circles( + ax, + peaks_x, + peaks_y, + peak_intensities, + central_idx, + center, + *, + marker_scaled=True, + marker_size=8.0, + marker_size_min=4.0, + marker_size_max=16.0, + selected_peak_color="red", + central_beam_color="red", + show_central_beam=True, + central_size=5.0, + peak_linewidth=2.0, + central_linewidth=1.5, +): + """Draw Bragg-peak markers as circles in DATA coordinates (radius in detector + pixels). + + Unlike ``_plot_bragg_peaks_on_ax`` (which sizes markers in fixed points**2, so they + do not track the figure size), these circles are in data units and therefore cover a + constant fraction of the diffraction pattern at any panel size -- matching the widget + canvas, where the overlay sizes markers in data pixels scaled to the display (js + ``drawDot`` / peak rings). Open circles for detected peaks (skipping the central-beam + peak) plus a filled dot at the calibrated beam ``center`` (row, col). + """ + from matplotlib.patches import Circle + + if _has_peak_positions(peaks_x, peaks_y): + px = np.asarray(peaks_x) + py = np.asarray(peaks_y) + idxs = [ + i for i in range(len(px)) + if i != central_idx and np.isfinite(px[i]) and np.isfinite(py[i]) + ] + ints = None if peak_intensities is None else np.asarray(peak_intensities, dtype=float) + use_scaled = marker_scaled and ints is not None and len(idxs) > 0 + if use_scaled: + vals = ints[idxs] + imin, imax = float(np.nanmin(vals)), float(np.nanmax(vals)) + rng = (imax - imin) if imax > imin else 1.0 + for i in idxs: + if use_scaled: + norm = (ints[i] - imin) / rng + if not np.isfinite(norm): + norm = 0.5 + r = marker_size_min + norm * (marker_size_max - marker_size_min) + else: + r = marker_size + ax.add_patch(Circle( + (px[i], py[i]), radius=r, fill=False, + edgecolor=selected_peak_color, linewidth=peak_linewidth, zorder=5, + )) + + if show_central_beam and center is not None: + cy, cx = center + ax.add_patch(Circle( + (cx, cy), radius=central_size, facecolor=central_beam_color, + edgecolor="k", linewidth=central_linewidth, zorder=10, + )) + + +class ScanMaskEditor: + """Interactive, persistent circular scan-mask editor. + + The horizontal X control maps directly to the scan-column coordinate. The + vertical Y control is visually inverted relative to the array row index so + moving the slider upward moves the probe marker upward on an ``origin="upper"`` + scan image. + """ + + SCHEMA_VERSION = 2 + GEOMETRIES = ("circle", "ellipse", "square", "rectangle") + + def __init__( + self, + analysis, + *, + initial_x=None, + initial_y=None, + initial_radius=None, + initial_geometry="circle", + initial_size_x=None, + initial_size_y=None, + reference_image=None, + state_path=None, + overlay_alpha=0.28, + crosshair_width=2, + crosshair_size=12, + autosave=False, + display_widget=True, + ): + self.analysis = analysis + self.scan_shape = tuple(int(v) for v in analysis.dataset_cartesian.shape[:2]) + self.state_path = None if state_path is None else Path(state_path) + self.autosave = bool(autosave) + self.overlay_alpha = float(overlay_alpha) + self.crosshair_width = float(crosshair_width) + self.crosshair_size = float(crosshair_size) + self._syncing = False + self._dirty = False + self._saved = False + self._loaded = False + + rows, columns = self.scan_shape + center_row = rows // 2 if initial_y is None else int(initial_y) + center_column = columns // 2 if initial_x is None else int(initial_x) + radius = ( + max(1, min(rows, columns) // 3) + if initial_radius is None + else int(initial_radius) + ) + geometry = str(initial_geometry).lower() + size_x = radius if initial_size_x is None else int(initial_size_x) + size_y = radius if initial_size_y is None else int(initial_size_y) + loaded_mask = None + if self.state_path is not None and self.state_path.is_file(): + state = self._read_state(self.state_path) + center_row = state["center_row"] + center_column = state["center_column"] + geometry = state["geometry"] + size_x = state["size_x"] + size_y = state["size_y"] + loaded_mask = state["mask"] + self._loaded = True + self._saved = True + + self._validate_geometry( + center_row, center_column, geometry, size_x, size_y + ) + self.reference_image = self._resolve_reference_image(reference_image) + self._preview_mask = ( + loaded_mask.copy() + if loaded_mask is not None + else self._geometry_mask( + center_row, center_column, geometry, size_x, size_y + ) + ) + + maximum_radius = int(np.ceil(np.hypot(rows - 1, columns - 1))) + 1 + self.x_slider = widgets.IntSlider( + value=center_column, + min=0, + max=columns - 1, + step=1, + description="X", + continuous_update=True, + readout=False, + style={"description_width": "18px"}, + layout=widgets.Layout(width="375px"), + ) + # Slider value increases upward, while array row indices increase downward. + self.y_slider = widgets.IntSlider( + value=rows - 1 - center_row, + min=0, + max=rows - 1, + step=1, + description="Y", + orientation="vertical", + continuous_update=True, + readout=False, + style={"description_width": "18px"}, + layout=widgets.Layout(height="281px", width="52px"), + ) + self.geometry_selector = widgets.Dropdown( + options=[ + ("Circular", "circle"), + ("Elliptical", "ellipse"), + ("Square", "square"), + ("Rectangular", "rectangle"), + ], + value=geometry, + description="Shape", + style={"description_width": "42px"}, + layout=widgets.Layout(width="155px"), + ) + self.size_x_slider = widgets.IntSlider( + value=size_x, + min=1, + max=maximum_radius, + step=1, + description="Radius", + continuous_update=True, + readout=False, + style={"description_width": "55px"}, + layout=widgets.Layout(width="375px"), + ) + self.size_y_slider = widgets.IntSlider( + value=size_y, + min=1, + max=maximum_radius, + step=1, + description="Y radius", + continuous_update=True, + readout=False, + style={"description_width": "68px"}, + layout=widgets.Layout(width="375px"), + ) + # Historical public attribute retained for callers that customize it. + self.radius_slider = self.size_x_slider + self.x_input = widgets.BoundedIntText( + value=center_column, + min=0, + max=columns - 1, + description="X column", + style={"description_width": "62px"}, + layout=widgets.Layout(width="150px"), + ) + self.y_input = widgets.BoundedIntText( + value=center_row, + min=0, + max=rows - 1, + description="Y row", + style={"description_width": "52px"}, + layout=widgets.Layout(width="140px"), + ) + self.size_x_input = widgets.BoundedIntText( + value=size_x, + min=1, + max=maximum_radius, + description="Radius", + style={"description_width": "48px"}, + layout=widgets.Layout(width="130px"), + ) + self.size_y_input = widgets.BoundedIntText( + value=size_y, + min=1, + max=maximum_radius, + description="Y radius", + style={"description_width": "58px"}, + layout=widgets.Layout(width="140px"), + ) + self.radius_input = self.size_x_input + + self.apply_button = widgets.Button( + description="Apply", + icon="check", + button_style="primary", + tooltip="Commit the preview mask to BraggPeaksPolymer.scan_mask", + layout=widgets.Layout(width="72px"), + ) + self.save_button = widgets.Button( + description="Apply & Save", + icon="save", + button_style="success", + tooltip="Commit and save this mask for the next notebook run", + disabled=self.state_path is None, + layout=widgets.Layout(width="105px"), + ) + self.center_button = widgets.Button( + description="Center", + icon="crosshairs", + tooltip="Center the circle", + layout=widgets.Layout(width="75px"), + ) + self.full_button = widgets.Button( + description="Full", + icon="expand", + tooltip="Include the entire scan", + layout=widgets.Layout(width="70px"), + ) + self.reset_button = widgets.Button( + description="Reset", + icon="undo", + tooltip="Restore the loaded/default state", + layout=widgets.Layout(width="72px"), + ) + self.status = widgets.HTML(layout=widgets.Layout(width="500px")) + self.output = widgets.Output( + layout=widgets.Layout( + width="438px", + height="356px", + max_width="438px", + overflow="hidden", + ) + ) + + self.figure, self.ax = plt.subplots(figsize=(4.0, 3.25)) + finite = self.reference_image[np.isfinite(self.reference_image)] + if finite.size: + vmin, vmax = np.percentile(finite, [1.0, 99.0]) + if not vmax > vmin: + vmin, vmax = float(np.min(finite)), float(np.max(finite) + 1.0) + else: + vmin, vmax = 0.0, 1.0 + self.image_artist = self.ax.imshow( + self.reference_image, + cmap="gray", + origin="upper", + vmin=vmin, + vmax=vmax, + interpolation="nearest", + ) + self.mask_artist = self.ax.imshow( + np.ma.masked_where(~self._preview_mask, self._preview_mask), + cmap="Reds", + origin="upper", + alpha=self.overlay_alpha, + vmin=0, + vmax=1, + interpolation="nearest", + ) + self.boundary_artist = None + self.circle_artist = None + self._replace_boundary_artist() + (self.center_artist,) = self.ax.plot( + center_column, + center_row, + marker="+", + color="#ff3030", + markersize=self.crosshair_size, + markeredgewidth=self.crosshair_width, + ) + self.ax.set( + title="Scan-mask editor", + xlabel="X — scan column", + ylabel="Y — scan row", + xlim=(-0.5, columns - 0.5), + ylim=(rows - 0.5, -0.5), + ) + self.ax.title.set_fontsize(10) + self.ax.xaxis.label.set_fontsize(9) + self.ax.yaxis.label.set_fontsize(9) + self.ax.tick_params(labelsize=8) + self.figure.tight_layout() + + self.x_slider.observe(self._on_x_slider, names="value") + self.y_slider.observe(self._on_y_slider, names="value") + self.geometry_selector.observe(self._on_geometry, names="value") + self.size_x_slider.observe(self._on_size_x_slider, names="value") + self.size_y_slider.observe(self._on_size_y_slider, names="value") + self.x_input.observe(self._on_x_input, names="value") + self.y_input.observe(self._on_y_input, names="value") + self.size_x_input.observe(self._on_size_x_input, names="value") + self.size_y_input.observe(self._on_size_y_input, names="value") + self.apply_button.on_click(lambda _: self.apply()) + self.save_button.on_click(lambda _: self.save()) + self.center_button.on_click( + lambda _: self.set_mask(x=columns // 2, y=rows // 2) + ) + self.full_button.on_click(lambda _: self._set_full_scan()) + self.reset_button.on_click(lambda _: self._restore_initial()) + + self._initial_geometry = ( + center_column, center_row, geometry, size_x, size_y + ) + self._initial_mask = self._preview_mask.copy() + # A loaded/default mask is immediately usable by Run All. Slider edits + # remain previews until Apply, preventing repeated inference-cache invalidation. + self.analysis.scan_mask = self._preview_mask.copy() + self._render() + self._refresh_status("Loaded saved mask" if self._loaded else "Default mask applied") + + position_row = widgets.HBox( + [self.geometry_selector, self.x_input, self.y_input], + layout=widgets.Layout(width="500px"), + ) + self.size_input_row = widgets.HBox( + [self.size_x_input, self.size_y_input], + layout=widgets.Layout(width="500px"), + ) + toolbar = widgets.HBox( + [ + self.apply_button, + self.save_button, + self.center_button, + self.full_button, + self.reset_button, + ], + layout=widgets.Layout(flex_flow="row wrap"), + ) + plot_row = widgets.HBox( + [self.y_slider, self.output], + layout=widgets.Layout( + align_items="center", width="500px", overflow="hidden" + ), + ) + self.size_x_row = widgets.HBox( + [ + widgets.Box(layout=widgets.Layout(width="52px")), + self.size_x_slider, + ], + layout=widgets.Layout(align_items="center", width="500px"), + ) + self.size_y_row = widgets.HBox( + [ + widgets.Box(layout=widgets.Layout(width="52px")), + self.size_y_slider, + ], + layout=widgets.Layout(align_items="center", width="500px"), + ) + x_row = widgets.HBox( + [ + widgets.Box(layout=widgets.Layout(width="52px")), + self.x_slider, + ], + layout=widgets.Layout(align_items="center", width="500px"), + ) + self.widget = widgets.VBox( + [ + toolbar, + position_row, + self.size_input_row, + self.size_x_row, + self.size_y_row, + plot_row, + x_row, + self.status, + ], + layout=widgets.Layout(width="500px", max_width="500px"), + ) + self._refresh_geometry_controls() + # Prevent the inline Matplotlib backend from appending a second copy of + # the figure after the widget cell. The explicitly displayed Output copy + # remains live and continues to update. + plt.close(self.figure) + if display_widget: + display(self.widget) + + @property + def x(self): + """Horizontal scan-column coordinate.""" + return int(self.x_slider.value) + + @property + def y(self): + """Vertical scan-row coordinate.""" + return int(self.scan_shape[0] - 1 - self.y_slider.value) + + @property + def radius(self): + """Circle radius / square half-width compatibility value.""" + return self.size_x + + @property + def geometry(self): + return str(self.geometry_selector.value) + + @property + def size_x(self): + """Horizontal radius or half-width in scan pixels.""" + return int(self.size_x_slider.value) + + @property + def size_y(self): + """Vertical radius or half-height in scan pixels.""" + if self.geometry in {"circle", "square"}: + return self.size_x + return int(self.size_y_slider.value) + + @property + def mask(self): + """Current preview mask.""" + return self._preview_mask.copy() + + @property + def applied_mask(self): + return None if self.analysis.scan_mask is None else self.analysis.scan_mask.copy() + + def _effective_mask(self): + """Return the committed mask for ndarray-style compatibility.""" + mask = self.analysis.scan_mask + return self._preview_mask if mask is None else np.asarray(mask, dtype=bool) + + @property + def shape(self): + return self.scan_shape + + @property + def dtype(self): + return np.dtype(bool) + + @property + def size(self): + return int(np.prod(self.scan_shape)) + + @property + def ndim(self): + return 2 + + def __array__(self, dtype=None, copy=None): + array = np.asarray(self._effective_mask(), dtype=dtype) + if copy: + array = array.copy() + return array + + def __len__(self): + return self.scan_shape[0] + + def sum(self, *args, **kwargs): + """NumPy-compatible sum for legacy ``mask_arr = editor`` cells.""" + return self._effective_mask().sum(*args, **kwargs) + + def astype(self, *args, **kwargs): + return self._effective_mask().astype(*args, **kwargs) + + def copy(self): + return self._effective_mask().copy() + + def __getitem__(self, key): + if not isinstance(key, str): + return self._effective_mask()[key] + # Compatibility with the historical returned dictionary. Its x0/y0 + # names represented row/column respectively despite their labels. + values = { + "mask": self.mask, + "x0": self.y, + "y0": self.x, + "r": self.radius, + "center_row": self.y, + "center_column": self.x, + "geometry": self.geometry, + "size_x": self.size_x, + "size_y": self.size_y, + } + return values[key] + + def get(self, key, default=None): + try: + return self[key] + except KeyError: + return default + + def _validate_geometry(self, row, column, geometry, size_x, size_y): + rows, columns = self.scan_shape + if not 0 <= row < rows or not 0 <= column < columns: + raise ValueError( + f"Mask center (row={row}, column={column}) is outside scan shape " + f"{self.scan_shape}." + ) + if geometry not in self.GEOMETRIES: + raise ValueError( + f"Unknown mask geometry {geometry!r}; choose one of " + f"{', '.join(self.GEOMETRIES)}." + ) + if size_x < 1 or size_y < 1: + raise ValueError("Mask half-sizes must be at least one scan pixel.") + maximum_radius = int(np.ceil(np.hypot(rows - 1, columns - 1))) + 1 + if size_x > maximum_radius or size_y > maximum_radius: + raise ValueError( + f"Mask half-size ({size_x}, {size_y}) exceeds the supported maximum " + f"{maximum_radius} for scan shape {self.scan_shape}." + ) + + def _geometry_mask(self, row, column, geometry, size_x, size_y): + yy, xx = np.ogrid[: self.scan_shape[0], : self.scan_shape[1]] + dy = yy - row + dx = xx - column + if geometry == "circle": + return dy**2 + dx**2 <= size_x**2 + if geometry == "ellipse": + return (dx / size_x) ** 2 + (dy / size_y) ** 2 <= 1 + if geometry == "square": + return (np.abs(dx) <= size_x) & (np.abs(dy) <= size_x) + if geometry == "rectangle": + return (np.abs(dx) <= size_x) & (np.abs(dy) <= size_y) + raise ValueError(f"Unknown mask geometry {geometry!r}.") + + def _resolve_reference_image(self, reference_image): + if reference_image is None: + virtual_images = getattr(self.analysis.dataset_cartesian, "virtual_images", {}) + if "virtual_image" in virtual_images: + reference_image = virtual_images["virtual_image"] + reference_image = getattr( + reference_image, + "array", + getattr(reference_image, "data", reference_image), + ) + else: + dataset = self.analysis.dataset_cartesian + array = getattr(dataset, "array", None) + if array is not None: + reference_image = np.asarray(array).mean(axis=(-2, -1)) + else: + reference_image = ( + dataset.tensor.float().mean(dim=(-2, -1)).detach().cpu().numpy() + ) + reference_image = np.asarray(reference_image, dtype=float) + if reference_image.shape != self.scan_shape: + raise ValueError( + f"reference_image shape {reference_image.shape} must match " + f"scan shape {self.scan_shape}." + ) + return reference_image + + def _read_state(self, path): + try: + with np.load(path, allow_pickle=False) as state: + version = int(state["schema_version"]) + shape = tuple(int(v) for v in state["scan_shape"]) + if version not in {1, self.SCHEMA_VERSION}: + raise ValueError( + f"Unsupported scan-mask schema {version}; expected " + f"1 or {self.SCHEMA_VERSION}." + ) + if shape != self.scan_shape: + raise ValueError( + f"Saved scan-mask shape {shape} does not match current " + f"scan shape {self.scan_shape}." + ) + mask = np.asarray(state["mask"], dtype=bool) + if mask.shape != self.scan_shape: + raise ValueError( + f"Saved mask array shape {mask.shape} does not match " + f"scan shape {self.scan_shape}." + ) + if version == 1: + geometry = "circle" + size_x = size_y = int(state["radius"]) + else: + geometry = str(state["geometry"].item()) + size_x = int(state["size_x"]) + size_y = int(state["size_y"]) + return { + "center_row": int(state["center_row"]), + "center_column": int(state["center_column"]), + "geometry": geometry, + "size_x": size_x, + "size_y": size_y, + "mask": mask, + } + except (OSError, KeyError) as exc: + raise ValueError(f"Could not load scan-mask state from {path}: {exc}") from exc + + def set_mask( + self, + *, + x=None, + y=None, + geometry=None, + size_x=None, + size_y=None, + ): + self._syncing = True + try: + if geometry is not None: + geometry = str(geometry).lower() + if geometry not in self.GEOMETRIES: + raise ValueError( + f"Unknown mask geometry {geometry!r}; choose one of " + f"{', '.join(self.GEOMETRIES)}." + ) + self.geometry_selector.value = geometry + if x is not None: + self.x_slider.value = int(x) + self.x_input.value = int(x) + if y is not None: + self.y_slider.value = self.scan_shape[0] - 1 - int(y) + self.y_input.value = int(y) + if size_x is not None: + self.size_x_slider.value = int(size_x) + self.size_x_input.value = int(size_x) + if size_y is not None: + self.size_y_slider.value = int(size_y) + self.size_y_input.value = int(size_y) + if self.geometry in {"circle", "square"}: + self.size_y_slider.value = self.size_x + self.size_y_input.value = self.size_x + finally: + self._syncing = False + self._refresh_geometry_controls() + self._update_preview() + return self + + def set_circle(self, *, x=None, y=None, radius=None): + """Compatibility helper that explicitly selects circular geometry.""" + return self.set_mask( + x=x, + y=y, + geometry="circle", + size_x=radius, + size_y=radius, + ) + + def apply(self): + self.analysis.scan_mask = self._preview_mask.copy() + self._dirty = False + self._refresh_status("Mask applied") + if self.autosave and self.state_path is not None: + self.save(apply_first=False) + return self.applied_mask + + def save(self, path=None, *, apply_first=True): + path = self.state_path if path is None else Path(path) + if path is None: + raise ValueError("No scan-mask state_path was configured.") + if apply_first: + self.apply() + path.parent.mkdir(parents=True, exist_ok=True) + sampling = np.asarray(self.analysis.dataset_cartesian.sampling[:2], dtype=float) + units = np.asarray( + [str(value) for value in self.analysis.dataset_cartesian.units[:2]], + dtype="U32", + ) + with tempfile.NamedTemporaryFile( + mode="wb", suffix=".npz", dir=path.parent, delete=False + ) as stream: + temporary_path = Path(stream.name) + np.savez_compressed( + stream, + schema_version=np.asarray(self.SCHEMA_VERSION, dtype=np.int64), + mask=self._preview_mask.astype(bool), + mask_type=np.asarray(self.geometry), + geometry=np.asarray(self.geometry), + center_row=np.asarray(self.y, dtype=np.int64), + center_column=np.asarray(self.x, dtype=np.int64), + radius=np.asarray(self.radius, dtype=np.int64), + size_x=np.asarray(self.size_x, dtype=np.int64), + size_y=np.asarray(self.size_y, dtype=np.int64), + scan_shape=np.asarray(self.scan_shape, dtype=np.int64), + sampling=sampling, + units=units, + ) + temporary_path.replace(path) + self.state_path = path + self.save_button.disabled = False + self._saved = True + self._dirty = False + self._refresh_status(f"Applied and saved to {path}") + return path + + def close(self): + plt.close(self.figure) + + def _restore_initial(self): + x, y, geometry, size_x, size_y = self._initial_geometry + self.set_mask( + x=x, + y=y, + geometry=geometry, + size_x=size_x, + size_y=size_y, + ) + self._preview_mask = self._initial_mask.copy() + self._dirty = True + self._render() + self._refresh_status("Initial state restored; click Apply") + + def _set_full_scan(self): + rows, columns = self.scan_shape + if self.geometry == "circle": + size_x = size_y = ( + int(np.ceil(np.hypot(rows - 1, columns - 1))) + 1 + ) + elif self.geometry == "ellipse": + size_x, size_y = columns, rows + elif self.geometry == "square": + size_x = size_y = max(rows, columns) + else: + size_x, size_y = columns, rows + self.set_mask( + x=columns // 2, + y=rows // 2, + size_x=size_x, + size_y=size_y, + ) + + def _on_x_slider(self, change): + if self._syncing: + return + self._syncing = True + self.x_input.value = int(change["new"]) + self._syncing = False + self._update_preview() + + def _on_y_slider(self, change): + if self._syncing: + return + self._syncing = True + self.y_input.value = self.scan_shape[0] - 1 - int(change["new"]) + self._syncing = False + self._update_preview() + + def _on_geometry(self, change): + if self._syncing: + return + self._syncing = True + try: + if change["new"] in {"circle", "square"}: + self.size_y_slider.value = self.size_x + self.size_y_input.value = self.size_x + finally: + self._syncing = False + self._refresh_geometry_controls() + self._update_preview() + + def _on_size_x_slider(self, change): + if self._syncing: + return + self._syncing = True + try: + self.size_x_input.value = int(change["new"]) + if self.geometry in {"circle", "square"}: + self.size_y_slider.value = int(change["new"]) + self.size_y_input.value = int(change["new"]) + finally: + self._syncing = False + self._update_preview() + + def _on_size_y_slider(self, change): + if self._syncing: + return + self._syncing = True + self.size_y_input.value = int(change["new"]) + self._syncing = False + self._update_preview() + + def _on_x_input(self, change): + if self._syncing: + return + self._syncing = True + self.x_slider.value = int(change["new"]) + self._syncing = False + self._update_preview() + + def _on_y_input(self, change): + if self._syncing: + return + self._syncing = True + self.y_slider.value = self.scan_shape[0] - 1 - int(change["new"]) + self._syncing = False + self._update_preview() + + def _on_size_x_input(self, change): + if self._syncing: + return + self._syncing = True + try: + self.size_x_slider.value = int(change["new"]) + if self.geometry in {"circle", "square"}: + self.size_y_slider.value = int(change["new"]) + self.size_y_input.value = int(change["new"]) + finally: + self._syncing = False + self._update_preview() + + def _on_size_y_input(self, change): + if self._syncing: + return + self._syncing = True + self.size_y_slider.value = int(change["new"]) + self._syncing = False + self._update_preview() + + def _update_preview(self): + self._preview_mask = self._geometry_mask( + self.y, self.x, self.geometry, self.size_x, self.size_y + ) + self._dirty = True + self._saved = False + self._render() + self._refresh_status("Preview changed; click Apply or Apply & Save") + + def _refresh_geometry_controls(self): + labels = { + "circle": ("Radius", None), + "ellipse": ("X radius", "Y radius"), + "square": ("Half-size", None), + "rectangle": ("Half-width", "Half-height"), + } + x_label, y_label = labels[self.geometry] + self.size_x_slider.description = x_label + self.size_x_input.description = x_label + if y_label is None: + self.size_y_slider.layout.display = "none" + self.size_y_input.layout.display = "none" + self.size_y_row.layout.display = "none" + else: + self.size_y_slider.description = y_label + self.size_y_input.description = y_label + self.size_y_slider.layout.display = "" + self.size_y_input.layout.display = "" + self.size_y_row.layout.display = "" + + def _replace_boundary_artist(self): + if self.boundary_artist is not None: + self.boundary_artist.remove() + style = { + "fill": False, + "edgecolor": "#ff3030", + "linewidth": 1.05, + "linestyle": (0, (1.2, 5.5)), + "alpha": 0.9, + } + if self.geometry in {"circle", "ellipse"}: + size_y = self.size_x if self.geometry == "circle" else self.size_y + artist = Ellipse( + (self.x, self.y), + width=2 * self.size_x, + height=2 * size_y, + **style, + ) + else: + size_y = self.size_x if self.geometry == "square" else self.size_y + artist = Rectangle( + (self.x - self.size_x - 0.5, self.y - size_y - 0.5), + width=2 * self.size_x + 1, + height=2 * size_y + 1, + **style, + ) + self.ax.add_patch(artist) + self.boundary_artist = artist + # Compatibility name retained even when the selected geometry is not circular. + self.circle_artist = artist + + def _render(self): + self.mask_artist.set_data( + np.ma.masked_where(~self._preview_mask, self._preview_mask) + ) + self._replace_boundary_artist() + self.center_artist.set_data([self.x], [self.y]) + self.figure.canvas.draw_idle() + with self.output: + clear_output(wait=True) + display(self.figure) + + def _refresh_status(self, message): + count = int(self._preview_mask.sum()) + total = int(self._preview_mask.size) + physical = "" + try: + row_sampling, column_sampling = ( + float(v) for v in self.analysis.dataset_cartesian.sampling[:2] + ) + row_unit, column_unit = ( + str(v) for v in self.analysis.dataset_cartesian.units[:2] + ) + if np.isclose(row_sampling, column_sampling) and row_unit == column_unit: + if self.geometry == "circle": + physical = ( + f" · radius ≈ {self.size_x * row_sampling:.4g} {row_unit}" + ) + else: + physical = ( + f" · half-size ≈ " + f"{self.size_x * column_sampling:.4g} × " + f"{self.size_y * row_sampling:.4g} {row_unit}" + ) + except (TypeError, ValueError): + pass + save_state = "saved" if self._saved else "not saved" + self.status.value = ( + f"{message}
" + f"{self.geometry.title()} · X column {self.x} · Y row {self.y} · " + f"half-size {self.size_x} × {self.size_y} px" + f"{physical} · {count:,}/{total:,} positions " + f"({100.0 * count / total:.1f}%) · {save_state}" + ) + + +# TODO: Likely dataset4dSTEM rather than dataset4d input class +# Bragg peaks from crystalline vs polymer +# +# TODO: "BraggPeaksPolymer" vs "BraggPeaksCrystal" +class BraggPeaksPolymer(AutoSerialize): + """ + + """ + + _token = object() + + def __init__( + self, + dataset_cartesian: Dataset4dstem, + compute_parameters: callable = None, + normalize_data: callable = None, + normalization_strategy: NormalizationStrategy | str | dict | None = None, + model: MultiChannelCNN2d = None, + final_shape: Tuple[int, int] = (256, 256), + device: str = 'cpu', + normalize_parameter_lower_percentile: float = 1.0, + normalize_parameter_upper_percentile: float = 99.0, + _token: object | None = None, + ): + if _token is not self._token: + raise RuntimeError( + "Use BraggPeaks.from_data() or .from_file() to instantiate this class." + ) + + self._dataset_cartesian = dataset_cartesian + self._device = device + self._final_shape = final_shape + self.normalize_parameter_lower_percentile = normalize_parameter_lower_percentile + self.normalize_parameter_upper_percentile = normalize_parameter_upper_percentile + if (compute_parameters is None) != (normalize_data is None): + raise ValueError( + "compute_parameters and normalize_data must be supplied together." + ) + if normalization_strategy is not None and compute_parameters is not None: + raise ValueError( + "Pass normalization_strategy or the legacy callback pair, not both." + ) + if compute_parameters is not None: + warnings.warn( + "compute_parameters and normalize_data are deprecated; pass a " + "normalization_strategy instead.", + DeprecationWarning, + stacklevel=2, + ) + normalization_strategy = LegacyNormalizationAdapter( + compute_parameters, + normalize_data, + normalize_parameter_lower_percentile, + normalize_parameter_upper_percentile, + ) + self.compute_parameters = compute_parameters + self.normalize_data = normalize_data + self._normalization_strategy = ( + resolve_normalization_strategy(normalization_strategy) + if normalization_strategy is not None + else None + ) + self._normalization_is_explicit = normalization_strategy is not None + # To be set by class methods + # self.resized_cartesian_data = None + self.peak_coordinates_cartesian = None + self.peak_intensities = None + self.image_centers = None + # Calibration parameters cached by preprocess() (lazy: applied downstream by + # the polar transforms, the raw 4D data is left untouched). + self.ellipse_params = None # (a, b, theta_deg) + self.ellipse_center = None # (row, col) of the mean-DP ellipse fit + self.descan_origin = None # (2, Ry, Rx) plane-fitted CoM background + self.origin_com_measured = None # (2, Ry, Rx) raw per-pattern CoM + self.detector_rotation_deg = None # r->q rotation (clockwise, degrees) + self.detector_transpose = None # detector transpose flag + self.sampling_inv_A = None # detector-pixel sampling in 1/A + self.polar_data = None + self.polar_peaks = None + self.max_radius = None + self.num_radial_bins = None + self.num_annular_bins = None + self.orient_corr = None + self.orient_corr_pairs = None + # Cached dataset-level normalization stats (median, iqr). Computed once by + # find_peaks_model / ensure_normalization_params and reused for live inference + # so single-DP predictions reproduce the full-scan results exactly. + self._normalization_parameters = None + # Deprecated cache aliases retained for serialized historical objects. + self._norm_median = None + self._norm_iqr = None + # True once BatchNorm running stats have been adapted to this dataset (for + # eval-mode single-DP inference); see adapt_batchnorm / infer_peaks_single. + self._bn_adapted = False + # Set when an angular detector calibration must be converted to reciprocal + # length. None means that the documented 300 kV default has not yet been + # accepted or overridden by the user. + self._accelerating_voltage_kv = None + # Cache of the most recent train-mode chunk output for live inference + # (bn_mode="train_batch"): (chunk_start, chunk_size, outs). Lets neighbouring + # cursor positions in the same find_peaks_model chunk reuse one forward pass. + self._live_chunk_cache = None + # Scan mask (region of interest) remembered from find_peaks_model / process_polar, + # so normalization + BN adaptation restrict to the sample ROI (see scan_mask). + self._scan_mask = None + + if model is None: + # Setup model + input_channels = 1 # 1 for a greyscale image, 3 for RGB, 4 for RGBA, etc. + k_size = 3 + # k_size = 7 + num_layers = 4 + start_filters = 32 + num_per_layer = 3 + # num_per_layer = 2 + use_skip_connections = True + dtype = torch.float32 + # The immutable paper checkpoint was trained with dropout disabled. + dropout = 0.0 + model = MultiChannelCNN2d( + in_channels=input_channels, + out_channels=2, + start_filters=start_filters, + num_layers=num_layers, + num_per_layer=num_per_layer, + use_skip_connections=use_skip_connections, + dtype=dtype, + dropout=dropout, + final_activations=["sigmoid", "sigmoid"], + conv_kernel_size=k_size, + ) + self._model = model + + @property + def model(self) -> MultiChannelCNN2d: + return self._model + + @model.setter + def model(self, model): + self._model = model + self._invalidate_inference_caches() + + @property + def normalization_strategy(self): + return self._normalization_strategy + + @normalization_strategy.setter + def normalization_strategy(self, strategy): + self._set_normalization_strategy(strategy, explicit=True) + + def _set_normalization_strategy(self, strategy, *, explicit): + resolved = ( + resolve_normalization_strategy(strategy) if strategy is not None else None + ) + if resolved != getattr(self, "_normalization_strategy", None): + self._normalization_strategy = resolved + self._invalidate_inference_caches() + self._normalization_is_explicit = explicit + + def _invalidate_inference_caches(self): + self._normalization_parameters = None + self._norm_median = None + self._norm_iqr = None + self._bn_adapted = False + self._live_chunk_cache = None + + def _require_normalization_strategy(self): + if self._normalization_strategy is None: + raise RuntimeError( + "No inference normalization is configured. Load a registered model, " + "or pass normalization_strategy (or the legacy compute_parameters and " + "normalize_data callbacks) when using a custom checkpoint." + ) + return self._normalization_strategy + + @property + def device(self) -> str: + return self._device + + @device.setter + def device(self, device): + self._device = device + + @property + def dataset_cartesian(self) -> Dataset4dstem: + return self._dataset_cartesian + + @dataset_cartesian.setter + def dataset_cartesian(self, dataset_cartesian): + self._dataset_cartesian = dataset_cartesian + + @property + def final_shape(self) -> str: + return self._final_shape + + @final_shape.setter + def final_shape(self, final_shape): + self._final_shape = final_shape + + @property + def scan_mask(self): + """Boolean (Ry, Rx) region-of-interest mask, or None for the whole scan. + + Remembered from ``find_peaks_model`` (and settable directly) so that + ``ensure_normalization_params`` / ``adapt_batchnorm`` estimate their statistics + from the sample ROI rather than off-sample regions (vacuum, edges, beam stop). + """ + return self._scan_mask + + @scan_mask.setter + def scan_mask(self, mask): + if mask is None: + new_mask = None + else: + new_mask = np.asarray(mask, dtype=bool) + Ry, Rx = int(self._dataset_cartesian.shape[0]), int(self._dataset_cartesian.shape[1]) + if new_mask.shape != (Ry, Rx): + raise ValueError( + f"scan_mask shape {new_mask.shape} must match scan shape ({Ry}, {Rx})" + ) + # Only invalidate the lazily-cached stats if the mask actually changed, so + # re-running find_peaks_model with the same mask doesn't needlessly recompute. + changed = not ( + (self._scan_mask is None and new_mask is None) + or ( + self._scan_mask is not None + and new_mask is not None + and np.array_equal(self._scan_mask, new_mask) + ) + ) + self._scan_mask = new_mask + if changed: + self._invalidate_inference_caches() + + @classmethod + def from_file( + cls, + file_path: str, + device: str = "cpu", + compute_parameters: callable = None, + normalize_data: callable = None, + normalization_strategy: NormalizationStrategy | str | dict | None = None, + file_type: str | None = None, + normalize_parameter_lower_percentile: float = 1.0, + normalize_parameter_upper_percentile: float = 99.0, + ) -> "BraggPeaksPolymer": + dataset_cartesian = Dataset4dstem.from_file(file_path, file_type=file_type) + return cls.from_data( + dataset_cartesian=dataset_cartesian, + device=device, + compute_parameters=compute_parameters, + normalize_data=normalize_data, + normalization_strategy=normalization_strategy, + normalize_parameter_lower_percentile=normalize_parameter_lower_percentile, + normalize_parameter_upper_percentile=normalize_parameter_upper_percentile, + ) + + @classmethod + def from_data( + cls, + dataset_cartesian: Dataset4dstem, + device: str = "cpu", + compute_parameters: callable = None, + normalize_data: callable = None, + normalization_strategy: NormalizationStrategy | str | dict | None = None, + normalize_parameter_lower_percentile: float = 1.0, + normalize_parameter_upper_percentile: float = 99.0, + ) -> "BraggPeaksPolymer": + return cls( + dataset_cartesian=dataset_cartesian, + _token=cls._token, + device=device, + compute_parameters=compute_parameters, + normalize_data=normalize_data, + normalization_strategy=normalization_strategy, + normalize_parameter_lower_percentile=normalize_parameter_lower_percentile, + normalize_parameter_upper_percentile=normalize_parameter_upper_percentile, + ) + + def pixels_to_inv_A(self, accelerating_voltage_kv: float = None): + """Return the detector-pixel sampling in inverse angstroms. + + Angular calibrations in mrad require the electron wavelength. If no voltage + has previously been supplied, 300 kV is assumed with an explicit warning. + Supplying a voltage stores it for subsequent reciprocal-space operations. + """ + unit = str(self.dataset_cartesian.units[2]).strip().lower() + sampling = self.dataset_cartesian.sampling[2] + + if unit == "mrad": + if accelerating_voltage_kv is None: + accelerating_voltage_kv = self._accelerating_voltage_kv + if accelerating_voltage_kv is None: + accelerating_voltage_kv = 300.0 + warnings.warn( + "Detector calibration is in mrad; assuming an accelerating " + "voltage of 300 kV for conversion to 1/Å. Pass " + "accelerating_voltage_kv to find_peaks_model() to override it.", + UserWarning, + stacklevel=2, + ) + if not np.isfinite(accelerating_voltage_kv) or accelerating_voltage_kv <= 0: + raise ValueError("accelerating_voltage_kv must be a positive finite value") + + self._accelerating_voltage_kv = float(accelerating_voltage_kv) + wavelength_angstrom = electron_wavelength_angstrom( + self._accelerating_voltage_kv * 1e3 + ) + return sampling / (1e3 * wavelength_angstrom) + + _, sampling_angstrom_conversion_factor = parse_reciprocal_units( + self.dataset_cartesian.units[2] + ) + return sampling * sampling_angstrom_conversion_factor + + def preprocess( + self, + accelerating_voltage_kv: float | None = None, + *, + center_source: str = "descent", + fit_ellipse: bool = True, + ellipse_fit_method: str = "angular_variance", + ellipse_threshold: float | None = None, + ellipse_radial_min: float | None = None, + ellipse_radial_max: float | None = None, + ellipse_device: str | None = None, + estimate_descan: bool = True, + descan_fit_method: str = "plane", + estimate_detector_rotation: bool = True, + scan_mask: ArrayLike = None, + center_device: str | None = None, + com_device: str | None = None, + com_batch_size: int | None = None, + store_metadata: bool = True, + show: bool = False, + verbose: bool = True, + ): + """Calibrate the 4D-STEM scan (centers, ellipticity, descan, detector rotation). + + This mirrors the lazy design of the rest of the class: it *measures and caches* + calibration parameters rather than re-warping the raw diffraction data (the ML + peak-finder runs on raw patterns; centers/ellipticity are applied downstream by + ``process_polar`` / the polar transforms). After ``preprocess`` you can call + ``process_polar(center_ellipse_params=bp.ellipse_params)`` and the cached + ``image_centers`` will be reused. + + Steps performed (each individually toggleable). Order matters: centering runs + first so the ellipse is measured on an already-centered mean DP -- fitting the + ellipse on the raw mean DP smears the ring by the descan drift and biases the + fit toward the central beam. + + 1. **Descan / detector rotation** (``estimate_descan`` / + ``estimate_detector_rotation``) -- a ``CenterOfMassOriginModel`` measures the + per-pattern centre of mass, fits a smooth background across scan positions + (``descan_fit_method``), and estimates the r->q detector rotation + transpose. + Results are cached on ``self.descan_origin`` (2, Ry, Rx), + ``self.detector_rotation_deg``, ``self.detector_transpose`` and (optionally) + ``dataset_cartesian.metadata["r_to_q_rotation_cw_deg"]``. + 2. **Image centers** -- ``self.image_centers`` (2, Ry, Rx), the per-pattern + origins consumed by the polar transforms. ``center_source`` selects the + estimator: ``"descent"`` / ``"grid"`` / ``"peaks"`` use + ``find_central_beams_4d`` (angular-uniformity, the pipeline default), + ``"com"`` uses the raw centre of mass, ``"descan"`` uses the plane-fitted + (descanned) origin field. + 3. **Ellipticity** (``fit_ellipse``), fit LAST -- a diffuse-ring fit on a + *centered* mean DP (each pattern shifted so its central beam sits at the + detector center, then averaged; see ``_centered_dp_mean``). The + ``"angular_variance"`` method searches ``(b/a, theta)`` to minimise the + annulus' azimuthal variance. The ``"ridge"`` method extracts the ring radius + independently at each azimuth, uses its first two harmonics to initialise the + center and ellipse, and accepts a robust joint refinement only when it + improves held-out angular sectors over a circle. Both methods ignore the + central beam itself. The centered mean DP is cached on + ``self.dp_mean_centered``; ``ellipse_radial_min`` / + ``ellipse_radial_max`` bound the ring band (auto-detected from the radial + profile when None). Stored as + ``self.ellipse_params = (a, b, theta_deg)`` and (optionally) into + ``dataset_cartesian.metadata["ellipticity"]``. ``ellipse_threshold`` is kept + for backward compatibility but is unused by the ring fit. + 4. **Reciprocal sampling** -- caches ``self.sampling_inv_A`` via + ``pixels_to_inv_A`` (accepts ``accelerating_voltage_kv`` for mrad detectors). + + Parameters + ---------- + accelerating_voltage_kv : float, optional + Beam voltage for mrad->1/A conversion (see ``pixels_to_inv_A``). + center_source : {"descent", "grid", "peaks", "com", "descan"} + Estimator backing ``self.image_centers``. Default "descent". + fit_ellipse : bool + Fit ellipticity from the mean DP. Default True. + ellipse_fit_method : {"angular_variance", "ridge"} + Angular-variance search or robust diffuse-ring ridge refinement. + Default ``"angular_variance"`` during ridge-method validation. + ellipse_threshold : float, optional + Binarisation threshold for ``fit_probe_ellipse`` (Otsu if None). + estimate_descan : bool + Run the CoM + background-fit descan estimate. Default True. + descan_fit_method : {"plane", "constant"} + Background model for the descan fit. Default "plane". + estimate_detector_rotation : bool + Estimate the r->q detector rotation + transpose (requires the CoM model, + so it forces ``estimate_descan``). Default True. + scan_mask : ArrayLike, optional + Boolean (Ry, Rx) ROI passed to ``find_central_beams_4d``. + center_device, com_device : str, optional + Device overrides for the angular-uniformity finder and the CoM model + respectively (both default to ``self.device``). Note the CoM model loads + the whole 4D tensor onto its device at once. + com_batch_size : int, optional + Batch size for the CoM origin calculation (whole scan if None). + store_metadata : bool + Write ellipticity / rotation into ``dataset_cartesian.metadata``. Default True. + show : bool + Show the ellipse-fit overlay. Default False. + verbose : bool + Print a short calibration summary. Default True. + + Returns + ------- + dict + The calibration parameters that were computed. + """ + center_source = center_source.lower() + valid_sources = ("descent", "grid", "peaks", "com", "descan") + if center_source not in valid_sources: + raise ValueError(f"center_source must be one of {valid_sources}, got {center_source!r}") + ellipse_fit_method = str(ellipse_fit_method).lower() + if ellipse_fit_method not in {"angular_variance", "ridge"}: + raise ValueError( + "ellipse_fit_method must be 'angular_variance' or 'ridge', " + f"got {ellipse_fit_method!r}" + ) + + Ry, Rx, Qy, Qx = self._dataset_cartesian.shape + need_com = ( + estimate_descan + or estimate_detector_rotation + or center_source in ("com", "descan") + ) + + results: dict = {} + + # 1. Descan (CoM + background fit) and detector rotation come FIRST: the + # per-pattern central-beam CoM and its smooth drift model are what let us + # build a properly centered mean DP for the ellipse fit in step 3. + self.descan_origin = None + self.origin_com_measured = None + self.detector_rotation_deg = None + self.detector_transpose = None + com_model = None + if need_com: + from quantem.diffractive_imaging.origin_models import CenterOfMassOriginModel + + com_dev = com_device if com_device is not None else self.device + com_model = CenterOfMassOriginModel.from_dataset( + self._dataset_cartesian, device=com_dev + ) + com_model.calculate_origin(max_batch_size=com_batch_size) + measured = com_model.origin_measured.detach().cpu().numpy().reshape(Ry, Rx, 2) + self.origin_com_measured = np.moveaxis(measured, -1, 0) # (2, Ry, Rx) + results["origin_com_measured"] = self.origin_com_measured + + # Restrict the descan/rotation fit to the ROI: out-of-mask patterns + # (vacuum/substrate) have meaningless CoM that drags a global plane, leaving a + # uniform residual inside the ROI. Fit the plane over in-mask patterns only + # (sigma-clipped to reject hot/dead-pixel CoM outliers), push it back into the + # CoM model, and zero the residual outside the ROI so the detector-rotation + # curl isn't contaminated by junk patterns either. + roi = ( + np.asarray(scan_mask, dtype=bool) + if scan_mask is not None + else np.ones((Ry, Rx), dtype=bool) + ) + if estimate_descan or estimate_detector_rotation or center_source == "descan": + import torch + + fitted = self._fit_origin_roi(measured, roi, fit_method=descan_fit_method) + self.descan_origin = np.moveaxis(fitted, -1, 0) # (2, Ry, Rx) + results["descan_origin"] = self.descan_origin + + dev = com_model.device + com_model.origin_fitted = torch.as_tensor( + fitted.reshape(-1, 2), dtype=torch.float, device=dev + ) + meas_clean = measured.copy() + meas_clean[~roi] = fitted[~roi] # residual := 0 outside the ROI + com_model.origin_measured = torch.as_tensor( + meas_clean.reshape(-1, 2), dtype=torch.float, device=dev + ) + + if estimate_detector_rotation: + com_model.estimate_detector_rotation() + self.detector_rotation_deg = float(com_model.detector_rotation_deg) + self.detector_transpose = bool(com_model.detector_transpose) + results["detector_rotation_deg"] = self.detector_rotation_deg + results["detector_transpose"] = self.detector_transpose + if store_metadata: + self._dataset_cartesian.metadata["r_to_q_rotation_cw_deg"] = ( + self.detector_rotation_deg + ) + + # 2. Per-pattern image centers consumed by the polar transforms. Centering + # runs BEFORE the ellipse fit (ellipse_params intentionally None here) so + # the ellipse is measured on an already-centered mean DP, not the reverse. + if center_source in ("descent", "grid", "peaks"): + self.image_centers = self.find_central_beams_4d( + scan_mask=scan_mask, + center_method=center_source, + ellipse_params=None, + center_device=center_device, + ) + elif center_source == "com": + self.image_centers = self.origin_com_measured.copy() + else: # "descan" + self.image_centers = self.descan_origin.copy() + results["image_centers"] = self.image_centers + + # 3. Ellipticity LAST, fit on a mean DP that has been centered so the central + # beam sits at the detector center and the diffraction ring is concentric. + # Both supported methods fit the diffuse ring rather than the probe blob; + # the ridge method may additionally remove a small residual center offset. + self.ellipse_params = None + self.ellipse_center = None + self.dp_mean_centered = None + self.ellipse_fit_diagnostics = None + if fit_ellipse: + self.dp_mean_centered = self._centered_dp_mean( + self.image_centers, com_model=com_model + ) + Qy, Qx = self._dataset_cartesian.shape[-2:] + center = ((Qy - 1) / 2.0, (Qx - 1) / 2.0) # _centered_dp_mean puts the beam here + fit_function = ( + self._fit_ellipse_from_ridge + if ellipse_fit_method == "ridge" + else self._fit_ellipse_from_ring + ) + a_axis, b_axis, theta_deg, ring_band = fit_function( + self.dp_mean_centered, + center, + radial_min=ellipse_radial_min, + radial_max=ellipse_radial_max, + device=ellipse_device if ellipse_device is not None else "cpu", + show=show, + verbose=verbose, + ) + self.ellipse_params = (float(a_axis), float(b_axis), float(theta_deg)) + refined_center = self.ellipse_fit_diagnostics.get( + "center_refined", center + ) + self.ellipse_center = tuple(float(value) for value in refined_center) + if ( + ellipse_fit_method == "ridge" + and self.ellipse_fit_diagnostics["accepted"] + ): + center_delta = np.asarray(self.ellipse_center) - np.asarray(center) + self.image_centers = np.asarray( + self.image_centers, dtype=float + ).copy() + valid_centers = ( + (self.image_centers[0] != 0) + | (self.image_centers[1] != 0) + ) + self.image_centers[0, valid_centers] += center_delta[0] + self.image_centers[1, valid_centers] += center_delta[1] + results["image_centers"] = self.image_centers + self.ellipse_ring_band = ring_band + results["ellipse_params"] = self.ellipse_params + results["ellipse_center"] = self.ellipse_center + results["ellipse_ring_band"] = ring_band + results["ellipse_fit_diagnostics"] = self.ellipse_fit_diagnostics + results["ellipse_fit_method"] = ellipse_fit_method + if store_metadata: + self._dataset_cartesian.metadata["ellipticity"] = self.ellipse_params + + # 4. Reciprocal-space sampling (pixels -> 1/A). + try: + self.sampling_inv_A = float(self.pixels_to_inv_A(accelerating_voltage_kv)) + results["sampling_inv_A"] = self.sampling_inv_A + except Exception as exc: # calibration/units may be unavailable + self.sampling_inv_A = None + if verbose: + print(f"preprocess: reciprocal calibration skipped ({exc})") + + if verbose: + print(f"preprocess: device={self.device}, scan=({Ry}, {Rx}), detector=({Qy}, {Qx})") + print(f" image_centers <- {center_source} shape {self.image_centers.shape}") + if self.ellipse_params is not None: + a, b, th = self.ellipse_params + print(f" ellipticity a={a:.3f} b={b:.3f} theta={th:.2f} deg (a/b={a / b:.4f})") + if self.descan_origin is not None: + print(f" descan {descan_fit_method}-fit CoM background") + if self.detector_rotation_deg is not None: + print( + f" r->q rotation {self.detector_rotation_deg:.2f} deg " + f"(transpose={self.detector_transpose})" + ) + if self.sampling_inv_A is not None: + print(f" sampling {self.sampling_inv_A:.5g} 1/A per pixel") + + return results + + def _fit_origin_roi(self, measured, mask, fit_method="plane", clip_sigma=5.0, n_iter=2): + """Fit the smooth CoM-origin background over the ROI (scan mask) only. + + ``measured`` is the per-pattern CoM origin ``(Ry, Rx, 2)``; ``mask`` is the + ``(Ry, Rx)`` scan ROI. Fitting over the whole scan lets out-of-ROI patterns + (vacuum/substrate, whose CoM is meaningless) drag the plane, leaving a uniform + residual inside the ROI. Each component is fit independently: ``"plane"`` does an + ordinary least-squares ``z = a + b*row + c*col`` with ``n_iter`` sigma-clip passes + to reject hot/dead-pixel CoM outliers; ``"constant"`` uses the ROI mean. Returns + the fitted field ``(Ry, Rx, 2)`` evaluated at every scan position. + """ + measured = np.asarray(measured, dtype=float) + Ry, Rx, ncomp = measured.shape + m0 = np.asarray(mask, dtype=bool) + yy, xx = np.mgrid[0:Ry, 0:Rx].astype(float) + fitted = np.empty_like(measured) + for c in range(ncomp): + z = measured[..., c] + if fit_method == "constant": + ref = z[m0] if m0.any() else z + fitted[..., c] = float(np.nanmean(ref)) + continue + use = m0 & np.isfinite(z) + plane = np.full((Ry, Rx), float(np.nanmean(z[use])) if use.any() else 0.0) + for _ in range(max(1, n_iter)): + if int(use.sum()) < 3: + break + A = np.stack([np.ones(int(use.sum())), xx[use], yy[use]], axis=1) + coef, *_ = np.linalg.lstsq(A, z[use], rcond=None) + plane = coef[0] + coef[1] * xx + coef[2] * yy + resid = z - plane + s = float(np.std(resid[use])) + if s == 0: + break + new_use = m0 & np.isfinite(z) & (np.abs(resid) < clip_sigma * s) + if int(new_use.sum()) == int(use.sum()) or int(new_use.sum()) < 8: + break + use = new_use + fitted[..., c] = plane + return fitted + + def _centered_dp_mean(self, image_centers, com_model=None): + """Mean diffraction pattern with every pattern shifted so its central beam + lands at the detector center. + + Averaging the raw patterns smears the diffraction ring by the descan drift and + leaves the central beam off-center, which biases an ellipse fit toward the + central beam. Aligning each pattern first yields a sharp, concentric ring. + + When a ``CenterOfMassOriginModel`` is available (it holds the 4D tensor on-device), + each ROI pattern is SUB-PIXEL shifted (bilinear) by ``image_centers`` -- the + angular-uniformity BEAM center -- so the beam lands exactly on the detector center, + and the mean is accumulated in batches. We center by ``image_centers`` rather than + the CoM/descan origin because the CoM is the centroid of the whole pattern: any + ring or background asymmetry pulls it a few px off the actual beam, which would + leave the beam off-center in the mean. Otherwise (no CoM model) we fall back to + translating the plain mean DP by the average center offset. + + Sub-pixel matters: plain integer rolls leave a per-pattern residual of up to 0.5 px + that does NOT average out when the center spread is narrow (all patterns round to + the same integer), so the mean beam ends up biased off-center by a fraction of a + pixel. Bilinear splatting removes that bias. + + We deliberately avoid ``CenterOfMassOriginModel.shift_origin_to`` here: it + materialises a full second copy of the 4D stack *plus* a per-pattern sampling + grid, which OOMs on large scans. The bilinear accumulator needs only a + ``(Qy, Qx)`` buffer plus one batch of patterns at a time. + """ + Qy, Qx = self._dataset_cartesian.shape[-2:] + center = ((Qy - 1) / 2.0, (Qx - 1) / 2.0) + ic = np.asarray(image_centers, dtype=float) # (2, Ry, Rx); 0 outside the scan mask + valid = (ic[0] != 0) | (ic[1] != 0) + if com_model is not None and valid.any(): + import torch + + with torch.no_grad(): + flat = com_model.tensor.reshape(-1, Qy, Qx) + dev = flat.device + # Center by the authoritative per-pattern beam centers (image_centers, from + # the angular-uniformity finder), NOT the CoM plane: the CoM is the centroid + # of the whole pattern, so ring/background asymmetry pulls it a few px off + # the beam, and centering by it would leave the beam off-center. Only ROI + # (in-mask, non-zero) patterns are averaged. + oy = torch.as_tensor(ic[0].ravel(), dtype=torch.float, device=dev) + ox = torch.as_tensor(ic[1].ravel(), dtype=torch.float, device=dev) + keep = ( + torch.as_tensor(valid.ravel(), device=dev) + .nonzero(as_tuple=False) + .squeeze(1) + ) + sy = center[0] - oy # continuous shift -> detector center, (y, x) + sx = center[1] - ox + acc = torch.zeros((Qy, Qx), dtype=torch.float32, device=dev) + batch = 256 + for bstart in range(0, int(keep.numel()), batch): + bidx = keep[bstart:bstart + batch] + chunk = flat[bidx].float() # (b, Qy, Qx), ROI patterns only + fyb = torch.floor(sy[bidx]) + fxb = torch.floor(sx[bidx]) + gy = sy[bidx] - fyb + gx = sx[bidx] - fxb + floor_pairs = torch.stack([fyb.long(), fxb.long()], dim=1) # (b, 2) + # bilinear weights for the 4 integer-shift corners around the fraction + corner_w = { + (0, 0): (1 - gy) * (1 - gx), + (0, 1): (1 - gy) * gx, + (1, 0): gy * (1 - gx), + (1, 1): gy * gx, + } + # descan drift is smooth -> few distinct floor shifts per batch: + # weight-sum each group, then roll each of the 4 corners. + uniq, inv = torch.unique(floor_pairs, dim=0, return_inverse=True) + for k in range(int(uniq.shape[0])): + m = inv == k + cg = chunk[m] # (g, Qy, Qx) + fy = int(uniq[k, 0]) + fx = int(uniq[k, 1]) + for (dy, dx), wt in corner_w.items(): + s = (cg * wt[m][:, None, None]).sum(0) + acc += torch.roll(s, shifts=(fy + dy, fx + dx), dims=(0, 1)) + dp = (acc / max(int(keep.numel()), 1)).detach().cpu().numpy() + return np.asarray(dp, dtype=float) + # Fallback: translate the raw mean DP so the average beam center is centered. + # image_centers is 0 outside the scan mask (find_central_beams_4d), so average + # over valid (non-zero) positions only to avoid a bias toward the origin. + from scipy.ndimage import shift as ndi_shift + + dp = np.asarray(self._dataset_cartesian.get_dp_mean().array, dtype=float) + valid = (image_centers[0] != 0) | (image_centers[1] != 0) + if not valid.any(): + valid = np.ones_like(image_centers[0], dtype=bool) + dy = center[0] - float(image_centers[0][valid].mean()) + dx = center[1] - float(image_centers[1][valid].mean()) + return ndi_shift(dp, (dy, dx), order=1, mode="constant", cval=0.0) + + def _fit_ellipse_from_ring( + self, + dp, + center, + *, + radial_min=None, + radial_max=None, + radial_step=1.0, + num_annular_bins=180, + ratio_range=(0.85, 1.18), + n_ratio=12, + n_theta=24, + refine=True, + max_ring_candidates=3, + min_fit_improvement=0.005, + max_fit_score=0.25, + min_angular_coverage=0.55, + device="cpu", + show=False, + verbose=False, + ): + """Fit ring ellipticity ``(a, b, theta_deg)`` by minimising the azimuthal variance + of a diffraction-ring annulus at a FIXED center -- Karen Ehrhardt's angular- + uniformity criterion (see ``quantem.diffraction.polar_transform``). + + Unlike a probe-blob fit (``fit_probe_ellipse``) this samples an annulus out at the + ring radius and never touches the central beam, so a smeared / off-center / doubled + central beam does not bias the result. Only the axis ratio ``b/a`` and orientation + ``theta`` are identifiable from a single ring, so the returned ``(a, b)`` are + normalised to the ring radius (``a ~ R0``); downstream consumers (``polar_transform`` + / ``find_central_beams_4d``) use only ``b/a`` and ``theta``. + + Parameters + ---------- + dp : ndarray + Centered mean diffraction pattern (beam at ``center``). + center : (float, float) + Fixed origin ``(y, x)`` in detector pixels. + radial_min, radial_max : float, optional + Ring band in pixels. If either is None the band is auto-detected from the + circular median radial profile and several prominent candidates are fitted + and quality-ranked. + ratio_range, n_ratio, n_theta, refine : + Coarse grid over ``b/a`` and ``theta`` (degrees), then a clipped local + refinement pass. + max_ring_candidates : int + Maximum prominent radial-profile peaks evaluated when the ring band is + selected automatically. + min_fit_improvement, max_fit_score, min_angular_coverage : float + Quality gates. Fits that do not improve held-out angular alignment, retain + excessive raw angular variance, lack ring coverage, or hit a ratio boundary + fall back to a circular correction with a warning. + + Returns + ------- + (a, b, theta_deg, (radial_min, radial_max)) + """ + from quantem.diffraction.polar_transform import polar_transform + + dp = np.asarray(dp, dtype=float) + Qy, Qx = dp.shape + origin = np.asarray(center, dtype=float) + + def _polar(image, ellipse_params, rmin, rmax): + # polar_transform returns (n_phi, n_r) when scan_pos is given. + return np.asarray( + polar_transform( + image, + origin_array=origin, + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=float(rmin), + radial_max=float(rmax), + radial_step=radial_step, + scan_pos=(0, 0), + device=device, + show_progress=False, + ), + dtype=float, + ) + + # Log compression plus global winsorisation strongly reduces the leverage of + # isolated Bragg spots without erasing the broad diffuse calibration ring. + fit_dp = np.log1p(np.clip(dp, 0.0, None)) + finite_fit = fit_dp[np.isfinite(fit_dp)] + if finite_fit.size: + fit_dp = np.minimum(fit_dp, np.percentile(finite_fit, 99.5)) + + # 1. Find several plausible diffuse-ring bands. A median angular profile is much + # less likely than a mean profile to select a sparse constellation of Bragg + # spots. Candidate fits are quality-ranked below rather than trusting the + # single strongest radial feature. + r_hi = float(min(Qy, Qx) / 2.0 - 1.0) + explicit_band = radial_min is not None and radial_max is not None + candidate_bands = [] + if radial_min is None or radial_max is None: + from scipy.ndimage import uniform_filter1d + + prof = _polar(fit_dp, (1.0, 1.0, 0.0), 0.0, r_hi) + radial_profile = uniform_filter1d(np.median(prof, axis=0), size=5) + r_axis = np.arange(radial_profile.size) * radial_step + r_exclude = max(6.0, 0.06 * r_hi) + i0 = int(r_exclude / radial_step) + search_profile = radial_profile.copy() + search_profile[:i0] = np.min(search_profile) + prominence_floor = max( + 1e-9, 0.03 * float(np.ptp(search_profile[i0:])) + ) + peak_indices, properties = find_peaks( + search_profile, + prominence=prominence_floor, + distance=max(3, int(round(6.0 / radial_step))), + ) + valid = ( + (peak_indices >= i0) + & (r_axis[peak_indices] <= 0.92 * r_hi) + ) + peak_indices = peak_indices[valid] + prominences = properties["prominences"][valid] + if not peak_indices.size: + peak_indices = np.asarray( + [i0 + int(np.argmax(search_profile[i0:]))] + ) + prominences = np.asarray([1.0]) + order = np.argsort(prominences)[::-1][:max_ring_candidates] + for index in peak_indices[order]: + r0 = float(r_axis[index]) + half = max(6.0, 0.20 * r0) + band_min = ( + max(r_exclude, r0 - half) + if radial_min is None + else float(radial_min) + ) + band_max = ( + min(r_hi, r0 + half) + if radial_max is None + else float(radial_max) + ) + if band_max > band_min: + candidate_bands.append((band_min, band_max, r0)) + else: + candidate_bands.append( + (float(radial_min), float(radial_max), + 0.5 * (float(radial_min) + float(radial_max))) + ) + + # Deduplicate overlapping candidates created by broad/shouldered peaks. + unique_bands = [] + for band in candidate_bands: + if not any(abs(band[2] - other[2]) < 3.0 for other in unique_bands): + unique_bands.append(band) + candidate_bands = unique_bands + + fit_angles = (np.arange(num_annular_bins) // 6) % 2 == 0 + validation_angles = ~fit_angles + + def _robust_score(ellipse_params, band, angle_mask): + polar = _polar( + fit_dp, ellipse_params, band[0], band[1] + )[angle_mask] + if not polar.size: + return np.inf + # Per-radius clipping removes angularly isolated hot pixels. Per-angle + # normalisation then scores radial alignment rather than polymer texture. + upper = np.percentile(polar, 90.0, axis=0, keepdims=True) + polar = np.minimum(polar, upper) + polar = polar - np.percentile( + polar, 10.0, axis=1, keepdims=True + ) + polar = np.clip(polar, 0.0, None) + scale = np.percentile(polar, 90.0, axis=1, keepdims=True) + valid_scale = scale[:, 0] > 1e-9 + if np.count_nonzero(valid_scale) < 4: + return np.inf + polar = polar[valid_scale] / (scale[valid_scale] + 1e-9) + reference = np.median(polar, axis=0) + return float( + np.median(np.abs(polar - reference), axis=0).sum() + / (np.abs(reference).sum() + 1e-9) + ) + + def _raw_score(ellipse_params, band): + polar = _polar(dp, ellipse_params, band[0], band[1]) + return float( + polar.std(axis=0).sum() + / (np.abs(polar.mean(axis=0)).sum() + 1e-6) + ) + + def _angular_coverage(ellipse_params, band): + polar = _polar(fit_dp, ellipse_params, band[0], band[1]) + contrast = np.percentile(polar, 95.0, axis=1) - np.percentile( + polar, 20.0, axis=1 + ) + reference = np.percentile(contrast, 90.0) + if not np.isfinite(reference) or reference <= 1e-9: + return 0.0 + return float(np.mean(contrast >= 0.15 * reference)) + + def _search(ratios, thetas, band): + best = (np.inf, 1.0, 0.0) + for th in thetas: + for rat in ratios: + s = _robust_score( + (1.0, float(rat), float(th)), band, fit_angles + ) + if s < best[0]: + best = (s, float(rat), float(th)) + return best + + # 2. Fit every candidate, clip refinement to the declared search range, and + # validate on held-out angular blocks. + coarse_ratios = np.linspace(ratio_range[0], ratio_range[1], n_ratio) + coarse_thetas = np.linspace(0.0, 180.0, n_theta, endpoint=False) + diagnostics = [] + ratio_step = ( + (ratio_range[1] - ratio_range[0]) / max(n_ratio - 1, 1) + ) + for band in candidate_bands: + best = _search(coarse_ratios, coarse_thetas, band) + if refine: + _, rat0, th0 = best + dth = 180.0 / n_theta + fine_ratios = np.unique(np.clip( + np.linspace(rat0 - ratio_step, rat0 + ratio_step, 11), + ratio_range[0], + ratio_range[1], + )) + fine_thetas = ( + np.linspace(th0 - dth, th0 + dth, 11) % 180.0 + ) + fine = _search(fine_ratios, fine_thetas, band) + best = min(best, fine, key=lambda item: item[0]) + fit_score, ratio, theta = best + circle_validation = _robust_score( + (1.0, 1.0, 0.0), band, validation_angles + ) + ellipse_validation = _robust_score( + (1.0, ratio, theta), band, validation_angles + ) + improvement = ( + (circle_validation - ellipse_validation) + / max(abs(circle_validation), 1e-9) + ) + raw_score = _raw_score((1.0, ratio, theta), band) + coverage = _angular_coverage((1.0, ratio, theta), band) + boundary_limited = ( + ratio <= ratio_range[0] + 0.25 * ratio_step + or ratio >= ratio_range[1] - 0.25 * ratio_step + ) + accepted = ( + np.isfinite(fit_score) + and improvement >= min_fit_improvement + and raw_score <= max_fit_score + and coverage >= min_angular_coverage + and not boundary_limited + ) + diagnostics.append({ + "band": (float(band[0]), float(band[1])), + "r0": float(band[2]), + "ratio_b_over_a": float(ratio), + "theta_deg": float(theta % 180.0), + "fit_score": float(fit_score), + "raw_score": float(raw_score), + "validation_improvement": float(improvement), + "angular_coverage": float(coverage), + "boundary_limited": bool(boundary_limited), + "accepted": bool(accepted), + }) + + accepted_candidates = [item for item in diagnostics if item["accepted"]] + if accepted_candidates: + selected = min( + accepted_candidates, + key=lambda item: ( + item["raw_score"], + -item["validation_improvement"], + ), + ) + fit_accepted = True + else: + selected = min( + diagnostics, + key=lambda item: ( + item["boundary_limited"], + item["raw_score"], + -item["validation_improvement"], + ), + ) + fit_accepted = False + + radial_min, radial_max = selected["band"] + r0 = selected["r0"] + ratio = selected["ratio_b_over_a"] if fit_accepted else 1.0 + theta_deg = selected["theta_deg"] if fit_accepted else 0.0 + score = selected["raw_score"] + + # 3. Normalise (a, b) to the selected ring radius; only b/a and theta are + # identifiable. A rejected fit deliberately becomes a circular correction. + a_axis, b_axis = r0, r0 * ratio + # Canonicalise so a is the MAJOR semi-axis (a/b >= 1): the (a, b, theta) and + # (b, a, theta+90) parametrisations describe the same ellipse, so pick the one + # with a >= b for an unambiguous a/b >= 1 readout. + if b_axis > a_axis: + a_axis, b_axis = b_axis, a_axis + theta_deg += 90.0 + theta_deg = float(theta_deg % 180.0) + self.ellipse_fit_diagnostics = { + "method": "angular_variance", + "accepted": fit_accepted, + "selected": selected, + "candidates": diagnostics, + "explicit_band": explicit_band, + "center_initial": tuple(float(v) for v in center), + "center_refined": tuple(float(v) for v in center), + "rejection_reasons": [], + "quality_thresholds": { + "min_fit_improvement": float(min_fit_improvement), + "max_fit_score": float(max_fit_score), + "min_angular_coverage": float(min_angular_coverage), + "ratio_range": tuple(float(v) for v in ratio_range), + }, + } + if not fit_accepted: + reasons = [] + if selected["boundary_limited"]: + reasons.append("ratio search boundary") + if selected["validation_improvement"] < min_fit_improvement: + reasons.append( + f"held-out improvement {selected['validation_improvement']:.3g}" + ) + if selected["raw_score"] > max_fit_score: + reasons.append(f"raw score {selected['raw_score']:.3g}") + if selected["angular_coverage"] < min_angular_coverage: + reasons.append( + f"angular coverage {selected['angular_coverage']:.1%}" + ) + message = ( + "Ellipse fit rejected; using a circular correction" + + (f" ({', '.join(reasons)})." if reasons else ".") + ) + self.ellipse_fit_diagnostics["rejection_reasons"] = reasons + warnings.warn(message, RuntimeWarning, stacklevel=2) + if verbose: + bands_text = ", ".join( + f"{item['r0']:.1f}" for item in diagnostics + ) + print( + f" ellipse ring candidates: r0=[{bands_text}] px; " + f"selected band=[{radial_min:.1f}, {radial_max:.1f}] px" + ) + print( + f" ellipse ring fit: {'accepted' if fit_accepted else 'rejected'} " + f"a/b={a_axis / b_axis:.4f} theta={theta_deg:.2f} deg " + f"(score={score:.4g}, held-out improvement=" + f"{selected['validation_improvement']:.2%}, " + f"coverage={selected['angular_coverage']:.1%})" + ) + + if show: + import matplotlib.pyplot as plt + + circ = _polar(dp, (1.0, 1.0, 0.0), radial_min, radial_max) + corr = _polar( + dp, (a_axis, b_axis, theta_deg), radial_min, radial_max + ) + fig, axes = plt.subplots(1, 3, figsize=(13, 4)) + axes[0].imshow(dp, cmap="magma") + axes[0].plot([origin[1]], [origin[0]], "c+", ms=10) + axes[0].set_title("centered mean DP") + axes[1].imshow(circ, aspect="auto", cmap="magma") + axes[1].set_title("polar: circular (before)") + axes[2].imshow(corr, aspect="auto", cmap="magma") + axes[2].set_title(f"polar: ellipse-corrected\na/b={a_axis / b_axis:.4f}, θ={theta_deg:.1f}°") + for ax in axes[1:]: + ax.set_xlabel("radius (band)") + ax.set_ylabel("φ bin") + plt.tight_layout() + plt.show() + + return float(a_axis), float(b_axis), float(theta_deg), (float(radial_min), float(radial_max)) + + def _fit_ellipse_from_ridge( + self, + dp, + center, + *, + radial_min=None, + radial_max=None, + radial_step=1.0, + num_annular_bins=180, + ratio_range=(0.85, 1.0), + center_search_radius=2.5, + max_ring_candidates=3, + min_angular_coverage=0.55, + min_validation_improvement=0.05, + max_validation_residual=2.5, + device="cpu", + show=False, + verbose=False, + ): + """Jointly refine ring center and ellipticity from a robust radial ridge. + + The diffuse-ring radius is measured independently at each azimuth after + log compression and hot-spot clipping. A first/second-harmonic model + initializes center and ellipse terms, followed by bounded robust geometric + least squares. Fits are accepted only when held-out azimuthal sectors improve + over an independently refined circle. + """ + from scipy.optimize import least_squares + from quantem.diffraction.polar_transform import polar_transform + + dp = np.asarray(dp, dtype=float) + origin = np.asarray(center, dtype=float) + qy, qx = dp.shape + r_hi = float(min(qy, qx) / 2.0 - 1.0) + fit_dp = np.log1p(np.clip(dp, 0.0, None)) + finite = fit_dp[np.isfinite(fit_dp)] + if finite.size: + fit_dp = np.minimum(fit_dp, np.percentile(finite, 99.5)) + + def polar_at(image, candidate_center, rmin, rmax): + return np.asarray( + polar_transform( + image, + origin_array=np.asarray(candidate_center, dtype=float), + ellipse_params=(1.0, 1.0, 0.0), + num_annular_bins=num_annular_bins, + radial_min=float(rmin), + radial_max=float(rmax), + radial_step=radial_step, + scan_pos=(0, 0), + device=device, + show_progress=False, + ), + dtype=float, + ) + + # Candidate rings from an angular median profile; sparse Bragg spots largely + # disappear in the median rather than becoming the selected calibration ring. + explicit_band = radial_min is not None and radial_max is not None + if explicit_band: + bands = [( + float(radial_min), + float(radial_max), + 0.5 * (float(radial_min) + float(radial_max)), + )] + else: + full = polar_at(fit_dp, origin, 0.0, r_hi) + profile = gaussian_filter1d(np.median(full, axis=0), 2.0) + r_axis = np.arange(profile.size, dtype=float) * radial_step + exclude = max(6.0, 0.06 * r_hi) + start = int(np.ceil(exclude / radial_step)) + prominence = max(1e-9, 0.03 * np.ptp(profile[start:])) + indices, properties = find_peaks( + profile, + prominence=prominence, + distance=max(3, int(round(6.0 / radial_step))), + ) + valid_peaks = ( + (indices >= start) & (r_axis[indices] <= 0.92 * r_hi) + ) + indices = indices[valid_peaks] + prominences = properties["prominences"][valid_peaks] + if not indices.size: + indices = np.asarray([start + np.argmax(profile[start:])]) + prominences = np.ones(1) + order = np.argsort(prominences)[::-1][:max_ring_candidates] + bands = [] + for index in indices[order]: + r0 = float(r_axis[index]) + half = max(6.0, 0.20 * r0) + low = max(exclude, r0 - half) if radial_min is None else float(radial_min) + high = min(r_hi, r0 + half) if radial_max is None else float(radial_max) + if high > low and not any(abs(r0 - old[2]) < 3.0 for old in bands): + bands.append((low, high, r0)) + if not bands: + fallback_r0 = float(np.clip(r_axis[start], exclude, 0.92 * r_hi)) + fallback_half = max(6.0, 0.20 * fallback_r0) + bands = [( + max(exclude, fallback_r0 - fallback_half), + min(r_hi, fallback_r0 + fallback_half), + fallback_r0, + )] + + phi = np.linspace(0.0, 2.0 * np.pi, num_annular_bins, endpoint=False) + block_fit = (np.arange(num_annular_bins) // 6) % 2 == 0 + + def extract_ridge(band): + polar = polar_at(fit_dp, origin, band[0], band[1]) + polar = np.minimum( + polar, np.percentile(polar, 90.0, axis=0, keepdims=True) + ) + smooth = gaussian_filter1d(polar, 1.25, axis=1, mode="nearest") + baseline = np.percentile(smooth, 20.0, axis=1, keepdims=True) + signal = np.clip(smooth - baseline, 0.0, None) + peak_index = np.argmax(signal, axis=1) + ridge = np.empty(num_annular_bins, dtype=float) + confidence = np.empty(num_annular_bins, dtype=float) + for index in range(num_annular_bins): + lo = max(0, peak_index[index] - 2) + hi = min(signal.shape[1], peak_index[index] + 3) + weights = signal[index, lo:hi] + bins = np.arange(lo, hi, dtype=float) + ridge[index] = ( + np.average(bins, weights=weights) + if weights.sum() > 1e-12 + else float(peak_index[index]) + ) + noise = ( + 1.4826 * np.median(np.abs(np.diff(smooth[index]))) + + 1e-9 + ) + confidence[index] = signal[index, peak_index[index]] / noise + ridge = band[0] + ridge * radial_step + valid = np.isfinite(ridge) & (confidence >= 2.0) + if np.count_nonzero(valid) >= 12: + design = np.column_stack([ + np.ones(num_annular_bins), + np.cos(phi), + np.sin(phi), + np.cos(2 * phi), + np.sin(2 * phi), + ]) + weights = np.clip(confidence / 10.0, 0.05, 1.0) + beta = np.zeros(5) + for _ in range(5): + use = valid & np.isfinite(weights) + root_weight = np.sqrt(weights[use]) + beta = np.linalg.lstsq( + design[use] * root_weight[:, None], + ridge[use] * root_weight, + rcond=None, + )[0] + residual = ridge - design @ beta + scale = 1.4826 * np.median( + np.abs(residual[use] - np.median(residual[use])) + ) + 1e-6 + robust = np.minimum(1.0, 1.5 * scale / (np.abs(residual) + 1e-9)) + weights = np.clip(confidence / 10.0, 0.05, 1.0) * robust + valid &= np.abs(ridge - design @ beta) <= max(3.0, 4.0 * scale) + else: + beta = np.asarray([np.nan] * 5) + weights = np.zeros_like(confidence) + return ridge, confidence, valid, beta, weights + + def ellipse_residual(parameters, x, y, weights): + cy, cx, axis_a, ratio, theta = parameters + dx, dy = x - cx, y - cy + cosine, sine = np.cos(theta), np.sin(theta) + major = dx * cosine + dy * sine + minor = -dx * sine + dy * cosine + geometric = ( + np.sqrt( + (major / axis_a) ** 2 + + (minor / (axis_a * ratio)) ** 2 + ) + - 1.0 + ) * axis_a + return geometric * np.sqrt(np.clip(weights, 1e-3, None)) + + def circle_residual(parameters, x, y, weights): + cy, cx, radius = parameters + return ( + np.hypot(x - cx, y - cy) - radius + ) * np.sqrt(np.clip(weights, 1e-3, None)) + + evaluated = [] + for band in bands: + ridge, confidence, valid, harmonic, weights = extract_ridge(band) + coverage = float(np.mean(valid)) + if np.count_nonzero(valid) < 12: + evaluated.append({ + "band": tuple(float(v) for v in band[:2]), + "r0": float(band[2]), + "accepted": False, + "angular_coverage": coverage, + "rejection_reasons": ["insufficient ridge points"], + }) + continue + + x = origin[1] + ridge * np.cos(phi) + y = origin[0] + ridge * np.sin(phi) + fit = valid & block_fit + validate = valid & ~block_fit + if np.count_nonzero(validate) < 6: + fit = valid + validate = valid + + r0, c1, s1, c2, s2 = harmonic + center_initial = np.asarray([ + origin[0] + np.clip(s1, -center_search_radius, center_search_radius), + origin[1] + np.clip(c1, -center_search_radius, center_search_radius), + ]) + second = float(np.hypot(c2, s2)) + ratio_initial = np.clip( + (max(r0, 1.0) - second) / (max(r0, 1.0) + second), + ratio_range[0], + ratio_range[1], + ) + theta_initial = 0.5 * np.arctan2(s2, c2) + center_low = origin - center_search_radius + center_high = origin + center_search_radius + axis_low = max(2.0, 0.65 * band[0]) + axis_high = min(r_hi * 1.5, 1.45 * band[1]) + + circle = least_squares( + circle_residual, + [*center_initial, np.median(ridge[fit])], + args=(x[fit], y[fit], weights[fit]), + bounds=([ + center_low[0], center_low[1], axis_low + ], [ + center_high[0], center_high[1], axis_high + ]), + loss="soft_l1", + f_scale=1.0, + ) + ellipse = least_squares( + ellipse_residual, + [ + *center_initial, + np.clip(np.max(ridge[fit]), axis_low, axis_high), + ratio_initial, + theta_initial, + ], + args=(x[fit], y[fit], weights[fit]), + bounds=([ + center_low[0], center_low[1], axis_low, + ratio_range[0], -np.pi, + ], [ + center_high[0], center_high[1], axis_high, + ratio_range[1], np.pi, + ]), + loss="soft_l1", + f_scale=1.0, + ) + circle_validation = np.median(np.abs(circle_residual( + circle.x, x[validate], y[validate], np.ones(np.count_nonzero(validate)) + ))) + ellipse_validation = np.median(np.abs(ellipse_residual( + ellipse.x, x[validate], y[validate], np.ones(np.count_nonzero(validate)) + ))) + improvement = float( + (circle_validation - ellipse_validation) + / max(circle_validation, 1e-9) + ) + center_boundary = bool(np.any( + np.isclose(ellipse.x[:2], center_low, atol=0.05) + | np.isclose(ellipse.x[:2], center_high, atol=0.05) + )) + ratio_boundary = bool( + ellipse.x[3] <= ratio_range[0] + 0.005 + or ellipse.x[3] >= ratio_range[1] - 0.001 + ) + reasons = [] + if coverage < min_angular_coverage: + reasons.append(f"angular coverage {coverage:.1%}") + if improvement < min_validation_improvement: + reasons.append(f"held-out improvement {improvement:.2%}") + if ellipse_validation > max_validation_residual: + reasons.append( + f"held-out residual {ellipse_validation:.2f} px" + ) + if center_boundary: + reasons.append("center search boundary") + if ratio_boundary: + reasons.append("ratio search boundary") + if not reasons: + ellipse = least_squares( + ellipse_residual, + ellipse.x, + args=(x[valid], y[valid], weights[valid]), + bounds=([ + center_low[0], center_low[1], axis_low, + ratio_range[0], -np.pi, + ], [ + center_high[0], center_high[1], axis_high, + ratio_range[1], np.pi, + ]), + loss="soft_l1", + f_scale=1.0, + ) + center_boundary = bool(np.any( + np.isclose(ellipse.x[:2], center_low, atol=0.05) + | np.isclose(ellipse.x[:2], center_high, atol=0.05) + )) + ratio_boundary = bool( + ellipse.x[3] <= ratio_range[0] + 0.005 + or ellipse.x[3] >= ratio_range[1] - 0.001 + ) + if center_boundary: + reasons.append("center search boundary after full refit") + if ratio_boundary: + reasons.append("ratio search boundary after full refit") + evaluated.append({ + "band": tuple(float(v) for v in band[:2]), + "r0": float(band[2]), + "accepted": not reasons, + "angular_coverage": coverage, + "center_initial": tuple(float(v) for v in center_initial), + "center_refined": tuple(float(v) for v in ellipse.x[:2]), + "a_pixels": float(ellipse.x[2]), + "b_pixels": float(ellipse.x[2] * ellipse.x[3]), + "ratio_b_over_a": float(ellipse.x[3]), + "theta_deg": float(np.rad2deg(ellipse.x[4]) % 180.0), + "circle_validation_residual": float(circle_validation), + "ellipse_validation_residual": float(ellipse_validation), + "validation_improvement": improvement, + "center_boundary": center_boundary, + "ratio_boundary": ratio_boundary, + "ridge_point_count": int(np.count_nonzero(valid)), + "rejection_reasons": reasons, + "_ridge": ridge, + "_valid": valid, + "_x": x, + "_y": y, + }) + + accepted = [item for item in evaluated if item["accepted"]] + candidates_with_fit = [item for item in evaluated if "a_pixels" in item] + if accepted: + selected = min( + accepted, + key=lambda item: ( + item["ellipse_validation_residual"], + -item["validation_improvement"], + ), + ) + fit_accepted = True + elif candidates_with_fit: + selected = min( + candidates_with_fit, + key=lambda item: item["ellipse_validation_residual"], + ) + fit_accepted = False + else: + selected = evaluated[0] + fit_accepted = False + + public_candidates = [ + {key: value for key, value in item.items() if not key.startswith("_")} + for item in evaluated + ] + if fit_accepted: + a_axis = selected["a_pixels"] + b_axis = selected["b_pixels"] + theta_deg = selected["theta_deg"] + refined_center = selected["center_refined"] + else: + a_axis = b_axis = selected["r0"] + theta_deg = 0.0 + refined_center = tuple(float(v) for v in origin) + reasons = selected.get("rejection_reasons", ["no valid ridge fit"]) + warnings.warn( + "Ridge ellipse fit rejected; using a circular correction " + f"({', '.join(reasons)}).", + RuntimeWarning, + stacklevel=2, + ) + + selected_public = { + key: value for key, value in selected.items() if not key.startswith("_") + } + self.ellipse_fit_diagnostics = { + "method": "ridge", + "accepted": fit_accepted, + "selected": selected_public, + "candidates": public_candidates, + "explicit_band": explicit_band, + "center_initial": tuple(float(v) for v in origin), + "center_refined": tuple(float(v) for v in refined_center), + "rejection_reasons": ( + [] if fit_accepted else selected_public.get("rejection_reasons", []) + ), + } + + if verbose: + print( + " ridge ellipse candidates: " + + ", ".join(f"{item['r0']:.1f}" for item in public_candidates) + + " px" + ) + print( + f" ridge ellipse fit: {'accepted' if fit_accepted else 'rejected'} " + f"a/b={a_axis / b_axis:.4f} theta={theta_deg:.2f} deg " + f"center=({refined_center[0]:.2f}, {refined_center[1]:.2f})" + ) + + if show and "_ridge" in selected: + fig, axes = plt.subplots(1, 3, figsize=(13, 4)) + axes[0].imshow(np.log1p(np.clip(dp, 0, None)), cmap="magma") + axes[0].scatter( + selected["_x"][selected["_valid"]], + selected["_y"][selected["_valid"]], + s=5, + c="cyan", + alpha=0.7, + label="ridge inliers", + ) + axes[0].add_patch(Ellipse( + (refined_center[1], refined_center[0]), + 2 * a_axis, + 2 * b_axis, + angle=theta_deg, + fill=False, + color="lime", + linewidth=1.5, + label="ridge fit", + )) + axes[0].legend(fontsize=8) + axes[0].set_title("diffuse-ring ridge and robust ellipse") + before = polar_at(dp, origin, selected["band"][0], selected["band"][1]) + after = np.asarray( + polar_transform( + dp, + origin_array=np.asarray(refined_center), + ellipse_params=(a_axis, b_axis, theta_deg), + num_annular_bins=num_annular_bins, + radial_min=selected["band"][0], + radial_max=selected["band"][1], + radial_step=radial_step, + scan_pos=(0, 0), + device=device, + show_progress=False, + ) + ) + axes[1].imshow(before, aspect="auto", cmap="magma") + axes[1].set_title("circular polar before") + axes[2].imshow(after, aspect="auto", cmap="magma") + axes[2].set_title("ridge-refined polar after") + fig.tight_layout() + plt.show() + + band = selected["band"] + return ( + float(a_axis), + float(b_axis), + float(theta_deg), + (float(band[0]), float(band[1])), + ) + + def resize_data(self, device:str = "cuda:0"): + print(device) + Ry, Rx, Qy, Qx = self._dataset_cartesian.shape + scale_factor = (self._final_shape[0] * self._final_shape[1]) / (Qy * Qx) + resized_data = np.zeros((Ry, Rx, self._final_shape[0], self._final_shape[1])) + for i in tqdm(range(Ry), desc='rows'): + inp = torch.tensor(self._dataset_cartesian[i].array, dtype=torch.float32).to(device) + inp = torch.nn.functional.interpolate(inp[None, ...], size=self._final_shape, mode='bilinear', align_corners=False) * scale_factor + resized_data[i, :, :, :] = inp.squeeze().detach().cpu().numpy() + self.resized_cartesian_data = resized_data + + def resize_images(self, images, device: str = "cuda:0", initial_chunk_size: int = 100, show_progress=False): + # Handle Dataset objects - extract array + if hasattr(images, 'array'): + images = images.array + elif isinstance(images, Dataset3d): + # If it's a Dataset3d, get the underlying array + images = np.array([images[i].array for i in range(images.shape[0])]) + + N, Qy, Qx = images.shape + scale_factor = (self._final_shape[0] * self._final_shape[1]) / (Qy * Qx) + resized_data = np.zeros((N, self._final_shape[0], self._final_shape[1])) + + chunk_size = initial_chunk_size + i = 0 + + with tqdm(total=N, desc='images', disable=not show_progress) as pbar: + while i < N: + try: + # Determine the end index for this chunk + end_idx = min(i + chunk_size, N) + chunk = images[i:end_idx] + + # Process chunk on GPU + inp = torch.tensor(chunk, dtype=torch.float32).to(device) + inp = torch.nn.functional.interpolate( + inp.unsqueeze(1), # Add channel dimension + size=self._final_shape, + mode='bilinear', + align_corners=False + ) * scale_factor + + resized_data[i:end_idx, :, :] = inp.squeeze(1).detach().cpu().numpy() + + # Clear GPU cache + del inp + if 'cuda' in device: + torch.cuda.empty_cache() + + # Update progress and move to next chunk + pbar.update(end_idx - i) + i = end_idx + + except RuntimeError as e: + if 'out of memory' in str(e): + # Clear cache and reduce chunk size + if 'cuda' in device: + torch.cuda.empty_cache() + + chunk_size = max(1, chunk_size // 2) + print(f"\nGPU OOM! Reducing chunk size to {chunk_size}") + + if chunk_size == 1: + # If even single image fails, fall back to CPU + print("Falling back to CPU processing") + device = "cpu" + else: + raise e + + return resized_data + + def set_model_weights( + self, + path_to_weights: str = None, + *, + model_id: str = PAPER_MODEL_ID, + version: str | None = None, + latest: bool = False, + local_model_dir: str | None = None, + cache_dir: str | None = None, + ) -> "BraggPeaksPolymer": + """Load explicit weights or a checksum-verified named model. + + Explicit paths retain the historical behavior. Without a path, the + immutable paper model is selected; ``latest=True`` is opt-in. + """ + if path_to_weights is None: + resolution = resolve_polymer_model( + model_id=model_id, + version=version, + latest=latest, + local_model_dir=local_model_dir, + cache_dir=cache_dir, + ) + self._model = build_polymer_model(resolution.specification) + if not self._normalization_is_explicit: + normalization_config = resolution.specification.get( + "experimental_normalization" + ) + if normalization_config is None: + raise RuntimeError( + f"Registered model {resolution.model_id!r} does not declare " + "experimental_normalization." + ) + self._set_normalization_strategy(normalization_config, explicit=False) + path_to_weights = str(resolution.weights_path) + self.model_resolution = resolution + self._model.load_state_dict( + torch.load(path_to_weights, weights_only=True, map_location=self.device) + ) + self._model.to(self.device) + self._invalidate_inference_caches() + return self + + def detect_ice( + self, + *, + params=None, + scan_mask=None, + intensity_threshold_global=None, + return_debug=False, + ): + """Detect ice peaks from this analysis's polar peaks and intensities.""" + + from quantem.diffraction.polymer_ice import IceFlaggerParams, detect_ice + + if self.polar_peaks is None or self.peak_intensities is None: + raise RuntimeError( + "detect_ice() requires polar_peaks and peak_intensities to be computed first." + ) + return detect_ice( + self.polar_peaks, + self.peak_intensities, + params=IceFlaggerParams() if params is None else params, + scan_mask=self.scan_mask if scan_mask is None else scan_mask, + intensity_threshold_global=intensity_threshold_global, + return_debug=return_debug, + polar_data=getattr(self, "polar_data", None), + # process_polar(two_fold_symmetry=True) folded theta to [0, 180). + theta_period_deg=180.0 if getattr(self, "two_fold_symmetry", False) else 360.0, + ) + + def measure_ice_peak_widths(self, *, params=None, scan_mask=None, **kwargs): + """Radial/annular widths of this analysis's peaks, for tuning the sharpness gate.""" + + from quantem.diffraction.polymer_ice import IceFlaggerParams, collect_peak_widths + + if self.polar_peaks is None or self.peak_intensities is None or getattr(self, "polar_data", None) is None: + raise RuntimeError( + "measure_ice_peak_widths() requires polar_peaks, peak_intensities and polar_data." + ) + return collect_peak_widths( + self.polar_peaks, + self.peak_intensities, + self.polar_data, + params=IceFlaggerParams() if params is None else params, + scan_mask=self.scan_mask if scan_mask is None else scan_mask, + **kwargs, + ) + + def plot_q_intensity_density(self, **kwargs): + """Plot q/intensity density from this analysis's aligned peak vectors.""" + + from quantem.diffraction.polymer_ice import plot_q_intensity_density + + if self.polar_peaks is None or self.peak_intensities is None: + raise RuntimeError( + "plot_q_intensity_density() requires polar_peaks and peak_intensities." + ) + return plot_q_intensity_density( + self.polar_peaks, self.peak_intensities, **kwargs + ) + + def _postprocess_single(self, position_map, intensity_map, sigma=1.0, threshold=0.25, show=False): + """Process a single 2D image""" + # Find peaks with subpixel-refinement + peak_coords, peak_position_signal_intensities, refinement_success = detect_blobs( + position_map, + sigma=sigma, # Sigma for Gaussian smoothing used in processing + threshold=threshold, # Threshold for strength of peak position signal to be valid peak + ) + + # If no peaks found, return empty lists + if len(peak_coords) == 0: + return np.array([]), np.array([]) + + # map_coordinates expects coordinates in (row, col) = (y, x) order + # peak_coords is already in [row, col] format from detect_blobs + interpolated_intensities = map_coordinates( + intensity_map, + peak_coords.T, # Transpose to get [[all_y], [all_x]] + order=1, # 1 = bilinear interpolation + mode='nearest' # How to handle edges + ) + + # Optional: filter out peaks that were not successfully refined + if np.any(refinement_success): + pass + + if show: + # Peak positions only + fig, ax = plt.subplots(figsize=(10, 8)) + ax.imshow(position_map, cmap='gray', alpha=0.8) + ax.set_title("Input Position Map with Marked Peaks") + ax.scatter(peak_coords[:, 1], peak_coords[:, 0], s=10, c='r', label="Peaks") + ax.legend() + plt.tight_layout() + plt.show() + + # Peak positions with color representing intensity + fig, ax = plt.subplots(figsize=(10, 8)) + im = ax.imshow(position_map, cmap='gray', alpha=0.8) + scatter = ax.scatter( + peak_coords[:, 1], # x coordinates + peak_coords[:, 0], # y coordinates + c=interpolated_intensities, # color by intensity + s=10, + cmap='turbo', + edgecolors='black', # white border for visibility + linewidths=2, + alpha=0.9, + marker='o' + ) + cbar = plt.colorbar(scatter, ax=ax) + cbar.set_label('Intensity', fontsize=12) + ax.set_title('Peak Positions and Intensities', fontsize=14) + ax.axis('off') + plt.tight_layout() + plt.show() + + return peak_coords, interpolated_intensities + + def ensure_normalization_params( + self, + device: str = None, + n_normalize_samples: int = 1000, + scan_mask: ArrayLike = None, + recompute: bool = False, + ): + """Fit and cache the configured inference-normalization parameters. + + These are estimated once from a random sample of valid diffraction patterns and + reused by both ``find_peaks_model`` (whole-scan) and ``infer_peaks_single`` + (live). Caching guarantees live single-DP inference reproduces the full-scan + peaks exactly (same normalization). Parameters are intentionally opaque. + """ + strategy = self._require_normalization_strategy() + if not recompute and self._normalization_parameters is not None: + return self._normalization_parameters + + device = device or self.device + Ry, Rx, _, _ = self.dataset_cartesian.shape + # Restrict to the stored ROI when no mask is passed explicitly (fall back to the + # whole scan only if none is set); estimate stats from the sample region. + if scan_mask is None: + scan_mask = self._scan_mask + if scan_mask is None: + scan_mask = np.ones((Ry, Rx), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + valid_positions = np.argwhere(scan_mask) + n_valid = len(valid_positions) + + n_normalize_samples = min(n_normalize_samples, n_valid) + sample_indices = np.random.choice(n_valid, size=n_normalize_samples, replace=False) + + stats_patterns = np.array([ + self.dataset_cartesian[ry, rx].array + for ry, rx in valid_positions[sample_indices] + ]) + + stats_patterns_resized = self.resize_images(stats_patterns, device=device) + parameters = strategy.fit(stats_patterns_resized) + self._normalization_parameters = parameters + if isinstance(parameters, tuple) and len(parameters) == 2: + self._norm_median, self._norm_iqr = parameters + else: + self._norm_median = self._norm_iqr = None + return parameters + + def adapt_batchnorm( + self, + device: str = None, + n_samples: int = 1000, + scan_mask: ArrayLike = None, + chunk_size: int = 100, + recompute: bool = False, + ): + """Adapt the model's BatchNorm running statistics to THIS dataset, then eval. + + The model trains on synthetic data, so its stored BatchNorm running stats do not + match the experimental scan; plain ``eval()`` inference then under-detects. + ``find_peaks_model`` sidesteps this by running in train mode (per-chunk batch + stats). For deterministic single-DP inference (``infer_peaks_single`` / the live + widget), we instead estimate the running stats *once* from a representative sample + of this dataset and freeze them: reset the BatchNorm buffers, run a sample through + the model in train mode with ``momentum=None`` (so the buffers accumulate the + cumulative mean/var over the sample), then switch to eval. Uses the same input + normalization pipeline (resize + ``normalize_data`` with the cached median/iqr). + Idempotent unless ``recompute=True``. Leaves the model in eval mode. + """ + if self._bn_adapted and not recompute: + return + import torch.nn as nn + + device = device or self.device + parameters = self.ensure_normalization_params( + device=device, n_normalize_samples=max(n_samples, 1000), scan_mask=scan_mask + ) + strategy = self._require_normalization_strategy() + + Ry, Rx, _, _ = self.dataset_cartesian.shape + # Restrict the adaptation sample to the stored ROI when none is passed. + if scan_mask is None: + scan_mask = self._scan_mask + if scan_mask is None: + scan_mask = np.ones((Ry, Rx), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + valid_positions = np.argwhere(scan_mask) + n_valid = len(valid_positions) + n_samples = min(n_samples, n_valid) + sample_indices = np.random.choice(n_valid, size=n_samples, replace=False) + sample_positions = valid_positions[sample_indices] + + # Temporarily switch BatchNorm layers to cumulative-average mode so the running + # buffers become the exact mean/var over the sample (not an EMA of the last batch). + self.model.to(device) + bn_layers = [m for m in self.model.modules() if isinstance(m, nn.modules.batchnorm._BatchNorm)] + saved_momentum = [m.momentum for m in bn_layers] + for m in bn_layers: + m.reset_running_stats() + m.momentum = None # cumulative moving average + self.model.train() + try: + with torch.no_grad(): + for i in range(0, n_samples, chunk_size): + chunk = np.array([ + self.dataset_cartesian[ry, rx].array + for ry, rx in sample_positions[i : i + chunk_size] + ]) + resized = self.resize_images(chunk, device=device, initial_chunk_size=chunk_size) + ins = torch.tensor(resized, dtype=torch.float32).to(device) + ins_batch = strategy.transform(ins, parameters)[:, None, ...] + self.model(ins_batch) # updates BN running stats only + finally: + for m, mom in zip(bn_layers, saved_momentum): + m.momentum = mom + self.model.eval() + self._bn_adapted = True + + def prepare_inference(self, device: str = None, n_samples: int = 1000, scan_mask: ArrayLike = None): + """Convenience: compute input-normalization stats + adapt BatchNorm in one call. + + Run after the model weights are loaded to ready the object for deterministic + eval-mode single-DP inference (``infer_peaks_single``). + """ + self.ensure_normalization_params(device=device, n_normalize_samples=n_samples, scan_mask=scan_mask) + self.adapt_batchnorm(device=device, n_samples=n_samples, scan_mask=scan_mask) + + def _infer_train_batch_output( + self, ry, rx, *, device, parameters, chunk_size=100, scan_mask=None + ): + """Model output ``(2, H, W)`` for the DP at (ry, rx), computed exactly as + ``find_peaks_model`` does. + + The DP is run inside its train-mode ``find_peaks_model`` chunk, so BatchNorm + normalizes it with the same ~``chunk_size`` real-DP statistics (the train-mode + test-time domain adaptation). This reproduces the precomputed detection for that + position -- unlike the eval + ``adapt_batchnorm`` path, whose global running stats + differ from the chunk-local stats and over-detect on this OOD scan. + + The chunk is the same slice of ``np.argwhere(scan_mask)`` (row-major) that + ``find_peaks_model`` would place (ry, rx) in; the resulting output is cached so + neighbouring cursor positions in the same chunk reuse one forward pass. + """ + Ry, Rx, _, _ = self.dataset_cartesian.shape + if scan_mask is None: + scan_mask = self._scan_mask + if scan_mask is None: + scan_mask = np.ones((Ry, Rx), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + valid = np.argwhere(scan_mask) # row-major: matches find_peaks_model's iteration + match = np.where((valid[:, 0] == ry) & (valid[:, 1] == rx))[0] + if len(match): + qi = int(match[0]) + start = (qi // chunk_size) * chunk_size + chunk_positions = valid[start : start + chunk_size] + local_i = qi - start + else: + # (ry, rx) is outside the ROI -- find_peaks_model never processes it. Still give + # a faithful readout by running it at the head of a representative ROI chunk. + start = -1 # never matches a real chunk_start -> not cacheable across positions + head = valid[: max(0, chunk_size - 1)] + chunk_positions = ( + np.concatenate([[[ry, rx]], head], axis=0) if len(head) else np.array([[ry, rx]]) + ) + local_i = 0 + + cache = self._live_chunk_cache + if start >= 0 and cache is not None and cache[0] == start and cache[1] == chunk_size: + return cache[2][local_i] + + chunk = np.array([self.dataset_cartesian[r, c].array for r, c in chunk_positions]) + resized = self.resize_images(chunk, device=device, initial_chunk_size=len(chunk)) + ins = torch.tensor(resized, dtype=torch.float32).to(device) + ins_batch = self._require_normalization_strategy().transform( + ins, parameters + )[:, None, ...] + self.model.to(device) + self.model.train() # per-chunk BatchNorm stats, exactly like find_peaks_model + with torch.no_grad(): + outs = self.model(ins_batch).detach().cpu().numpy() # (n, 2, H, W) + if start >= 0: + self._live_chunk_cache = (start, chunk_size, outs) + return outs[local_i] + + def infer_peaks_single( + self, + ry: int, + rx: int, + *, + device: str = None, + sigma_peak_blur: float = 1.0, + threshold_peak: float = 0.5, + n_normalize_samples: int = 1000, + bn_mode: str = "train_batch", + chunk_size: int = 100, + scan_mask: ArrayLike = None, + ): + """Run the model on the single diffraction pattern at (ry, rx). + + Live counterpart of ``find_peaks_model`` for one scan position: resize -> + normalize (cached median/iqr) -> model forward -> decode -> rescale to detector + pixels. Returns a dict with keys ``"y_pixels"``, ``"x_pixels"``, ``"intensities"`` + (empty arrays when no peaks are found), matching the columns/units of + ``peak_coordinates_cartesian`` / ``peak_intensities``. + + ``bn_mode`` selects the BatchNorm regime: + + - ``"train_batch"`` (default): run the DP inside its train-mode ``find_peaks_model`` + chunk so it gets the same per-chunk domain adaptation. Output **matches the + precomputed find_peaks_model detection** for that position. Deterministic given + the ROI + chunk_size. + - ``"eval_adapt"``: eval mode using dataset-adapted BatchNorm running stats (see + ``adapt_batchnorm``, lazy + cached). Faster (single-DP forward) but an + approximation that over-detects on this out-of-distribution scan. + """ + device = device or self.device + parameters = self.ensure_normalization_params( + device=device, n_normalize_samples=n_normalize_samples, scan_mask=scan_mask + ) + + if bn_mode == "train_batch": + out = self._infer_train_batch_output( + ry, rx, device=device, parameters=parameters, + chunk_size=chunk_size, scan_mask=scan_mask, + ) + elif bn_mode == "eval_adapt": + # Domain-adapt BatchNorm to this dataset once, then infer in eval mode. + self.adapt_batchnorm(device=device, n_samples=n_normalize_samples, scan_mask=scan_mask) + dp = np.asarray(self.dataset_cartesian[ry, rx].array) + resized = self.resize_images(dp[None], device=device, initial_chunk_size=1) + ins = torch.tensor(resized, dtype=torch.float32).to(device) + ins_batch = self._require_normalization_strategy().transform( + ins, parameters + )[:, None, ...] + self.model.to(device) + self.model.eval() + with torch.no_grad(): + out = self.model(ins_batch).detach().cpu().numpy()[0] # (2, H, W) + else: + raise ValueError( + f"bn_mode must be 'train_batch' or 'eval_adapt', got {bn_mode!r}" + ) + + peak_coords, peak_ints = self._postprocess_single( + out[0], out[1], sigma=sigma_peak_blur, threshold=threshold_peak + ) + if len(peak_coords) == 0: + empty = np.array([]) + return {"y_pixels": empty, "x_pixels": empty, "intensities": empty} + + # Rescale from model-input pixels back to original detector pixels (matches + # the whole-scan rescale in find_peaks_model). + scale = self.dataset_cartesian.shape[2] / self.final_shape[0] + coords = np.asarray(peak_coords) * scale # (N, 2) = [row=y, col=x] + return { + "y_pixels": coords[:, 0], + "x_pixels": coords[:, 1], + "intensities": np.asarray(peak_ints), + } + + def find_peaks_model( + self, + device: str = "cuda:0", + scan_mask: ArrayLike = None, + n_normalize_samples: int = 1000, + initial_chunk_size: int = 100, + sigma_peak_blur: float = 1.0, + threshold_peak: float = 0.5, + show_plots=False, + accelerating_voltage_kv: float = None, + ): + """Detect peaks throughout the scan with the trained model. + + Parameters + ---------- + accelerating_voltage_kv + Electron accelerating voltage used to convert detector sampling from + mrad to inverse angstroms. For mrad data, the default is 300 kV and an + explicit warning is emitted. Ignored for reciprocal-length calibration. + """ + Ry, Rx, Qy, Qx = self.dataset_cartesian.shape + total_positions = Ry * Rx + + # Resolve this once per run, both to avoid repeated unit parsing and to retain + # the selected voltage for later polar-coordinate operations. + sampling_inv_A = self.pixels_to_inv_A(accelerating_voltage_kv) + + # Remember the ROI so later normalization / BN adaptation (and the live widget) + # restrict to the sample region. Storing the user-provided value (None stays the + # whole scan); the setter invalidates cached stats only if the mask changed. + self.scan_mask = scan_mask + + # ============================================ + # Handle scan_mask + # ============================================ + if scan_mask is None: + scan_mask = np.ones((Ry, Rx), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + if scan_mask.shape != (Ry, Rx): + raise ValueError(f"scan_mask shape {scan_mask.shape} must match scan shape ({Ry}, {Rx})") + + # Get list of valid positions + valid_positions = np.argwhere(scan_mask) # Returns array of (ry, rx) pairs + n_valid = len(valid_positions) + + peaks = Vector.from_shape( + shape=(Ry, Rx), + fields=["y_pixels", "x_pixels", "y_invA", "x_invA"], + name="peaks_vector", + units=["Pixels", "Pixels", "1/Å", "1/Å"], + ) + intensities = Vector.from_shape( + shape=(Ry, Rx), + fields=["intensities", "intensities_sampled_from_dp"], + name="intensities_vector", + units=["Normalized", "Normalized"], + ) + + # ============================================ + # 1. Compute normalization parameters (only from valid positions) + # ============================================ + # recompute=True to preserve the original per-call semantics (find_peaks_model + # always recomputed the sample stats); the cache still serves infer/adapt. + parameters = self.ensure_normalization_params( + device=device, + n_normalize_samples=n_normalize_samples, + scan_mask=scan_mask, + recompute=True, + ) + + # Run in TRAIN mode on purpose. The model trains on synthetic data; on the + # (out-of-distribution) experimental scan, train-mode BatchNorm normalizes each + # chunk with the experimental data's own statistics — test-time domain adaptation + # that detects far better than eval mode (which would impose the synthetic-training + # population stats on real data). Set it explicitly so a prior eval() / adapt_batchnorm + # (e.g. from the live widget) can't leave the shared model in eval mode. The live + # single-DP path (infer_peaks_single) instead uses adapt_batchnorm + eval. + self.model.train() + + # ============================================ + # 2. Process only valid positions with chunking + # ============================================ + chunk_size = initial_chunk_size + pos_idx = 0 + + with tqdm(total=n_valid, desc="Processing patterns") as pbar: + while pos_idx < n_valid: + try: + # ---------------------------------------- + # 2a. Determine chunk boundaries + # ---------------------------------------- + end_pos_idx = min(pos_idx + chunk_size, n_valid) + actual_chunk_size = end_pos_idx - pos_idx + + # ---------------------------------------- + # 2b. Extract chunk data (only valid positions) + # ---------------------------------------- + chunk_data = [] + chunk_positions = [] + + for i in range(pos_idx, end_pos_idx): + ry, rx = valid_positions[i] + chunk_data.append(self.dataset_cartesian[ry, rx].array) + chunk_positions.append((ry, rx)) + + chunk_array = np.array(chunk_data) + + # ---------------------------------------- + # 2c. Resize chunk + # ---------------------------------------- + # self.model.to(device) + chunk_resized = self.resize_images( + chunk_array, + device=device, + initial_chunk_size=actual_chunk_size + ) + + # ---------------------------------------- + # 2d. Normalize and run model + # ---------------------------------------- + ins = torch.tensor(chunk_resized, dtype=torch.float32).to(device) + dps_norm = self._require_normalization_strategy().transform( + ins, parameters + ) + ins_batch = dps_norm[:, None, ...] + + with torch.no_grad(): + outs = self.model(ins_batch).detach().cpu().numpy() + + # ---------------------------------------- + # 2e. Post-process each pattern in chunk + # ---------------------------------------- + for k in range(outs.shape[0]): + ry, rx = chunk_positions[k] + + peak_coords, peak_intensities = self._postprocess_single( + outs[k, 0], + outs[k, 1], + show=show_plots, + sigma=sigma_peak_blur, + threshold=threshold_peak, + ) + + if len(peak_coords) > 0: + peak_intensity_averages = sample_average_from_image( + ins_batch[k].squeeze().detach().cpu().numpy(), + peak_coords + ) + peak_intensities_data = np.column_stack([ + peak_intensities, + peak_intensity_averages, + ]) + + peak_coords_original = peak_coords * ( + self.dataset_cartesian.shape[2] / self.final_shape[0] + ) + + peak_data = np.column_stack([ + peak_coords_original, + peak_coords_original * sampling_inv_A + ]) + + peaks[ry, rx] = peak_data + intensities[ry, rx] = peak_intensities_data + + # ---------------------------------------- + # 2f. Memory cleanup + # ---------------------------------------- + del ins, dps_norm, ins_batch, outs, chunk_array, chunk_resized + if 'cuda' in device: + torch.cuda.empty_cache() + + # ---------------------------------------- + # 2g. Update progress and move to next chunk + # ---------------------------------------- + pbar.update(actual_chunk_size) + pos_idx = end_pos_idx + + except RuntimeError as e: + if 'out of memory' in str(e): + if 'cuda' in device: + torch.cuda.empty_cache() + + chunk_size = max(1, chunk_size // 2) + print(f"\nGPU OOM! Reducing chunk size to {chunk_size}") + + if chunk_size == 1: + print("Falling back to CPU processing") + device = "cpu" + else: + raise e + + print('Done!') + self.peak_coordinates_cartesian = peaks + self.peak_intensities = intensities + + @staticmethod + def _save_object(filepath, obj): + # Wrap in a 0-d object array so np.save pickles the WHOLE object. A canon Vector is + # array-like, so np.save(vector) would otherwise flatten it to an (Ry, Rx) object + # array of cells that loses the Vector's fields/units and can't be reconstructed + # (breaks polar_transform_peaks, which needs a Vector). The 0-d wrapper round-trips + # via the size-1 .item() unwrap in the load_* methods. + arr = np.empty((), dtype=object) + arr[()] = obj + np.save(filepath, arr, allow_pickle=True) + + def save_cartesian_peaks(self, filepath): + self._save_object(filepath, self.peak_coordinates_cartesian) + + def load_cartesian_peaks(self, filepath): + peak_coordinates_cartesian = np.load(filepath, allow_pickle=True) + if isinstance(peak_coordinates_cartesian, np.ndarray) and peak_coordinates_cartesian.dtype == object and peak_coordinates_cartesian.size == 1: + peak_coordinates_cartesian = peak_coordinates_cartesian.item() + self.peak_coordinates_cartesian = peak_coordinates_cartesian + + def save_polar_peaks(self, filepath): + self._save_object(filepath, self.polar_peaks) + + def save_polar_data(self, filepath): + self._save_object(filepath, self.polar_data) + + def load_polar_peaks(self, filepath): + polar_peaks = np.load(filepath, allow_pickle=True) + if isinstance(polar_peaks, np.ndarray) and polar_peaks.dtype == object and polar_peaks.size == 1: + polar_peaks = polar_peaks.item() + self.polar_peaks = polar_peaks + + def load_polar_data(self, filepath): + obj = np.load(filepath, allow_pickle=True) + if isinstance(obj, np.ndarray) and obj.dtype == object and obj.shape == (): + obj = obj.item() + self.polar_data = obj + + # Populate attributes expected elsewhere + r_grid = self.polar_data['r_invA'] + self.max_radius_invA = float(np.max(r_grid)) + self.num_radial_bins = int(r_grid.shape[0]) + self.num_annular_bins = int(r_grid.shape[1]) + + def save_peak_intensities(self, filepath): + self._save_object(filepath, self.peak_intensities) + + def load_peak_intensities(self, filepath): + peak_intensities = np.load(filepath, allow_pickle=True) + if isinstance(peak_intensities, np.ndarray) and peak_intensities.dtype == object and peak_intensities.size == 1: + peak_intensities = peak_intensities.item() + self.peak_intensities = peak_intensities + + def save_image_centers(self, filepath): + np.save(filepath, self.image_centers) + + def load_image_centers(self, filepath): + image_centers = np.load(filepath, allow_pickle=True) + if isinstance(image_centers, np.ndarray) and image_centers.dtype == object and image_centers.size == 1: + image_centers = image_centers.item() + self.image_centers = image_centers + + def process_polar( + self, + scan_mask: ArrayLike = None, + two_fold_symmetry: bool = True, + center_method: str = "descent", + center_radial_min: float = 4.0, + center_radial_max: float | None = None, + center_radial_step: float = 1.0, + center_num_annular_bins: int = 180, + center_n_phi: int = 120, + center_kpow: float = 0.0, + center_ellipse_params: tuple[float, float, float] | None = None, + center_device: str | None = None, + center_batch_size: int = 16, + center_local_margin: int = 40, + fallback_to_peaks: bool = True, + ): + """Find image centers, then return polar transforms of data and peaks. + + ``center_method`` defaults to Karen Ehrhardt's angular-uniformity descent + method. Use ``center_method="grid"`` for the slower coarse-to-fine + search, or ``center_method="peaks"`` to force the previous peak-based + central-beam heuristic. + """ + self.image_centers = self.find_central_beams_4d( + scan_mask=scan_mask, + center_method=center_method, + radial_min=center_radial_min, + radial_max=center_radial_max, + radial_step=center_radial_step, + num_annular_bins=center_num_annular_bins, + n_phi=center_n_phi, + kpow=center_kpow, + ellipse_params=center_ellipse_params, + center_device=center_device, + center_batch_size=center_batch_size, + local_margin=center_local_margin, + fallback_to_peaks=fallback_to_peaks, + ) + self.polar_peaks = self.polar_transform_peaks( + cartesian_peaks=self.peak_coordinates_cartesian, + centers=self.image_centers, + scan_mask=scan_mask, + two_fold_symmetry=two_fold_symmetry, + ellipse_params=center_ellipse_params, + ) + self.polar_data = self.polar_transform_4d( + self.dataset_cartesian, + centers=self.image_centers, + scan_mask=scan_mask, + two_fold_symmetry=two_fold_symmetry, + ellipse_params=center_ellipse_params, + ) + + def find_central_beams_4d( + self, + scan_mask: ArrayLike = None, + intensity_threshold=0.3, + distance_weight=0.5, + sampling_radius=2, + debug=False, + use_tqdm=True, + center_method: str = "descent", + radial_min: float = 4.0, + radial_max: float | None = None, + radial_step: float = 1.0, + num_annular_bins: int = 180, + n_phi: int = 120, + kpow: float = 0.0, + ellipse_params: tuple[float, float, float] | None = None, + center_device: str | None = None, + center_batch_size: int = 16, + local_margin: int = 40, + fallback_to_peaks: bool = True, + ): + """ + Fast central beam finding for entire 4D dataset. + + Parameters: + ----------- + scan_mask : ArrayLike, optional + Boolean mask (Ry, Rx) indicating which positions to process + use_tqdm : bool + Show progress bar + + Returns: + -------- + centers : ndarray, shape (2, scan_y, scan_x) + Center coordinates (y, x) for each scan position + """ + center_method = center_method.lower() + if center_method == "peaks": + return self._find_central_beams_from_peaks_4d( + scan_mask=scan_mask, + intensity_threshold=intensity_threshold, + distance_weight=distance_weight, + sampling_radius=sampling_radius, + debug=debug, + use_tqdm=use_tqdm, + ) + if center_method not in ("descent", "grid"): + raise ValueError( + "center_method must be 'descent', 'grid', or 'peaks', " + f"got {center_method!r}." + ) + + scan_y, scan_x, _det_y, _det_x = self.dataset_cartesian.shape + if scan_mask is None: + scan_mask_arr = np.ones((scan_y, scan_x), dtype=bool) + else: + scan_mask_arr = np.asarray(scan_mask, dtype=bool) + if scan_mask_arr.shape != (scan_y, scan_x): + raise ValueError( + f"scan_mask shape {scan_mask_arr.shape} must match {(scan_y, scan_x)}" + ) + + device = center_device if center_device is not None else self.device + try: + origins = find_origin_angular_uniformity( + self.dataset_cartesian, + method=center_method, + ellipse_params=ellipse_params, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + num_annular_bins=num_annular_bins, + n_phi=n_phi, + kpow=kpow, + device=device, + batch_size=center_batch_size, + local_margin=local_margin, + ) + except Exception as exc: + if not fallback_to_peaks: + raise + warnings.warn( + "Angular-uniformity center finding failed; falling back to " + f"peak-based central-beam heuristic. Original error: {exc}", + stacklevel=2, + ) + return self._find_central_beams_from_peaks_4d( + scan_mask=scan_mask_arr, + intensity_threshold=intensity_threshold, + distance_weight=distance_weight, + sampling_radius=sampling_radius, + debug=debug, + use_tqdm=use_tqdm, + ) + + if origins.shape != (scan_y, scan_x, 2): + raise ValueError( + f"Origin finder returned shape {origins.shape}; expected {(scan_y, scan_x, 2)}." + ) + centers = np.moveaxis(np.asarray(origins, dtype=float), -1, 0) + centers[:, ~scan_mask_arr] = 0.0 + return centers + + def _find_central_beams_from_peaks_4d( + self, + scan_mask: ArrayLike = None, + intensity_threshold=0.3, + distance_weight=0.5, + sampling_radius=2, + debug=False, + use_tqdm=True, + ): + """Previous central-beam heuristic based on detected peak locations.""" + scan_y, scan_x, det_y, det_x = self.dataset_cartesian.shape + centers = np.zeros((2, scan_y, scan_x)) + + # Handle scan_mask + if scan_mask is None: + scan_mask = np.ones((scan_y, scan_x), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + + iterator = tqdm(range(scan_y), disable=not use_tqdm, desc="Finding centers") + + for i in iterator: + for j in range(scan_x): + if not scan_mask[i, j]: + continue + if self.peak_coordinates_cartesian[i, j] is None: + print(f"None at i={i}, j={j}") + centers[:, i, j] = find_central_beam_from_peaks( + peak_coords=self.peak_coordinates_cartesian[i, j], + peak_intensities=None, + image_shape=(det_y, det_x), + intensity_threshold=intensity_threshold, + distance_weight=distance_weight, + debug=debug, + image=self.dataset_cartesian[i, j].array.squeeze(), + sampling_radius=sampling_radius + ) + return centers + + def polar_transform_peaks( + self, + cartesian_peaks, + centers, + scan_mask: ArrayLike = None, + two_fold_symmetry=True, + ellipse_params: tuple[float, float, float] | None = None, + use_tqdm: bool=True, + ): + """Transform detected Cartesian peak coordinates with Karen's polar convention. + + Peaks are preserved one-to-one. With two-fold symmetry, theta is folded + modulo pi while partner detections remain separate rows. + """ + return karen_polar_transform_peaks( + cartesian_peaks, + centers, + scan_mask=scan_mask, + sampling_conversion_factor=self.pixels_to_inv_A(), + two_fold_rotation_symmetry=two_fold_symmetry, + ellipse_params=ellipse_params, + use_tqdm=use_tqdm, + ) + + def polar_transform_4d( + self, + data, + centers, + scan_mask: ArrayLike = None, + num_r=None, + num_theta=360, + two_fold_symmetry=True, + ellipse_params: tuple[float, float, float] | None = None, + device: str | None = None, + batch_size: int = 128, + use_tqdm: bool=True, + ): + """ + Perform polar transform on the last two axes of a 4D array. + + Parameters: + ----------- + data : ndarray, shape (N, M, H, W) + 4D input array where H, W are the axes to transform + centers : ndarray, shape (2, N, M) + Center of each diffraction pattern (usually determined by central beam) + scan_mask : ArrayLike, optional + Boolean mask (N, M) indicating which positions to process + num_r : int, optional + Number of radial bins. If None, uses max radius across all patterns + num_theta : int, optional + Number of angular bins (default: 360) + two_fold_symmetry : bool, optional + If True, applies 2-fold symmetry by summing opposite angles (default: True). + Samples the full [0, 2π] range but folds it to [0, π] by summing + theta and theta+π positions. + use_tqdm : bool, optional + Whether to show progress bar (default: True) + + Returns: + -------- + polar_data : dict + Dictionary containing polar-transformed data with keys: + - 'r_pixels': radial coordinates in pixels + - 'theta': angular coordinates in radians [0, π] if two_fold_symmetry, else [0, 2π] + - 'r_invA': radial coordinates in 1/Å + - 'intensity': transformed intensity data + + Notes: + ------ + Also sets the following attributes on self: + - self.max_radius_pixels : maximum radius in pixels + - self.max_radius_invA : maximum radius in 1/Å + - self.num_radial_bins : number of radial bins + - self.num_annular_bins : number of angular bins (after symmetry folding) + - self.two_fold_symmetry : whether 2-fold symmetry was used + """ + N, M, H, W = data.shape + + # Handle scan_mask + if scan_mask is None: + scan_mask = np.ones((N, M), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + if scan_mask.shape != (N, M): + raise ValueError(f"scan_mask shape {scan_mask.shape} must match {(N, M)}") + if not np.any(scan_mask): + raise ValueError("scan_mask must include at least one scan position.") + + centers = np.asarray(centers, dtype=float) + if centers.shape == (2, N, M): + centers_karen = np.moveaxis(centers, 0, -1) + centers_bragg = centers + elif centers.shape == (N, M, 2): + centers_karen = centers + centers_bragg = np.moveaxis(centers, -1, 0) + else: + raise ValueError( + f"centers must have shape {(2, N, M)} or {(N, M, 2)}, got {centers.shape}" + ) + if two_fold_symmetry and num_theta % 2 != 0: + raise ValueError("num_theta must be even when two_fold_symmetry=True.") + + # Calculate consistent max_radius across entire dataset (only from masked positions) + valid_centers_0 = centers_bragg[0][scan_mask] + valid_centers_1 = centers_bragg[1][scan_mask] + dist_to_origin_sq = (valid_centers_0**2 + valid_centers_1**2).min() + dist_to_corner_sq = ((H-1 - valid_centers_0)**2 + (W-1 - valid_centers_1)**2).max() + max_radius_pixels = np.sqrt(max(dist_to_origin_sq, dist_to_corner_sq)) + + if num_r is None: + num_r = int(np.ceil(max_radius_pixels)) + num_r = max(1, int(num_r)) + radial_step = max_radius_pixels / num_r if max_radius_pixels > 0 else 1.0 + + # Calculate maximum radius in inverse angstroms + max_radius_invA = max_radius_pixels * self.pixels_to_inv_A() + + polar_full = karen_polar_transform( + data, + origin_array=centers_karen, + ellipse_params=ellipse_params, + num_annular_bins=num_theta, + radial_min=0.0, + radial_max=max_radius_pixels, + radial_step=radial_step, + two_fold_rotation_symmetry=False, + device=device if device is not None else self.device, + batch_size=batch_size, + show_progress=use_tqdm, + ) + polar_intensity_full = np.asarray(polar_full.array, dtype=np.float32).transpose(0, 1, 3, 2) + polar_intensity_full[~scan_mask] = 0.0 + + # Pre-calculate coordinate arrays in both units using Karen's radial bins. + num_r_actual = polar_intensity_full.shape[2] + r_pixels = np.arange(num_r_actual, dtype=float) * radial_step + theta_full = np.linspace(0, 2*np.pi, polar_intensity_full.shape[-1], endpoint=False) + r_grid_full, theta_grid_full = np.meshgrid(r_pixels, theta_full, indexing='ij') + + # Apply 2-fold symmetry if requested + if two_fold_symmetry: + # Fold to [0, π] + num_theta_folded = polar_intensity_full.shape[-1] // 2 + theta_folded = np.linspace(0, np.pi, num_theta_folded, endpoint=False) + + # Create output arrays + r_grid, theta_grid = np.meshgrid(r_pixels, theta_folded, indexing='ij') + r_invA_grid = r_grid * self.pixels_to_inv_A() + polar_intensity = ( + polar_intensity_full[:, :, :, :num_theta_folded] + + polar_intensity_full[:, :, :, num_theta_folded:] + ) + + num_annular_bins = num_theta_folded + else: + # Use full range + theta_grid = theta_grid_full + r_grid = r_grid_full + r_invA_grid = r_grid * self.pixels_to_inv_A() + polar_intensity = polar_intensity_full + num_annular_bins = num_theta + + # Store metadata + self.max_radius_pixels = max_radius_pixels + self.max_radius_invA = max_radius_invA + self.num_radial_bins = num_r_actual + self.num_annular_bins = num_annular_bins + self.two_fold_symmetry = two_fold_symmetry + + polar_data = { + "r_pixels": r_grid, + "theta": theta_grid, + "r_invA": r_invA_grid, + "intensity": polar_intensity, + } + + return polar_data + + def visualize_peak_detection(self, n_images=10, indices=None, images_per_row=5, figsize_per_image=(3.2, 3), vmax_polar=20, vmax_cartesian=None): + """ + Visualize peak detection results for multiple diffraction patterns. + + Parameters: + ----------- + self : BraggPeaksPolymer + BraggPeaksPolymer object with processed data + n_images : int + Number of images to display (ignored if indices is provided) + indices : list of tuples, optional + List of (ind_y, ind_x) coordinates to visualize. If None, random indices are selected. + images_per_row : int + Number of images per row (default: 5) + figsize_per_image : tuple + Size of each subplot (width, height) + vmax_polar : float + Maximum value for polar data colormap + vmax_cartesian : float + Maximum value for cartesian data colormap + + Returns: + -------- + fig, axes : matplotlib figure and axes + """ + + # Generate or validate indices + if indices is None: + Ry, Rx = self.dataset_cartesian.shape[:2] + # Generate random indices + flat_indices = np.random.choice(Ry * Rx, size=min(n_images, Ry * Rx), replace=False) + indices = [(idx // Rx, idx % Rx) for idx in flat_indices] + else: + n_images = len(indices) + + # Calculate grid dimensions + n_rows = int(np.ceil(n_images / images_per_row)) + n_cols = 5 # 5 types of visualizations per pattern + actual_cols = images_per_row * n_cols + + # Create figure + fig_width = figsize_per_image[0] * actual_cols + fig_height = figsize_per_image[1] * n_rows + fig, axes = plt.subplots(n_rows, actual_cols, figsize=(fig_width, fig_height)) + + # Handle single row case + if n_rows == 1: + axes = axes.reshape(1, -1) + + # Column titles (only for first row) + col_titles = [ + "Polar Transform", + "Polar + Peaks", + "Cartesian + Peaks", + "Cartesian Original", + "Cartesian Normalized" + ] + + # Process each image + for img_idx, (ind_y, ind_x) in enumerate(indices): + row = img_idx // images_per_row + col_offset = (img_idx % images_per_row) * n_cols + + # Check if peaks exist for this pattern + has_peaks = (self.peak_coordinates_cartesian[ind_y, ind_x] is not None and + len(self.peak_coordinates_cartesian[ind_y, ind_x]) > 0) + + # 1. Polar Transform + ax = axes[row, col_offset] + print(self.polar_data["intensity"][ind_y, ind_x].shape) + im = ax.matshow(self.polar_data["intensity"][ind_y, ind_x], cmap='turbo', vmax=vmax_polar) + if row == 0: + ax.set_title(col_titles[0], fontsize=10, pad=10) + ax.text(0.05, 0.95, f'({ind_y},{ind_x})', transform=ax.transAxes, + fontsize=8, va='top', ha='left', color='white', + bbox=dict(boxstyle='round', facecolor='black', alpha=0.5)) + ax.set_axis_off() + + # 2. Polar Transform with Peaks + ax = axes[row, col_offset + 1] + ax.matshow(self.polar_data["intensity"][ind_y, ind_x], cmap='turbo', vmax=vmax_polar) + if has_peaks and self.polar_peaks[ind_y, ind_x] is not None and len(self.polar_peaks[ind_y, ind_x]) > 0: + # Convert radial coordinates to bin indices + r_coords = self.polar_peaks[ind_y, ind_x][:, 0] + theta_coords = self.polar_peaks[ind_y, ind_x][:, 1] + + # Convert theta from radians to angular bins (0 to num_annular_bins) + theta_period = np.pi if getattr(self, "two_fold_symmetry", False) else 2 * np.pi + theta_bins = theta_coords * (self.num_annular_bins / theta_period) + + ax.scatter(theta_bins, r_coords, c='red', s=15, alpha=0.8, edgecolors='white', linewidths=0.5) + if row == 0: + ax.set_title(col_titles[1], fontsize=10, pad=10) + ax.set_axis_off() + + # 3. Cartesian with Peaks and Center + img = self.dataset_cartesian[ind_y, ind_x].array + lower_q = 0.01 + upper_q = 0.99 + vmin, vmax = np.quantile(img[np.isfinite(img)], [lower_q, upper_q]) + if vmax_cartesian is None: + vmax_cartesian = vmax + ax = axes[row, col_offset + 2] + ax.matshow(img, cmap="gray", vmin=vmin, vmax=vmax_cartesian) + # ax.matshow(self.dataset_cartesian[ind_y, ind_x].array, cmap="gray", vmax=vmax_cartesian) + # ax.matshow(self.resized_cartesian_data[ind_y, ind_x], cmap="gray", vmax=vmax_cartesian) + if has_peaks: + ax.scatter(self.peak_coordinates_cartesian[ind_y, ind_x][:, 1], + self.peak_coordinates_cartesian[ind_y, ind_x][:, 0], + c='red', s=15, alpha=0.8, edgecolors='white', linewidths=0.5) + ax.scatter(self.image_centers[1, ind_y, ind_x], + self.image_centers[0, ind_y, ind_x], + c='red', s=500, marker='x', linewidths=2) + if row == 0: + ax.set_title(col_titles[2], fontsize=10, pad=10) + ax.set_axis_off() + + # 4. Original Cartesian + ax = axes[row, col_offset + 3] + im = ax.matshow(img, cmap="gray", vmin=vmin, vmax=vmax_cartesian) + # im = ax.matshow(self.dataset_cartesian[ind_y, ind_x].array, cmap="gray", vmax=vmax_cartesian) + # im = ax.matshow(self.resized_cartesian_data[ind_y, ind_x], cmap="gray", vmax=vmax_cartesian) + if row == 0: + ax.set_title(col_titles[3], fontsize=10, pad=10) + ax.set_axis_off() + plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04) + + # # 5. Normalized Cartesian + # ax = axes[row, col_offset + 4] + # im = ax.matshow(self.normalized_dps_array[ind_y, ind_x], cmap="gray") + # if row == 0: + # ax.set_title(col_titles[4], fontsize=10, pad=10) + # ax.set_axis_off() + # plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04) + + # Hide unused subplots + total_plots = n_images + for idx in range(total_plots, n_rows * images_per_row): + row = idx // images_per_row + col_offset = (idx % images_per_row) * n_cols + for col in range(n_cols): + axes[row, col_offset + col].set_visible(False) + + fig.tight_layout() + return fig, axes + + def estimate_peak_windows( + self, + num_bins=200, + q_min=None, + q_max=None, + n_peaks=5, + height_percentile=10, + prominence_factor=0.1, + width_factor=2.0, + min_width=0.05, + smoothing_sigma=2.0, + intensity_field='intensities', + mode='intensity', + log_scale=False, + ): + """ + Automatically detect the top N most prominent peaks and estimate their windows. + + Parameters + ---------- + num_bins : int + Number of radial bins + q_min : float, optional + Minimum q value for binning + q_max : float, optional + Maximum q value for binning + n_peaks : int + Number of top peaks to detect + height_percentile : float + Percentile threshold for peak height (peaks below this are ignored) + prominence_factor : float + Factor of max intensity for minimum peak prominence + width_factor : float + Multiplier for estimating peak window width from FWHM + min_width : float + Minimum window width in 1/Å + smoothing_sigma : float + Gaussian smoothing sigma for noise reduction before peak detection + mode : {'intensity', 'count'} + Radial profile to detect peaks on: intensity-weighted histogram of peak q + ('intensity', default) or the number of detected peaks per bin ('count'). + log_scale : bool + If True, detect peaks on log1p(profile) so small peaks are not dominated + by large ones. + + Returns + ------- + peak_centers : array + q-values for peak centers (shape: n_peaks) + peak_windows : array + Window boundaries for each peak (shape: n_peaks, 2) + Each row is [q_min, q_max] for that peak + peak_info : dict + Additional information about detected peaks including: + - 'heights': peak heights + - 'prominences': peak prominences + - 'widths': estimated peak widths (FWHM) + """ + + if mode not in ('intensity', 'count'): + raise ValueError(f"mode must be 'intensity' or 'count', got {mode!r}") + + # Get radial profile (intensity-weighted or peak-count) + all_r = _vector_field_flat(self.polar_peaks, "r_invA") + + if q_min is None: + q_min = 0 + if q_max is None: + q_max = np.max(all_r) + + r_bins = np.linspace(q_min, q_max, num_bins + 1) + if mode == 'intensity': + all_intensity = _vector_field_flat(self.peak_intensities, intensity_field) + profile, _ = np.histogram(all_r, bins=r_bins, weights=all_intensity) + else: # 'count' + profile, _ = np.histogram(all_r, bins=r_bins) + r_centers = (r_bins[:-1] + r_bins[1:]) / 2 + + # Optional log compression so small peaks are not dominated by large ones + if log_scale: + profile = np.log1p(profile) + + # Smooth the data to reduce noise + if smoothing_sigma > 0: + intensity_smooth = gaussian_filter1d(profile, smoothing_sigma) + else: + intensity_smooth = profile + + # Calculate thresholds + height_threshold = np.percentile(intensity_smooth, height_percentile) + prominence_threshold = prominence_factor * np.max(intensity_smooth) + + # Find peaks + peaks_indices, properties = find_peaks( + intensity_smooth, + height=height_threshold, + prominence=prominence_threshold, + distance=int(min_width / (r_centers[1] - r_centers[0])) # Minimum separation + ) + + if len(peaks_indices) == 0: + print("No peaks found with current parameters!") + return np.array([]), np.array([]).reshape(0, 2), {} + + # Sort by prominence and take top N + prominences = properties['prominences'] + sorted_indices = np.argsort(prominences)[::-1][:n_peaks] + top_peak_indices = peaks_indices[sorted_indices] + top_peak_indices = np.sort(top_peak_indices) # Re-sort by position + + # Get peak centers + peak_centers = r_centers[top_peak_indices] + + # Calculate peak widths (FWHM) + widths_data = peak_widths(intensity_smooth, top_peak_indices, rel_height=0.5) + fwhm_bins = widths_data[0] # Width in bins + fwhm_invA = fwhm_bins * (r_centers[1] - r_centers[0]) # Convert to 1/Å + + # Estimate windows: center ± width_factor * FWHM/2, with minimum width + half_widths = np.maximum(width_factor * fwhm_invA / 2, min_width / 2) + peak_windows = np.column_stack([ + peak_centers - half_widths, + peak_centers + half_widths + ]) + + # Clip windows to data range + peak_windows[:, 0] = np.maximum(peak_windows[:, 0], q_min) + peak_windows[:, 1] = np.minimum(peak_windows[:, 1], q_max) + + # Collect additional info + peak_info = { + 'heights': intensity_smooth[top_peak_indices], + 'prominences': prominences[sorted_indices], + 'widths_fwhm': fwhm_invA, + 'intensity_profile': intensity_smooth, + 'profile': intensity_smooth, + 'r_centers': r_centers, + 'mode': mode, + 'log_scale': log_scale, + } + + # Print summary + print(f"Detected {len(peak_centers)} peaks:") + _to_d = lambda q: (1.0 / q if q > 0 else float('inf')) # d-spacing (Å) = 1 / q (1/Å) + for i, (center, window, height, prom, width) in enumerate(zip( + peak_centers, peak_windows, peak_info['heights'], + peak_info['prominences'], peak_info['widths_fwhm'] + )): + print(f" Peak {i+1}: center={center:.3f} 1/Å (d={_to_d(center):.2f} Å), " + f"window=[{window[0]:.3f}, {window[1]:.3f}] 1/Å " + f"(d=[{_to_d(window[1]):.2f}, {_to_d(window[0]):.2f}] Å), " + f"height={height:.1f}, prominence={prom:.1f}, FWHM={width:.3f} 1/Å") + + return peak_centers, peak_windows, peak_info + + def peak_radial_intensity_plot( + self, + num_bins=200, + q_min=None, + q_max=None, + ROI_xs=None, + ROI_ys=None, + peak_centers=None, + peak_windows=None, + vlines=None, + vline_colors=None, + vline_labels=None, + window_alpha=0.3, + window_color='red', + fill_alpha=0.5, + fill_color=None, + plot=True, + return_data=False, + intensity_field='intensities', + log_scale=False, + show_d_spacing=False, + ): + """ + Create radial intensity line plot summarizing polar peaks. + + Parameters + ---------- + num_bins : int + Number of radial bins + q_min : float, optional + Minimum q value for binning + q_max : float, optional + Maximum q value for binning + ROI_xs : tuple, optional + X range for region of interest (not yet implemented) + ROI_ys : tuple, optional + Y range for region of interest (not yet implemented) + peak_centers : array, optional + 1D array of peak center positions to mark with vertical lines + peak_windows : array, optional + 2D array (N, 2) of [q_min, q_max] for each peak window to highlight + vlines : list of lists/arrays, optional + Additional vertical lines to plot. Each element is a list/array of x-positions. + vline_colors : list of colors, optional + Colors for each group of vertical lines + vline_labels : list of str, optional + Labels for each group of vertical lines (for legend) + window_alpha : float + Transparency for peak window background highlighting (0-1) + window_color : str or color + Color for peak window background highlighting + fill_alpha : float + Transparency for filled area under curve within windows (0-1) + fill_color : str or color, optional + Color for filled area under curve. If None, uses window_color + plot : bool + Whether to display the plot + return_data : bool + Whether to return the binned data + + Returns + ------- + r_centers : array (optional) + Radial bin centers + intensity_sum : array (optional) + Integrated intensity per bin + """ + all_r = _vector_field_flat(self.polar_peaks, "r_invA") + all_intensity = _vector_field_flat(self.peak_intensities, intensity_field) + + if q_min is None: + q_min = 0 + if q_max is None: + q_max = np.max(all_r) + r_bins = np.linspace(q_min, q_max, num_bins + 1) + + # Histogram the data + intensity_sum, _ = np.histogram(all_r, bins=r_bins, weights=all_intensity) + counts, _ = np.histogram(all_r, bins=r_bins) + + # Bin centers + r_centers = (r_bins[:-1] + r_bins[1:]) / 2 + + # Use window_color for fill if not specified + if fill_color is None: + fill_color = window_color + + if plot: + # Create line plot + fig, ax = plt.subplots() + ax.plot(r_centers, intensity_sum, linewidth=2, label='Intensity', color='black') + ax.set_xlabel('Radial Distance (1/Å)', fontsize=12) + ax.set_ylabel('Integrated Intensity', fontsize=12) + ax.set_title('Radial Intensity Profile (All Patterns)', fontsize=14) + ax.grid(True, alpha=0.3) + + fill_base = 0 + if log_scale: + ax.set_yscale('log') + _pos = intensity_sum[intensity_sum > 0] + fill_base = (_pos.min() if _pos.size else 1e-9) + + if show_d_spacing: + # top axis: real-space d-spacing (Å) = 1 / q (1/Å) + secax = ax.secondary_xaxis( + 'top', + functions=(lambda q: 1.0 / np.clip(q, 1e-12, None), + lambda d: 1.0 / np.clip(d, 1e-12, None)), + ) + secax.set_xlabel('d-spacing (Å)', fontsize=12) + + # Add peak windows as filled regions and fill under curve + if peak_windows is not None: + peak_windows = np.atleast_2d(peak_windows) + for i, (q_min_win, q_max_win) in enumerate(peak_windows): + # Background window highlight + ax.axvspan(q_min_win, q_max_win, alpha=window_alpha, + color=window_color, zorder=0, + label='Peak windows' if i == 0 else None) + + # Fill under the curve within this window + # Find indices within the window + mask = (r_centers >= q_min_win) & (r_centers <= q_max_win) + if np.any(mask): + r_window = r_centers[mask] + intensity_window = intensity_sum[mask] + ax.fill_between(r_window, fill_base, intensity_window, + alpha=fill_alpha, color=fill_color, + label='Peak intensity' if i == 0 else None, + zorder=1) + + # Add peak centers as vertical lines + if peak_centers is not None: + peak_centers = np.atleast_1d(peak_centers) + for i, center in enumerate(peak_centers): + ax.axvline(center, color=window_color, linestyle='-', + linewidth=2, alpha=0.8, + label='Peak centers' if i == 0 else None, zorder=2) + + # Add additional vertical lines if provided + if vlines is not None: + # Convert to list of lists if needed + if not isinstance(vlines[0], (list, np.ndarray)): + vlines = [vlines] + + # Default colors if not provided + if vline_colors is None: + default_colors = plt.cm.tab10(np.linspace(0, 1, len(vlines))) + vline_colors = default_colors + + # Ensure vline_colors is a list + if not isinstance(vline_colors, list): + vline_colors = [vline_colors] + + # Check length match + if len(vline_colors) != len(vlines): + raise ValueError( + f"Number of vline_colors ({len(vline_colors)}) must match " + f"number of vline groups ({len(vlines)})" + ) + + # Plot each group of vertical lines + for i, (vline_group, color) in enumerate(zip(vlines, vline_colors)): + # Get label if provided + label = vline_labels[i] if vline_labels is not None and i < len(vline_labels) else None + + # Plot each line in the group + for j, x_pos in enumerate(vline_group): + # Only add label to first line in group (for legend) + line_label = label if j == 0 else None + ax.axvline(x_pos, color=color, linestyle='--', + linewidth=1.5, alpha=0.7, label=line_label, zorder=2) + + # Add legend + ax.legend() + elif peak_centers is not None or peak_windows is not None: + # Add legend for peak markers if present + ax.legend() + + fig.tight_layout() + plt.show() + + if return_data: + return r_centers, intensity_sum + + def peak_radial_count_plot( + self, + num_bins=200, + q_min=None, + q_max=None, + ROI_xs=None, + ROI_ys=None, + peak_centers=None, + peak_windows=None, + vlines=None, + vline_colors=None, + vline_labels=None, + window_alpha=0.3, + window_color='red', + fill_alpha=0.5, + fill_color=None, + plot=True, + return_data=False, + log_scale=False, + show_d_spacing=False, + ): + """ + Create radial peak count line plot summarizing polar peaks. + + Parameters + ---------- + num_bins : int + Number of radial bins + q_min : float, optional + Minimum q value for binning + q_max : float, optional + Maximum q value for binning + ROI_xs : tuple, optional + X range for region of interest (not yet implemented) + ROI_ys : tuple, optional + Y range for region of interest (not yet implemented) + peak_centers : array, optional + 1D array of peak center positions to mark with vertical lines + peak_windows : array, optional + 2D array (N, 2) of [q_min, q_max] for each peak window to highlight + vlines : list of lists/arrays, optional + Additional vertical lines to plot. Each element is a list/array of x-positions. + vline_colors : list of colors, optional + Colors for each group of vertical lines + vline_labels : list of str, optional + Labels for each group of vertical lines (for legend) + window_alpha : float + Transparency for peak window background highlighting (0-1) + window_color : str or color + Color for peak window background highlighting + fill_alpha : float + Transparency for filled area under curve within windows (0-1) + fill_color : str or color, optional + Color for filled area under curve. If None, uses window_color + plot : bool + Whether to display the plot + return_data : bool + Whether to return the binned data + + Returns + ------- + r_centers : array (optional) + Radial bin centers + peak_counts : array (optional) + Number of peaks per bin + """ + all_r = _vector_field_flat(self.polar_peaks, "r_invA") + + if q_min is None: + q_min = 0 + if q_max is None: + q_max = np.max(all_r) + r_bins = np.linspace(q_min, q_max, num_bins + 1) + + # Histogram the data - counts only, no weights + peak_counts, _ = np.histogram(all_r, bins=r_bins) + + # Bin centers + r_centers = (r_bins[:-1] + r_bins[1:]) / 2 + + # Use window_color for fill if not specified + if fill_color is None: + fill_color = window_color + + if plot: + # Create line plot + fig, ax = plt.subplots() + ax.plot(r_centers, peak_counts, linewidth=2, label='Peak Count', color='black') + ax.set_xlabel('Radial Distance (1/Å)', fontsize=12) + ax.set_ylabel('Number of Peaks', fontsize=12) + ax.set_title('Radial Peak Count Profile (All Patterns)', fontsize=14) + ax.grid(True, alpha=0.3) + + fill_base = 0 + if log_scale: + ax.set_yscale('log') + _pos = peak_counts[peak_counts > 0] + fill_base = (_pos.min() if _pos.size else 1e-9) + + if show_d_spacing: + # top axis: real-space d-spacing (Å) = 1 / q (1/Å) + secax = ax.secondary_xaxis( + 'top', + functions=(lambda q: 1.0 / np.clip(q, 1e-12, None), + lambda d: 1.0 / np.clip(d, 1e-12, None)), + ) + secax.set_xlabel('d-spacing (Å)', fontsize=12) + + # Add peak windows as filled regions and fill under curve + if peak_windows is not None: + peak_windows = np.atleast_2d(peak_windows) + for i, (q_min_win, q_max_win) in enumerate(peak_windows): + # Background window highlight + ax.axvspan(q_min_win, q_max_win, alpha=window_alpha, + color=window_color, zorder=0, + label='Peak windows' if i == 0 else None) + + # Fill under the curve within this window + # Find indices within the window + mask = (r_centers >= q_min_win) & (r_centers <= q_max_win) + if np.any(mask): + r_window = r_centers[mask] + counts_window = peak_counts[mask] + ax.fill_between(r_window, fill_base, counts_window, + alpha=fill_alpha, color=fill_color, + label='Peak counts' if i == 0 else None, + zorder=1) + + # Add peak centers as vertical lines + if peak_centers is not None: + peak_centers = np.atleast_1d(peak_centers) + for i, center in enumerate(peak_centers): + ax.axvline(center, color=window_color, linestyle='-', + linewidth=2, alpha=0.8, + label='Peak centers' if i == 0 else None, zorder=2) + + # Add additional vertical lines if provided + if vlines is not None: + # Convert to list of lists if needed + if not isinstance(vlines[0], (list, np.ndarray)): + vlines = [vlines] + + # Default colors if not provided + if vline_colors is None: + default_colors = plt.cm.tab10(np.linspace(0, 1, len(vlines))) + vline_colors = default_colors + + # Ensure vline_colors is a list + if not isinstance(vline_colors, list): + vline_colors = [vline_colors] + + # Check length match + if len(vline_colors) != len(vlines): + raise ValueError( + f"Number of vline_colors ({len(vline_colors)}) must match " + f"number of vline groups ({len(vlines)})" + ) + + # Plot each group of vertical lines + for i, (vline_group, color) in enumerate(zip(vlines, vline_colors)): + # Get label if provided + label = vline_labels[i] if vline_labels is not None and i < len(vline_labels) else None + + # Plot each line in the group + for j, x_pos in enumerate(vline_group): + # Only add label to first line in group (for legend) + line_label = label if j == 0 else None + ax.axvline(x_pos, color=color, linestyle='--', + linewidth=1.5, alpha=0.7, label=line_label, zorder=2) + + # Add legend + ax.legend() + elif peak_centers is not None or peak_windows is not None: + # Add legend for peak markers if present + ax.legend() + + fig.tight_layout() + plt.show() + + if return_data: + return r_centers, peak_counts + + def make_orientation_histogram( + self, + radial_ranges: np.ndarray = None, + orientation_map=None, + orientation_ind: int = 0, + orientation_growth_angles: np.array = 0.0, + orientation_separate_bins: bool = False, + orientation_flip_sign: bool = False, + orientation_offset_degrees: float = 0.0, + upsample_factor: float = 4.0, + theta_step_deg: float = 1.0, + sigma_x: float = 1.0, + sigma_y: float = 1.0, + sigma_theta: float = 3.0, + use_peak_sigma: bool = False, + peak_sigma_samples: int = 6, + normalize_intensity_image: bool = False, + normalize_intensity_stack: bool = True, + progress_bar: bool = True, + r_field: str = "r_invA", + theta_field: str = "theta", + intensity_field: str = "intensities", + # intensity_field: str = "intensities_sampled_from_dp", + ): + """ + Create a 3D or 4D orientation histogram from bragg peaks. + + Can generate histograms from either: + 1. Polar peak data with radial ranges + 2. Orientation map with Euler angles (for fiber textures) + + Parameters + ---------- + radial_ranges : np.ndarray, optional + Size (N x 2) array for N radial bins, or (2,) for a single bin. + orientation_map : OrientationMap, optional + Class containing Euler angles to generate a flowline map. + orientation_ind : int + Index of the orientation map (default 0) + orientation_growth_angles : np.array + Angles to place into histogram, relative to orientation. + orientation_separate_bins : bool + Whether to place multiple angles into multiple radial bins. + orientation_flip_sign : bool + Flip the direction of theta + orientation_offset_degrees : float + Offset for orientation angles in degrees + upsample_factor : float + Upsample factor for output histogram + theta_step_deg : float + Step size along annular direction in degrees + sigma_x : float + Smoothing in x direction before upsample + sigma_y : float + Smoothing in y direction before upsample + sigma_theta : float + Smoothing in annular direction (units of bins, periodic) + use_peak_sigma : bool + Spread signal along annular direction using measured peak width + peak_sigma_samples : int + Number of samples for peak sigma spreading + normalize_intensity_image : bool + Normalize to max peak intensity = 1, per image + normalize_intensity_stack : bool + Normalize to max peak intensity = 1, all images + progress_bar : bool + Enable progress bar + r_field : str + Name of radial coordinate field + theta_field : str + Name of angular coordinate field + intensity_field : str + Name of intensity field + + Returns + ------- + orient_hist : np.ndarray + 4D array containing Bragg peak intensity histogram + [radial_bin, x_probe, y_probe, theta] + """ + # Coordinates + theta = np.arange(0, 180, theta_step_deg) * np.pi / 180.0 + dtheta = theta[1] - theta[0] + dtheta_deg = dtheta * 180 / np.pi + num_theta_bins = np.size(theta) + + # Setup for peak sigma spreading + if use_peak_sigma: + v_sigma = np.linspace(-2, 2, 2 * peak_sigma_samples + 1) + w_sigma = np.exp(-(v_sigma**2) / 2) + + if orientation_map is None: + # Input bins + radial_ranges = np.array(radial_ranges) + if radial_ranges.ndim == 1: + radial_ranges = radial_ranges[None, :] + radial_ranges_2 = radial_ranges**2 + num_radii = radial_ranges.shape[0] + size_input = self.polar_peaks.shape + else: + orientation_growth_angles = np.atleast_1d(orientation_growth_angles) + num_angles = orientation_growth_angles.shape[0] + size_input = [orientation_map.num_x, orientation_map.num_y] + if orientation_separate_bins is False: + num_radii = 1 + else: + num_radii = num_angles + + size_output = np.round( + np.array(size_input).astype("float") * upsample_factor + ).astype("int") + + # Output init + orient_hist = np.zeros([num_radii, size_output[0], size_output[1], num_theta_bins]) + + # Loop over all probe positions + for a0 in range(num_radii): + t = "Generating histogram " + str(a0) + for rx, ry in tqdmnd( + *size_input, desc=t, unit=" probe positions", disable=not progress_bar + ): + x = (rx + 0.5) * upsample_factor - 0.5 + y = (ry + 0.5) * upsample_factor - 0.5 + x = np.clip(x, 0, size_output[0] - 2) + y = np.clip(y, 0, size_output[1] - 2) + xF = np.floor(x).astype("int") + yF = np.floor(y).astype("int") + dx = x - xF + dy = y - yF + + add_data = False + + if orientation_map is None: + p_r = _vector_field_cell(self.polar_peaks, r_field, rx, ry) + p_theta = _vector_field_cell(self.polar_peaks, theta_field, rx, ry) + + if p_r is not None and len(p_r) > 0: + r2 = p_r**2 + sub = np.logical_and( + r2 >= radial_ranges_2[a0, 0], + r2 < radial_ranges_2[a0, 1] + ) + if np.any(sub): + intensity_data = _vector_field_cell( + self.peak_intensities, intensity_field, rx, ry + ) + if intensity_data is not None and len(intensity_data) > 0: + add_data = True + intensity = intensity_data[sub] + + # Get theta values + theta_radians = p_theta[sub] + if orientation_flip_sign: + theta_radians *= -1 + # Add offset + theta_radians += orientation_offset_degrees * np.pi / 180 + theta_radians = np.mod(theta_radians, np.pi) + t = theta_radians / dtheta + + # Spread signal using peak sigma if requested + if use_peak_sigma: + # Try to get sigma values if available + if 'sigma_theta' in self.polar_peaks.fields: + theta_std = _vector_field_cell( + self.polar_peaks, "sigma_theta", rx, ry + )[sub] / dtheta + t = (t[:, None] + theta_std[:, None] * v_sigma[None, :]).ravel() + intensity = (intensity[:, None] * w_sigma[None, :]).ravel() + else: + if orientation_map.corr[rx, ry, orientation_ind] > 0: + if orientation_separate_bins is False: + if orientation_flip_sign: + t = ( + np.array( + [ + ( + -orientation_map.angles[ + rx, ry, orientation_ind, 0 + ] + - orientation_map.angles[ + rx, ry, orientation_ind, 2 + ] + ) + / dtheta + ] + ) + + orientation_growth_angles + ) + else: + t = ( + np.array( + [ + ( + orientation_map.angles[ + rx, ry, orientation_ind, 0 + ] + + orientation_map.angles[ + rx, ry, orientation_ind, 2 + ] + ) + / dtheta + ] + ) + + orientation_growth_angles + ) + # Add offset + t += orientation_offset_degrees / dtheta_deg + intensity = ( + np.ones(num_angles) + * orientation_map.corr[rx, ry, orientation_ind] + ) + add_data = True + else: + if orientation_flip_sign: + t = ( + np.array( + [ + ( + -orientation_map.angles[ + rx, ry, orientation_ind, 0 + ] + - orientation_map.angles[ + rx, ry, orientation_ind, 2 + ] + ) + / dtheta + ] + ) + + orientation_growth_angles[a0] + ) + else: + t = ( + np.array( + [ + ( + orientation_map.angles[ + rx, ry, orientation_ind, 0 + ] + + orientation_map.angles[ + rx, ry, orientation_ind, 2 + ] + ) + / dtheta + ] + ) + + orientation_growth_angles[a0] + ) + # Add offset + t += orientation_offset_degrees / dtheta_deg + intensity = orientation_map.corr[rx, ry, orientation_ind] + add_data = True + + if add_data: + tF = np.floor(t).astype("int") + dt = t - tF + + orient_hist[a0, xF, yF, :] = orient_hist[a0, xF, yF, :] + np.bincount( + np.mod(tF, num_theta_bins), + weights=(1 - dx) * (1 - dy) * (1 - dt) * intensity, + minlength=num_theta_bins, + ) + orient_hist[a0, xF, yF, :] = orient_hist[a0, xF, yF, :] + np.bincount( + np.mod(tF + 1, num_theta_bins), + weights=(1 - dx) * (1 - dy) * (dt) * intensity, + minlength=num_theta_bins, + ) + + orient_hist[a0, xF + 1, yF, :] = orient_hist[ + a0, xF + 1, yF, : + ] + np.bincount( + np.mod(tF, num_theta_bins), + weights=(dx) * (1 - dy) * (1 - dt) * intensity, + minlength=num_theta_bins, + ) + orient_hist[a0, xF + 1, yF, :] = orient_hist[ + a0, xF + 1, yF, : + ] + np.bincount( + np.mod(tF + 1, num_theta_bins), + weights=(dx) * (1 - dy) * (dt) * intensity, + minlength=num_theta_bins, + ) + + orient_hist[a0, xF, yF + 1, :] = orient_hist[ + a0, xF, yF + 1, : + ] + np.bincount( + np.mod(tF, num_theta_bins), + weights=(1 - dx) * (dy) * (1 - dt) * intensity, + minlength=num_theta_bins, + ) + orient_hist[a0, xF, yF + 1, :] = orient_hist[ + a0, xF, yF + 1, : + ] + np.bincount( + np.mod(tF + 1, num_theta_bins), + weights=(1 - dx) * (dy) * (dt) * intensity, + minlength=num_theta_bins, + ) + + orient_hist[a0, xF + 1, yF + 1, :] = orient_hist[ + a0, xF + 1, yF + 1, : + ] + np.bincount( + np.mod(tF, num_theta_bins), + weights=(dx) * (dy) * (1 - dt) * intensity, + minlength=num_theta_bins, + ) + orient_hist[a0, xF + 1, yF + 1, :] = orient_hist[ + a0, xF + 1, yF + 1, : + ] + np.bincount( + np.mod(tF + 1, num_theta_bins), + weights=(dx) * (dy) * (dt) * intensity, + minlength=num_theta_bins, + ) + + # Smoothing / interpolation + if (sigma_x is not None) or (sigma_y is not None) or (sigma_theta is not None): + if num_radii > 1: + print("Interpolating orientation matrices ...", end="") + else: + print("Interpolating orientation matrix ...", end="") + if sigma_x is not None and sigma_x > 0: + orient_hist = gaussian_filter1d( + orient_hist, + sigma_x * upsample_factor, + mode="nearest", + axis=1, + truncate=3.0, + ) + if sigma_y is not None and sigma_y > 0: + orient_hist = gaussian_filter1d( + orient_hist, + sigma_y * upsample_factor, + mode="nearest", + axis=2, + truncate=3.0, + ) + if sigma_theta is not None and sigma_theta > 0: + orient_hist = gaussian_filter1d( + orient_hist, sigma_theta / dtheta_deg, mode="wrap", axis=3, truncate=2.0 + ) + print(" done.") + + # Normalization + if normalize_intensity_stack is True: + stack_max = np.max(orient_hist) + if stack_max > 0: + orient_hist = orient_hist / stack_max + elif normalize_intensity_image is True: + for a0 in range(num_radii): + image_max = np.max(orient_hist[a0, :, :, :]) + if image_max > 0: + orient_hist[a0, :, :, :] /= image_max + + return orient_hist + + def calculate_orientation_correlation( + self, + orient_hist, + radius_max=None, + pairs="all", + backend="auto", + device=None, + mode_batch_size=None, + pair_batch_size=None, + max_memory_fraction=0.6, + dtype="float32", + workers=None, + zero_policy="nan", + return_numpy=True, + store_result=True, + progress_bar=True, + ): + """ + Calculate distance-angle correlations from an orientation histogram. + + This method is mathematically equivalent to constructing the full + ``(dx, dy, relative_theta)`` correlation volume, but processes angular + Fourier modes in batches and performs the radial integration before the + angular inverse transform. This substantially reduces peak memory and + allows the FFT work to run on a GPU. + + Parameters + ---------- + orient_hist : numpy.ndarray or torch.Tensor + Histogram with shape ``(radial_bin, scan_x, scan_y, theta)``. + A three-dimensional ``(scan_x, scan_y, theta)`` input is treated as + a single radial bin. + radius_max : int, optional + Maximum spatial separation in orientation-histogram pixels. + Defaults to half of the smaller scan dimension. + pairs : {"all", "autocorrelation"} or sequence of tuple[int, int] + Radial-bin pairs to correlate. ``"all"`` uses upper-triangular + ordering; ``"autocorrelation"`` calculates only ``(i, i)``. + backend : {"auto", "numpy", "torch"} + ``"auto"`` uses PyTorch when CUDA is available and NumPy otherwise. + device : str or torch.device, optional + PyTorch device. Defaults to CUDA when available, otherwise CPU. + mode_batch_size, pair_batch_size : int, optional + Angular-frequency and radial-pair batch sizes. CUDA mode batching is + automatically sized from available memory when omitted. + max_memory_fraction : float + Fraction of currently free CUDA memory available to automatic + batching. + dtype : {"float32", "float64"} + Real computation dtype. ``float32`` is recommended for CUDA. + workers : int, optional + Number of SciPy FFT workers for the NumPy backend. + zero_policy : {"nan", "zero", "raise"} + Handling for radial distances with no normalization signal. + return_numpy : bool + Convert PyTorch output to a NumPy array before returning. + store_result : bool + Store output in ``self.orient_corr`` and its radial-bin mapping in + ``self.orient_corr_pairs``. + progress_bar : bool + Display progress over angular-mode and radial-pair batches. + + Returns + ------- + numpy.ndarray or torch.Tensor + Array with shape + ``(num_pairs, num_theta // 2 + 1, radius_max + 1)`` in multiples of + a random distribution. A value of 1 indicates random association. + + Notes + ----- + The full ``pairs="all"`` output uses upper-triangular radial-bin pair + ordering. Use ``self.orient_corr_pairs`` to label the first output axis. + """ + orient_corr, pair_indices = _calculate_orientation_correlation( + orient_hist, + radius_max=radius_max, + pairs=pairs, + backend=backend, + device=device, + mode_batch_size=mode_batch_size, + pair_batch_size=pair_batch_size, + max_memory_fraction=max_memory_fraction, + dtype=dtype, + workers=workers, + zero_policy=zero_policy, + return_numpy=return_numpy, + progress_bar=progress_bar, + ) + if store_result: + self.orient_corr = orient_corr + self.orient_corr_pairs = pair_indices + return orient_corr + + def plot_orientation_correlation( + self, + orient_corr=None, + *, + pair_indices=None, + pixel_size=1.0, + pixel_units="scan pixels", + probability_range=(0.5, 2.0), + cmap="correlation", + figsize=None, + show_metrics=True, + return_metrics=False, + slope_weight_scale=None, + ): + """Plot distance-orientation correlations using Matplotlib. + + The 50% boundary is halfway between the correlation at zero separation + and the random-association baseline of one. Its intercepts give the + radial and annular 50% distances. The signed slope is fitted separately + to the primary correlation-equals-one boundary between positive + correlation and anticorrelation. + + Parameters + ---------- + slope_weight_scale : float, optional + 1/e decay length, in ``pixel_units``, of the exponential weighting + applied to the slope fit. The correlation-equals-one boundary + saturates with distance, so an unweighted straight line over the + whole lobe is dominated by the flat tail: it biases the slope low + and pushes the fitted intercept off the measured boundary at zero + separation. Weighting towards short distances makes ``slope`` the + near-origin tangent instead. Defaults to + ``SLOPE_WEIGHT_FRACTION`` times the fitted distance span. Pass + ``numpy.inf`` to restore the previous unweighted full-lobe fit. + + Notes + ----- + Each metrics entry reports ``slope_fit_intercept_degrees`` (compare it + against the boundary at zero separation to check the fit), + ``slope_fit_effective_point_count`` (Kish effective sample size, which + falls as the weighting sharpens), and the resolved + ``slope_weight_scale``. ``slope_fit_r_squared`` is weighted with the + same weights, so it is not comparable to an unweighted R-squared. + """ + from matplotlib.colors import LinearSegmentedColormap, LogNorm + from matplotlib.lines import Line2D + + def crossing(coordinates, profile, level): + profile = np.asarray(profile, dtype=float) + coordinates = np.asarray(coordinates, dtype=float) + if not np.isfinite(profile[0]): + return np.nan + initial_side = profile[0] - level + if initial_side == 0: + return float(coordinates[0]) + for point in range(1, len(profile)): + before, after = profile[point - 1], profile[point] + if not np.isfinite(before) or not np.isfinite(after): + continue + before_side = before - level + after_side = after - level + if before_side == 0: + return float(coordinates[point - 1]) + if before_side * after_side <= 0: + if before == after: + return float(coordinates[point]) + fraction = -before_side / (after_side - before_side) + return float( + coordinates[point - 1] + + fraction * (coordinates[point] - coordinates[point - 1]) + ) + return np.nan + + values = self.orient_corr if orient_corr is None else orient_corr + if values is None: + raise RuntimeError( + "No orientation correlation is available. Run " + "calculate_orientation_correlation() first or pass orient_corr." + ) + values = np.asarray(values) + if values.ndim != 3: + raise ValueError( + "orient_corr must have shape (pair, relative_angle, distance)." + ) + if values.shape[0] == 0: + raise ValueError("orient_corr must contain at least one radial-bin pair.") + if values.shape[1] < 2 or values.shape[2] < 2: + raise ValueError( + "orient_corr requires at least two angle and two distance samples." + ) + labels = self.orient_corr_pairs if pair_indices is None else pair_indices + if labels is not None: + labels = np.asarray(labels) + if labels.shape != (values.shape[0], 2): + raise ValueError( + f"pair_indices must have shape ({values.shape[0]}, 2)." + ) + + panel_count = values.shape[0] + column_count = min(3, max(1, panel_count)) + row_count = int(np.ceil(panel_count / column_count)) + if figsize is None: + figsize = (4.5 * column_count, 3.8 * row_count) + fig, axes = plt.subplots( + row_count, + column_count, + figsize=figsize, + squeeze=False, + constrained_layout=True, + ) + lower, upper = map(float, probability_range) + if not 0 < lower < upper: + raise ValueError("probability_range must satisfy 0 < lower < upper.") + if cmap == "correlation": + cmap = LinearSegmentedColormap.from_list( + "quantem_correlation", + [ + (0.00, "#002b9a"), + (0.32, "#1769e8"), + (0.50, "#b8b8b8"), + (0.68, "#f23838"), + (1.00, "#9e0015"), + ], + ) + distance_max = (values.shape[2] - 1) * float(pixel_size) + distances = np.arange(values.shape[2], dtype=float) * float(pixel_size) + angles = np.linspace(0.0, 180.0, values.shape[1]) + image = None + metrics = [] + for index, ax in enumerate(axes.flat): + if index >= panel_count: + ax.set_visible(False) + continue + image = ax.imshow( + values[index], + origin="lower", + aspect="auto", + extent=(0, distance_max, 0, 180), + norm=LogNorm(vmin=lower, vmax=upper), + cmap=cmap, + ) + if labels is None: + title = f"Ring pair {index}" + pair = (index, index) + else: + pair = tuple(int(value) for value in labels[index]) + title = ( + f"Autocorrelation of Ring {pair[0]}" + if pair[0] == pair[1] + else f"Correlation of Rings {pair[0]} and {pair[1]}" + ) + ax.set( + title=title, + xlabel=f"distance ({pixel_units})", + ylabel="relative orientation (degrees)", + ) + panel = np.asarray(values[index], dtype=float) + origin_probability = panel[0, 0] + half_probability = ( + 1.0 + 0.5 * (origin_probability - 1.0) + if np.isfinite(origin_probability) + and origin_probability > 0.0 + and not np.isclose(origin_probability, 1.0) + else np.nan + ) + radial_distance = np.nan + annular_distance = np.nan + slope = np.nan + slope_fit_r_squared = np.nan + slope_fit_point_count = 0 + slope_fit_intercept = np.nan + slope_fit_effective_count = np.nan + weight_scale = np.nan + fit_distances = np.array([]) + fit_angles = np.array([]) + if np.isfinite(half_probability): + radial_distance = crossing( + distances, panel[0, :], half_probability + ) + annular_distance = crossing( + angles, panel[:, 0], half_probability + ) + if np.isfinite(radial_distance): + ax.scatter( + [radial_distance], + [0], + marker="o", + s=45, + facecolor="white", + edgecolor="black", + linewidth=0.8, + zorder=5, + ) + if np.isfinite(annular_distance): + ax.scatter( + [0], + [annular_distance], + marker="D", + s=40, + facecolor="white", + edgecolor="black", + linewidth=0.8, + zorder=5, + ) + + # The slope belongs to the gray probability/random = 1 boundary, not + # the half-maximum contour used for the two distance intercepts. + baseline_boundary = np.array( + [ + crossing(angles, panel[:, radius], 1.0) + for radius in range(panel.shape[1]) + ] + ) + baseline_radial_intercept = crossing(distances, panel[0, :], 1.0) + valid = np.isfinite(baseline_boundary) + if np.isfinite(baseline_radial_intercept): + valid &= distances <= baseline_radial_intercept + float(pixel_size) + + # Select the earliest contiguous run: this follows the principal + # red/blue lobe from the angular axis and rejects remote closed loops. + valid_indices = np.flatnonzero(valid) + primary_indices = np.array([], dtype=int) + if valid_indices.size: + angular_jump = np.abs( + np.diff(baseline_boundary[valid_indices]) + ) + maximum_step = max(15.0, 4.0 * (angles[1] - angles[0])) + split_points = np.flatnonzero( + (np.diff(valid_indices) > 1) | (angular_jump > maximum_step) + ) + 1 + segments = np.split( + valid_indices, split_points + ) + primary_indices = next( + (segment for segment in segments if len(segment) >= 2), + np.array([], dtype=int), + ) + if primary_indices.size >= 5: + # A connected correlation=1 contour can rise away from the + # origin, turn around at large distance, and return as part of + # the same loop. Fitting that entire loop can reverse the sign + # of the visually obvious near-origin boundary. Stop at the + # first sustained turning point while tolerating isolated + # pixel-scale contour noise. + boundary_run = baseline_boundary[primary_indices] + smoothing_sigma = min(3.0, max(0.75, len(boundary_run) / 100.0)) + boundary_smooth = gaussian_filter1d( + boundary_run, sigma=smoothing_sigma, mode="nearest" + ) + boundary_gradient = np.gradient(boundary_smooth) + initial_count = min(20, max(3, len(boundary_run) // 10)) + initial_trend = float( + np.median(boundary_gradient[:initial_count]) + ) + if not np.isclose(initial_trend, 0.0): + reversal = boundary_gradient * np.sign(initial_trend) < 0 + persistence = min(5, max(2, len(boundary_run) // 20)) + sustained = np.convolve( + reversal.astype(int), + np.ones(persistence, dtype=int), + mode="valid", + ) + turning_points = np.flatnonzero(sustained == persistence) + if turning_points.size and turning_points[0] >= 2: + primary_indices = primary_indices[ + : turning_points[0] + 1 + ] + if primary_indices.size: + fit_distances = distances[primary_indices] + fit_values = baseline_boundary[primary_indices] + + # The correlation=1 boundary saturates: it climbs steeply near the + # origin and flattens at large separation. An unweighted straight + # line over the whole lobe is therefore dominated by the flat tail, + # which drags the intercept off the measured boundary at d=0 (up to + # ~7 degrees on real data, visibly landing in the blue region) and + # biases the slope low. Weight the fit towards short distances so + # `slope` is the near-origin tangent, which is the physically + # meaningful quantity. + fit_span = float(fit_distances.max() - fit_distances.min()) + if slope_weight_scale is None: + weight_scale = SLOPE_WEIGHT_FRACTION * fit_span + else: + weight_scale = float(slope_weight_scale) + if not np.isfinite(weight_scale) or weight_scale <= 0: + # np.inf (or a non-positive scale) restores the legacy + # unweighted fit over the full lobe. + fit_weights = np.ones_like(fit_distances) + weight_scale = np.inf + else: + fit_weights = np.exp( + -(fit_distances - fit_distances.min()) / weight_scale + ) + + root_weights = np.sqrt(fit_weights) + design = ( + np.vstack([fit_distances, np.ones_like(fit_distances)]).T + * root_weights[:, None] + ) + fit_slope, fit_intercept = np.linalg.lstsq( + design, fit_values * root_weights, rcond=None + )[0] + fit_angles = fit_intercept + fit_slope * fit_distances + slope = float(fit_slope) + slope_fit_intercept = float(fit_intercept) + slope_fit_point_count = int(fit_distances.size) + # Kish effective sample size: how many points the weighting really + # uses, so a too-aggressive scale is visible rather than silent. + slope_fit_effective_count = float( + fit_weights.sum() ** 2 / np.sum(fit_weights**2) + ) + weight_mean = float( + np.average(fit_values, weights=fit_weights) + ) + residual_sum_squares = float( + np.sum(fit_weights * (fit_values - fit_angles) ** 2) + ) + total_sum_squares = float( + np.sum(fit_weights * (fit_values - weight_mean) ** 2) + ) + slope_fit_r_squared = ( + 1.0 - residual_sum_squares / total_sum_squares + if total_sum_squares > 0 + else np.nan + ) + + if fit_distances.size: + # Draw only where the weighting actually constrains the line. Past + # ~3 decay lengths the boundary has saturated away from this + # tangent, and extending the line there would misrepresent the fit. + if np.isfinite(weight_scale): + drawn = fit_distances <= ( + fit_distances.min() + 3.0 * weight_scale + ) + if drawn.sum() < 2: + drawn = np.zeros_like(fit_distances, dtype=bool) + drawn[: min(2, drawn.size)] = True + else: + drawn = np.ones_like(fit_distances, dtype=bool) + visible_fit = drawn & (fit_angles >= 0) & (fit_angles <= 180) + ax.plot( + fit_distances[visible_fit], + fit_angles[visible_fit], + color="#ffe600", + linestyle="-", + linewidth=2.5, + zorder=6, + ) + ax.legend( + handles=[ + Line2D( + [0], + [0], + marker="o", + color="none", + markerfacecolor="white", + markeredgecolor="black", + label="50% radial intercept", + ), + Line2D( + [0], + [0], + marker="D", + color="none", + markerfacecolor="white", + markeredgecolor="black", + label="50% annular intercept", + ), + Line2D( + [0], + [0], + color="#ffe600", + linewidth=2.5, + label="near-origin baseline fit", + ), + ], + loc="lower right", + fontsize=7, + framealpha=0.82, + ) + + panel_metrics = { + "pair": pair, + "title": title, + "half_probability": float(half_probability), + "radial_distance": float(radial_distance), + "annular_distance_degrees": float(annular_distance), + "slope_degrees_per_unit": float(slope), + "slope_fit_r_squared": float(slope_fit_r_squared), + "slope_fit_point_count": slope_fit_point_count, + "slope_fit_intercept_degrees": float(slope_fit_intercept), + "slope_fit_effective_point_count": float( + slope_fit_effective_count + ), + "slope_weight_scale": float(weight_scale), + "slope_contour_probability": 1.0, + "distance_units": pixel_units, + } + metrics.append(panel_metrics) + if show_metrics: + radial_text = ( + f"{radial_distance:.2f} {pixel_units}" + if np.isfinite(radial_distance) + else "not resolved" + ) + annular_text = ( + f"{annular_distance:.2f} degrees" + if np.isfinite(annular_distance) + else "not resolved" + ) + slope_text = ( + f"{slope:.2f} degrees/{pixel_units}" + if np.isfinite(slope) + else "not resolved" + ) + ax.text( + 0.98, + 0.98, + "50% radial distance = " + + radial_text + + "\n50% annular distance = " + + annular_text + + "\nslope = " + + slope_text, + transform=ax.transAxes, + ha="right", + va="top", + fontsize=8, + bbox={ + "boxstyle": "round,pad=0.35", + "facecolor": "white", + "edgecolor": "black", + "alpha": 0.82, + }, + ) + if image is not None: + fig.colorbar( + image, + ax=[ax for ax in axes.flat if ax.get_visible()], + label="probability / random", + ) + if return_metrics: + return fig, axes, metrics + return fig, axes + + def plot_interactive_image_map(self, ry=None, rx=None, intensity_map=None, vmax_cartesian=None, vmin_cartesian=None, + map_cmap='viridis', map_title='Intensity Map', dp_cmap="gray", + norm_upper_quantile=None, norm_power=1.0, + show_polar=True, vmax_polar=None, crosshair_color='r', figsize=None, + crosshair_width=2, crosshair_size=15, gaussian_filter_sigma=None): + """ + Interactive plot for browsing diffraction patterns with optional intensity map. + + Parameters + ---------- + intensity_map : array, optional + 2D array to display as reference map. Can be upsampled relative to dataset. + If None, shows mean intensity at original resolution. + Upsample factor is automatically detected from array dimensions. + vmax_cartesian : float + Maximum value for diffraction pattern display + vmin_cartesian : float + Minimum value for diffraction pattern display + map_cmap : str + Colormap for the intensity map + map_title : str + Title for the intensity map panel + dp_cmap : str + Colormap for diffraction patterns + norm_upper_quantile : float, optional + Upper quantile for normalization (0-1). If None, not used. + norm_power : float + Power law normalization exponent + show_polar : bool + Whether to show the polar transformed data panel + vmax_polar : float, optional + Maximum value for polar pattern display. If None, uses vmax_cartesian. + """ + + Ry, Rx = self.dataset_cartesian.shape[:2] + + # Check polar data availability + if show_polar and not (hasattr(self, 'polar_data') and self.polar_data is not None): + print("Warning: polar_data not found. Set show_polar=False or run polar_transform_4d first.") + show_polar = False + + intensity_map, upsample_factor = _resolve_intensity_map( + self.dataset_cartesian, + intensity_map, + (Ry, Rx), + validate=True, + announce_upsample=intensity_map is not None, + ) + + # Compute intensity map display limits + _is_rgb_map, vmin_intensity_map, vmax_intensity_map = _intensity_display_limits( + intensity_map + ) + + vmax_polar = vmax_polar or vmax_cartesian + slider_Ry, slider_Rx = Ry * upsample_factor, Rx * upsample_factor + + # ---- Create figure and axes once ---- + if show_polar: + if figsize is None: + figsize=(15, 4) + fig, (ax_map, ax_diff, ax_polar) = plt.subplots(1, 3, figsize=figsize) + else: + if figsize is None: + figsize=(12, 5) + fig, (ax_map, ax_diff) = plt.subplots(1, 2, figsize=figsize) + ax_polar = None + + # Initialize image objects + if vmin_intensity_map is None: + im_map = ax_map.imshow(intensity_map, cmap=map_cmap) + else: + im_map = ax_map.imshow(intensity_map, cmap=map_cmap, + vmin=vmin_intensity_map, vmax=vmax_intensity_map) + line_marker, = ax_map.plot([], [], color=crosshair_color, marker='+', markersize=crosshair_size, markeredgewidth=crosshair_width) + ax_map.set_title(map_title) + ax_map.set_xlabel('Rx (upsampled)' if upsample_factor > 1 else 'Rx') + ax_map.set_ylabel('Ry (upsampled)' if upsample_factor > 1 else 'Ry') + cbar_map = plt.colorbar(im_map, ax=ax_map) + + # Diffraction pattern (initialize with zeros) + im_diff = ax_diff.imshow(np.zeros((10, 10)), cmap=dp_cmap, vmin=vmin_cartesian, vmax=vmax_cartesian) + ax_diff.set_title('Diffraction Pattern') + ax_diff.set_xticks([]) + ax_diff.set_yticks([]) + cbar_diff = plt.colorbar(im_diff, ax=ax_diff) + + # Polar transform + if show_polar: + # ax_polar.set_aspect('equal', adjustable='box') + im_polar = ax_polar.imshow(np.zeros((10, 10)), cmap=dp_cmap, vmax=vmax_polar, aspect='auto') + ax_polar.set_title('Polar Transform') + ax_polar.set_xlabel('Radius (bins)') + ax_polar.set_ylabel('Theta (bins)') + cbar_polar = plt.colorbar(im_polar, ax=ax_polar) + + plt.tight_layout() + plt.close(fig) + + # ---- Interactive display callback (updates only) ---- + def show_pattern(ry_slider, rx_slider): + ry_data = ry_slider // upsample_factor + rx_data = rx_slider // upsample_factor + + # Update marker + line_marker.set_data([rx_slider], [ry_slider]) + + # Update diffraction pattern + dp_data = _normalized_dp( + self.dataset_cartesian, + ry_data, + rx_data, + norm_upper_quantile=norm_upper_quantile, + norm_power=norm_power, + copy_data=False, + ) + im_polar_data = self.polar_data['intensity'][ry_data, rx_data].T if show_polar else None + if gaussian_filter_sigma is not None: + dp_data = gaussian_filter(dp_data, gaussian_filter_sigma) + if show_polar: + im_polar_data = gaussian_filter(im_polar_data, gaussian_filter_sigma) + + im_diff.set_data(dp_data) + ax_diff.set_title(f'Diffraction Pattern (Ry={ry_data}, Rx={rx_data})') + + # Update polar transform + if show_polar: + im_polar.set_data(im_polar_data) + ax_polar.set_title(f'Polar Transform (Ry={ry_data}, Rx={rx_data})') + + clear_output(wait=True) + display(fig) + + # Create widgets + if ry is None: + ry = slider_Ry//2 + if rx is None: + rx = slider_Rx//2 + ry_slider = IntSlider(min=0, max=slider_Ry-1, value=ry, description='Ry:', continuous_update=False) + rx_slider = IntSlider(min=0, max=slider_Rx-1, value=rx, description='Rx:', continuous_update=False) + + controls = VBox([HBox([ry_slider, rx_slider])]) + interactive_plot = interactive_output(show_pattern, {'ry_slider': ry_slider, 'rx_slider': rx_slider}) + display(controls, interactive_plot) + + def save_diffraction_figures(self, ry, rx, intensity_map=None, prefix='diffraction', save_dir='.', + vmax_cartesian=None, vmin_cartesian=None, + map_cmap='viridis', map_title='Intensity Map', dp_cmap="gray", + norm_upper_quantile=None, norm_power=1.0, + show_polar=True, vmax_polar=None, crosshair_color='r', + figsize_individual=None, figsize_combined=None, crosshair_width=2, crosshair_size=15, + gaussian_filter_sigma=None): + """ + Save diffraction pattern figures for a specific scan position. + + Parameters + ---------- + ry : int + Y position in original dataset coordinates + rx : int + X position in original dataset coordinates + intensity_map : array, optional + 2D array to display as reference map. If None, shows mean intensity. + prefix : str + Filename prefix for saved files + save_dir : str + Directory path for saving files + vmax_cartesian : float + Maximum value for diffraction pattern display + vmin_cartesian : float + Minimum value for diffraction pattern display + map_cmap : str + Colormap for the intensity map + map_title : str + Title for the intensity map panel + dp_cmap : str + Colormap for diffraction patterns + norm_upper_quantile : float, optional + Upper quantile for normalization (0-1). If None, not used. + norm_power : float + Power law normalization exponent + show_polar : bool + Whether to save the polar transformed data + vmax_polar : float, optional + Maximum value for polar pattern display. If None, uses vmax_cartesian. + """ + + from pathlib import Path + + Ry, Rx = self.dataset_cartesian.shape[:2] + + # Validate coordinates + if not (0 <= ry < Ry and 0 <= rx < Rx): + raise ValueError(f"Coordinates ({ry}, {rx}) out of bounds for dataset shape ({Ry}, {Rx})") + + # Check polar data availability + if show_polar and not (hasattr(self, 'polar_data') and self.polar_data is not None): + print("Warning: polar_data not found. Skipping polar transform save.") + show_polar = False + + intensity_map, upsample_factor = _resolve_intensity_map( + self.dataset_cartesian, + intensity_map, + (Ry, Rx), + validate=True, + ) + + # Compute intensity map display limits + _is_rgb_map, vmin_intensity_map, vmax_intensity_map = _intensity_display_limits( + intensity_map + ) + + vmax_polar = vmax_polar or vmax_cartesian + + # Create save directory + save_path = Path(save_dir) + try: + save_path.mkdir(parents=True, exist_ok=True) + except Exception as e: + print(f"Error creating directory: {e}") + return + + # Calculate marker positions + marker_ry = ry * upsample_factor + marker_rx = rx * upsample_factor + + try: + # Save intensity map + if figsize_individual is None: + figsize_individual = (6, 6) + fig_map, ax = plt.subplots(figsize=figsize_individual) + if vmin_intensity_map is None: + im = ax.imshow(intensity_map, cmap=map_cmap) + else: + im = ax.imshow(intensity_map, cmap=map_cmap, + vmin=vmin_intensity_map, vmax=vmax_intensity_map) + ax.plot(marker_rx, marker_ry, color=crosshair_color, marker='+', markersize=crosshair_size, markeredgewidth=crosshair_width) + ax.set_title(map_title) + ax.set_xlabel('Rx (upsampled)' if upsample_factor > 1 else 'Rx') + ax.set_ylabel('Ry (upsampled)' if upsample_factor > 1 else 'Ry') + filename = save_path / f'{prefix}_ry{ry}_rx{rx}_intensity_map.pdf' + fig_map.savefig(filename, format='pdf', bbox_inches='tight', pad_inches=0) + plt.close(fig_map) + print(f'✓ Saved: {filename}') + + # Save diffraction pattern + fig_diff, ax = plt.subplots(figsize=figsize_individual) + dp_data = _normalized_dp( + self.dataset_cartesian, + ry, + rx, + norm_upper_quantile=norm_upper_quantile, + norm_power=norm_power, + ) + polar_im_data = self.polar_data['intensity'][ry, rx].T if show_polar else None + if gaussian_filter_sigma is not None: + dp_data = gaussian_filter(dp_data, gaussian_filter_sigma) + if show_polar: + polar_im_data = gaussian_filter(polar_im_data, gaussian_filter_sigma) + + im = ax.imshow(dp_data, cmap=dp_cmap, vmin=vmin_cartesian, vmax=vmax_cartesian) + ax.set_title(f'Diffraction Pattern (Ry={ry}, Rx={rx})') + ax.set_xticks([]) + ax.set_yticks([]) + filename = save_path / f'{prefix}_ry{ry}_rx{rx}_diffraction.pdf' + fig_diff.savefig(filename, format='pdf', bbox_inches='tight', pad_inches=0) + plt.close(fig_diff) + print(f'✓ Saved: {filename}') + + # Save polar transform + if show_polar: + fig_polar, ax = plt.subplots(figsize=figsize_individual) + im = ax.imshow(polar_im_data, cmap=dp_cmap, vmax=vmax_polar, aspect='auto') + # ax.set_aspect('equal', adjustable='box') + ax.set_title(f'Polar Transform (Ry={ry}, Rx={rx})') + ax.set_xlabel('Radius (bins)') + ax.set_ylabel('Theta (bins)') + filename = save_path / f'{prefix}_ry{ry}_rx{rx}_polar.pdf' + fig_polar.savefig(filename, format='pdf', bbox_inches='tight', pad_inches=0) + plt.close(fig_polar) + print(f'✓ Saved: {filename}') + + # Save combined figure + if show_polar: + if figsize_combined is None: + figsize_combined = (15, 4) + fig_combined, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=figsize_combined) + else: + if figsize_combined is None: + figszie_combined = (12, 5) + fig_combined, (ax1, ax2) = plt.subplots(1, 2, figsize=figsize_combined) + ax3 = None + + # Plot intensity map + if vmin_intensity_map is None: + im1 = ax1.imshow(intensity_map, cmap=map_cmap) + else: + im1 = ax1.imshow(intensity_map, cmap=map_cmap, + vmin=vmin_intensity_map, vmax=vmax_intensity_map) + ax1.plot(marker_rx, marker_ry, color=crosshair_color, marker='+', markersize=crosshair_size, markeredgewidth=crosshair_width) + ax1.set_title(map_title) + ax1.set_xlabel('Rx (upsampled)' if upsample_factor > 1 else 'Rx') + ax1.set_ylabel('Ry (upsampled)' if upsample_factor > 1 else 'Ry') + + # Plot diffraction pattern + im2 = ax2.imshow(dp_data, cmap=dp_cmap, vmin=vmin_cartesian, vmax=vmax_cartesian) + ax2.set_title(f'Diffraction Pattern (Ry={ry}, Rx={rx})') + ax2.set_xticks([]) + ax2.set_yticks([]) + + # Plot polar transform + if show_polar: + im3 = ax3.imshow(polar_im_data, + cmap=dp_cmap, vmax=vmax_polar, aspect='auto') + # ax3.set_aspect('equal', adjustable='box') + ax3.set_title(f'Polar Transform (Ry={ry}, Rx={rx})') + ax3.set_xlabel('Radius (bins)') + ax3.set_ylabel('Theta (bins)') + + plt.tight_layout() + filename = save_path / f'{prefix}_ry{ry}_rx{rx}_combined.pdf' + fig_combined.savefig(filename, format='pdf', bbox_inches='tight', pad_inches=0) + plt.close(fig_combined) + print(f'✓ Saved: {filename}') + + print(f'\nAll figures saved successfully to: {save_path}') + + except Exception as e: + print(f"Error saving figures: {e}") + + def show_widget(self, **kwargs): + """Open the interactive polymer 4D-STEM viewer (``quantem.widget``). + + Thin wrapper over ``quantem.widget.show_polymer_4DSTEM``: drag the map to update + the Current/Lamellar/Backbone/pi-pi DP panels and the polar view, with detected + peaks overlaid when ``find_peaks_model`` has run. All ``**kwargs`` are forwarded + to the factory (e.g. ``intensity_map``, ``map_cmap``, ``dp_cmap``, ``show_polar``, + ``title``). + """ + from quantem.widget import show_polymer_4DSTEM + return show_polymer_4DSTEM(self, **kwargs) + + def plot_interactive_peak_map(self, radial_range=None, intensity_map=None, + ry=None, rx=None, + vmax_cartesian=7, vmin_cartesian=0, show_all_peaks=True, + selected_peak_color='red', other_peak_color='gray', + central_beam_color='red', + norm_upper_quantile=None, norm_power=1.0, + peak_intensity_mode='size', peak_size_range=(30, 300), + peak_cmap='hot', peak_vmin=None, peak_vmax=None, + show_polar=True, vmax_polar=None, two_fold_symmetry=True, + map_cmap="viridis", dp_cmap="gray", intensity_field='intensities', + crosshair_color='r', figsize=None, crosshair_width=2, crosshair_size=15, + crosshair_width_peaks=2, crosshair_scaling_peaks=1, crosshair_scaling_central_beam=1, + gaussian_filter_sigma=None, zoom=1): + """ + Interactive plot for browsing diffraction patterns with peak overlay. + Central beam (closest to image center) plotted in blue. + """ + if figsize is None: + if show_polar: + figsize = (15, 4) + else: + figsize = (12, 5) + Ry, Rx = self.peak_coordinates_cartesian.shape + + if show_polar and not (hasattr(self, 'polar_data') and self.polar_data is not None): + print("Warning: polar_data not found. Set show_polar=False or run polar_transform_4d first.") + show_polar = False + + # Setup intensity map + if intensity_map is not None: + intensity_map, upsample_factor = _resolve_intensity_map( + self.dataset_cartesian, + intensity_map, + (Ry, Rx), + validate=False, + ) + map_title = f'Custom Map ({radial_range[0]:.2f}-{radial_range[1]:.2f} 1/Å)' if radial_range else 'Custom Map' + else: + intensity_map, upsample_factor = _resolve_intensity_map( + self.dataset_cartesian, + intensity_map, + (Ry, Rx), + validate=False, + ) + map_title = f'Peak Map ({radial_range[0]:.2f}-{radial_range[1]:.2f} 1/Å)' if radial_range else 'Peak Map' + + is_rgb_map, vmin_intensity_map, vmax_intensity_map = _intensity_display_limits( + intensity_map + ) + + vmax_polar = vmax_polar or vmax_cartesian + + # Peak plotting function + def plot_peaks_on_ax( + ax, + peaks_x, + peaks_y, + peaks_r_invA, + peak_intensities, + central_idx, + ry_data, + rx_data, + center=None, + ): + _plot_bragg_peaks_on_ax( + ax, + peaks_x, + peaks_y, + peaks_r_invA, + peak_intensities, + central_idx, + radial_range=radial_range, + show_all_peaks=show_all_peaks, + selected_peak_color=selected_peak_color, + other_peak_color=other_peak_color, + central_beam_color=central_beam_color, + peak_intensity_mode=peak_intensity_mode, + peak_size_range=peak_size_range, + peak_cmap=peak_cmap, + peak_vmin=peak_vmin, + peak_vmax=peak_vmax, + crosshair_width_peaks=crosshair_width_peaks, + crosshair_scaling_peaks=crosshair_scaling_peaks, + crosshair_scaling_central_beam=crosshair_scaling_central_beam, + add_colorbar=True, + center=center, + ) + + # Interactive callback + def show_pattern(ry_slider, rx_slider): + ry_data = ry_slider // upsample_factor + rx_data = rx_slider // upsample_factor + fig, axes = plt.subplots(1, 3 if show_polar else 2, figsize=figsize) + ax1, ax2 = axes[0], axes[1] + ax3 = axes[2] if show_polar else None + + # Intensity map + if vmin_intensity_map is None: + im1 = ax1.imshow(intensity_map, cmap=map_cmap) + else: + im1 = ax1.imshow(intensity_map, cmap=map_cmap, vmin=vmin_intensity_map, vmax=vmax_intensity_map) + ax1.scatter(rx_slider, ry_slider, facecolor='none', edgecolor=crosshair_color, marker='o', s=crosshair_size, linewidth=crosshair_width) + ax1.set_title(map_title) + ax1.set_xlabel('Rx (upsampled)' if upsample_factor > 1 else 'Rx') + ax1.set_ylabel('Ry (upsampled)' if upsample_factor > 1 else 'Ry') + if not is_rgb_map: + plt.colorbar(im1, ax=ax1) + + # Create inset axes for the zoomed view + axins = inset_axes(ax1, width="30%", height="30%", loc='upper right', + borderpad=1.5) + + # Calculate 9x9 region with selected pixel at center (4 pixels margin each side) + margin = 4 + ry_min = max(0, ry_slider - margin) + ry_max = min(intensity_map.shape[0], ry_slider + margin + 1) + rx_min = max(0, rx_slider - margin) + rx_max = min(intensity_map.shape[1], rx_slider + margin + 1) + + # Extract and display the zoomed region + zoomed_region = intensity_map[ry_min:ry_max, rx_min:rx_max] + + if vmin_intensity_map is None: + axins.imshow(zoomed_region, cmap=map_cmap, + extent=[rx_min, rx_max, ry_max, ry_min], + interpolation='nearest') + else: + axins.imshow(zoomed_region, cmap=map_cmap, + extent=[rx_min, rx_max, ry_max, ry_min], + vmin=vmin_intensity_map, vmax=vmax_intensity_map, + interpolation='nearest') + + # Draw border around the selected (central) pixel + pixel_border = Rectangle((rx_slider, ry_slider), 1, 1, + linewidth=2, edgecolor=crosshair_color, + facecolor='none', zorder=10) + axins.add_patch(pixel_border) + + # Set limits and styling + axins.set_xlim(rx_min, rx_max) + axins.set_ylim(ry_max, ry_min) + axins.set_xticks([]) + axins.set_yticks([]) + axins.set_title('9×9 zoom', fontsize=8, pad=2) + + # Optional: Add a rectangle on main plot showing zoomed region + rect = Rectangle((rx_min, ry_min), rx_max-rx_min, ry_max-ry_min, + linewidth=1.5, edgecolor=crosshair_color, + facecolor='none', linestyle='--', alpha=0.7) + ax1.add_patch(rect) + + + # Diffraction pattern + dp_data = _normalized_dp( + self.dataset_cartesian, + ry_data, + rx_data, + norm_upper_quantile=norm_upper_quantile, + norm_power=norm_power, + ) + im_polar_data = self.polar_data['intensity'][ry_data, rx_data].T if show_polar else None + if gaussian_filter_sigma is not None: + dp_data = gaussian_filter(dp_data, gaussian_filter_sigma) + if show_polar: + im_polar_data = gaussian_filter(im_polar_data, gaussian_filter_sigma) + + peaks_r_invA = _vector_field_cell(self.polar_peaks, "r_invA", ry_data, rx_data) + peaks_y = _vector_field_cell( + self.peak_coordinates_cartesian, "y_pixels", ry_data, rx_data + ) + peaks_x = _vector_field_cell( + self.peak_coordinates_cartesian, "x_pixels", ry_data, rx_data + ) + peak_ints = _vector_field_cell( + self.peak_intensities, intensity_field, ry_data, rx_data + ) + has_peak_positions = _has_peak_positions(peaks_x, peaks_y) + center = _display_center( + getattr(self, "image_centers", None), ry_data, rx_data, dp_data.shape + ) + central_idx = _central_peak_index( + peaks_x, peaks_y, peaks_r_invA, center, + max_dist=_central_beam_max_dist(dp_data.shape), + ) + ( + dp_data, + peaks_x, + peaks_y, + peaks_r_invA, + peak_ints, + central_idx, + display_center, + ) = _zoom_peak_overlay( + dp_data, + peaks_x, + peaks_y, + peaks_r_invA, + peak_ints, + central_idx, + zoom, + center, + ) + + im2 = ax2.imshow(dp_data, cmap=dp_cmap, vmax=vmax_cartesian, vmin=vmin_cartesian) + ax2.set_xticks([]) + ax2.set_yticks([]) + + plot_peaks_on_ax( + ax2, + peaks_x, + peaks_y, + peaks_r_invA, + peak_ints, + central_idx, + ry_data, + rx_data, + center=display_center, + ) + ax2.set_xlim(-0.5, dp_data.shape[1] - 0.5) + ax2.set_ylim(dp_data.shape[0] - 0.5, -0.5) + + title = f'Diffraction Pattern (Ry={ry_data}, Rx={rx_data})' + if radial_range: + title += f'\n{radial_range[0]:.2f}-{radial_range[1]:.2f} 1/Å' + if not has_peak_positions: + title += '\nNo peaks at this scan position' + ax2.set_title(title) + + # Polar transform + if show_polar: + im3 = ax3.imshow(im_polar_data, + cmap=dp_cmap, vmax=vmax_polar, aspect='auto') + # ax3.set_aspect('equal', adjustable='box') + ax3.set_xlabel('Radius (bins)') + ax3.set_ylabel('Theta (bins)') + ax3.set_title(f'Polar (Ry={ry_data}, Rx={rx_data})') + + if hasattr(self, 'polar_peaks') and self.polar_peaks is not None: + polar_r = _vector_field_cell( + self.polar_peaks, "r_invA", ry_data, rx_data + ) + polar_theta = _vector_field_cell( + self.polar_peaks, "theta", ry_data, rx_data + ) + if polar_r is not None and len(polar_r) > 0: + r_bins, theta_bins = _polar_peak_bins( + polar_r, + polar_theta, + self.max_radius_invA, + self.num_radial_bins, + self.num_annular_bins, + two_fold_symmetry, + ) + plot_peaks_on_ax(ax3, r_bins, theta_bins, polar_r, peak_ints, central_idx, ry_data, rx_data) + + plt.tight_layout() + plt.show() + + # Widgets + if ry is None: + ry = Ry*upsample_factor//2 + if rx is None: + rx = Rx*upsample_factor//2 + ry_slider = IntSlider(min=0, max=Ry*upsample_factor-1, value=ry, description='Ry:', continuous_update=False) + rx_slider = IntSlider(min=0, max=Rx*upsample_factor-1, value=rx, description='Rx:', continuous_update=False) + interactive_plot = interactive_output(show_pattern, {'ry_slider': ry_slider, 'rx_slider': rx_slider}) + display(VBox([HBox([ry_slider, rx_slider]), interactive_plot])) + + def save_peak_figures(self, ry, rx, intensity_map=None, + map_title="", prefix='peaks', save_dir='.', + vmax_cartesian=7, vmin_cartesian=0, + selected_peak_color='red', + central_beam_color='red', + norm_upper_quantile=None, norm_power=1.0, + peak_intensity_mode='size', peak_size_range=(30, 300), + peak_cmap='hot', peak_vmin=None, peak_vmax=None, + show_polar=True, vmax_polar=None, two_fold_symmetry=True, + map_cmap="viridis", dp_cmap="gray", intensity_field='intensities', + crosshair_color='r', figsize_individual=None, figsize_combined=None, + crosshair_width=2, crosshair_size=15, crosshair_width_peaks=2, + crosshair_scaling_peaks=1, crosshair_scaling_central_beam=1, peak_marker="o", + peak_marker_facecolors='none', peak_marker_size=None, gaussian_filter_sigma=None, + zoom=1, peak_alpha=1.0, central_linewidth=None, + peaks_x=None, peaks_y=None, peak_ints=None, peaks_r_invA=None, + central_idx=None, show_central_beam=True, + save_intensity_map=True, save_diffraction=True, save_polar=None, + dpi=400): + """ + Save peak-annotated diffraction figures for a specific scan position. + Central beam (closest to image center) plotted in blue. + + Peaks are read from the precomputed ``peak_coordinates_cartesian`` / + ``peak_intensities`` / ``polar_peaks`` by default. Pass ``peaks_x`` / ``peaks_y`` + / ``peak_ints`` (and optionally ``peaks_r_invA`` / ``central_idx``) to inject + peaks directly instead — e.g. from live single-DP inference, where no scan-wide + peak arrays exist. ``save_intensity_map`` / ``save_diffraction`` / ``save_polar`` + select which figures to write (``save_polar=None`` follows ``show_polar``); this + lets a caller save the context map once and the DP per panel. + """ + + override_peaks = peaks_x is not None + if self.peak_coordinates_cartesian is not None: + Ry, Rx = self.peak_coordinates_cartesian.shape + else: + Ry, Rx = int(self.dataset_cartesian.shape[0]), int(self.dataset_cartesian.shape[1]) + + if not (0 <= ry < Ry and 0 <= rx < Rx): + raise ValueError(f"Coordinates ({ry}, {rx}) out of bounds") + + if save_polar is not None: + show_polar = bool(save_polar) + if show_polar and not (hasattr(self, 'polar_data') and self.polar_data is not None): + print("Warning: polar_data not found. Skipping polar save.") + show_polar = False + + intensity_map, upsample_factor = _resolve_intensity_map( + self.dataset_cartesian, + intensity_map, + (Ry, Rx), + validate=False, + ) + + _is_rgb_map, vmin_intensity_map, vmax_intensity_map = _intensity_display_limits( + intensity_map + ) + + vmax_polar = vmax_polar or vmax_cartesian + + # Peak plotting function + def plot_peaks_on_ax(ax, peaks_x, peaks_y, peaks_r_invA, peak_intensities, central_idx, center=None): + _plot_bragg_peaks_on_ax( + ax, + peaks_x, + peaks_y, + peaks_r_invA, + peak_intensities, + central_idx, + selected_peak_color=selected_peak_color, + central_beam_color=central_beam_color, + peak_intensity_mode=peak_intensity_mode, + peak_size_range=peak_size_range, + peak_cmap=peak_cmap, + peak_vmin=peak_vmin, + peak_vmax=peak_vmax, + crosshair_width_peaks=crosshair_width_peaks, + crosshair_scaling_peaks=crosshair_scaling_peaks, + crosshair_scaling_central_beam=crosshair_scaling_central_beam, + peak_marker=peak_marker, + peak_marker_facecolors=peak_marker_facecolors, + peak_marker_size=peak_marker_size, + peak_alpha=peak_alpha, + central_alpha=peak_alpha, + central_linewidth=( + crosshair_width_peaks if central_linewidth is None else central_linewidth + ), + center=center, + show_central_beam=show_central_beam, + ) + + # Create save directory + save_path = Path(save_dir) + save_path.mkdir(parents=True, exist_ok=True) + + # Get peaks data once (injected overrides win; otherwise read precomputed). + if override_peaks: + peaks_x = np.asarray(peaks_x) + peaks_y = np.asarray(peaks_y) + peak_ints = None if peak_ints is None else np.asarray(peak_ints) + peaks_r_invA = None if peaks_r_invA is None else np.asarray(peaks_r_invA) + else: + peaks_y = _vector_field_cell(self.peak_coordinates_cartesian, "y_pixels", ry, rx) + peaks_x = _vector_field_cell(self.peak_coordinates_cartesian, "x_pixels", ry, rx) + peak_ints = _vector_field_cell(self.peak_intensities, intensity_field, ry, rx) + peaks_r_invA = ( + _vector_field_cell(self.polar_peaks, "r_invA", ry, rx) + if getattr(self, 'polar_peaks', None) is not None + else None + ) + dp_data = _normalized_dp( + self.dataset_cartesian, + ry, + rx, + norm_upper_quantile=norm_upper_quantile, + norm_power=norm_power, + ) + polar_im_data = self.polar_data['intensity'][ry, rx].T if show_polar else None + if gaussian_filter_sigma is not None: + dp_data = gaussian_filter(dp_data, gaussian_filter_sigma) + if show_polar: + polar_im_data = gaussian_filter(polar_im_data, gaussian_filter_sigma) + + center = _display_center(getattr(self, "image_centers", None), ry, rx, dp_data.shape) + if central_idx is None: + central_idx = _central_peak_index( + peaks_x, peaks_y, peaks_r_invA, center, + max_dist=_central_beam_max_dist(dp_data.shape), + ) + + ( + dp_data, + peaks_x, + peaks_y, + peaks_r_invA, + peak_ints, + central_idx, + display_center, + ) = _zoom_peak_overlay( + dp_data, + peaks_x, + peaks_y, + peaks_r_invA, + peak_ints, + central_idx, + zoom, + center, + ) + + # Save intensity map + if figsize_individual is None: + figsize_individual = (6, 6) + if save_intensity_map: + fig_map, ax = plt.subplots(figsize=figsize_individual) + if vmin_intensity_map is None: + im = ax.imshow(intensity_map, cmap=map_cmap) + else: + im = ax.imshow(intensity_map, cmap=map_cmap, vmin=vmin_intensity_map, vmax=vmax_intensity_map) + + ry_slider = ry * upsample_factor + rx_slider = rx * upsample_factor + + ax.scatter(rx_slider, ry_slider, facecolor='none', edgecolor=crosshair_color, marker='o', s=crosshair_size, linewidth=crosshair_width) + + # Add inset + from mpl_toolkits.axes_grid1.inset_locator import inset_axes + axins = inset_axes(ax, width="30%", height="30%", loc='upper right', borderpad=1.5) + + margin = 4 + ry_min = max(0, ry_slider - margin) + ry_max = min(intensity_map.shape[0], ry_slider + margin + 1) + rx_min = max(0, rx_slider - margin) + rx_max = min(intensity_map.shape[1], rx_slider + margin + 1) + + zoomed_region = intensity_map[ry_min:ry_max, rx_min:rx_max] + + if vmin_intensity_map is None: + axins.imshow(zoomed_region, cmap=map_cmap, extent=[rx_min, rx_max, ry_max, ry_min], interpolation='nearest') + else: + axins.imshow(zoomed_region, cmap=map_cmap, extent=[rx_min, rx_max, ry_max, ry_min], + vmin=vmin_intensity_map, vmax=vmax_intensity_map, interpolation='nearest') + + pixel_border = Rectangle((rx_slider, ry_slider), 1, 1, linewidth=2, edgecolor=crosshair_color, + facecolor='none', zorder=10) + axins.add_patch(pixel_border) + + axins.set_xlim(rx_min, rx_max) + axins.set_ylim(ry_max, ry_min) + axins.set_xticks([]) + axins.set_yticks([]) + axins.set_title('9×9 zoom', fontsize=8, pad=2) + + rect = Rectangle((rx_min, ry_min), rx_max-rx_min, ry_max-ry_min, + linewidth=1.5, edgecolor=crosshair_color, facecolor='none', linestyle='--', alpha=0.7) + ax.add_patch(rect) + + ax.set_title(map_title) + ax.set_xlabel('Rx (upsampled)' if upsample_factor > 1 else 'Rx') + ax.set_ylabel('Ry (upsampled)' if upsample_factor > 1 else 'Ry') + fig_map.savefig(save_path / f'{prefix}_ry{ry}_rx{rx}_intensity_map.pdf', format='pdf', bbox_inches='tight', pad_inches=0, dpi=dpi) + plt.close(fig_map) + print(f'✓ Saved: {prefix}_ry{ry}_rx{rx}_intensity_map.pdf') + + # Save diffraction pattern with peaks + if save_diffraction: + fig_diff, ax = plt.subplots(figsize=figsize_individual) + im = ax.imshow(dp_data, cmap=dp_cmap, vmax=vmax_cartesian, vmin=vmin_cartesian) + ax.set_xticks([]) + ax.set_yticks([]) + if peaks_x is not None: + plot_peaks_on_ax(ax, peaks_x, peaks_y, peaks_r_invA, peak_ints, central_idx, center=display_center) + ax.set_xlim(-0.5, dp_data.shape[1] - 0.5) + ax.set_ylim(dp_data.shape[0] - 0.5, -0.5) + ax.set_title(f'Diffraction Pattern (Ry={ry}, Rx={rx})') + fig_diff.savefig(save_path / f'{prefix}_ry{ry}_rx{rx}_diffraction.pdf', format='pdf', bbox_inches='tight', pad_inches=0, dpi=dpi) + plt.close(fig_diff) + print(f'✓ Saved: {prefix}_ry{ry}_rx{rx}_diffraction.pdf') + + # Save polar transform with peaks + if show_polar: + fig_polar, ax = plt.subplots(figsize=figsize_individual) + im = ax.imshow(polar_im_data, cmap=dp_cmap, vmax=vmax_polar, aspect='auto') + ax.set_title(f'Polar (Ry={ry}, Rx={rx})') + ax.set_xlabel('Radius (bins)') + ax.set_ylabel('Theta (bins)') + + if hasattr(self, 'polar_peaks') and self.polar_peaks is not None: + polar_r = _vector_field_cell(self.polar_peaks, "r_invA", ry, rx) + polar_theta = _vector_field_cell(self.polar_peaks, "theta", ry, rx) + if polar_r is not None and len(polar_r) > 0: + r_bins, theta_bins = _polar_peak_bins( + polar_r, + polar_theta, + self.max_radius_invA, + self.num_radial_bins, + self.num_annular_bins, + two_fold_symmetry, + ) + + # Find central beam for polar + polar_central_idx = np.argmin(polar_r) + # Use the full (unzoomed) intensities: polar_r / r_bins / theta_bins are read + # from the full polar_peaks, whereas `peak_ints` may have been subset by + # zoom > 1 for the Cartesian panel (length mismatch -> IndexError otherwise). + polar_peak_ints = _vector_field_cell( + self.peak_intensities, intensity_field, ry, rx + ) + if polar_r is not None and len(polar_r) > 0: + plot_peaks_on_ax(ax, r_bins, theta_bins, polar_r, polar_peak_ints, polar_central_idx) + fig_polar.savefig(save_path / f'{prefix}_ry{ry}_rx{rx}_polar.pdf', format='pdf', bbox_inches='tight', pad_inches=0, dpi=dpi) + plt.close(fig_polar) + print(f'✓ Saved: {prefix}_ry{ry}_rx{rx}_polar.pdf') + + def save_peak_animation( + self, + path, + *, + region=None, + step=1, + bidirectional=True, + fps=10, + intensity_map=None, + map_title="", + map_cmap="viridis", + crosshair_color="r", + crosshair_size=80, + crosshair_width=2, + dp_cmap="gray", + vmin_cartesian=0, + vmax_cartesian=7, + norm_upper_quantile=None, + norm_power=1.0, + gaussian_filter_sigma=None, + zoom=1, + show_peaks=True, + selected_peak_color="red", + central_beam_color="red", + show_central_beam=True, + peak_intensity_mode="size", + peak_size_range=(30, 300), + peak_marker_size=None, + crosshair_width_peaks=2, + crosshair_scaling_central_beam=1, + peak_alpha=1.0, + central_linewidth=None, + intensity_field="intensities", + live_inference=False, + infer_device=None, + sigma_peak_blur=1.0, + threshold_peak=0.5, + panels=None, + figsize=None, + dpi=100, + progress=True, + ): + """Render a snaking-cursor animation to an animated GIF. + + Walks a boustrophedon (snake) path over the scan and, for each position, + renders one combined frame: the real-space intensity map with a cursor + crosshair at the current position (left) beside one or more diffraction-pattern + panels with detected Bragg peaks overlaid (right). Frames are assembled into a + looping GIF. This reuses the same rendering primitives as + :meth:`save_peak_figures` so frames match the per-position saved figures. + + Parameters + ---------- + path : str | pathlib.Path + Output ``.gif`` path. + region : tuple[int, int, int, int] | None + ``(ry0, ry1, rx0, rx1)`` half-open scan bounds to snake over; ``None`` + covers the whole scan. + step : int + Stride between visited positions (>= 1). + bidirectional : bool + Snake/boustrophedon path (alternate row direction). ``False`` scans every + row left->right. + fps : float + Playback frames per second. + intensity_map : np.ndarray | None + Real-space map to display (computed once). ``None`` uses the mean-intensity + virtual image. May be scalar ``(H, W)`` or RGB ``(H, W, 3|4)``. + panels : list[dict] | None + One dict per diffraction-pattern panel to draw beside the map, each holding + that panel's display settings (any of: ``title``, ``dp_cmap``, + ``vmin_cartesian``, ``vmax_cartesian``, ``norm_upper_quantile``, + ``norm_power``, ``gaussian_filter_sigma``, ``zoom``, ``selected_peak_color``, + ``central_beam_color``, ``show_central_beam``, ``peak_intensity_mode``, + ``peak_size_range``, ``peak_marker_size``, ``crosshair_width_peaks``, + ``crosshair_scaling_central_beam``, ``peak_alpha``, ``central_linewidth``). + Missing keys fall back to the corresponding top-level argument. ``None`` + (default) draws a single panel from the top-level arguments. + live_inference : bool + Run the model per position via :meth:`infer_peaks_single` instead of reading + precomputed ``peak_coordinates_cartesian`` (slow over large regions). + figsize : tuple | None + Figure size. ``None`` auto-sizes to ``(5 * (1 + n_panels), 5)``. + + Returns + ------- + pathlib.Path + The written GIF path. + """ + from PIL import Image + + # A single top-level panel spec unless the caller passes an explicit list. + if panels is None: + panels = [dict( + title=None, + dp_cmap=dp_cmap, + vmin_cartesian=vmin_cartesian, + vmax_cartesian=vmax_cartesian, + norm_upper_quantile=norm_upper_quantile, + norm_power=norm_power, + gaussian_filter_sigma=gaussian_filter_sigma, + zoom=zoom, + selected_peak_color=selected_peak_color, + central_beam_color=central_beam_color, + show_central_beam=show_central_beam, + peak_intensity_mode=peak_intensity_mode, + peak_size_range=peak_size_range, + peak_marker_size=peak_marker_size, + crosshair_width_peaks=crosshair_width_peaks, + crosshair_scaling_central_beam=crosshair_scaling_central_beam, + peak_alpha=peak_alpha, + central_linewidth=central_linewidth, + )] + n_panels = len(panels) + if n_panels == 0: + raise ValueError("panels must contain at least one DP panel spec") + + Ry, Rx = int(self.dataset_cartesian.shape[0]), int(self.dataset_cartesian.shape[1]) + base_shape = (int(self.dataset_cartesian.shape[2]), int(self.dataset_cartesian.shape[3])) + + # Resolve the real-space map ONCE; _mean_intensity_map rescans every DP, so + # rebuilding it per frame would be quadratic in scan size. + intensity_map, upsample_factor = _resolve_intensity_map( + self.dataset_cartesian, intensity_map, (Ry, Rx), validate=False, + ) + is_rgb_map, map_vmin, map_vmax = _intensity_display_limits(intensity_map) + + # Boustrophedon path over the requested region (mirrors Show4DSTEM.raster). + if region is None: + ry0, ry1, rx0, rx1 = 0, Ry, 0, Rx + else: + ry0, ry1, rx0, rx1 = region + ry0, ry1 = max(0, int(ry0)), min(Ry, int(ry1)) + rx0, rx1 = max(0, int(rx0)), min(Rx, int(rx1)) + if ry1 <= ry0 or rx1 <= rx0: + raise ValueError(f"Empty region {region!r} for scan shape ({Ry}, {Rx})") + step = max(1, int(step)) + points = [] + for i, ry in enumerate(range(ry0, ry1, step)): + cols = list(range(rx0, rx1, step)) + if bidirectional and i % 2 == 1: + cols = cols[::-1] + points.extend((ry, rx) for rx in cols) + + has_precomputed = (not live_inference) and self.peak_coordinates_cartesian is not None + has_polar_peaks = getattr(self, "polar_peaks", None) is not None + + if figsize is None: + figsize = (5 * (1 + n_panels), 5) + fig, axes = plt.subplots(1, 1 + n_panels, figsize=figsize, dpi=dpi) + ax_map = axes[0] + dp_axes = axes[1:] + frames = [] + try: + for ry, rx in tqdm(points, desc="Rendering snake", disable=not progress): + # Peaks + beam center are fetched ONCE per position; each panel then + # applies its own normalization / zoom crop below. + peaks_x = peaks_y = peak_ints = peaks_r_invA = None + if show_peaks: + if live_inference: + res = self.infer_peaks_single( + ry, rx, device=infer_device, + sigma_peak_blur=sigma_peak_blur, threshold_peak=threshold_peak, + ) + peaks_x, peaks_y, peak_ints = ( + res["x_pixels"], res["y_pixels"], res["intensities"], + ) + elif has_precomputed: + peaks_y = _vector_field_cell(self.peak_coordinates_cartesian, "y_pixels", ry, rx) + peaks_x = _vector_field_cell(self.peak_coordinates_cartesian, "x_pixels", ry, rx) + if self.peak_intensities is not None: + peak_ints = _vector_field_cell(self.peak_intensities, intensity_field, ry, rx) + if has_polar_peaks: + peaks_r_invA = _vector_field_cell(self.polar_peaks, "r_invA", ry, rx) + + center = _display_center(getattr(self, "image_centers", None), ry, rx, base_shape) + # _plot_bragg_peaks_on_ax draws no rings when peaks_r_invA is None. When + # there is no polar transform, fall back to the pixel radius from center so + # the rings still render (r_invA is otherwise only used for radial filtering, + # which this call does not use). + if peaks_r_invA is None and _has_peak_positions(peaks_x, peaks_y): + peaks_r_invA = np.sqrt( + (np.asarray(peaks_x) - center[1]) ** 2 + + (np.asarray(peaks_y) - center[0]) ** 2 + ) + central_idx = _central_peak_index( + peaks_x, peaks_y, peaks_r_invA, center, + max_dist=_central_beam_max_dist(base_shape), + ) + + ax_map.clear() + if is_rgb_map: + ax_map.imshow(intensity_map) + elif map_vmin is None: + ax_map.imshow(intensity_map, cmap=map_cmap) + else: + ax_map.imshow(intensity_map, cmap=map_cmap, vmin=map_vmin, vmax=map_vmax) + ax_map.scatter( + rx * upsample_factor, ry * upsample_factor, + facecolor="none", edgecolor=crosshair_color, marker="o", + s=crosshair_size, linewidth=crosshair_width, zorder=10, + ) + ax_map.set_title(f"{map_title} Ry={ry}, Rx={rx}" if map_title else f"Ry={ry}, Rx={rx}") + ax_map.set_xticks([]) + ax_map.set_yticks([]) + + for ax, spec in zip(dp_axes, panels): + npow = spec.get("norm_power", norm_power) + npow = 1.0 if npow is None else npow + sigma = spec.get("gaussian_filter_sigma", gaussian_filter_sigma) + dp_p = _normalized_dp( + self.dataset_cartesian, ry, rx, + norm_upper_quantile=spec.get("norm_upper_quantile", norm_upper_quantile), + norm_power=npow, + ) + if sigma is not None: + dp_p = gaussian_filter(dp_p, sigma) + ( + dp_p, px, py, r_invA, pint, cidx, disp_center, + ) = _zoom_peak_overlay( + dp_p, peaks_x, peaks_y, peaks_r_invA, peak_ints, + central_idx, spec.get("zoom", zoom), center, + ) + ax.clear() + ax.imshow( + dp_p, cmap=spec.get("dp_cmap", dp_cmap), + vmin=spec.get("vmin_cartesian", vmin_cartesian), + vmax=spec.get("vmax_cartesian", vmax_cartesian), + ) + if show_peaks and px is not None: + cwp = spec.get("crosshair_width_peaks", crosshair_width_peaks) + clw = spec.get("central_linewidth", central_linewidth) + if ("marker_size" in spec) or ("central_size" in spec): + # Data-proportional circles (radius in detector px) so markers + # cover the same fraction of the pattern as the widget canvas + # (which scales its px marker radii to the display). + _draw_peaks_data_circles( + ax, px, py, pint, cidx, disp_center, + marker_scaled=spec.get("marker_scaled", True), + marker_size=spec.get("marker_size", 8.0), + marker_size_min=spec.get("marker_size_min", 4.0), + marker_size_max=spec.get("marker_size_max", 16.0), + selected_peak_color=spec.get("selected_peak_color", selected_peak_color), + central_beam_color=spec.get("central_beam_color", central_beam_color), + show_central_beam=spec.get("show_central_beam", show_central_beam), + central_size=spec.get("central_size", 5.0), + peak_linewidth=cwp, + central_linewidth=(cwp if clw is None else clw), + ) + else: + palpha = spec.get("peak_alpha", peak_alpha) + _plot_bragg_peaks_on_ax( + ax, px, py, r_invA, pint, cidx, + selected_peak_color=spec.get("selected_peak_color", selected_peak_color), + central_beam_color=spec.get("central_beam_color", central_beam_color), + peak_intensity_mode=spec.get("peak_intensity_mode", peak_intensity_mode), + peak_size_range=spec.get("peak_size_range", peak_size_range), + peak_marker_size=spec.get("peak_marker_size", peak_marker_size), + crosshair_width_peaks=cwp, + crosshair_scaling_central_beam=spec.get( + "crosshair_scaling_central_beam", crosshair_scaling_central_beam + ), + peak_alpha=palpha, + central_alpha=palpha, + central_linewidth=(cwp if clw is None else clw), + center=disp_center, + show_central_beam=spec.get("show_central_beam", show_central_beam), + ) + ax.set_xlim(-0.5, dp_p.shape[1] - 0.5) + ax.set_ylim(dp_p.shape[0] - 0.5, -0.5) + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_title(spec.get("title") or "") + + fig.canvas.draw() + rgba = np.asarray(fig.canvas.buffer_rgba()) + frames.append(Image.fromarray(rgba[..., :3].copy())) + finally: + plt.close(fig) + + if not frames: + raise ValueError("Snake path is empty; check region / step.") + + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + duration_ms = max(10, int(round(1000.0 / max(0.1, fps)))) + frames[0].save( + str(path), save_all=True, append_images=frames[1:], + duration=duration_ms, loop=0, optimize=True, disposal=2, + ) + if progress: + print(f"✓ Saved {len(frames)}-frame animation: {path.resolve()}") + return path + + def edit_scan_mask( + self, + *, + initial_x=None, + initial_y=None, + initial_radius=None, + initial_geometry="circle", + initial_size_x=None, + initial_size_y=None, + reference_image=None, + state_path=None, + overlay_alpha=0.28, + crosshair_width=2, + crosshair_size=12, + autosave=False, + display_widget=True, + ): + """Create an interactive scan-mask editor. + + X is the horizontal scan-column coordinate and Y is the vertical + scan-row coordinate. Circle, ellipse, square, and rectangle geometries + are available. Sizes are radii for round geometries and half-sizes for + rectangular geometries. A saved ``state_path`` is loaded automatically. + """ + return ScanMaskEditor( + self, + initial_x=initial_x, + initial_y=initial_y, + initial_radius=initial_radius, + initial_geometry=initial_geometry, + initial_size_x=initial_size_x, + initial_size_y=initial_size_y, + reference_image=reference_image, + state_path=state_path, + overlay_alpha=overlay_alpha, + crosshair_width=crosshair_width, + crosshair_size=crosshair_size, + autosave=autosave, + display_widget=display_widget, + ) + + def create_interactive_circular_mask( + self, + initial_x0=None, + initial_y0=None, + initial_r=None, + reference_image=None, + overlay_alpha=0.3, + crosshair_width=2, + crosshair_size=15, + state_path=None, + autosave=False, + display_widget=True, + ): + """Compatibility wrapper for :meth:`edit_scan_mask`. + + Historically ``initial_x0`` represented the array row and + ``initial_y0`` represented the array column. New code should use + ``edit_scan_mask(initial_x=column, initial_y=row, ...)``. + """ + return self.edit_scan_mask( + initial_x=initial_y0, + initial_y=initial_x0, + initial_radius=initial_r, + reference_image=reference_image, + state_path=state_path, + overlay_alpha=overlay_alpha, + crosshair_width=crosshair_width, + crosshair_size=crosshair_size, + autosave=autosave, + display_widget=display_widget, + ) + + def plot_peak_histogram_map( + self, + intensity_threshold=None, + intensity_percentile=None, + figsize=(8, 6), + cmap='viridis', + return_values=False, + intensity_field='intensities', + ): + """ + Plot 2D map showing the number of peaks found at each scan position. + + Parameters: + ----------- + intensity_threshold : float, optional + Absolute intensity threshold. Only count peaks above this value. + intensity_percentile : float, optional + Percentile threshold (0-100). Overrides intensity_threshold. + figsize : tuple + Figure size (width, height) + cmap : str + Colormap to use + return_values : bool + If True, return figure, axes, and count_map + + Returns: + -------- + fig, ax, count_map : (optional) matplotlib figure, axes, and count array + """ + Ry, Rx = self.peak_coordinates_cartesian.shape + + # Convert percentile to threshold if needed + if intensity_percentile is not None: + all_intensities = [ + values + for i in range(Ry) + for j in range(Rx) + if len( + values := _vector_field_cell( + self.peak_intensities, intensity_field, i, j + ) + ) + ] + if all_intensities: + intensity_threshold = np.percentile(np.concatenate(all_intensities), intensity_percentile) + + # Build count map + count_map = np.zeros((Ry, Rx)) + for i in range(Ry): + for j in range(Rx): + peaks = self.peak_coordinates_cartesian[i, j].array + if len(peaks) == 0: + continue + + if intensity_threshold is None: + count_map[i, j] = len(peaks) + else: + intensities = _vector_field_cell( + self.peak_intensities, intensity_field, i, j + ) + if len(intensities): + count_map[i, j] = np.sum(intensities >= intensity_threshold) + + # Plot + fig, ax = plt.subplots(figsize=figsize) + im = ax.imshow(count_map, cmap=cmap, origin='lower') + + cbar = plt.colorbar(im, ax=ax) + cbar.set_label('Number of Peaks', fontsize=12) + + # Integer colorbar ticks + max_count = int(np.max(count_map)) + if max_count > 0: + ticks = np.arange(0, max_count + 1, max(1, max_count // 5)) + cbar.set_ticks(ticks) + + # Title + title = 'Peak Count per Scan Position' + if intensity_threshold is not None: + title += f'\n(intensity ≥ {intensity_threshold:.3f})' + ax.set_title(title, fontsize=14) + ax.set_xlabel('Scan X', fontsize=12) + ax.set_ylabel('Scan Y', fontsize=12) + + plt.tight_layout() + plt.show() + + if return_values: + return fig, ax, count_map + + def plot_peak_count_map(self, q_ranges, figsize_per_map=(5, 4), cmap='viridis', return_values=False): + """ + Plot 2D maps showing the number of peaks in specified q-ranges. + + Parameters: + ----------- + q_ranges : list of tuples or single tuple + Either a single (q_min, q_max) tuple or a list of tuples for multiple ranges. + Example: (2.8, 3.2) or [(0.3, 0.7), (2.8, 3.2), (5.0, 5.4)] + figsize_per_map : tuple + Size of each subplot (width, height) + cmap : str + Colormap to use + + Returns: + -------- + fig, axes : matplotlib figure and axes + count_maps : list of ndarrays + The count maps for each q-range + """ + # Handle single range or list of ranges + if isinstance(q_ranges, tuple): + q_ranges = [q_ranges] + + Ry, Rx = self.peak_coordinates_cartesian.shape + n_ranges = len(q_ranges) + + # Create figure + n_cols = min(3, n_ranges) # Max 3 columns + n_rows = int(np.ceil(n_ranges / n_cols)) + fig, axes = plt.subplots(n_rows, n_cols, + figsize=(figsize_per_map[0]*n_cols, figsize_per_map[1]*n_rows)) + + # Handle single subplot case + if n_ranges == 1: + axes = np.array([axes]) + axes = axes.flatten() + + count_maps = [] + + for idx, (q_min, q_max) in enumerate(q_ranges): + # Create count map + count_map = np.zeros((Ry, Rx)) + + for i in range(Ry): + for j in range(Rx): + peaks_r_invA = _vector_field_cell(self.polar_peaks, "r_invA", i, j) + if peaks_r_invA is not None and len(peaks_r_invA) > 0: + # Get radial distances in 1/Å + distances = peaks_r_invA + # Count peaks in range + mask = (distances >= q_min) & (distances < q_max) + count_map[i, j] = np.sum(mask) + + count_maps.append(count_map) + + # Calculate max_count early for use in both colorbar and statistics + max_count = int(np.max(count_map)) + + # Plot as a true integer-count map. Avoid show_2d's default quantile + # normalization here: count maps are discrete, not continuous images. + boundaries = np.arange(-0.5, max_count + 1.5, 1) + norm = BoundaryNorm(boundaries, ncolors=plt.get_cmap(cmap).N, clip=True) + im = axes[idx].imshow( + count_map, + cmap=cmap, + norm=norm, + interpolation='nearest', + origin='upper', + ) + _dlo = (1.0 / q_max) if q_max > 0 else float('inf') # d-spacing (Å) = 1 / q (1/Å) + _dhi = (1.0 / q_min) if q_min > 0 else float('inf') + axes[idx].set_title( + f'Peak Count\n{q_min:.2f} - {q_max:.2f} 1/Å\n' + f'd = {_dlo:.2f} - {_dhi:.2f} Å', + fontsize=14) + axes[idx].set_xlabel('Scan X', fontsize=12) + axes[idx].set_ylabel('Scan Y', fontsize=12) + axes[idx].set_xticks([]) + axes[idx].set_yticks([]) + cbar = plt.colorbar(im, ax=axes[idx], ticks=np.arange(max_count + 1)) + cbar.set_label('Number of Peaks', fontsize=10) + + # Print statistics + total_peaks = int(np.sum(count_map)) + positions_with_peaks = np.sum(count_map > 0) + print(f"Range {q_min:.2f}-{q_max:.2f} 1/Å:") + print(f" Total peaks: {total_peaks}") + print(f" Positions with peaks: {positions_with_peaks}/{Ry*Rx}") + print(f" Max peaks at one position: {max_count}") + print(f" Mean peaks per position: {np.mean(count_map):.2f}") + print() + + # Hide unused subplots + for idx in range(n_ranges, len(axes)): + axes[idx].set_visible(False) + + plt.tight_layout() + plt.show() + + if return_values: + return fig, axes, count_maps + + def make_flowline_map( + self, + orient_hist, + thresh_seed=0.2, + thresh_grow=0.05, + thresh_collision=0.001, + sep_seeds=None, + sep_xy=6.0, + sep_theta=5.0, + sort_seeds="intensity", + linewidth=2.0, + step_size=0.5, + min_steps=4, + max_steps=1000, + sigma_x=1.0, + sigma_y=1.0, + sigma_theta=2.0, + progress_bar: bool = True, + ): + """ + Create an 3D or 4D orientation flowline map - essentially a pixelated "stream map" which represents diffraction data. + + Args: + orient_hist (array): Histogram of all orientations with coordinates + [radial_bin x_probe y_probe theta] + We assume theta bin ranges from 0 to 180 degrees and is periodic. + thresh_seed (float): Threshold for seed generation in histogram. + thresh_grow (float): Threshold for flowline growth in histogram. + thresh_collision (float): Threshold for termination of flowline growth in histogram. + sep_seeds (float): Initial seed separation in bins - set to None to use default value, + which is equal to 0.5*sep_xy. + sep_xy (float): Search radius for flowline direction in x and y. + sep_theta = (float): Search radius for flowline direction in theta. + sort_seeds (str): How to sort the initial seeds for growth: + None - no sorting + 'intensity' - sort by histogram intensity + 'random' - random order + linewidth (float): Thickness of the flowlines in pixels. + step_size (float): Step size for flowline growth in pixels. + min_steps (int): Minimum number of steps for a flowline to be drawn. + max_steps (int): Maximum number of steps for a flowline to be drawn. + sigma_x (float): Weighted sigma in x direction for direction update. + sigma_y (float): Weighted sigma in y direction for direction update. + sigma_theta (float): Weighted sigma in theta for direction update. + progress_bar (bool): Enable progress bar + + Returns: + orient_flowlines (array): 4D array containing flowlines + [radial_bin x_probe y_probe theta] + """ + + # Ensure sep_xy and sep_theta are arrays + sep_xy = np.atleast_1d(sep_xy) + sep_theta = np.atleast_1d(sep_theta) + + # number of radial bins + num_radii = orient_hist.shape[0] + if num_radii > 1 and len(sep_xy) == 1: + sep_xy = np.ones(num_radii) * sep_xy + if num_radii > 1 and len(sep_theta) == 1: + sep_theta = np.ones(num_radii) * sep_theta + + # Default seed separation + if sep_seeds is None: + sep_seeds = np.round(np.min(sep_xy) / 2 + 0.5).astype("int") + else: + sep_seeds = np.atleast_1d(sep_seeds).astype("int") + if num_radii > 1 and len(sep_seeds) == 1: + sep_seeds = (np.ones(num_radii) * sep_seeds).astype("int") + + # coordinates + theta = np.linspace(0, np.pi, orient_hist.shape[3], endpoint=False) + dtheta = theta[1] - theta[0] + size_3D = np.array( + [ + orient_hist.shape[1], + orient_hist.shape[2], + orient_hist.shape[3], + ] + ) + + # initialize weighting array + vx = np.arange(-np.ceil(2 * sigma_x), np.ceil(2 * sigma_x) + 1) + vy = np.arange(-np.ceil(2 * sigma_y), np.ceil(2 * sigma_y) + 1) + vt = np.arange(-np.ceil(2 * sigma_theta), np.ceil(2 * sigma_theta) + 1) + ay, ax, at = np.meshgrid(vy, vx, vt) + k = ( + np.exp(ax**2 / (-2 * sigma_x**2)) + * np.exp(ay**2 / (-2 * sigma_y**2)) + * np.exp(at**2 / (-2 * sigma_theta**2)) + ) + k = k / np.sum(k) + vx = vx[:, None, None].astype("int") + vy = vy[None, :, None].astype("int") + vt = vt[None, None, :].astype("int") + + # initalize flowline array + orient_flowlines = np.zeros_like(orient_hist) + + # initialize output + xy_t_int = np.zeros((max_steps + 1, 4)) + xy_t_int_rev = np.zeros((max_steps + 1, 4)) + + # Loop over radial bins + for a0 in range(num_radii): + # initialize collision check array + cr = np.arange(-np.ceil(sep_xy[a0]), np.ceil(sep_xy[a0]) + 1) + ct = np.arange(-np.ceil(sep_theta[a0]), np.ceil(sep_theta[a0]) + 1) + ay, ax, at = np.meshgrid(cr, cr, ct) + c_mask = ( + (ax**2 + ay**2) / sep_xy[a0] ** 2 + at**2 / sep_theta[a0] ** 2 + <= (1 + 1 / sep_xy[a0]) ** 2 + )[None, :, :, :] + cx = cr[None, :, None, None].astype("int") + cy = cr[None, None, :, None].astype("int") + ct = ct[None, None, None, :].astype("int") + + # Find all seed locations + orient = orient_hist[a0, :, :, :] + sub_seeds = np.logical_and( + np.logical_and( + orient >= np.roll(orient, 1, axis=2), + orient >= np.roll(orient, -1, axis=2), + ), + orient >= thresh_seed, + ) + + # Separate seeds + if sep_seeds > 0: + for a1 in range(sep_seeds - 1): + sub_seeds[a1::sep_seeds, :, :] = False + sub_seeds[:, a1::sep_seeds, :] = False + + # Index seeds + x_inds, y_inds, t_inds = np.where(sub_seeds) + if sort_seeds is not None: + if sort_seeds == "intensity": + inds_sort = np.argsort(orient[sub_seeds])[::-1] + elif sort_seeds == "random": + inds_sort = np.random.permutation(np.count_nonzero(sub_seeds)) + x_inds = x_inds[inds_sort] + y_inds = y_inds[inds_sort] + t_inds = t_inds[inds_sort] + + # for a1 in tqdmnd(range(0,40), desc="Drawing flowlines",unit=" seeds", disable=not progress_bar): + t = "Drawing flowlines " + str(a0) + for a1 in tqdmnd( + range(0, x_inds.shape[0]), desc=t, unit=" seeds", disable=not progress_bar + ): + # initial coordinate and intensity + xy0 = np.array((x_inds[a1], y_inds[a1])) + t0 = theta[t_inds[a1]] + + # init theta + inds_theta = np.mod( + np.round(t0 / dtheta).astype("int") + vt, orient.shape[2] + ) + orient_crop = ( + k + * orient[ + np.clip( + np.round(xy0[0]).astype("int") + vx, 0, orient.shape[0] - 1 + ), + np.clip( + np.round(xy0[1]).astype("int") + vy, 0, orient.shape[1] - 1 + ), + inds_theta, + ] + ) + theta_crop = theta[inds_theta] + t0 = np.sum(orient_crop * theta_crop) / np.sum(orient_crop) + + # forward direction + t = t0 + v0 = np.array((np.cos(t), -np.sin(t))) + v = v0 * step_size + xy = xy0 + int_val = self.get_intensity(orient, xy0[0], xy0[1], t0 / dtheta) + xy_t_int[0, 0:2] = xy0 + xy_t_int[0, 2] = t / dtheta + xy_t_int[0, 3] = int_val + # main loop + grow = True + count = 0 + while grow is True: + count += 1 + + # update position and intensity + xy = xy + v + int_val = self.get_intensity(orient, xy[0], xy[1], t / dtheta) + + # check for collision + flow_crop = orient_flowlines[ + a0, + np.clip(np.round(xy[0]).astype("int") + cx, 0, orient.shape[0] - 1), + np.clip(np.round(xy[1]).astype("int") + cy, 0, orient.shape[1] - 1), + np.mod(np.round(t / dtheta).astype("int") + ct, orient.shape[2]), + ] + int_flow = np.max(flow_crop[c_mask]) + + if ( + xy[0] < 0 + or xy[1] < 0 + or xy[0] > orient.shape[0] + or xy[1] > orient.shape[1] + or int_val < thresh_grow + or int_flow > thresh_collision + ): + grow = False + else: + # update direction + inds_theta = np.mod( + np.round(t / dtheta).astype("int") + vt, orient.shape[2] + ) + orient_crop = ( + k + * orient[ + np.clip( + np.round(xy[0]).astype("int") + vx, + 0, + orient.shape[0] - 1, + ), + np.clip( + np.round(xy[1]).astype("int") + vy, + 0, + orient.shape[1] - 1, + ), + inds_theta, + ] + ) + theta_crop = theta[inds_theta] + t = np.sum(orient_crop * theta_crop) / np.sum(orient_crop) + # v = np.array((np.cos(t), np.sin(t))) * step_size + # v = np.array((np.sin(t), np.cos(t))) * step_size + # v = np.array((-np.sin(t), np.cos(t))) * step_size + + xy_t_int[count, 0:2] = xy + xy_t_int[count, 2] = t / dtheta + xy_t_int[count, 3] = int_val + + if count > max_steps - 1: + grow = False + + # reverse direction + t = t0 + np.pi + v0 = np.array((np.cos(t), -np.sin(t))) + v = v0 * step_size + xy = xy0 + int_val = self.get_intensity(orient, xy0[0], xy0[1], t0 / dtheta) + xy_t_int_rev[0, 0:2] = xy0 + xy_t_int_rev[0, 2] = t / dtheta + xy_t_int_rev[0, 3] = int_val + # main loop + grow = True + count_rev = 0 + while grow is True: + count_rev += 1 + + # update position and intensity + xy = xy + v + int_val = self.get_intensity(orient, xy[0], xy[1], t / dtheta) + + # check for collision + flow_crop = orient_flowlines[ + a0, + np.clip(np.round(xy[0]).astype("int") + cx, 0, orient.shape[0] - 1), + np.clip(np.round(xy[1]).astype("int") + cy, 0, orient.shape[1] - 1), + np.mod(np.round(t / dtheta).astype("int") + ct, orient.shape[2]), + ] + int_flow = np.max(flow_crop[c_mask]) + + if ( + xy[0] < 0 + or xy[1] < 0 + or xy[0] > orient.shape[0] + or xy[1] > orient.shape[1] + or int_val < thresh_grow + or int_flow > thresh_collision + ): + grow = False + else: + # update direction + inds_theta = np.mod( + np.round(t / dtheta).astype("int") + vt, orient.shape[2] + ) + orient_crop = ( + k + * orient[ + np.clip( + np.round(xy[0]).astype("int") + vx, + 0, + orient.shape[0] - 1, + ), + np.clip( + np.round(xy[1]).astype("int") + vy, + 0, + orient.shape[1] - 1, + ), + inds_theta, + ] + ) + theta_crop = theta[inds_theta] + t = np.sum(orient_crop * theta_crop) / np.sum(orient_crop) + np.pi + v = np.array((np.cos(t), -np.sin(t))) * step_size + v = np.array((np.cos(t), -np.sin(t))) * step_size + # v = np.array((-np.sin(t), np.cos(t))) * step_size + + xy_t_int_rev[count_rev, 0:2] = xy + xy_t_int_rev[count_rev, 2] = t / dtheta + xy_t_int_rev[count_rev, 3] = int_val + + if count_rev > max_steps - 1: + grow = False + + # write into output array + if count + count_rev > min_steps: + if count > 0: + orient_flowlines[a0, :, :, :] = self.set_intensity( + orient_flowlines[a0, :, :, :], xy_t_int[1:count, :] + ) + if count_rev > 1: + orient_flowlines[a0, :, :, :] = self.set_intensity( + orient_flowlines[a0, :, :, :], xy_t_int_rev[1:count_rev, :] + ) + + # normalize to step size + orient_flowlines = orient_flowlines * step_size + + # linewidth + if linewidth > 1.0: + s = linewidth - 1.0 + + orient_flowlines = gaussian_filter1d(orient_flowlines, s, axis=1, truncate=3.0) + orient_flowlines = gaussian_filter1d(orient_flowlines, s, axis=2, truncate=3.0) + orient_flowlines = orient_flowlines * (s**2) + + return orient_flowlines + + + def make_flowline_rainbow_image( + self, + orient_flowlines, + int_range=[0, 0.2], + sym_rotation_order=2, + theta_offset=np.pi, + greyscale=False, + greyscale_max=True, + white_background=False, + power_scaling=1.0, + sum_radial_bins=False, + plot_images=True, + figsize=None, + ): + """ + Generate RGB output images from the flowline arrays. + + Args: + orient_flowline (array): Histogram of all orientations with coordinates [x y radial_bin theta] + We assume theta bin ranges from 0 to 180 degrees and is periodic. + int_range (float) 2 element array giving the intensity range + sym_rotation_order (int): rotational symmety for colouring + theta_offset (float): Offset the anglular coloring by this value in radians. + Default pi rotates the hue mapping by 90deg (nematic + sym=2) so color tracks the drawn flowline direction: + cyan for vertical lines, red for horizontal. + greyscale (bool): Set to False for color output, True for greyscale output. + greyscale_max (bool): If output is greyscale, use max instead of mean for overlapping flowlines. + white_background (bool): For either color or greyscale output, switch to white background (from black). + power_scaling (float): Power law scaling for flowline intensity output. + sum_radial_bins (bool): Sum all radial bins (alternative is to output separate images). + plot_images (bool): Plot the outputs for quick visualization. + figsize (2-tuple): Size of output figure. + + Returns: + im_flowline (array): 3D or 4D array containing flowline images + """ + + # init array + size_input = orient_flowlines.shape + size_output = np.array([size_input[0], size_input[1], size_input[2], 3]) + im_flowline = np.zeros(size_output) + theta_offset = np.atleast_1d(theta_offset) + + if greyscale is True: + for a0 in range(size_input[0]): + if greyscale_max is True: + im = np.max(orient_flowlines[a0, :, :, :], axis=2) + else: + im = np.mean(orient_flowlines[a0, :, :, :], axis=2) + + sig = np.clip((im - int_range[0]) / (int_range[1] - int_range[0]), 0, 1) + + if power_scaling != 1: + sig = sig**power_scaling + + if white_background is False: + im_flowline[a0, :, :, :] = sig[:, :, None] + else: + im_flowline[a0, :, :, :] = 1 - sig[:, :, None] + + else: + # Color basis + c0 = np.array([1.0, 0.0, 0.0]) + c1 = np.array([0.0, 0.7, 0.0]) + c2 = np.array([0.0, 0.3, 1.0]) + + # angles + theta = np.linspace(0, np.pi, size_input[3], endpoint=False) + # Negate so the hue handedness matches the drawn flowline direction in the + # displayed (y-down) map: red horizontal, cyan vertical, "/" yellow, "\" purple. + theta_color = -theta * sym_rotation_order + + if size_input[0] > 1 and len(theta_offset) == 1: + theta_offset = np.ones(size_input[0]) * theta_offset + + for a0 in range(size_input[0]): + # color projections + b0 = np.maximum( + 1 + - np.abs( + np.mod(theta_offset[a0] + theta_color + np.pi, 2 * np.pi) - np.pi + ) + ** 2 + / (np.pi * 2 / 3) ** 2, + 0, + ) + b1 = np.maximum( + 1 + - np.abs( + np.mod( + theta_offset[a0] + theta_color - np.pi * 2 / 3 + np.pi, + 2 * np.pi, + ) + - np.pi + ) + ** 2 + / (np.pi * 2 / 3) ** 2, + 0, + ) + b2 = np.maximum( + 1 + - np.abs( + np.mod( + theta_offset[a0] + theta_color - np.pi * 4 / 3 + np.pi, + 2 * np.pi, + ) + - np.pi + ) + ** 2 + / (np.pi * 2 / 3) ** 2, + 0, + ) + + sig = np.clip( + (orient_flowlines[a0, :, :, :] - int_range[0]) + / (int_range[1] - int_range[0]), + 0, + 1, + ) + if power_scaling != 1: + sig = sig**power_scaling + + im_flowline[a0, :, :, :] = ( + np.sum(sig * b0[None, None, :], axis=2)[:, :, None] * c0[None, None, :] + + np.sum(sig * b1[None, None, :], axis=2)[:, :, None] + * c1[None, None, :] + + np.sum(sig * b2[None, None, :], axis=2)[:, :, None] + * c2[None, None, :] + ) + + # clip limits + im_flowline[a0, :, :, :] = np.clip(im_flowline[a0, :, :, :], 0, 1) + + # contrast flip + if white_background is True: + im = rgb_to_hsv(im_flowline[a0]) + im_v = im[:, :, 2] + im[:, :, 1] = im_v + im[:, :, 2] = 1 + im_flowline[a0] = hsv_to_rgb(im) + + if sum_radial_bins is True: + if white_background is False: + im_flowline = np.clip(np.sum(im_flowline, axis=0), 0, 1)[None, :, :, :] + else: + # im_flowline = np.clip(np.sum(im_flowline,axis=0)+1-im_flowline.shape[0],0,1)[None,:,:,:] + im_flowline = np.min(im_flowline, axis=0)[None, :, :, :] + + if plot_images is True: + if figsize is None: + fig, ax = plt.subplots( + im_flowline.shape[0], 1, figsize=(10, im_flowline.shape[0] * 10) + ) + else: + fig, ax = plt.subplots(im_flowline.shape[0], 1, figsize=figsize) + + if im_flowline.shape[0] > 1: + for a0 in range(im_flowline.shape[0]): + ax[a0].imshow(im_flowline[a0]) + # ax[a0].axis('off') + plt.subplots_adjust(wspace=0, hspace=0.02) + else: + ax.imshow(im_flowline[0]) + # ax.axis('off') + plt.show() + + return im_flowline + + + def make_flowline_rainbow_legend( + self, + im_size=np.array([256, 256]), + sym_rotation_order=2, + theta_offset_degrees=0.0, + white_background=False, + return_image=False, + radial_range=np.array([0.45, 0.9]), + plot_legend=True, + figsize=(4, 4), + ): + """ + This function generates a legend for a the rainbow colored flowline maps, and returns it as an RGB image. + + Parameters + ---------- + im_size (np.array): + Size of legend image in pixels. + sym_rotation_order (int): + rotational symmety for colouring + theta_offset_degrees (float): + Offset the anglular coloring by this value in degrees. + Rotation is Q with respect to R, in the positive (counter clockwise) direction. + white_background (bool): + For either color or greyscale output, switch to white background (from black). + return_image (bool): + Return the image array. + radial_range (np.array): + Inner and outer radius for the legend ring. + plot_legend (bool): + Plot the generated legend. + figsize (tuple or list): + Size of the plotted legend. + + Returns + ---------- + + im_legend (array): + Image array for the legend. + """ + + # Coordinates + x = np.linspace(-1, 1, im_size[0]) + y = np.linspace(-1, 1, im_size[1]) + ya, xa = np.meshgrid(y, x) + # TODO: Can replace with squared term? ra2? Faster + # ra = np.sqrt(xa**2 + ya**2) + ra2 = xa**2 + ya**2 + ta = np.arctan2(ya, xa) + np.deg2rad(theta_offset_degrees) + ta_sym = ta * sym_rotation_order + + # mask + mask = np.logical_and(ra2 > radial_range[0]**2, ra2 < radial_range[1]**2) + # mask = np.logical_and(ra > radial_range[0], ra < radial_range[1]) + + # rgb image + z = mask * np.exp(1j * ta_sym) + # hue_offset = 0 + amp = np.abs(z) + vmin = np.min(amp) + vmax = np.max(amp) + ph = np.angle(z) # + hue_offset + h = np.mod(ph / (2 * np.pi), 1) + s = 0.85 * np.ones_like(h) + v = (amp - vmin) / (vmax - vmin) + im_legend = hsv_to_rgb(np.dstack((h, s, v))) + + if white_background is True: + im_legend[im_legend.sum(2) == 0] = 1 + + # plotting + if plot_legend: + fig, ax = plt.subplots(1, 1, figsize=figsize) + ax.imshow(im_legend) + ax.invert_yaxis() + # ax.set_axis_off() + ax.axis("off") + + if return_image: + return im_legend + + + def make_flowline_combined_image( + self, + orient_flowlines, + int_range=[0, 0.2], + cvals=np.array( + [ + [0.0, 0.7, 0.0], + [1.0, 0.0, 0.0], + [0.0, 0.7, 1.0], + ] + ), + white_background=False, + power_scaling=1.0, + sum_radial_bins=True, + plot_images=True, + figsize=None, + ): + """ + Generate RGB output images from the flowline arrays. + + Args: + orient_flowline (array): Histogram of all orientations with coordinates [x y radial_bin theta] + We assume theta bin ranges from 0 to 180 degrees and is periodic. + int_range (float) 2 element array giving the intensity range + cvals (array): Nx3 size array containing RGB colors for different radial ibns. + white_background (bool): For either color or greyscale output, switch to white background (from black). + power_scaling (float): Power law scaling for flowline intensities. + sum_radial_bins (bool): Sum outputs over radial bins. + plot_images (bool): Plot the output images for quick visualization. + figsize (2-tuple): Size of output figure. + + Returns: + im_flowline (array): flowline images + """ + + # init array + size_input = orient_flowlines.shape + size_output = np.array([size_input[0], size_input[1], size_input[2], 3]) + im_flowline = np.zeros(size_output) + cvals = np.array(cvals) + + # Generate all color images + for a0 in range(size_input[0]): + sig = np.clip( + (np.sum(orient_flowlines[a0, :, :, :], axis=2) - int_range[0]) + / (int_range[1] - int_range[0]), + 0, + 1, + ) + if power_scaling != 1: + sig = sig**power_scaling + + if white_background: + im_flowline[a0, :, :, :] = 1 - sig[:, :, None] * ( + 1 - cvals[a0, :][None, None, :] + ) + else: + im_flowline[a0, :, :, :] = sig[:, :, None] * cvals[a0, :][None, None, :] + + # # contrast flip + # if white_background is True: + # im = rgb_to_hsv(im_flowline[a0,:,:,:]) + # # im_s = im[:,:,1] + # im_v = im[:,:,2] + # v_range = [np.min(im_v), np.max(im_v)] + # print(v_range) + + # im[:,:,1] = im_v + # im[:,:,2] = 1 + # im_flowline[a0,:,:,:] = hsv_to_rgb(im) + + if sum_radial_bins is True: + if white_background is False: + im_flowline = np.clip(np.sum(im_flowline, axis=0), 0, 1)[None, :, :, :] + else: + # im_flowline = np.clip(np.sum(im_flowline,axis=0)+1-im_flowline.shape[0],0,1)[None,:,:,:] + im_flowline = np.min(im_flowline, axis=0)[None, :, :, :] + + if plot_images is True: + if figsize is None: + fig, ax = plt.subplots( + im_flowline.shape[0], 1, figsize=(10, im_flowline.shape[0] * 10) + ) + else: + fig, ax = plt.subplots(im_flowline.shape[0], 1, figsize=figsize) + + if im_flowline.shape[0] > 1: + for a0 in range(im_flowline.shape[0]): + ax[a0].imshow(im_flowline[a0]) + ax[a0].axis("off") + plt.subplots_adjust(wspace=0, hspace=0.02) + else: + ax.imshow(im_flowline[0]) + ax.axis("off") + plt.show() + + return im_flowline + + def get_intensity( + self, + orient, + x, + y, + t + ): + # utility function to get histogram intensites + + x = np.clip(x, 0, orient.shape[0] - 2) + y = np.clip(y, 0, orient.shape[1] - 2) + + xF = np.floor(x).astype("int") + yF = np.floor(y).astype("int") + tF = np.floor(t).astype("int") + dx = x - xF + dy = y - yF + dt = t - tF + t1 = np.mod(tF, orient.shape[2]) + t2 = np.mod(tF + 1, orient.shape[2]) + + int_vals = ( + orient[xF, yF, t1] * ((1 - dx) * (1 - dy) * (1 - dt)) + + orient[xF, yF, t2] * ((1 - dx) * (1 - dy) * (dt)) + + orient[xF, yF + 1, t1] * ((1 - dx) * (dy) * (1 - dt)) + + orient[xF, yF + 1, t2] * ((1 - dx) * (dy) * (dt)) + + orient[xF + 1, yF, t1] * ((dx) * (1 - dy) * (1 - dt)) + + orient[xF + 1, yF, t2] * ((dx) * (1 - dy) * (dt)) + + orient[xF + 1, yF + 1, t1] * ((dx) * (dy) * (1 - dt)) + + orient[xF + 1, yF + 1, t2] * ((dx) * (dy) * (dt)) + ) + + return int_vals + + + def set_intensity( + self, + orient, + xy_t_int + ): + # utility function to set flowline intensites + + xF = np.floor(xy_t_int[:, 0]).astype("int") + yF = np.floor(xy_t_int[:, 1]).astype("int") + tF = np.floor(xy_t_int[:, 2]).astype("int") + dx = xy_t_int[:, 0] - xF + dy = xy_t_int[:, 1] - yF + dt = xy_t_int[:, 2] - tF + + inds_1D = np.ravel_multi_index( + [xF, yF, tF], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (1 - dx) * ( + 1 - dy + ) * (1 - dt) + inds_1D = np.ravel_multi_index( + [xF, yF, tF + 1], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (1 - dx) * ( + 1 - dy + ) * (dt) + inds_1D = np.ravel_multi_index( + [xF, yF + 1, tF], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (1 - dx) * ( + dy + ) * (1 - dt) + inds_1D = np.ravel_multi_index( + [xF, yF + 1, tF + 1], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (1 - dx) * ( + dy + ) * (dt) + inds_1D = np.ravel_multi_index( + [xF + 1, yF, tF], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (dx) * ( + 1 - dy + ) * (1 - dt) + inds_1D = np.ravel_multi_index( + [xF + 1, yF, tF + 1], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (dx) * ( + 1 - dy + ) * (dt) + inds_1D = np.ravel_multi_index( + [xF + 1, yF + 1, tF], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (dx) * (dy) * ( + 1 - dt + ) + inds_1D = np.ravel_multi_index( + [xF + 1, yF + 1, tF + 1], orient.shape[0:3], mode=["clip", "clip", "wrap"] + ) + orient.ravel()[inds_1D] = orient.ravel()[inds_1D] + xy_t_int[:, 3] * (dx) * (dy) * ( + dt + ) + + return orient + + def interactive_probe_selector(self, probe_map=None, figsize=(14, 8), cmap='viridis'): + """ + Interactive GUI to select probe positions and view diffraction patterns. + + Parameters + ---------- + probe_map : ndarray, optional + 2D array to display as the probe position map. If None, uses mean diffraction intensity. + figsize : tuple + Figure size (width, height) + cmap : str + Colormap for the probe map + + Returns + ------- + selected_positions : list of tuples + List of (ry, rx) coordinates of selected positions + """ + from matplotlib.widgets import Button + from matplotlib.patches import Circle + import ipywidgets as widgets + from IPython.display import display, clear_output + + Ry, Rx = self.dataset_cartesian.shape[:2] + + # Create default probe map if not provided + if probe_map is None: + probe_map = np.mean(self.dataset_cartesian.array, axis=(2, 3)) + + # Storage for selected positions + selected_positions = [] + markers = [] + + # Create figure + fig = plt.figure(figsize=figsize) + gs = fig.add_gridspec(2, 3, width_ratios=[2, 2, 1], height_ratios=[1, 1], + hspace=0.3, wspace=0.3) + + # Probe map axis + ax_probe = fig.add_subplot(gs[:, 0]) + im_probe = ax_probe.imshow(probe_map, cmap=cmap, origin='lower', + interpolation='nearest', aspect='auto') + ax_probe.set_title('Probe Position Map\n(Click to add point)', fontsize=12) + ax_probe.set_xlabel('Rx') + ax_probe.set_ylabel('Ry') + plt.colorbar(im_probe, ax=ax_probe, label='Intensity') + + # Diffraction pattern axes + ax_dp1 = fig.add_subplot(gs[0, 1]) + ax_dp2 = fig.add_subplot(gs[1, 1]) + ax_dp1.set_title('Diffraction Pattern 1') + ax_dp2.set_title('Diffraction Pattern 2') + ax_dp1.axis('off') + ax_dp2.axis('off') + + # Text area for position list + ax_list = fig.add_subplot(gs[:, 2]) + ax_list.axis('off') + ax_list.set_title('Selected Positions', fontsize=11, fontweight='bold') + + # Add clear all button + ax_button = plt.axes([0.7, 0.02, 0.1, 0.04]) + btn_clear = Button(ax_button, 'Clear All') + + def update_display(): + """Update the position list and diffraction patterns.""" + # Clear position list + ax_list.clear() + ax_list.axis('off') + ax_list.set_title('Selected Positions', fontsize=11, fontweight='bold') + + # Display positions + y_pos = 0.95 + for idx, (ry, rx) in enumerate(selected_positions): + text = f"{idx+1}. ({ry}, {rx})" + ax_list.text(0.1, y_pos, text, fontsize=10, transform=ax_list.transAxes, + verticalalignment='top') + y_pos -= 0.08 + + # Update diffraction patterns + if len(selected_positions) >= 1: + ry, rx = selected_positions[-1] + dp = self.dataset_cartesian[ry, rx].array + ax_dp1.clear() + ax_dp1.imshow(dp, cmap='gray') + ax_dp1.set_title(f'DP at ({ry}, {rx})') + ax_dp1.axis('off') + + if len(selected_positions) >= 2: + ry, rx = selected_positions[-2] + dp = self.dataset_cartesian[ry, rx].array + ax_dp2.clear() + ax_dp2.imshow(dp, cmap='gray') + ax_dp2.set_title(f'DP at ({ry}, {rx})') + ax_dp2.axis('off') + + fig.canvas.draw_idle() + + def onclick(event): + """Handle click events on probe map.""" + if event.inaxes == ax_probe and event.button == 1: # Left click + rx = int(np.round(event.xdata)) + ry = int(np.round(event.ydata)) + + # Check bounds + if 0 <= ry < Ry and 0 <= rx < Rx: + selected_positions.append((ry, rx)) + + # Add marker + marker = Circle((rx, ry), radius=0.5, color='red', + fill=True, zorder=10) + ax_probe.add_patch(marker) + markers.append(marker) + + # Add label + label = ax_probe.text(rx, ry, str(len(selected_positions)), + color='white', fontsize=8, ha='center', + va='center', fontweight='bold', zorder=11) + markers.append(label) + + update_display() + + def clear_all(event): + """Clear all selected positions.""" + selected_positions.clear() + for marker in markers: + marker.remove() + markers.clear() + ax_dp1.clear() + ax_dp1.axis('off') + ax_dp2.clear() + ax_dp2.axis('off') + update_display() + + # Connect events + fig.canvas.mpl_connect('button_press_event', onclick) + btn_clear.on_clicked(clear_all) + + plt.show() + + return selected_positions + + + def visualize_selected_patterns(self, positions, ncols=4, figsize_per_pattern=(3, 3), + cmap='gray', vmax=None): + """ + Display diffraction patterns at selected probe positions in a grid. + + Parameters + ---------- + positions : list of tuples + List of (ry, rx) coordinates + ncols : int + Number of columns in the grid + figsize_per_pattern : tuple + Size of each subplot (width, height) + cmap : str + Colormap for diffraction patterns + vmax : float, optional + Maximum value for colormap normalization + + Returns + ------- + fig, axes : matplotlib figure and axes + """ + n_positions = len(positions) + nrows = int(np.ceil(n_positions / ncols)) + + fig, axes = plt.subplots(nrows, ncols, + figsize=(figsize_per_pattern[0]*ncols, + figsize_per_pattern[1]*nrows)) + + # Handle single subplot case + if n_positions == 1: + axes = np.array([axes]) + axes = axes.flatten() + + for idx, (ry, rx) in enumerate(positions): + dp = self.dataset_cartesian[ry, rx].array + + im = axes[idx].imshow(dp, cmap=cmap, vmax=vmax) + axes[idx].set_title(f'({ry}, {rx})', fontsize=10) + axes[idx].axis('off') + plt.colorbar(im, ax=axes[idx], fraction=0.046, pad=0.04) + + # Hide unused subplots + for idx in range(n_positions, len(axes)): + axes[idx].set_visible(False) + + plt.tight_layout() + plt.show() + + return fig, axes + + + def interactive_probe_selector_widget(self, probe_map=None, cmap='viridis'): + """ + Enhanced interactive GUI using ipywidgets for fine-tuning positions. + + Parameters + ---------- + probe_map : ndarray, optional + 2D array to display as the probe position map + cmap : str + Colormap for the probe map + + Returns + ------- + selected_positions : list of tuples + List of (ry, rx) coordinates of selected positions + """ + import ipywidgets as widgets + from IPython.display import display, clear_output + + Ry, Rx = self.dataset_cartesian.shape[:2] + + # Create default probe map if not provided + if probe_map is None: + probe_map = np.mean(self.dataset_cartesian.array, axis=(2, 3)) + + # Storage + selected_positions = [] + + # Create output widgets + output_plot = widgets.Output() + output_list = widgets.Output() + + def update_plot(): + """Update the main plot with markers.""" + with output_plot: + clear_output(wait=True) + fig, ax = plt.subplots(figsize=(8, 6)) + + im = ax.imshow(probe_map, cmap=cmap, origin='lower', + interpolation='nearest', aspect='auto') + ax.set_title('Probe Position Map (Click to add point)', fontsize=12) + ax.set_xlabel('Rx') + ax.set_ylabel('Ry') + plt.colorbar(im, ax=ax, label='Intensity') + + # Add markers + for idx, (ry, rx) in enumerate(selected_positions): + circle = Circle((rx, ry), radius=0.5, color='red', + fill=True, zorder=10) + ax.add_patch(circle) + ax.text(rx, ry, str(idx+1), color='white', fontsize=8, + ha='center', va='center', fontweight='bold', zorder=11) + + def onclick(event): + if event.inaxes == ax and event.button == 1: + rx = int(np.round(event.xdata)) + ry = int(np.round(event.ydata)) + if 0 <= ry < Ry and 0 <= rx < Rx: + selected_positions.append((ry, rx)) + update_plot() + update_list() + + fig.canvas.mpl_connect('button_press_event', onclick) + plt.show() + + def update_list(): + """Update the position list with controls.""" + with output_list: + clear_output(wait=True) + + if not selected_positions: + print("No positions selected") + return + + for idx, (ry, rx) in enumerate(selected_positions): + print(f"--- Position {idx+1} ---") + + # Create sliders for fine-tuning + ry_slider = widgets.IntSlider( + value=ry, min=0, max=Ry-1, step=1, + description=f'Ry {idx+1}:', continuous_update=False + ) + rx_slider = widgets.IntSlider( + value=rx, min=0, max=Rx-1, step=1, + description=f'Rx {idx+1}:', continuous_update=False + ) + + def make_update(i): + def update_position(change): + selected_positions[i] = (ry_slider.value, rx_slider.value) + update_plot() + return update_position + + ry_slider.observe(make_update(idx), names='value') + rx_slider.observe(make_update(idx), names='value') + + # Delete button + delete_btn = widgets.Button(description=f'Delete {idx+1}', + button_style='danger') + + def make_delete(i): + def delete_position(b): + del selected_positions[i] + update_plot() + update_list() + return delete_position + + delete_btn.on_click(make_delete(idx)) + + display(widgets.HBox([ry_slider, rx_slider, delete_btn])) + + # Clear all button + clear_btn = widgets.Button(description='Clear All', button_style='warning') + def clear_all(b): + selected_positions.clear() + update_plot() + update_list() + clear_btn.on_click(clear_all) + + display(clear_btn) + + # Layout + ui = widgets.VBox([ + widgets.HBox([output_plot, output_list]) + ]) + + display(ui) + update_plot() + update_list() + + return selected_positions diff --git a/src/quantem/diffraction/grain_clustering.py b/src/quantem/diffraction/grain_clustering.py new file mode 100644 index 000000000..cef24cbfa --- /dev/null +++ b/src/quantem/diffraction/grain_clustering.py @@ -0,0 +1,1080 @@ +"""Signal-level grain clustering for 4D-STEM polymer orientation data. + +This module assigns individual diffraction *signals* (Bragg arcs) to *grains* by a +global, order-independent procedure, as an alternative to greedy seed-and-grow +region growing. + +Design (see the full proposal for rationale): + +* The clustered unit is a **signal**, not a probe/pixel: a single detected peak at a + probe position, carrying ``(pos, theta, r, intensity, window)``. Because one signal + carries exactly one label, "each signal belongs to at most one grain" holds by + construction, and several grains may coexist at one probe (their signals get + different labels). + +* A **window** (radial range) is a hard, immutable *signal class* (backbone, lamellar, + pi-pi, ...). Clustering happens strictly within a window; ``r`` never drifts across + a window boundary. + +* Within a window, both **orientation** (circular, 180-deg period for 2-fold polymer + texture) and **radius** are "smooth within a grain, a discontinuity splits grains". + Two signals at the same orientation but distinct *quantized* radii are different + grains. A merge across a probe boundary is allowed only if *both* the orientation + jump and the (relative) radius jump are below tolerance. + +* Grains are spatially coherent because adjacency exists **only between neighbouring + probes** -- never in pure feature space. Hence identical orientation in + *disconnected* regions stays separate, and a missing detection can be bridged by a + larger ``neighbor_dist``. + +Core algorithm (Stage A): build a region-adjacency graph over signals (nodes = signals, +edges = signal pairs at neighbouring probes), then agglomeratively merge the adjacent +region pair with the smallest *boundary discontinuity* (mean orientation/radius jump +across the shared boundary), stopping when no boundary is within tolerance. This is +average-linkage on the spatial graph with a discontinuity stop: order-independent +(driven by the global minimum cost, not by traversal), chaining-resistant (a boundary +statistic, not a single lucky edge), tolerant of gentle orientation gradients (bent +grains), and it separates quantized-radius grains. + +Only numpy + scipy + the standard library are required. +""" + +from __future__ import annotations + +import heapq +from dataclasses import dataclass, field +from typing import Optional, Sequence + +import numpy as np +from scipy.spatial import cKDTree + +__all__ = [ + "SignalTable", + "GrainInfo", + "GrainResult", + "extract_signals", + "cluster_signals_into_grains", + "refine_grains_crf", + "circular_distance_deg", + "orientation_to_rgb", + "grain_rgb_overlay", + "orientation_legend_image", + "plot_grain_map", +] + + +# -------------------------------------------------------------------------------------- +# data structures +# -------------------------------------------------------------------------------------- +@dataclass +class SignalTable: + """A flat table of detected signals across the probe grid. + + Attributes + ---------- + pos : (N, 2) int array + Probe position ``(rx, ry)`` of each signal. + theta : (N,) float array + Orientation angle in degrees, folded to ``[0, 180)``. + r : (N,) float array + Scattering-vector magnitude ``|q|`` of each signal. + intensity : (N,) float array + Peak intensity. + window : (N,) int array + Radial-window / signal-class id (immutable class label). + map_shape : (Rx, Ry) + Probe-grid shape. + """ + + pos: np.ndarray + theta: np.ndarray + r: np.ndarray + intensity: np.ndarray + window: np.ndarray + map_shape: tuple + + def __post_init__(self): + self.pos = np.asarray(self.pos, dtype=np.int64).reshape(-1, 2) + self.theta = np.asarray(self.theta, dtype=np.float64).reshape(-1) + self.r = np.asarray(self.r, dtype=np.float64).reshape(-1) + self.intensity = np.asarray(self.intensity, dtype=np.float64).reshape(-1) + self.window = np.asarray(self.window, dtype=np.int64).reshape(-1) + n = self.pos.shape[0] + if not (len(self.theta) == len(self.r) == len(self.intensity) == len(self.window) == n): + raise ValueError("SignalTable field lengths are inconsistent") + + def __len__(self) -> int: + return self.pos.shape[0] + + +@dataclass +class GrainInfo: + label: int + window: int + signal_ids: np.ndarray + n_signals: int + theta_mean: float + theta_std: float + r_mean: float + intensity_median: float + centroid: tuple + + +@dataclass +class GrainResult: + """Result of :func:`cluster_signals_into_grains`. + + ``labels`` (length N, -1 = outlier) is the authoritative output. ``label_map`` is a + convenience raster ``(num_windows, Rx, Ry)`` for visualisation; where a probe holds + several same-window signals in different grains it keeps the highest-intensity one. + """ + + labels: np.ndarray + n_grains: int + grains: list + label_map: np.ndarray + params: dict = field(default_factory=dict) + confidence: Optional[np.ndarray] = None # (N,) max posterior, filled by Stage B + margin: Optional[np.ndarray] = None # (N,) energy gap top1-top2, filled by Stage B + + +# -------------------------------------------------------------------------------------- +# orientation geometry +# -------------------------------------------------------------------------------------- +def circular_distance_deg(a, b, period: float = 180.0): + """Circular distance between angles (degrees), default 180-deg period (2-fold).""" + d = np.abs(np.asarray(a, float) - np.asarray(b, float)) % period + return np.minimum(d, period - d) + + +# -------------------------------------------------------------------------------------- +# ingestion adapter (matches make_orientation_histogram conventions) +# -------------------------------------------------------------------------------------- +def extract_signals( + bragg_peaks, + radial_ranges, + *, + r_field: Optional[str] = None, + theta_field: str = "theta", + intensity_field: Optional[str] = None, + flip_sign: bool = False, + offset_deg: float = 0.0, +) -> SignalTable: + """Build a :class:`SignalTable` from a ``BraggPeaksPolymer``-style object. + + Replicates the orientation convention of + ``BraggPeaksPolymer.make_orientation_histogram`` (Karen's polar transform): the + **stored polar** ``theta`` (radians) is used directly -- optionally sign-flipped + (``flip_sign`` <-> ``orientation_flip_sign``), offset (``offset_deg`` <-> + ``orientation_offset_degrees``), then folded ``mod pi`` to ``[0, 180)`` degrees. + Reading the stored ``theta`` inherits Karen's sign by construction rather than + re-deriving it from ``qx``/``qy``. ``r`` is the stored polar magnitude; radial + windows gate on ``r**2`` against ``radial_ranges`` (each row ``[r_min, r_max]``). + + Two ingestion APIs are supported and auto-detected: + + * A real quantem ``Vector`` (``bp.polar_peaks`` / ``bp.peak_intensities``): per-cell + 1-D field arrays are read with ``vec.select_fields(field)[rx, ry].array[:, 0]`` + -- exactly the access ``make_orientation_histogram`` uses. Default fields are then + ``r_invA`` / ``theta`` / ``intensities``. + * A legacy field-indexable container (``polar[field][rx, ry]`` -> 1-D array), used by + the unit tests. Default fields are then ``r`` / ``theta`` / ``intensity``. + + ``r_field`` / ``intensity_field`` default to ``None`` and resolve per API above; + pass explicit names to override. + + Parameters + ---------- + bragg_peaks : object + Provides ``.polar_peaks`` (a ``Vector`` or field-indexable container with + ``.shape == (Rx, Ry)``) and ``.peak_intensities``; or is itself such a + container (then intensities are read from it too). + """ + polar = getattr(bragg_peaks, "polar_peaks", bragg_peaks) + inten_src = getattr(bragg_peaks, "peak_intensities", polar) + + # quantem Vector rejects string field indexing (needs select_fields); auto-detect it. + is_vector = hasattr(polar, "select_fields") + if r_field is None: + r_field = "r_invA" if is_vector else "r" + if intensity_field is None: + intensity_field = "intensities" if is_vector else "intensity" + + radial_ranges = np.atleast_2d(np.asarray(radial_ranges, dtype=float)) + rr2 = radial_ranges ** 2 + Rx, Ry = polar.shape + offset_rad = np.deg2rad(offset_deg) + + if is_vector: + R_v = polar.select_fields(r_field) + TH_v = polar.select_fields(theta_field) + II_v = inten_src.select_fields(intensity_field) + + def r_cell(i, j): + return np.asarray(R_v[i, j].array[:, 0], dtype=float) + + def th_cell(i, j): + return np.asarray(TH_v[i, j].array[:, 0], dtype=float) + + def i_cell(i, j): + return np.asarray(II_v[i, j].array[:, 0], dtype=float) + else: + R, TH, II = polar[r_field], polar[theta_field], inten_src[intensity_field] + + def _legacy(grid, i, j): + a = grid[i, j] + return np.empty(0) if a is None else np.asarray(a, dtype=float) + + def r_cell(i, j): + return _legacy(R, i, j) + + def th_cell(i, j): + return _legacy(TH, i, j) + + def i_cell(i, j): + return _legacy(II, i, j) + + pos_l, th_l, r_l, i_l, w_l = [], [], [], [], [] + for rx in range(Rx): + for ry in range(Ry): + p_r = r_cell(rx, ry) + if len(p_r) == 0: + continue + p_th = th_cell(rx, ry) + inten = i_cell(rx, ry) + r2 = p_r ** 2 + ang = -p_th if flip_sign else p_th # do not mutate the source array + ang = np.degrees(np.mod(ang + offset_rad, np.pi)) # -> [0, 180) + for w, (lo2, hi2) in enumerate(rr2): + sub = (r2 >= lo2) & (r2 < hi2) + if not np.any(sub): + continue + n = int(sub.sum()) + pos_l.append(np.column_stack([np.full(n, rx), np.full(n, ry)])) + th_l.append(ang[sub]) + r_l.append(p_r[sub]) + i_l.append(inten[sub]) + w_l.append(np.full(n, w)) + + if not pos_l: + empty_i = np.zeros((0, 2), dtype=np.int64) + empty_f = np.zeros((0,), dtype=float) + return SignalTable(empty_i, empty_f, empty_f, empty_f, + empty_f.astype(np.int64), (Rx, Ry)) + + return SignalTable( + np.concatenate(pos_l, axis=0), + np.concatenate(th_l), + np.concatenate(r_l), + np.concatenate(i_l), + np.concatenate(w_l), + (Rx, Ry), + ) + + +# -------------------------------------------------------------------------------------- +# adjacency +# -------------------------------------------------------------------------------------- +def _signal_edges(pos: np.ndarray, neighbor_dist: int): + """All signal pairs whose probes are within Chebyshev ``neighbor_dist`` (excluding + same-probe pairs). Returns local index arrays ``(ii, jj)`` with ``ii < jj``. + """ + n = pos.shape[0] + if n < 2: + return np.empty(0, np.int64), np.empty(0, np.int64) + tree = cKDTree(pos.astype(float)) + pairs = tree.query_pairs(r=neighbor_dist, p=np.inf, output_type="ndarray") + if pairs.shape[0] == 0: + return np.empty(0, np.int64), np.empty(0, np.int64) + ii, jj = pairs[:, 0], pairs[:, 1] + # drop same-probe pairs (different signals at the identical probe are not neighbours) + diff = np.any(pos[ii] != pos[jj], axis=1) + return ii[diff], jj[diff] + + +# -------------------------------------------------------------------------------------- +# agglomerative boundary merge (Stage A core) +# -------------------------------------------------------------------------------------- +def _key(a: int, b: int): + return (a, b) if a < b else (b, a) + + +def _agglomerative_merge( + n: int, + ii: np.ndarray, + jj: np.ndarray, + d_theta: np.ndarray, + d_r_rel: np.ndarray, + d_i_rel: np.ndarray, + theta_tol: float, + r_tol_rel: float, + intensity_tol_rel: float, + probe_lin: np.ndarray, + enforce_one_per_probe: bool, +): + """Average-linkage agglomeration on the spatial signal graph with a discontinuity + stop. Returns a (n,) array of root ids (a flat clustering).""" + parent = np.arange(n, dtype=np.int64) + + def find(x: int) -> int: + root = x + while parent[root] != root: + root = parent[root] + while parent[x] != root: + parent[x], x = root, parent[x] + return root + + inv_theta = (1.0 / theta_tol) if np.isfinite(theta_tol) else 0.0 + inv_r = (1.0 / r_tol_rel) if np.isfinite(r_tol_rel) else 0.0 + inv_i = (1.0 / intensity_tol_rel) if np.isfinite(intensity_tol_rel) else 0.0 + + def cost_of(stats) -> float: + s_th, s_r, s_i, c = stats + return max((s_th / c) * inv_theta, (s_r / c) * inv_r, (s_i / c) * inv_i) + + # region-adjacency graph: edge stats keyed by current root pair, neighbour sets + edge_stats: dict = {} + nbrs: list = [set() for _ in range(n)] + for a, b, t, rr, di in zip(ii.tolist(), jj.tolist(), d_theta, d_r_rel, d_i_rel): + k = _key(a, b) + st = edge_stats.get(k) + if st is None: + edge_stats[k] = [float(t), float(rr), float(di), 1] + nbrs[a].add(b) + nbrs[b].add(a) + else: # parallel edges between the same singleton pair shouldn't happen, but be safe + st[0] += float(t); st[1] += float(rr); st[2] += float(di); st[3] += 1 + + probes = [{int(probe_lin[k])} for k in range(n)] if enforce_one_per_probe else None + + heap = [] + for (a, b), st in edge_stats.items(): + heapq.heappush(heap, (cost_of(st), a, b)) + + while heap: + cost, a, b = heapq.heappop(heap) + if cost >= 1.0: + break + ra, rb = find(a), find(b) + if ra == rb: + continue + k = _key(ra, rb) + st = edge_stats.get(k) + if st is None: + continue # no longer adjacent + cur = cost_of(st) + if cur > cost + 1e-12: # stale: re-push with the up-to-date cost + heapq.heappush(heap, (cur, ra, rb)) + continue + if cur >= 1.0: + continue + + # keep the region with the larger probe set as the survivor (small-to-large) + if enforce_one_per_probe: + if len(probes[rb]) > len(probes[ra]): + ra, rb = rb, ra + k = _key(ra, rb) + small, large = (probes[rb], probes[ra]) if len(probes[rb]) <= len(probes[ra]) else (probes[ra], probes[rb]) + if not small.isdisjoint(large): + # merging would put two signals of one probe in a single grain: forbid + edge_stats.pop(k, None) + nbrs[ra].discard(rb) + nbrs[rb].discard(ra) + continue + + # merge rb -> ra + parent[rb] = ra + edge_stats.pop(k, None) + nbrs[ra].discard(rb) + nbrs[rb].discard(ra) + if enforce_one_per_probe: + probes[ra] |= probes[rb] + probes[rb] = None + + for c in list(nbrs[rb]): + kbc = _key(rb, c) + stbc = edge_stats.pop(kbc) + nbrs[c].discard(rb) + if c == ra: + continue + kac = _key(ra, c) + stac = edge_stats.get(kac) + if stac is None: + edge_stats[kac] = stbc + nbrs[ra].add(c) + nbrs[c].add(ra) + else: + stac[0] += stbc[0]; stac[1] += stbc[1]; stac[2] += stbc[2]; stac[3] += stbc[3] + heapq.heappush(heap, (cost_of(edge_stats[kac]), ra, c)) + nbrs[rb] = set() + + roots = np.array([find(i) for i in range(n)], dtype=np.int64) + return roots + + +# -------------------------------------------------------------------------------------- +# main entry point +# -------------------------------------------------------------------------------------- +def cluster_signals_into_grains( + signals: SignalTable, + *, + theta_tol_deg: float = 10.0, + r_tol_rel: float = 0.10, + intensity_tol_rel: float = np.inf, + neighbor_dist: int = 1, + area_min: int = 3, + enforce_one_per_probe: bool = True, +) -> GrainResult: + """Cluster signals into grains, strictly within each radial window. + + Parameters + ---------- + signals : SignalTable + theta_tol_deg : float + Max orientation discontinuity (deg, circular) across a within-grain boundary. + Tie to the histogram angular resolution (a few x ``sigma_theta``). ``inf`` + disables orientation gating. + r_tol_rel : float + Max *relative* radius discontinuity ``|dr| / r_mean`` across a within-grain + boundary. Set below the inter-peak (quantized) radial gap so distinct radii in + the same window separate into different grains. ``inf`` disables radius gating. + intensity_tol_rel : float + Optional relative intensity discontinuity tolerance. Default ``inf`` (off): + intensity varies within real grains, so it is not gated. + neighbor_dist : int + Chebyshev probe radius for adjacency. 1 = 8-connectivity; 2 bridges single + missing detections so a dropout does not fragment a grain. + area_min : int + Grains with fewer than this many signals become outliers (label -1). + enforce_one_per_probe : bool + Forbid a grain from containing two signals at the same probe. + + Returns + ------- + GrainResult + """ + N = len(signals) + Rx, Ry = signals.map_shape + labels = np.full(N, -1, dtype=np.int64) + theta_tol = float(theta_tol_deg) + next_label = 0 + windows = np.unique(signals.window) + + for w in windows: + idx = np.nonzero(signals.window == w)[0] + if idx.size == 0: + continue + pos = signals.pos[idx] + theta = signals.theta[idx] + r = signals.r[idx] + inten = signals.intensity[idx] + probe_lin = pos[:, 0].astype(np.int64) * Ry + pos[:, 1].astype(np.int64) + + ii, jj = _signal_edges(pos, neighbor_dist) + if ii.size: + d_theta = circular_distance_deg(theta[ii], theta[jj]) + rbar = 0.5 * (r[ii] + r[jj]) + d_r_rel = np.abs(r[ii] - r[jj]) / np.where(rbar > 0, rbar, 1.0) + ibar = 0.5 * (inten[ii] + inten[jj]) + d_i_rel = np.abs(inten[ii] - inten[jj]) / np.where(ibar > 0, ibar, 1.0) + else: + d_theta = d_r_rel = d_i_rel = np.empty(0) + + roots = _agglomerative_merge( + idx.size, ii, jj, d_theta, d_r_rel, d_i_rel, + theta_tol, float(r_tol_rel), float(intensity_tol_rel), + probe_lin, enforce_one_per_probe, + ) + + # area filter + contiguous relabelling (per window), offset into the global space + uniq, counts = np.unique(roots, return_counts=True) + keep = {root: counts[i] >= area_min for i, root in enumerate(uniq)} + remap = {} + for root in uniq: + if keep[root]: + remap[root] = next_label + next_label += 1 + for local_i, root in enumerate(roots): + if keep[root]: + labels[idx[local_i]] = remap[root] + + n_grains = next_label + grains = _summarize(signals, labels, n_grains) + label_map = _rasterize(signals, labels, len(windows)) + + return GrainResult( + labels=labels, + n_grains=n_grains, + grains=grains, + label_map=label_map, + params=dict( + theta_tol_deg=theta_tol_deg, + r_tol_rel=r_tol_rel, + intensity_tol_rel=intensity_tol_rel, + neighbor_dist=neighbor_dist, + area_min=area_min, + enforce_one_per_probe=enforce_one_per_probe, + ), + ) + + +# -------------------------------------------------------------------------------------- +# postprocessing helpers +# -------------------------------------------------------------------------------------- +def _summarize(signals: SignalTable, labels: np.ndarray, n_grains: int) -> list: + grains = [] + for g in range(n_grains): + sids = np.nonzero(labels == g)[0] + if sids.size == 0: + continue + th = signals.theta[sids] + # circular mean / std on the doubled angle (2-fold) + ang2 = np.deg2rad(2.0 * th) + c, s = np.cos(ang2).mean(), np.sin(ang2).mean() + theta_mean = (np.rad2deg(np.arctan2(s, c)) / 2.0) % 180.0 + R = np.hypot(c, s) + theta_std = np.rad2deg(np.sqrt(max(0.0, -2.0 * np.log(max(R, 1e-12))))) / 2.0 + pos = signals.pos[sids] + grains.append( + GrainInfo( + label=g, + window=int(signals.window[sids[0]]), + signal_ids=sids, + n_signals=int(sids.size), + theta_mean=float(theta_mean), + theta_std=float(theta_std), + r_mean=float(signals.r[sids].mean()), + intensity_median=float(np.median(signals.intensity[sids])), + centroid=(float(pos[:, 0].mean()), float(pos[:, 1].mean())), + ) + ) + return grains + + +def _rasterize(signals: SignalTable, labels: np.ndarray, num_windows: int) -> np.ndarray: + Rx, Ry = signals.map_shape + label_map = np.full((num_windows, Rx, Ry), -1, dtype=np.int64) + order = np.argsort(signals.intensity, kind="stable") # higher intensity overwrites + for i in order: + g = labels[i] + if g < 0: + continue + w = int(signals.window[i]) + rx, ry = signals.pos[i] + label_map[w, rx, ry] = g + return label_map + + +def _apply_area_min(labels: np.ndarray, area_min: int) -> np.ndarray: + """Dissolve grains with < area_min signals to -1 and relabel 0..K-1 contiguously.""" + out = np.full_like(labels, -1) + valid = labels >= 0 + if not np.any(valid): + return out + uniq, counts = np.unique(labels[valid], return_counts=True) + keep = uniq[counts >= area_min] + for new, old in enumerate(np.sort(keep)): + out[labels == old] = new + return out + + +# -------------------------------------------------------------------------------------- +# Stage B: CRF + EM refinement (boundary precision, outliers, soft confidence) +# -------------------------------------------------------------------------------------- +def _circular_mean_deg(theta_deg: np.ndarray, weights: np.ndarray) -> float: + """Weighted circular mean on the doubled angle (2-fold), returned in [0, 180).""" + a = np.deg2rad(2.0 * np.asarray(theta_deg, float)) + w = np.asarray(weights, float) + c = float(np.sum(w * np.cos(a))) + s = float(np.sum(w * np.sin(a))) + return (np.rad2deg(np.arctan2(s, c)) / 2.0) % 180.0 + + +def refine_grains_crf( + signals: SignalTable, + init, + *, + theta_sigma_deg: float = 8.0, + r_sigma_rel: float = 0.06, + lam: float = 0.5, + model_radius: float = 3.0, + neighbor_dist: int = 1, + outlier_energy: float = 6.0, + max_iter: int = 8, + enforce_one_per_probe: bool = True, + area_min: int = 1, +) -> GrainResult: + """Refine a Stage-A clustering by minimising a contrast-sensitive CRF energy with + ICM, EM-style: each grain's orientation model is a *local* weighted circular mean of + its nearby members (recomputed from current labels), so the energy is + + E(x) = sum_s U(x_s) + lam * sum_{(i,j) in nbrs} w_ij * [x_i != x_j] + U(s, g) = dtheta(theta_s, theta_g_local(p_s))^2 / (2 sigma_theta^2) + + (drel r)^2 / (2 sigma_r^2) + U(s, outlier) = outlier_energy + w_ij = exp(-dtheta_ij^2/2sigma_theta^2 - drel_ij^2/2sigma_r^2) + + Because the grain model is evaluated *at the signal's location* (not a global mean), + continuously-flowing / bent grains are preserved exactly as in Stage A. A signal's + candidate labels are only grains present within ``model_radius`` probes, so it can + never jump to a spatially-distant grain; ``enforce_one_per_probe`` is upheld. + + Returns a :class:`GrainResult` with ``confidence`` (max posterior) and ``margin`` + (top1-top2 energy gap) populated. ``init`` may be a Stage-A ``GrainResult`` or an + ``(N,)`` label array; for best results run Stage A with ``area_min=1`` so no grains + are dissolved before refinement, then let this apply the final ``area_min``. + """ + init_labels = init.labels if isinstance(init, GrainResult) else np.asarray(init) + labels = init_labels.astype(np.int64).copy() + N = len(signals) + Rx, Ry = signals.map_shape + confidence = np.zeros(N) + margin = np.full(N, np.inf) + + two_sig_th2 = 2.0 * theta_sigma_deg ** 2 + two_sig_r2 = 2.0 * r_sigma_rel ** 2 + sig_d = max(model_radius / 2.0, 1e-6) + + for w in np.unique(signals.window): + idx = np.nonzero(signals.window == w)[0] + if idx.size == 0: + continue + pos = signals.pos[idx].astype(float) + theta = signals.theta[idx] + r = signals.r[idx] + probe_lin = (signals.pos[idx, 0] * Ry + signals.pos[idx, 1]).astype(np.int64) + lab = labels[idx].copy() + n = idx.size + + # contrast-sensitive pairwise adjacency + ii, jj = _signal_edges(signals.pos[idx], neighbor_dist) + adj = [[] for _ in range(n)] + if ii.size: + dth = circular_distance_deg(theta[ii], theta[jj]) + rbar = 0.5 * (r[ii] + r[jj]) + drr = np.abs(r[ii] - r[jj]) / np.where(rbar > 0, rbar, 1.0) + wij = np.exp(-(dth ** 2) / two_sig_th2 - (drr ** 2) / two_sig_r2) + for a, b, wv in zip(ii.tolist(), jj.tolist(), wij.tolist()): + adj[a].append((b, wv)) + adj[b].append((a, wv)) + + # local-model neighbours within model_radius (Chebyshev) + Gaussian weights + tree = cKDTree(pos) + ball = tree.query_ball_point(pos, r=model_radius, p=np.inf) + model_nbr, model_w = [], [] + for i in range(n): + nb = np.array([k for k in ball[i] if k != i], dtype=np.int64) + model_nbr.append(nb) + if nb.size: + d = np.linalg.norm(pos[nb] - pos[i], axis=1) + model_w.append(np.exp(-(d ** 2) / (2.0 * sig_d ** 2))) + else: + model_w.append(np.zeros(0)) + + probe_members = {} + for i in range(n): + probe_members.setdefault(int(probe_lin[i]), []).append(i) + + def pair_cost(i, g): + c = 0.0 + for (j, wv) in adj[i]: + if lab[j] != g: + c += wv + return lam * c + + def unary(i, g, nb, nbl, nbw): + mask = nbl == g + if not np.any(mask): + return None + th_g = _circular_mean_deg(theta[nb[mask]], nbw[mask]) + wsum = float(nbw[mask].sum()) + r_g = float(np.sum(r[nb[mask]] * nbw[mask]) / wsum) if wsum > 0 else float(r[nb[mask]].mean()) + dth = float(circular_distance_deg(theta[i], th_g)) + drr = abs(r[i] - r_g) / r_g if r_g > 0 else 0.0 + return (dth ** 2) / two_sig_th2 + (drr ** 2) / two_sig_r2 + + # ICM sweeps to convergence (EM model recomputed implicitly each evaluation) + for _ in range(max_iter): + changed = 0 + for i in range(n): + nb = model_nbr[i] + if nb.size == 0: + new = -1 + else: + nbl = lab[nb] + nbw = model_w[i] + cands = {int(x) for x in nbl if x >= 0} + forbidden = set() + if enforce_one_per_probe: + for m in probe_members[int(probe_lin[i])]: + if m != i and lab[m] >= 0: + forbidden.add(int(lab[m])) + best_lab, best_E = -1, outlier_energy + pair_cost(i, -1) + for g in cands: + if g in forbidden: + continue + U = unary(i, g, nb, nbl, nbw) + if U is None: + continue + E = U + pair_cost(i, g) + if E < best_E: + best_E, best_lab = E, g + new = best_lab + if new != lab[i]: + lab[i] = new + changed += 1 + if changed == 0: + break + + # posteriors from final neighbour labels + for i in range(n): + nb = model_nbr[i] + energies = [outlier_energy + pair_cost(i, -1)] + if nb.size: + nbl = lab[nb] + nbw = model_w[i] + for g in {int(x) for x in nbl if x >= 0}: + U = unary(i, g, nb, nbl, nbw) + if U is not None: + energies.append(U + pair_cost(i, g)) + e = np.sort(np.array(energies)) + p = np.exp(-(e - e[0])) + p /= p.sum() + confidence[idx[i]] = float(p[0]) + margin[idx[i]] = float(e[1] - e[0]) if e.size > 1 else np.inf + + labels[idx] = lab + + labels = _apply_area_min(labels, area_min) + n_grains = int(labels.max()) + 1 if labels.max() >= 0 else 0 + grains = _summarize(signals, labels, n_grains) + label_map = _rasterize(signals, labels, int(np.unique(signals.window).size)) + return GrainResult( + labels=labels, + n_grains=n_grains, + grains=grains, + label_map=label_map, + confidence=confidence, + margin=margin, + params=dict( + stage="B", + theta_sigma_deg=theta_sigma_deg, + r_sigma_rel=r_sigma_rel, + lam=lam, + model_radius=model_radius, + neighbor_dist=neighbor_dist, + outlier_energy=outlier_energy, + max_iter=max_iter, + area_min=area_min, + ), + ) + + +# -------------------------------------------------------------------------------------- +# visualization overlays +# -------------------------------------------------------------------------------------- +# flowline colour basis (matches make_flowline_rainbow_image so hues are consistent) +_FLOWLINE_C0 = np.array([1.0, 0.0, 0.0]) +_FLOWLINE_C1 = np.array([0.0, 0.7, 0.0]) +_FLOWLINE_C2 = np.array([0.0, 0.3, 1.0]) + + +def orientation_to_rgb(theta_deg, sym_rotation_order: int = 2, theta_offset: float = 0.0): + """Map orientation angle(s) (degrees) to RGB using the *flowline* colour basis. + + Reproduces ``make_flowline_rainbow_image``: ``theta_color = theta_offset + + sym_rotation_order * theta`` projected onto three colour vectors peaked at 0, 2pi/3, + 4pi/3, so grain hues match existing flowline plots. Output adds a trailing length-3 + axis. For ``sym_rotation_order=2`` (polymer 2-fold) angles theta and theta+180 map to + the same colour. + """ + th = np.deg2rad(np.asarray(theta_deg, dtype=float)) + tc = theta_offset + sym_rotation_order * th + denom = (np.pi * 2.0 / 3.0) ** 2 + + def proj(shift): + return np.maximum(1.0 - np.abs(np.mod(tc - shift + np.pi, 2 * np.pi) - np.pi) ** 2 / denom, 0.0) + + b0, b1, b2 = proj(0.0), proj(np.pi * 2.0 / 3.0), proj(np.pi * 4.0 / 3.0) + rgb = b0[..., None] * _FLOWLINE_C0 + b1[..., None] * _FLOWLINE_C1 + b2[..., None] * _FLOWLINE_C2 + return np.clip(rgb, 0.0, 1.0) + + +def _hsv_to_rgb(h, s, v): + h = np.asarray(h, float); s = np.asarray(s, float); v = np.asarray(v, float) + i = np.floor(h * 6.0).astype(int) + f = h * 6.0 - i + p = v * (1.0 - s) + q = v * (1.0 - f * s) + t = v * (1.0 - (1.0 - f) * s) + i = i % 6 + r = np.choose(i, [v, q, p, p, t, v]) + g = np.choose(i, [t, v, v, q, p, p]) + b = np.choose(i, [p, p, t, v, v, q]) + return np.stack([r, g, b], axis=-1) + + +def _qualitative_palette(n: int, seed: int = 0): + """n maximally-spaced distinct colours (golden-ratio hue spacing).""" + if n <= 0: + return np.zeros((0, 3)) + k = np.arange(n) + h = (0.61803398875 * (k + 1) + seed * 0.137) % 1.0 + return _hsv_to_rgb(h, np.full(n, 0.62), np.full(n, 0.97)) + + +def _rasterize_window(signals: SignalTable, result: GrainResult, window: int): + """Per-window maps (highest-intensity signal wins each probe): label, theta, confidence, + and a 'filled' mask (a signal of this window present regardless of label).""" + Rx, Ry = signals.map_shape + lab = np.full((Rx, Ry), -1, dtype=np.int64) + th = np.zeros((Rx, Ry)) + conf = np.zeros((Rx, Ry)) if result.confidence is not None else None + filled = np.zeros((Rx, Ry), dtype=bool) + idx = np.nonzero(signals.window == window)[0] + order = idx[np.argsort(signals.intensity[idx], kind="stable")] + for i in order: + rx, ry = signals.pos[i] + filled[rx, ry] = True + th[rx, ry] = signals.theta[i] + lab[rx, ry] = result.labels[i] + if conf is not None: + conf[rx, ry] = result.confidence[i] + return lab, th, conf, filled + + +def _boundary_mask(lab: np.ndarray, outline_background: bool = False) -> np.ndarray: + """Boundary pixels between two *distinct grains* (both labels >= 0). + + Grain<->outlier / grain<->empty transitions are NOT marked, so scattered outliers do + not leave black halos. With ``outline_background=True`` the assigned side of a + grain<->background edge is also outlined. + """ + b = np.zeros(lab.shape, dtype=bool) + up, dn = lab[:-1, :], lab[1:, :] + le, ri = lab[:, :-1], lab[:, 1:] + dv = (up != dn) & (up >= 0) & (dn >= 0) + dh = (le != ri) & (le >= 0) & (ri >= 0) + b[:-1, :] |= dv; b[1:, :] |= dv + b[:, :-1] |= dh; b[:, 1:] |= dh + if outline_background: + b[:-1, :] |= (up >= 0) & (dn < 0) + b[1:, :] |= (dn >= 0) & (up < 0) + b[:, :-1] |= (le >= 0) & (ri < 0) + b[:, 1:] |= (ri >= 0) & (le < 0) + return b + + +def grain_rgb_overlay( + signals: SignalTable, + result: GrainResult, + *, + window: int = 0, + mode: str = "orientation", + overlap: str = "dominant", + stripe_width: int = 2, + boundary: bool = True, + boundary_color=(0.0, 0.0, 0.0), + outline_background: bool = False, + background=(0.12, 0.12, 0.12), + outlier_color=None, + confidence_shading: bool = False, + sym_rotation_order: int = 2, + theta_offset: float = 0.0, + qualitative_seed: int = 0, + upsample: int = 1, +) -> np.ndarray: + """Build an RGB image of the grain clustering for one radial window. + + Visually distinct from flowlines: a *filled segmentation with hard grain boundaries* + rather than streamlines. Modes: + + * ``"orientation"`` -- each probe coloured by its own signal orientation (flowline + hue), so within-grain orientation gradients stay visible *and* grains are outlined + (the candidate to supersede flowlines: same orientation field + grain structure). + * ``"mean_orientation"`` -- each grain a flat colour = its circular-mean orientation. + * ``"grain"`` -- a distinct qualitative colour per grain id (partition only). + + ``overlap="stripe"`` renders probes carrying several grains as diagonally striped tiles + (one stripe colour per grain, ordered by intensity), so overlapping grains are visible in + one image; there ``upsample`` sets the tile size (auto-bumped to 8 if < 4) and + ``stripe_width`` the stripe period. ``overlap="dominant"`` (default) keeps the + highest-intensity grain per probe. + + ``confidence_shading`` (Stage B) dims low-confidence signals; ``outlier_color`` fills + rejected signals; ``upsample`` does nearest-neighbour zoom. Returns ``(Rx*u, Ry*u, 3)`` + in [0, 1]. + """ + Rx, Ry = signals.map_shape + if overlap == "stripe": + tile = upsample if upsample >= 4 else 8 + return _striped_overlay( + signals, result, window, mode=mode, tile=tile, stripe_width=stripe_width, + background=background, boundary=boundary, boundary_color=boundary_color, + outline_background=outline_background, confidence_shading=confidence_shading, + sym_rotation_order=sym_rotation_order, theta_offset=theta_offset, + qualitative_seed=qualitative_seed, + ) + if overlap != "dominant": + raise ValueError(f"unknown overlap {overlap!r}") + lab, th, conf, filled = _rasterize_window(signals, result, window) + assigned = lab >= 0 + + rgb = np.zeros((Rx, Ry, 3), float) + if background is not None: + rgb[:] = np.asarray(background, float) + + if mode == "orientation": + rgb[assigned] = orientation_to_rgb(th[assigned], sym_rotation_order, theta_offset) + elif mode == "mean_orientation": + mean_th = {int(g.label): g.theta_mean for g in result.grains} + gm = np.array([mean_th.get(int(l), 0.0) for l in lab[assigned]]) + rgb[assigned] = orientation_to_rgb(gm, sym_rotation_order, theta_offset) + elif mode == "grain": + palette = _qualitative_palette(max(result.n_grains, 1), qualitative_seed) + rgb[assigned] = palette[lab[assigned]] + else: + raise ValueError(f"unknown mode {mode!r}") + + if outlier_color is not None: + rgb[filled & ~assigned] = np.asarray(outlier_color, float) + + if confidence_shading and conf is not None: + factor = np.ones((Rx, Ry)) + factor[assigned] = np.clip(conf[assigned], 0.0, 1.0) + rgb = rgb * factor[..., None] + + if boundary: + rgb[_boundary_mask(lab, outline_background)] = np.asarray(boundary_color, float) + + if upsample > 1: + rgb = np.kron(rgb, np.ones((upsample, upsample, 1))) + return rgb + + +def _probe_signal_stacks(signals: SignalTable, result: GrainResult, window: int): + """Per probe, the list of assigned signals ``(label, theta, intensity, confidence)``, + de-duplicated by grain and sorted by descending intensity (the 'stack' at that probe).""" + idx = np.nonzero(signals.window == window)[0] + conf_arr = result.confidence + stacks: dict = {} + for i in idx: + lab = int(result.labels[i]) + if lab < 0: + continue + key = (int(signals.pos[i, 0]), int(signals.pos[i, 1])) + c = float(conf_arr[i]) if conf_arr is not None else 1.0 + stacks.setdefault(key, []).append((lab, float(signals.theta[i]), float(signals.intensity[i]), c)) + out = {} + for key, lst in stacks.items(): + lst.sort(key=lambda t: -t[2]) + seen, uniq = set(), [] + for t in lst: + if t[0] in seen: + continue + seen.add(t[0]) + uniq.append(t) + out[key] = uniq + return out + + +def _striped_overlay( + signals, result, window, *, mode, tile, stripe_width, background, boundary, + boundary_color, outline_background, confidence_shading, sym_rotation_order, + theta_offset, qualitative_seed, +): + """Render multi-grain probes as diagonally striped tiles (see ``grain_rgb_overlay``).""" + Rx, Ry = signals.map_shape + palette = _qualitative_palette(max(result.n_grains, 1), qualitative_seed) + mean_th = {int(g.label): g.theta_mean for g in result.grains} + + def color_for(label, theta): + if mode == "orientation": + return np.asarray(orientation_to_rgb(theta, sym_rotation_order, theta_offset), float) + if mode == "mean_orientation": + return np.asarray(orientation_to_rgb(mean_th.get(int(label), 0.0), sym_rotation_order, theta_offset), float) + if mode == "grain": + return np.asarray(palette[int(label)], float) + raise ValueError(f"unknown mode {mode!r}") + + stacks = _probe_signal_stacks(signals, result, window) + img = np.zeros((Rx * tile, Ry * tile, 3), float) + if background is not None: + img[:] = np.asarray(background, float) + dom = np.full((Rx, Ry), -1, dtype=np.int64) + iu, ju = np.mgrid[0:tile, 0:tile] + base = (iu + ju) // max(int(stripe_width), 1) + + for (rx, ry), stack in stacks.items(): + dom[rx, ry] = stack[0][0] + colors = [] + for (lab, theta, _inten, conf) in stack: + c = color_for(lab, theta) + if confidence_shading: + c = c * float(np.clip(conf, 0.0, 1.0)) + colors.append(c) + sub = img[rx * tile:(rx + 1) * tile, ry * tile:(ry + 1) * tile] + if len(colors) == 1: + sub[:] = colors[0] + else: + sidx = base % len(colors) + for k, c in enumerate(colors): + sub[sidx == k] = c + + if boundary: + bt = max(1, tile // 6) + bc = np.asarray(boundary_color, float) + for rx, ry in zip(*np.nonzero((dom[:-1, :] >= 0) & (dom[1:, :] >= 0) & (dom[:-1, :] != dom[1:, :]))): + y = (int(rx) + 1) * tile + img[max(0, y - bt):y + bt, int(ry) * tile:(int(ry) + 1) * tile] = bc + for rx, ry in zip(*np.nonzero((dom[:, :-1] >= 0) & (dom[:, 1:] >= 0) & (dom[:, :-1] != dom[:, 1:]))): + x = (int(ry) + 1) * tile + img[int(rx) * tile:(int(rx) + 1) * tile, max(0, x - bt):x + bt] = bc + return img + + +def orientation_legend_image(size: int = 128, sym_rotation_order: int = 2, theta_offset: float = 0.0): + """RGBA colour-wheel legend (orientation -> flowline hue), transparent outside a ring.""" + yy, xx = np.mgrid[0:size, 0:size].astype(float) + c = (size - 1) / 2.0 + dx, dy = xx - c, -(yy - c) + rad = np.hypot(dx, dy) / (size / 2.0) + ang = np.degrees(np.arctan2(dy, dx)) % 180.0 + rgb = orientation_to_rgb(ang, sym_rotation_order, theta_offset) + alpha = ((rad <= 1.0) & (rad >= 0.32)).astype(float) + return np.concatenate([rgb, alpha[..., None]], axis=-1) + + +def plot_grain_map( + signals: SignalTable, + result: GrainResult, + *, + window: int = 0, + mode: str = "orientation", + ax=None, + title=None, + legend: bool = True, + **overlay_kw, +): + """Plot a grain overlay (lazy matplotlib import). Returns the matplotlib Axes.""" + import matplotlib.pyplot as plt + + rgb = grain_rgb_overlay(signals, result, window=window, mode=mode, **overlay_kw) + if ax is None: + _, ax = plt.subplots(figsize=(6, 6)) + ax.imshow(rgb, origin="upper", interpolation="nearest") + ax.set_xticks([]); ax.set_yticks([]) + ax.set_title(title or f"grains (window {window}, mode={mode}, n={result.n_grains})") + if legend and mode in ("orientation", "mean_orientation"): + leg = ax.inset_axes([0.80, 0.80, 0.18, 0.18]) + leg.imshow( + orientation_legend_image( + sym_rotation_order=overlay_kw.get("sym_rotation_order", 2), + theta_offset=overlay_kw.get("theta_offset", 0.0), + ), + origin="upper", + interpolation="bilinear", + ) + leg.set_xticks([]); leg.set_yticks([]) + leg.patch.set_alpha(0.0) + return ax diff --git a/src/quantem/diffraction/orientation_correlation.py b/src/quantem/diffraction/orientation_correlation.py new file mode 100644 index 000000000..ec523aeb2 --- /dev/null +++ b/src/quantem/diffraction/orientation_correlation.py @@ -0,0 +1,594 @@ +"""Memory-efficient distance-angle correlations for orientation histograms.""" + +from __future__ import annotations + +from collections.abc import Sequence + +import numpy as np +import torch +from scipy import fft as scipy_fft +from tqdm.auto import tqdm + + +def _validate_and_shape_input(orient_hist): + is_torch = isinstance(orient_hist, torch.Tensor) + if orient_hist.ndim == 3: + orient_hist = orient_hist[None] + elif orient_hist.ndim != 4: + raise ValueError( + "orient_hist must have shape (x, y, theta) or " + "(radial_bin, x, y, theta)" + ) + + num_radii, size_x, size_y, num_theta = orient_hist.shape + if min(num_radii, size_x, size_y) < 1 or num_theta < 2: + raise ValueError( + "orient_hist must contain at least one radial bin and spatial pixel, " + "and at least two theta bins" + ) + return orient_hist, is_torch + + +def _resolve_pairs(pairs, num_radii): + if isinstance(pairs, str): + if pairs == "all": + pair_list = [ + (first, second) + for first in range(num_radii) + for second in range(first, num_radii) + ] + elif pairs == "autocorrelation": + pair_list = [(index, index) for index in range(num_radii)] + else: + raise ValueError( + "pairs must be 'all', 'autocorrelation', or a sequence of pairs" + ) + else: + pair_list = [] + for pair in pairs: + if len(pair) != 2: + raise ValueError("each entry in pairs must contain two indices") + first, second = int(pair[0]), int(pair[1]) + if not ( + 0 <= first < num_radii and 0 <= second < num_radii + ): + raise ValueError( + f"radial-bin pair {(first, second)} is outside " + f"[0, {num_radii})" + ) + pair_list.append((first, second)) + + if not pair_list: + raise ValueError("pairs must contain at least one radial-bin pair") + + return np.asarray(pair_list, dtype=np.int64) + + +def _radial_geometry(size_x, size_y, radius_max): + """Build two-point linear interpolation from spatial pixels to radial bins.""" + padded_x = max(2 * size_x, 2 * radius_max) + padded_y = max(2 * size_y, 2 * radius_max) + + x = np.mod(np.arange(padded_x) + padded_x / 2, padded_x) - padded_x / 2 + y = np.mod(np.arange(padded_y) + padded_y / 2, padded_y) - padded_y / 2 + yy, xx = np.meshgrid(y, x) + radius = np.sqrt(xx**2 + yy**2) + + lower_mask = radius <= radius_max + upper_mask = radius <= radius_max - 1 + lower_floor = np.floor(radius[lower_mask]).astype(np.int64) + upper_floor = np.floor(radius[upper_mask]).astype(np.int64) + + return { + "padded_shape": (padded_x, padded_y), + "point_indices": ( + np.flatnonzero(lower_mask), + np.flatnonzero(upper_mask), + ), + "radial_bins": (lower_floor, upper_floor + 1), + "radial_weights": ( + 1.0 - (radius[lower_mask] - lower_floor), + radius[upper_mask] - upper_floor, + ), + } + + +def _normalize_correlation( + radial_correlation, + correlation_spectrum, + num_modes, + num_theta, + zero_policy, +): + denominator = correlation_spectrum[:, 0, :].real / num_theta + + if isinstance(radial_correlation, torch.Tensor): + dtype = radial_correlation.dtype + maximum = torch.max(torch.abs(denominator)) + threshold = torch.finfo(dtype).eps * torch.clamp( + maximum, min=torch.finfo(dtype).tiny + ) + valid = torch.abs(denominator) > threshold + if zero_policy == "raise" and not bool(torch.all(valid).item()): + raise ZeroDivisionError( + "orientation correlation has radial distances with zero " + "normalization signal" + ) + safe_denominator = torch.where( + valid, denominator, torch.ones_like(denominator) + ) + output = ( + radial_correlation[:, :num_modes, :] + / safe_denominator[:, None, :] + ) + fill_value = float("nan") if zero_policy == "nan" else 0.0 + return output.masked_fill(~valid[:, None, :], fill_value) + + dtype = radial_correlation.dtype + maximum = float(np.max(np.abs(denominator), initial=0.0)) + threshold = np.finfo(dtype).eps * max(maximum, np.finfo(dtype).tiny) + valid = np.abs(denominator) > threshold + if zero_policy == "raise" and not np.all(valid): + raise ZeroDivisionError( + "orientation correlation has radial distances with zero " + "normalization signal" + ) + fill_value = np.nan if zero_policy == "nan" else 0.0 + output = np.full( + (radial_correlation.shape[0], num_modes, radial_correlation.shape[2]), + fill_value, + dtype=dtype, + ) + np.divide( + radial_correlation[:, :num_modes, :], + denominator[:, None, :], + out=output, + where=valid[:, None, :], + ) + return output + + +def _calculate_numpy( + orient_hist, + pair_indices, + geometry, + num_theta, + radius_max, + *, + dtype, + mode_batch_size, + pair_batch_size, + workers, + zero_policy, + progress_bar, +): + real_dtype = np.float32 if dtype == "float32" else np.float64 + complex_dtype = np.complex64 if dtype == "float32" else np.complex128 + histogram = np.asarray(orient_hist, dtype=real_dtype) + if not np.all(np.isfinite(histogram)): + raise ValueError("orient_hist contains NaN or infinite values") + + num_pairs = len(pair_indices) + num_modes = num_theta // 2 + 1 + num_distances = radius_max + 1 + mode_batch_size = min(mode_batch_size or 1, num_modes) + pair_batch_size = min(pair_batch_size or 4, num_pairs) + if mode_batch_size < 1 or pair_batch_size < 1: + raise ValueError("mode_batch_size and pair_batch_size must be at least 1") + + theta_spectrum = scipy_fft.rfft(histogram, axis=-1, workers=workers) + correlation_spectrum = np.empty( + (num_pairs, num_modes, num_distances), dtype=complex_dtype + ) + point_indices = geometry["point_indices"] + radial_bins = geometry["radial_bins"] + radial_weights = geometry["radial_weights"] + + total = ( + int(np.ceil(num_modes / mode_batch_size)) + * int(np.ceil(num_pairs / pair_batch_size)) + ) + progress = tqdm( + total=total, + desc="Calculate orientation correlations (CPU)", + unit="batch", + disable=not progress_bar, + ) + try: + for mode_start in range(0, num_modes, mode_batch_size): + mode_stop = min(mode_start + mode_batch_size, num_modes) + spatial_spectrum = scipy_fft.fft2( + np.moveaxis( + theta_spectrum[..., mode_start:mode_stop], -1, 1 + ), + s=geometry["padded_shape"], + axes=(-2, -1), + workers=workers, + ) + + for pair_start in range(0, num_pairs, pair_batch_size): + pair_stop = min(pair_start + pair_batch_size, num_pairs) + pair_batch = pair_indices[pair_start:pair_stop] + cross_spectrum = ( + spatial_spectrum[pair_batch[:, 0]] + * np.conj(spatial_spectrum[pair_batch[:, 1]]) + ) + spatial_correlation = scipy_fft.ifft2( + cross_spectrum, axes=(-2, -1), workers=workers + ).reshape(len(pair_batch), mode_stop - mode_start, -1) + radial_correlation = np.zeros( + ( + len(pair_batch), + mode_stop - mode_start, + num_distances, + ), + dtype=complex_dtype, + ) + + # NumPy does not provide batched bincount. Only the small + # pair/mode dimensions are looped; spatial work stays vectorized. + for pair_index in range(len(pair_batch)): + for mode_index in range(mode_stop - mode_start): + output = radial_correlation[pair_index, mode_index] + for points, bins, weights in zip( + point_indices, radial_bins, radial_weights + ): + values = ( + spatial_correlation[ + pair_index, mode_index, points + ] + * weights + ) + output += np.bincount( + bins, + weights=values.real, + minlength=num_distances, + ) + output += 1j * np.bincount( + bins, + weights=values.imag, + minlength=num_distances, + ) + + correlation_spectrum[ + pair_start:pair_stop, mode_start:mode_stop + ] = radial_correlation + progress.update() + finally: + progress.close() + + radial_correlation = scipy_fft.irfft( + correlation_spectrum, n=num_theta, axis=1, workers=workers + ) + output = _normalize_correlation( + radial_correlation, + correlation_spectrum, + num_modes, + num_theta, + zero_policy, + ) + return output.astype(real_dtype, copy=False) + + +def _calculate_torch( + orient_hist, + pair_indices, + geometry, + num_theta, + radius_max, + *, + device, + dtype, + mode_batch_size, + pair_batch_size, + max_memory_fraction, + zero_policy, + progress_bar, +): + device = torch.device( + device + if device is not None + else ("cuda" if torch.cuda.is_available() else "cpu") + ) + if device.type == "cuda" and not torch.cuda.is_available(): + raise RuntimeError( + f"CUDA device {device} was requested, but CUDA is not available" + ) + + real_dtype = torch.float32 if dtype == "float32" else torch.float64 + complex_dtype = ( + torch.complex64 if real_dtype == torch.float32 else torch.complex128 + ) + histogram = torch.as_tensor( + orient_hist, dtype=real_dtype, device=device + ) + if not bool(torch.all(torch.isfinite(histogram)).item()): + raise ValueError("orient_hist contains NaN or infinite values") + + num_radii = histogram.shape[0] + num_pairs = len(pair_indices) + num_modes = num_theta // 2 + 1 + num_distances = radius_max + 1 + pair_batch_was_requested = pair_batch_size is not None + pair_batch_size = min(pair_batch_size or 4, num_pairs) + if pair_batch_size < 1: + raise ValueError("pair_batch_size must be at least 1") + + if mode_batch_size is None: + if device.type == "cuda": + free_memory, total_memory = torch.cuda.mem_get_info(device) + # cudaMemGetInfo reports memory free on the *device*, which ignores + # any per-process cap from torch.cuda.set_per_process_memory_fraction. + # Budgeting off the device figure under a cap sizes batches for + # headroom this process may not allocate, and the fft2 below then + # raises OutOfMemoryError while the device still looks mostly free. + # Take whichever allowance is smaller. + try: + process_fraction = torch.cuda.get_per_process_memory_fraction(device) + except (AttributeError, RuntimeError, TypeError): + process_fraction = 1.0 + complex_bytes = 8 if complex_dtype == torch.complex64 else 16 + padded_pixels = int(np.prod(geometry["padded_shape"])) + # theta_spectrum is allocated after this estimate, so reserve room for + # it up front rather than discovering the shortfall mid-loop. + spectrum_bytes = ( + histogram.numel() // num_theta * (num_theta // 2 + 1) + ) * complex_bytes + + def available_bytes(): + budget = free_memory + if 0.0 < process_fraction < 1.0: + # Headroom is measured against *allocated*, not *reserved*: + # blocks the caching allocator holds but no tensor is using + # are reusable, and after a failed run they can account for + # most of the cap. Counting them as spent wrongly reports a + # zero budget. + remaining = ( + process_fraction * total_memory + - torch.cuda.memory_allocated(device) + ) + budget = min(budget, max(0, int(remaining))) + return max(0, budget - spectrum_bytes) + + memory_budget = int(available_bytes() * max_memory_fraction) + + def estimate_bytes_per_mode(batch_size): + return ( + padded_pixels + * complex_bytes + * (num_radii + 3 * batch_size) + ) + + bytes_per_mode = estimate_bytes_per_mode(pair_batch_size) + if not pair_batch_was_requested: + while pair_batch_size > 1 and bytes_per_mode > memory_budget: + pair_batch_size = max(1, pair_batch_size // 2) + bytes_per_mode = estimate_bytes_per_mode(pair_batch_size) + if bytes_per_mode > memory_budget: + # A previous failure can leave the cap saturated with cached + # blocks. Return them and re-measure before giving up. + torch.cuda.empty_cache() + free_memory, total_memory = torch.cuda.mem_get_info(device) + memory_budget = int(available_bytes() * max_memory_fraction) + if not pair_batch_was_requested: + while pair_batch_size > 1 and bytes_per_mode > memory_budget: + pair_batch_size = max(1, pair_batch_size // 2) + bytes_per_mode = estimate_bytes_per_mode(pair_batch_size) + if bytes_per_mode > memory_budget: + allocated_gib = torch.cuda.memory_allocated(device) / 1024**3 + cap_text = ( + f"{process_fraction * total_memory / 1024**3:.2f} GiB " + "(torch.cuda.set_per_process_memory_fraction)" + if 0.0 < process_fraction < 1.0 + else f"{total_memory / 1024**3:.2f} GiB (device total, no cap)" + ) + raise MemoryError( + "A single angular-mode batch is estimated to require " + f"{bytes_per_mode / 1024**3:.2f} GiB, but the CUDA memory " + f"budget is only {memory_budget / 1024**3:.2f} GiB " + f"(max_memory_fraction={max_memory_fraction}). This process " + f"is capped at {cap_text}, currently holds " + f"{allocated_gib:.2f} GiB live, and must also reserve " + f"{spectrum_bytes / 1024**3:.2f} GiB for the angular " + f"spectrum; {free_memory / 1024**3:.2f} GiB is free on the " + "device. Raise the per-process cap or max_memory_fraction, " + "or reduce radius_max or the orientation-histogram upsampling." + ) + mode_batch_size = max( + 1, + memory_budget // max(bytes_per_mode, 1), + ) + else: + mode_batch_size = 1 + mode_batch_size = min(int(mode_batch_size), num_modes) + if mode_batch_size < 1: + raise ValueError("mode_batch_size must be at least 1") + + point_indices = [ + torch.as_tensor(values, dtype=torch.long, device=device) + for values in geometry["point_indices"] + ] + radial_bins = [ + torch.as_tensor(values, dtype=torch.long, device=device) + for values in geometry["radial_bins"] + ] + radial_weights = [ + torch.as_tensor(values, dtype=real_dtype, device=device) + for values in geometry["radial_weights"] + ] + pair_indices = torch.as_tensor( + pair_indices, dtype=torch.long, device=device + ) + + theta_spectrum = torch.fft.rfft(histogram, dim=-1) + correlation_spectrum = torch.empty( + (num_pairs, num_modes, num_distances), + dtype=complex_dtype, + device=device, + ) + total = ( + int(np.ceil(num_modes / mode_batch_size)) + * int(np.ceil(num_pairs / pair_batch_size)) + ) + progress = tqdm( + total=total, + desc=f"Calculate orientation correlations ({device})", + unit="batch", + disable=not progress_bar, + ) + try: + for mode_start in range(0, num_modes, mode_batch_size): + mode_stop = min(mode_start + mode_batch_size, num_modes) + spatial_spectrum = torch.fft.fft2( + theta_spectrum[..., mode_start:mode_stop].movedim(-1, 1), + s=geometry["padded_shape"], + dim=(-2, -1), + ) + + for pair_start in range(0, num_pairs, pair_batch_size): + pair_stop = min(pair_start + pair_batch_size, num_pairs) + pair_batch = pair_indices[pair_start:pair_stop] + cross_spectrum = ( + spatial_spectrum.index_select(0, pair_batch[:, 0]) + * torch.conj( + spatial_spectrum.index_select(0, pair_batch[:, 1]) + ) + ) + spatial_correlation = torch.fft.ifft2( + cross_spectrum, dim=(-2, -1) + ).flatten(-2) + radial_correlation = torch.zeros( + ( + len(pair_batch), + mode_stop - mode_start, + num_distances, + ), + dtype=complex_dtype, + device=device, + ) + for points, bins, weights in zip( + point_indices, radial_bins, radial_weights + ): + radial_correlation.index_add_( + -1, + bins, + spatial_correlation.index_select(-1, points) * weights, + ) + + correlation_spectrum[ + pair_start:pair_stop, mode_start:mode_stop + ] = radial_correlation + progress.update() + finally: + progress.close() + + radial_correlation = torch.fft.irfft( + correlation_spectrum, n=num_theta, dim=1 + ) + return _normalize_correlation( + radial_correlation, + correlation_spectrum, + num_modes, + num_theta, + zero_policy, + ) + + +def calculate_orientation_correlation( + orient_hist, + radius_max: int | None = None, + pairs: str | Sequence[tuple[int, int]] = "all", + backend: str = "auto", + device=None, + mode_batch_size: int | None = None, + pair_batch_size: int | None = None, + max_memory_fraction: float = 0.6, + dtype: str = "float32", + workers: int | None = None, + zero_policy: str = "nan", + return_numpy: bool = True, + progress_bar: bool = True, +): + """ + Compute spatial-distance versus relative-angle correlations. + + The angular Fourier modes are streamed through batched 2D spatial + correlations and radially integrated before the angular inverse transform. + This is equivalent to constructing a full 3D correlation volume, while + requiring substantially less peak memory. + + Returns + ------- + orient_corr, pair_indices + Correlation values have shape + ``(num_pairs, num_theta // 2 + 1, radius_max + 1)`` and are normalized + in multiples of a random distribution. ``pair_indices`` maps the first + axis back to radial-bin pairs. + """ + if backend not in {"auto", "numpy", "torch"}: + raise ValueError("backend must be 'auto', 'numpy', or 'torch'") + if dtype not in {"float32", "float64"}: + raise ValueError("dtype must be 'float32' or 'float64'") + if zero_policy not in {"nan", "zero", "raise"}: + raise ValueError("zero_policy must be 'nan', 'zero', or 'raise'") + if not 0 < max_memory_fraction <= 1: + raise ValueError("max_memory_fraction must be in the interval (0, 1]") + + orient_hist, is_torch_input = _validate_and_shape_input(orient_hist) + num_radii, size_x, size_y, num_theta = orient_hist.shape + if radius_max is None: + radius_max = int(np.ceil(min(size_x, size_y) / 2)) + elif not isinstance(radius_max, (int, np.integer)): + raise TypeError("radius_max must be an integer or None") + radius_max = int(radius_max) + if radius_max < 0: + raise ValueError("radius_max must be non-negative") + + pair_indices = _resolve_pairs(pairs, num_radii) + geometry = _radial_geometry(size_x, size_y, radius_max) + if backend == "auto": + wants_cuda = device is None or str(device).startswith("cuda") + backend = ( + "torch" + if torch.cuda.is_available() and wants_cuda + else "numpy" + ) + + if backend == "numpy": + histogram = ( + orient_hist.detach().cpu().numpy() + if is_torch_input + else orient_hist + ) + output = _calculate_numpy( + histogram, + pair_indices, + geometry, + num_theta, + radius_max, + dtype=dtype, + mode_batch_size=mode_batch_size, + pair_batch_size=pair_batch_size, + workers=workers, + zero_policy=zero_policy, + progress_bar=progress_bar, + ) + else: + output = _calculate_torch( + orient_hist, + pair_indices, + geometry, + num_theta, + radius_max, + device=device, + dtype=dtype, + mode_batch_size=mode_batch_size, + pair_batch_size=pair_batch_size, + max_memory_fraction=max_memory_fraction, + zero_policy=zero_policy, + progress_bar=progress_bar, + ) + if return_numpy: + output = output.detach().cpu().numpy() + + return output, pair_indices diff --git a/src/quantem/diffraction/peak_detection.py b/src/quantem/diffraction/peak_detection.py new file mode 100644 index 000000000..bd3bdc6cb --- /dev/null +++ b/src/quantem/diffraction/peak_detection.py @@ -0,0 +1,752 @@ +import torch +import torch.nn.functional as F +import numpy as np +from typing import List, Tuple +from scipy.spatial import cKDTree +from scipy.ndimage import gaussian_filter, maximum_filter, grey_dilation, map_coordinates +from quantem.core.datastructures import Vector + + +def visualize_blobs(image: np.ndarray, blobs: np.ndarray): + """Visualize detected blobs""" + import matplotlib.pyplot as plt + from matplotlib.patches import Circle + + fig, ax = plt.subplots(1, 1, figsize=(10, 10)) + ax.imshow(image, cmap='gray') + + for y, x, r in blobs: + circle = Circle((x, y), r, color='red', fill=False, linewidth=2) + ax.add_patch(circle) + + ax.set_title(f'Detected {len(blobs)} blobs') + plt.show() + +def detect_blobs(image, sigma=1.0, threshold=None): + """ + Detect strict local maxima (greater than 8 nearest neighbors) with subpixel quadratic refinement. + + Parameters: + ----------- + image : 2D array + sigma : float, for Gaussian smoothing + threshold : float or None, minimum intensity for peak to be valid + + Returns: + -------- + peaks : Nx2 array of (row, col) subpixel coordinates + intensities : N array of signal intensities for peak position + success : N array of booleans (True if refinement succeeded) + """ + + smoothed = gaussian_filter(image, sigma=sigma) + local_max = maximum_filter(smoothed, size=3) + # Make strict: exclude plateaus by checking inequality with neighbors + # Use erosion to get image of maximum value in kernel convolution. + # Used to check if strictly greater than nearest 8 + # Footprint to exclude center pixel. Evaluates nearest 8. + footprint = np.array([[1, 1, 1], + [1, 0, 1], + [1, 1, 1]], dtype=bool) + max_neighbors = grey_dilation(smoothed, footprint=footprint) + peaks = (smoothed == local_max) & (smoothed > max_neighbors) + + # Remove borders and apply threshold + peaks[:, 0] = peaks[:, -1] = peaks[0, :] = peaks[-1, :] = False + if threshold is not None: + peaks &= (smoothed > threshold) + + # Get integer coordinates + peak_coords = np.argwhere(peaks) + # If no peaks, return empty lists + if len(peak_coords) == 0: + return np.array([]), np.array([]), np.array([]) + + # Subpixel refinement + refined_coords, success = refine_peaks_quadratic(smoothed, peak_coords) + # Get intensities of peak position signal + intensities = map_coordinates(smoothed, refined_coords.T, order=1) + + return refined_coords, intensities, success + +def refine_peaks_quadratic(smoothed, peak_coords): + """ + Refine peak positions to subpixel accuracy using 2D quadratic fitting. + + Parameters: + ----------- + smoothed : 2D array, image after Gaussian smoothing + peak_coords : Nx2 array of (row, col) integer peak positions + + Returns: + -------- + refined_coords : Nx2 array of (row, col) subpixel peak positions + success : N array of booleans, True if refinement succeeded + """ + refined = [] + success = [] + + for y, x in peak_coords: + # Skip peaks too close to border (need 3x3 neighborhood) + if y < 1 or y >= smoothed.shape[0]-1 or x < 1 or x >= smoothed.shape[1]-1: + refined.append([float(y), float(x)]) + success.append(False) + continue + + # Get 3x3 neighborhood + patch = smoothed[y-1:y+2, x-1:x+2] + + # Taylor expansion around the peak: + # f(x+dx, y+dy) ≈ f(x,y) + g·[dx,dy] + 0.5·[dx,dy]·H·[dx,dy] + # where g is gradient (1st power) and H is Hessian (2nd power) + + # First derivatives (gradient) using central differences + dy = (patch[2, 1] - patch[0, 1]) / 2.0 + dx = (patch[1, 2] - patch[1, 0]) / 2.0 + + # Second derivatives (Hessian) using finite differences + dyy = patch[2, 1] - 2*patch[1, 1] + patch[0, 1] + dxx = patch[1, 2] - 2*patch[1, 1] + patch[1, 0] + dxy = (patch[2, 2] - patch[2, 0] - patch[0, 2] + patch[0, 0]) / 4.0 + + # Build Hessian matrix + H = np.array([[dyy, dxy], + [dxy, dxx]]) + + # Gradient vector + g = np.array([dy, dx]) + + # At the peak, gradient should be zero: g + H·offset = 0 + # So: offset = -H^(-1)·g + try: + # Check if Hessian is negative definite (proper maximum) + eigenvalues = np.linalg.eigvalsh(H) + if np.all(eigenvalues < 0): # Both eigenvalues negative = local maximum + offset = -np.linalg.solve(H, g) + + # Sanity check: offset shouldn't be too large + # (if it is, the quadratic approximation is probably bad and should just use integer coords) + if np.all(np.abs(offset) <= 1.5): + refined.append([y + offset[0], x + offset[1]]) + success.append(True) + else: + # Offset too large, use integer position, as more accurate + refined.append([float(y), float(x)]) + success.append(False) + else: + # Not a proper maximum (saddle point or minimum) + refined.append([float(y), float(x)]) + success.append(False) + + except np.linalg.LinAlgError: + # Singular matrix (flat region), use integer position + refined.append([float(y), float(x)]) + success.append(False) + + return np.array(refined), np.array(success) + +def pair_peaks(peaks_experimental, peaks_reference, radius_max): + """ + Pair experimental Bragg peaks with reference peaks. + + Parameters: + - peaks_experimental: np.array, shape (n, 2) for n experimental peaks + - peaks_reference: np.array, shape (m, 2) for m reference peaks + - radius_max: float, maximum distance for a match + + Returns: + - matches: list of tuples (exp_index, ref_index, distance) + - unmatched_exp: list of indices of unmatched experimental peaks + """ + # Create KD-Tree for efficient nearest neighbor search + tree = cKDTree(peaks_reference) + + # Find nearest neighbors for all experimental peaks + distances, indices = tree.query(peaks_experimental, distance_upper_bound=radius_max) + + matches = [] + unmatched_exp = [] + + for exp_index, (dist, ref_index) in enumerate(zip(distances, indices)): + if dist <= radius_max: + matches.append((exp_index, ref_index, dist)) + else: + unmatched_exp.append(exp_index) + + return matches, unmatched_exp + +def angle_difference(angle1, angle2): + """Calculate the smallest difference between two angles in degrees with ML model coordinate system.""" + return np.mod(angle1 - angle2 + 180, 360) - 180 + +def pair_peaks_polar(peaks_experimental, peaks_reference, radius_max, angle_max=180, central_radius_threshold=5, filter_central_beam=False): + """ + Pair experimental Bragg peaks with reference peaks in polar coordinates. + + Parameters: + - peaks_experimental: np.array, shape (n, 2) for n experimental peaks (r, theta in degrees) + - peaks_reference: np.array, shape (m, 2) for m reference peaks (r, theta in degrees) + - radius_max: float, maximum radial distance for a match + - angle_max: float, maximum angular difference for a match (in degrees) + - central_radius_threshold: float, radius below which angles are ignored for matching + - filter_central_beam: bool, if True return central beam info + + Returns: + - matches: list of tuples (exp_index, ref_index, distance, delta_r, delta_phi) + - unmatched_exp: list of indices of unmatched experimental peaks + - unmatched_ref: list of indices of unmatched reference peaks + - central_beam_info_exp: dict with keys 'exp_index', 'match_index' (ref_index if matched), 'in_unmatched_exp' + - central_beam_info_ref: dict with keys 'ref_index', 'match_index' (exp_index if matched), 'in_unmatched_ref' + """ + matches = [] + unmatched_exp = list(range(len(peaks_experimental))) + unmatched_ref = list(range(len(peaks_reference))) + + # Find the central beams (smallest radius in both experimental and reference peaks) + central_beam_exp_index = np.argmin(peaks_experimental[:, 0]) if len(peaks_experimental) > 0 else None + central_beam_ref_index = np.argmin(peaks_reference[:, 0]) if len(peaks_reference) > 0 else None + + central_beam_info_exp = { + 'exp_index': central_beam_exp_index, + 'match_index': None, # ref_index if matched + 'in_unmatched_exp': None, # Index in unmatched_exp list if unmatched + } + central_beam_info_ref = { + 'ref_index': central_beam_ref_index, + 'match_index': None, # exp_index if matched + 'in_unmatched_ref': None, # Index in unmatched_ref list if unmatched + } + + for ref_index in unmatched_ref.copy(): + ref_peak = peaks_reference[ref_index] + best_match = None + best_distance = float('inf') + + for exp_index in unmatched_exp.copy(): + exp_peak = peaks_experimental[exp_index] + + delta_r = exp_peak[0] - ref_peak[0] + delta_phi = angle_difference(exp_peak[1], ref_peak[1]) + delta_x = exp_peak[0] * np.cos(exp_peak[1] * np.pi/180) - ref_peak[0] * np.cos(ref_peak[1] * np.pi/180) + delta_y = exp_peak[0] * np.sin(exp_peak[1] * np.pi/180) - ref_peak[0] * np.sin(ref_peak[1] * np.pi/180) + + # Check if either peak is within the central radius threshold + if exp_peak[0] <= central_radius_threshold or ref_peak[0] <= central_radius_threshold: + # For central peaks, only consider radial distance + distance = np.sqrt(delta_x**2 + delta_y**2) + else: + # Use a combination of radial and angular difference for matching + distance = np.sqrt(delta_x**2 + delta_y**2) + + if distance < radius_max and distance < best_distance and np.abs(delta_phi) < angle_max: + best_match = (exp_index, ref_index, distance, delta_r, delta_phi, delta_x, delta_y) + best_distance = distance + + if best_match: + # Check if this match involves the experimental central beam + if best_match[0] == central_beam_exp_index: + central_beam_info_exp['match_index'] = best_match[1] # Store the ref_index + + # Check if this match involves the reference central beam + if best_match[1] == central_beam_ref_index: + central_beam_info_ref['match_index'] = best_match[0] # Store the exp_index + + matches.append(best_match) + unmatched_exp.remove(best_match[0]) + unmatched_ref.remove(best_match[1]) + + # Update central beam info for unmatched cases + if central_beam_exp_index in unmatched_exp: + central_beam_info_exp['in_unmatched_exp'] = unmatched_exp.index(central_beam_exp_index) + + if central_beam_ref_index in unmatched_ref: + central_beam_info_ref['in_unmatched_ref'] = unmatched_ref.index(central_beam_ref_index) + + if filter_central_beam: + return matches, unmatched_exp, unmatched_ref, central_beam_info_exp, central_beam_info_ref + else: + return matches, unmatched_exp, unmatched_ref + +# def pair_peaks_polar(peaks_experimental, peaks_reference, radius_max, angle_max=180, central_radius_threshold=5, filter_central_beam=False): +# """ +# Pair experimental Bragg peaks with reference peaks in polar coordinates. + +# Parameters: +# - peaks_experimental: np.array, shape (n, 2) for n experimental peaks (r, theta in degrees) +# - peaks_reference: np.array, shape (m, 2) for m reference peaks (r, theta in degrees) +# - radius_max: float, maximum radial distance for a match +# - angle_max: float, maximum angular difference for a match (in degrees) +# - central_radius_threshold: float, radius below which angles are ignored for matching + +# Returns: +# - matches: list of tuples (exp_index, ref_index, distance, delta_r, delta_phi) +# - unmatched_exp: list of indices of unmatched experimental peaks +# - unmatched_ref: list of indices of unmatched reference peaks +# """ +# matches = [] +# unmatched_exp = list(range(len(peaks_experimental))) +# unmatched_ref = list(range(len(peaks_reference))) + +# for ref_index in unmatched_ref.copy(): +# ref_peak = peaks_reference[ref_index] +# best_match = None +# best_distance = float('inf') + +# for exp_index in unmatched_exp.copy(): +# exp_peak = peaks_experimental[exp_index] + +# delta_r = exp_peak[0] - ref_peak[0] +# delta_phi = angle_difference(exp_peak[1], ref_peak[1]) +# delta_x = exp_peak[0] * np.cos(exp_peak[1] * np.pi/180) - ref_peak[0] * np.cos(ref_peak[1] * np.pi/180) +# delta_y = exp_peak[0] * np.sin(exp_peak[1] * np.pi/180) - ref_peak[0] * np.sin(ref_peak[1] * np.pi/180) +# # Check if either peak is within the central radius threshold +# if exp_peak[0] <= central_radius_threshold or ref_peak[0] <= central_radius_threshold: +# # For central peaks, only consider radial distance +# # distance = abs(delta_r) +# distance = np.sqrt(delta_x**2 + delta_y**2) +# else: +# # Use a combination of radial and angular difference for matching +# distance = np.sqrt(delta_x**2 + delta_y**2) +# # distance = np.sqrt((delta_r / radius_max)**2 + (delta_phi / angle_max)**2) + +# if distance < radius_max and distance < best_distance and np.abs(delta_phi) < angle_max: # '1' represents a normalized distance threshold +# best_match = (exp_index, ref_index, distance, delta_r, delta_phi, delta_x, delta_y) +# best_distance = distance + +# if best_match: +# matches.append(best_match) +# unmatched_exp.remove(best_match[0]) +# unmatched_ref.remove(best_match[1]) + +# return matches, unmatched_exp, unmatched_ref + + +def get_peak_intensity_from_image(peak_coord, image, radius=2): + """ + Get average intensity in a circular region around a peak. + + Parameters: + ----------- + peak_coord : tuple or array + Peak coordinate (y, x) + image : ndarray + Original diffraction pattern + radius : int + Radius in pixels for sampling region + + Returns: + -------- + intensity : float + Average intensity in the circular region + """ + y, x = peak_coord + h, w = image.shape + + # Create coordinate grids + y_grid, x_grid = np.ogrid[:h, :w] + + # Calculate distance from peak + distances = np.sqrt((y_grid - y)**2 + (x_grid - x)**2) + + # Create circular mask + mask = distances <= radius + + # Get average intensity in the circular region + if np.sum(mask) > 0: + intensity = np.mean(image[mask]) + else: + # Fallback: just use the pixel value at the peak + intensity = image[int(np.clip(y, 0, h-1)), int(np.clip(x, 0, w-1))] + + return intensity + + +def find_central_beam_from_peaks(peak_coords, peak_intensities, image_shape, + intensity_threshold=0.5, distance_weight=0.3, + debug=False, image=None, sampling_radius=2, + vector_x_field=['x_pixels', 'x'], + vector_y_field=['y_pixels', 'y']): + """ + Find central beam from detected peaks with debugging visualization. + + Parameters: + ----------- + peak_coords : ndarray, shape (N, 2) or (N, 4), or Vector + Peak coordinates (y, x) or Vector with coordinate fields + peak_intensities : ndarray, shape (N,) or None + Peak intensities from model (ignored if image is provided) + image_shape : tuple + Shape of image (H, W) + intensity_threshold : float + Minimum intensity to consider (0-1) + distance_weight : float + Weight for distance vs intensity (0=only intensity, 1=only distance) + debug : bool + Show debug plots and print info + image : ndarray, optional + Original diffraction pattern for intensity sampling + sampling_radius : int + Radius in pixels for sampling intensity around each peak + vector_x_field : str or list of str + Field name(s) for x-coordinates. Default: ['x_pixels', 'x'] + vector_y_field : str or list of str + Field name(s) for y-coordinates. Default: ['y_pixels', 'y'] + + Returns: + -------- + center : tuple + (y, x) coordinates of central beam + """ + # Helper function to find first matching field + def find_field(field_options, available_fields): + fields = [field_options] if isinstance(field_options, str) else field_options + return next((f for f in fields if f in available_fields), None) + + # Check if None type passed (indicates no entries in Vector FieldView) + if peak_coords is None: + if debug: + print("⚠️ No peaks! Using image center.") + return (image_shape[0] / 2, image_shape[1] / 2) + + # Handle Vector input + if isinstance(peak_coords, Vector): + if debug: + print(f"Vector fields: {peak_coords.fields}") + + x_field = find_field(vector_x_field, peak_coords.fields) + y_field = find_field(vector_y_field, peak_coords.fields) + + if not (x_field and y_field): + raise ValueError( + f"Missing fields in Vector. Available: {peak_coords.fields}\n" + f"Looking for x in {[vector_x_field] if isinstance(vector_x_field, str) else vector_x_field}, " + f"y in {[vector_y_field] if isinstance(vector_y_field, str) else vector_y_field}" + ) + + if debug: + print(f"Using x='{x_field}', y='{y_field}'") + + vector_data = peak_coords.flatten() + if len(vector_data) > 0: + y_idx = peak_coords.fields.index(y_field) + x_idx = peak_coords.fields.index(x_field) + peak_coords = np.column_stack( + [vector_data[:, y_idx], vector_data[:, x_idx]] + ) + else: + peak_coords = np.empty((0, 2)) + # Handle ndarray input + elif isinstance(peak_coords, np.ndarray): + if peak_coords.ndim == 2 and peak_coords.shape[1] == 4: + if debug: + print("ndarray with 4 columns, using first 2 (y, x)") + peak_coords = peak_coords[:, :2] + elif peak_coords.ndim == 2 and peak_coords.shape[1] == 2: + pass # Already correct + elif peak_coords.ndim == 1 and len(peak_coords) == 0: + peak_coords = np.empty((0, 2)) + else: + raise ValueError(f"Array must be (N, 2) or (N, 4), got {peak_coords.shape}") + else: + raise TypeError(f"peak_coords must be Vector or ndarray, got {type(peak_coords)}") + + # Check for empty peaks + if len(peak_coords) == 0: + if debug: + print("⚠️ No peaks! Using image center.") + return (image_shape[0] / 2, image_shape[1] / 2) + + # Image center + center_y, center_x = image_shape[0] / 2, image_shape[1] / 2 + + # Determine which intensities to use + if image is not None: + # Sample intensities from actual diffraction pattern + sampled_intensities = np.array([ + get_peak_intensity_from_image(coord, image, radius=sampling_radius) + for coord in peak_coords + ]) + intensities_to_use = sampled_intensities + intensity_source = f"Sampled from DP (radius={sampling_radius}px)" + else: + # Use model-predicted intensities + intensities_to_use = peak_intensities + intensity_source = "Model predictions" + + # Normalize intensities to [0, 1] + max_intensity = np.max(intensities_to_use) + if max_intensity > 0: + intensities_norm = intensities_to_use / max_intensity + else: + intensities_norm = intensities_to_use + + if debug: + print(f"\n{'='*60}") + print(f"DEBUG: Central Beam Detection") + print(f"{'='*60}") + print(f"Number of peaks detected: {len(peak_coords)}") + print(f"Image shape: {image_shape}") + print(f"Image center: ({center_y:.1f}, {center_x:.1f})") + print(f"Intensity source: {intensity_source}") + print(f"Intensity threshold: {intensity_threshold}") + print(f"Distance weight: {distance_weight}") + print(f"\nAll peaks:") + for i, (coord, intensity, intensity_norm) in enumerate(zip(peak_coords, intensities_to_use, intensities_norm)): + print(f" Peak {i}: coord=({coord[0]:.2f}, {coord[1]:.2f}), " + f"intensity={intensity:.4f}, normalized={intensity_norm:.4f}") + + # Filter by intensity threshold + intensity_mask = intensities_norm > intensity_threshold + num_above_threshold = np.sum(intensity_mask) + + if debug: + print(f"\nPeaks above intensity threshold ({intensity_threshold}): {num_above_threshold}/{len(peak_coords)}") + + if num_above_threshold == 0: + if debug: + print("⚠️ No peaks above intensity threshold! Using all peaks.") + intensity_mask = np.ones(len(intensities_norm), dtype=bool) + + filtered_coords = peak_coords[intensity_mask] + filtered_intensities = intensities_to_use[intensity_mask] + filtered_intensities_norm = intensities_norm[intensity_mask] + + if debug: + print(f"\nFiltered peaks ({len(filtered_coords)}):") + for i, (coord, intensity, intensity_norm) in enumerate(zip(filtered_coords, filtered_intensities, filtered_intensities_norm)): + print(f" Peak {i}: coord=({coord[0]:.2f}, {coord[1]:.2f}), " + f"intensity={intensity:.4f}, normalized={intensity_norm:.4f}") + + # Calculate distance from image center + distances = np.sqrt( + (filtered_coords[:, 0] - center_y)**2 + + (filtered_coords[:, 1] - center_x)**2 + ) + + if debug: + print(f"\nDistances from center:") + for i, dist in enumerate(distances): + print(f" Peak {i}: {dist:.2f} pixels") + + # Normalize distances + if np.max(distances) > 0: + distances_norm = distances / np.max(distances) + else: + distances_norm = distances + + if debug: + print(f"\nNormalized values:") + print(f" Distance range: [{np.min(distances_norm):.3f}, {np.max(distances_norm):.3f}]") + print(f" Intensity range: [{np.min(filtered_intensities_norm):.3f}, {np.max(filtered_intensities_norm):.3f}]") + + # Score: high intensity, low distance wins + # Lower score is better + scores = (1 - filtered_intensities_norm) * (1 - distance_weight) + distances_norm * distance_weight + + if debug: + print(f"\nScores (lower is better):") + print(f" Formula: (1 - intensity_norm) * {1-distance_weight:.2f} + distance_norm * {distance_weight:.2f}") + for i, score in enumerate(scores): + print(f" Peak {i}: score={score:.4f} " + f"[intensity_term={(1-filtered_intensities_norm[i])*(1-distance_weight):.4f}, " + f"distance_term={distances_norm[i]*distance_weight:.4f}]") + + # Pick peak with best score + best_idx = np.argmin(scores) + central_beam_coords = filtered_coords[best_idx] + + # Map back to original peak index for reference + original_indices = np.where(intensity_mask)[0] + original_best_idx = original_indices[best_idx] + + if debug: + print(f"\n{'='*60}") + print(f"SELECTED CENTRAL BEAM:") + print(f" Peak index (filtered): {best_idx}") + print(f" Peak index (original): {original_best_idx}") + print(f" Coordinates: ({central_beam_coords[0]:.2f}, {central_beam_coords[1]:.2f})") + print(f" Intensity: {filtered_intensities[best_idx]:.4f}") + print(f" Normalized intensity: {filtered_intensities_norm[best_idx]:.4f}") + print(f" Distance from center: {distances[best_idx]:.2f} pixels") + print(f" Score: {scores[best_idx]:.4f}") + print(f"{'='*60}\n") + + # Visualization + if debug and image is not None: + import matplotlib.pyplot as plt + from matplotlib.patches import Circle + + fig, axes = plt.subplots(2, 2, figsize=(16, 14)) + + # Top-left: Diffraction pattern with sampling circles + ax = axes[0, 0] + ax.imshow(image, cmap='viridis') + ax.axhline(center_y, color='white', linestyle='--', alpha=0.5, linewidth=1, label='Image center') + ax.axvline(center_x, color='white', linestyle='--', alpha=0.5, linewidth=1) + + # Draw sampling circles for all peaks + for i, coord in enumerate(peak_coords): + circle = Circle((coord[1], coord[0]), sampling_radius, + fill=False, edgecolor='cyan', linewidth=1, alpha=0.5) + ax.add_patch(circle) + ax.text(coord[1] + sampling_radius + 2, coord[0], f'{i}', + color='cyan', fontsize=8, alpha=0.7) + + # Highlight filtered peaks + for i, coord in enumerate(filtered_coords): + circle = Circle((coord[1], coord[0]), sampling_radius, + fill=False, edgecolor='yellow', linewidth=2, alpha=0.8) + ax.add_patch(circle) + + # Highlight selected central beam + circle = Circle((central_beam_coords[1], central_beam_coords[0]), sampling_radius, + fill=False, edgecolor='red', linewidth=3) + ax.add_patch(circle) + ax.scatter(central_beam_coords[1], central_beam_coords[0], + s=500, c='red', marker='*', + edgecolors='yellow', linewidths=3, zorder=10, label='Selected central beam') + + ax.set_title(f'Diffraction Pattern with Sampling Circles (radius={sampling_radius}px)', + fontsize=12, fontweight='bold') + ax.set_xlabel('X (pixels)') + ax.set_ylabel('Y (pixels)') + ax.legend(loc='upper right', fontsize=9) + + # Top-right: Peaks colored by sampled intensity + ax = axes[0, 1] + ax.imshow(image, cmap='viridis', alpha=0.6) + + scatter = ax.scatter(filtered_coords[:, 1], filtered_coords[:, 0], + c=filtered_intensities, s=300, cmap='hot', marker='o', + edgecolors='black', linewidths=2, label='Filtered peaks') + + ax.scatter(central_beam_coords[1], central_beam_coords[0], + s=500, c='red', marker='*', + edgecolors='yellow', linewidths=3, label='Selected central beam', + zorder=10) + + plt.colorbar(scatter, ax=ax, label='Sampled Intensity') + ax.set_title('Peaks Colored by Sampled Intensity', fontsize=12, fontweight='bold') + ax.legend() + + # Bottom-left: Score visualization + ax = axes[1, 0] + ax.imshow(image, cmap='viridis', alpha=0.6) + + scatter = ax.scatter(filtered_coords[:, 1], filtered_coords[:, 0], + c=scores, s=300, cmap='RdYlGn_r', marker='o', + edgecolors='black', linewidths=2, + vmin=0, vmax=1, label='Filtered peaks (by score)') + + ax.scatter(central_beam_coords[1], central_beam_coords[0], + s=500, c='red', marker='*', + edgecolors='yellow', linewidths=3, label='Selected central beam', + zorder=10) + + # Add score labels + for i, (coord, score) in enumerate(zip(filtered_coords, scores)): + ax.annotate(f'{i}\n{score:.2f}', + xy=(coord[1], coord[0]), + xytext=(10, 10), textcoords='offset points', + fontsize=8, color='white', + bbox=dict(boxstyle='round,pad=0.3', facecolor='black', alpha=0.7)) + + plt.colorbar(scatter, ax=ax, label='Score (lower = better)') + ax.set_title('Peaks Colored by Score', fontsize=12, fontweight='bold') + ax.legend() + + # Bottom-right: Score breakdown + ax = axes[1, 1] + + x = np.arange(len(filtered_coords)) + width = 0.35 + + intensity_component = (1 - filtered_intensities_norm) * (1 - distance_weight) + distance_component = distances_norm * distance_weight + + bars1 = ax.bar(x - width/2, intensity_component, width, + label=f'Intensity term (weight={1-distance_weight:.2f})', + alpha=0.8, color='steelblue') + bars2 = ax.bar(x + width/2, distance_component, width, + label=f'Distance term (weight={distance_weight:.2f})', + alpha=0.8, color='coral') + + # Highlight selected peak + bars1[best_idx].set_color('darkblue') + bars1[best_idx].set_edgecolor('yellow') + bars1[best_idx].set_linewidth(3) + bars2[best_idx].set_color('darkred') + bars2[best_idx].set_edgecolor('yellow') + bars2[best_idx].set_linewidth(3) + + ax.set_xlabel('Peak Index (filtered)', fontsize=12) + ax.set_ylabel('Score Component', fontsize=12) + ax.set_title('Score Breakdown by Component', fontsize=14, fontweight='bold') + ax.set_xticks(x) + ax.legend(fontsize=10) + ax.grid(axis='y', alpha=0.3) + + # Add total score line + ax.plot(x, scores, 'ko-', linewidth=2, markersize=8, + label='Total score', zorder=5) + ax.scatter([best_idx], [scores[best_idx]], s=300, c='red', + marker='*', edgecolors='yellow', linewidths=2, + zorder=10, label='Selected') + ax.legend(fontsize=9) + + plt.tight_layout() + plt.show() + + # Additional info table + fig, ax = plt.subplots(1, 1, figsize=(12, 6)) + + peak_info = [] + for i in range(len(filtered_coords)): + peak_info.append({ + 'Peak': i, + 'Y': f"{filtered_coords[i, 0]:.1f}", + 'X': f"{filtered_coords[i, 1]:.1f}", + 'Intensity': f"{filtered_intensities[i]:.4f}", + 'Norm Int': f"{filtered_intensities_norm[i]:.3f}", + 'Distance': f"{distances[i]:.1f}", + 'Score': f"{scores[i]:.4f}", + 'Selected': '★' if i == best_idx else '' + }) + + # Create table + table_data = [[info[key] for key in ['Peak', 'Y', 'X', 'Intensity', 'Norm Int', 'Distance', 'Score', 'Selected']] + for info in peak_info] + + table = ax.table(cellText=table_data, + colLabels=['Peak', 'Y', 'X', 'Intensity', 'Norm Int', 'Distance', 'Score', ''], + cellLoc='center', + loc='center', + bbox=[0, 0, 1, 1]) + + table.auto_set_font_size(False) + table.set_fontsize(10) + table.scale(1, 2) + + # Color header + for j in range(8): + table[(0, j)].set_facecolor('#4472C4') + table[(0, j)].set_text_props(weight='bold', color='white') + + # Highlight selected row + for i in range(len(peak_info)): + if i == best_idx: + for j in range(8): + table[(i+1, j)].set_facecolor('#ffff99') + table[(i+1, j)].set_text_props(weight='bold') + + ax.axis('off') + title_text = f'Peak Summary ({intensity_source})\n' + title_text += f'distance_weight={distance_weight}, intensity_threshold={intensity_threshold}' + if image is not None: + title_text += f', sampling_radius={sampling_radius}px' + ax.set_title(title_text, fontsize=12, fontweight='bold', pad=20) + + plt.tight_layout() + plt.show() + + return (float(central_beam_coords[0]), float(central_beam_coords[1])) diff --git a/src/quantem/diffraction/polar_transform.py b/src/quantem/diffraction/polar_transform.py new file mode 100644 index 000000000..cbb9b5a89 --- /dev/null +++ b/src/quantem/diffraction/polar_transform.py @@ -0,0 +1,997 @@ +"""Karen Ehrhardt-derived polar transforms and angular-uniformity origin finding.""" + +from __future__ import annotations + +import warnings +from typing import Literal + +import numpy as np +import torch +import torch.nn.functional as F +from numpy.typing import NDArray +from tqdm import tqdm + + +OriginMethod = Literal["descent", "grid"] + +__all__ = [ + "OriginMethod", + "find_origin", + "find_origin_angular_descent", + "find_origin_angular_grid", + "polar_transform", + "polar_transform_peaks", +] + + +def find_origin( + data, + *, + method: OriginMethod = "descent", + ellipse_params: tuple[float, float, float] | None = None, + radial_min: float = 4.0, + radial_max: float | None = None, + radial_step: float = 1.0, + num_annular_bins: int = 180, + n_phi: int = 120, + two_fold_rotation_symmetry: bool = False, + kpow: float = 0.0, + device: str = "cpu", + batch_size: int = 16, + local_margin: int = 40, +) -> NDArray: + """Estimate diffraction-pattern origins as ``(scan_y, scan_x, 2)`` row/col pixels.""" + if method == "descent": + return find_origin_angular_descent( + data, + ellipse_params=ellipse_params, + radial_min=radial_min, + radial_max=radial_max, + n_phi=n_phi, + radial_step=radial_step, + kpow=kpow, + device=device, + ) + if method == "grid": + return find_origin_angular_grid( + data, + ellipse_params=ellipse_params, + num_annular_bins=num_annular_bins, + radial_min=radial_min, + radial_max=radial_max, + radial_step=radial_step, + two_fold_rotation_symmetry=two_fold_rotation_symmetry, + device=device, + batch_size=batch_size, + local_margin=local_margin, + ) + raise ValueError(f"method must be 'descent' or 'grid', got {method!r}.") + + +def polar_transform( + data, + origin_array: NDArray | torch.Tensor | None = None, + ellipse_params: tuple[float, float, float] | None = None, + num_annular_bins: int = 180, + radial_min: float = 0.0, + radial_max: float | None = None, + radial_step: float = 1.0, + two_fold_rotation_symmetry: bool = False, + name: str | None = None, + signal_units: str | None = None, + scan_pos: tuple[int, int] | None = None, + device: str = "cpu", + batch_size: int = 128, + show_progress: bool = True, +): + """Torch-native polar transform ported from Karen Ehrhardt's PDF workflow. + + The returned :class:`Polar4dstem` stores data as ``(scan_y, scan_x, phi, r)``. + ``two_fold_rotation_symmetry=True`` follows Karen's native behavior: sample + directly over ``0..pi``. Callers that need summed Friedel partners should + sample the full plane and fold explicitly. + """ + from quantem.core.datastructures.polar4dstem import Polar4dstem + + array, scan_y, scan_x, n_row, n_col = _as_4d_array(data) + + if isinstance(origin_array, torch.Tensor): + origin_array = origin_array.detach().cpu().numpy() + origin_array = np.asarray(origin_array, dtype=float) if origin_array is not None else None + if origin_array is None: + center = np.array([(n_row - 1) / 2.0, (n_col - 1) / 2.0], dtype=float) + origins = np.broadcast_to(center, (scan_y, scan_x, 2)).copy() + elif origin_array.shape == (2,): + origins = np.empty((scan_y, scan_x, 2), dtype=float) + origins[...] = origin_array + elif origin_array.shape == (scan_y, scan_x, 2): + origins = origin_array + else: + raise ValueError( + f"origin_array must have shape None, (2,), or {(scan_y, scan_x, 2)}, " + f"got {origin_array.shape}." + ) + + if scan_pos is not None: + iy, ix = scan_pos + dp = torch.as_tensor(array[iy, ix], dtype=torch.float32, device=device) + r0 = float(origins[iy, ix, 0]) + c0 = float(origins[iy, ix, 1]) + radial_max_eff = _resolve_radial_max( + n_row, n_col, origins[iy : iy + 1, ix : ix + 1], radial_min, radial_max, radial_step + ) + offset_row, offset_col, _, _ = _build_polar_sampling_offsets( + ellipse_params, + num_annular_bins, + radial_min, + radial_max_eff, + radial_step, + two_fold_rotation_symmetry, + device, + ) + col_norm = 2.0 * (offset_col + c0) / (n_col - 1) - 1.0 + row_norm = 2.0 * (offset_row + r0) / (n_row - 1) - 1.0 + grid = torch.stack([col_norm, row_norm], dim=-1).unsqueeze(0) + polar2d = F.grid_sample( + dp[None, None], + grid, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + ) + return polar2d.squeeze(0).squeeze(0).cpu().numpy() + + radial_max_eff = _resolve_radial_max( + n_row, n_col, origins, radial_min, radial_max, radial_step + ) + offset_row, offset_col, phi_bins, radial_bins = _build_polar_sampling_offsets( + ellipse_params, + num_annular_bins, + radial_min, + radial_max_eff, + radial_step, + two_fold_rotation_symmetry, + device, + ) + n_phi = phi_bins.numel() + n_r = radial_bins.numel() + + col_norm_scale = 2.0 / (n_col - 1) + row_norm_scale = 2.0 / (n_row - 1) + base_col_norm = offset_col * col_norm_scale + base_row_norm = offset_row * row_norm_scale + + n_pos = scan_y * scan_x + dp_view = torch.as_tensor(array.reshape(n_pos, n_row, n_col), dtype=torch.float32) + origins_t = torch.as_tensor(origins.reshape(n_pos, 2), dtype=torch.float32, device=device) + out = torch.empty((n_pos, n_phi, n_r), dtype=torch.float32, device=device) + + for start in tqdm( + range(0, n_pos, batch_size), + desc="Polar transform", + disable=(not show_progress) or n_pos < 8, + ): + end = min(start + batch_size, n_pos) + row_origins = origins_t[start:end, 0] + col_origins = origins_t[start:end, 1] + grid_col = base_col_norm.unsqueeze(0) + (col_origins * col_norm_scale - 1.0)[:, None, None] + grid_row = base_row_norm.unsqueeze(0) + (row_origins * row_norm_scale - 1.0)[:, None, None] + grids = torch.stack([grid_col, grid_row], dim=-1) + dp_batch = dp_view[start:end].to(device=device, dtype=torch.float32) + polars = F.grid_sample( + dp_batch.unsqueeze(1), + grids, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + ) + out[start:end] = polars.squeeze(1) + + out_np = out.reshape(scan_y, scan_x, n_phi, n_r).cpu().numpy() + phi_range = np.pi if two_fold_rotation_symmetry else 2.0 * np.pi + phi_step_deg = (phi_range / float(n_phi)) * (180.0 / np.pi) + + sampling = np.zeros(4, dtype=float) + origin = np.zeros(4, dtype=float) + sampling[0:2] = np.asarray(getattr(data, "sampling", np.ones(4)))[0:2] + sampling[2] = phi_step_deg + sampling[3] = float(np.asarray(getattr(data, "sampling", np.ones(4)))[-1]) * radial_step + origin[0:2] = np.asarray(getattr(data, "origin", np.zeros(4)))[0:2] + origin[2] = 0.0 + origin[3] = radial_min * float(np.asarray(getattr(data, "sampling", np.ones(4)))[-1]) + units_in = list(getattr(data, "units", ["pixels", "pixels", "pixels", "pixels"])) + metadata = dict(getattr(data, "metadata", {})) + metadata.update( + { + "polar_radial_min": float(radial_min), + "polar_radial_max": float(radial_max_eff), + "polar_radial_step": float(radial_step), + "polar_num_annular_bins": int(n_phi), + "polar_two_fold_rotation_symmetry": bool(two_fold_rotation_symmetry), + "polar_origin_row": float(origins[0, 0, 0]), + "polar_origin_col": float(origins[0, 0, 1]), + "polar_ellipse_params": tuple(ellipse_params) if ellipse_params is not None else None, + } + ) + return Polar4dstem( + array=out_np, + name=name if name is not None else f"{getattr(data, 'name', 'dataset')}_polar", + origin=origin, + sampling=sampling, + units=[units_in[0], units_in[1], "deg", units_in[-1]], + signal_units=signal_units if signal_units is not None else getattr(data, "signal_units", "arb. units"), + metadata=metadata, + _token=Polar4dstem._token, + ) + + +def polar_transform_peaks( + cartesian_vector, + centers: NDArray, + *, + scan_mask: NDArray | None = None, + x_field: str | list[str] = ["x_pixels", "x"], + y_field: str | list[str] = ["y_pixels", "y"], + sampling_conversion_factor: float | None = None, + two_fold_rotation_symmetry: bool = True, + ellipse_params: tuple[float, float, float] | None = None, + r_unit: str = "pixels", + theta_unit: str = "radians", + name_suffix: str = "_polar", + use_tqdm: bool = True, +): + """Transform Cartesian peak coordinates with Karen's polar convention. + + Peaks remain one-to-one with the input rows. Under two-fold symmetry, partner + peaks are folded to the same theta coordinate with ``theta % pi`` but are not + aggregated. + """ + from quantem.core.datastructures import Vector + + if isinstance(cartesian_vector, np.ndarray) and cartesian_vector.dtype == object: + cartesian_vector = cartesian_vector.item() + if not isinstance(cartesian_vector, Vector): + raise TypeError(f"Expected Vector, got {type(cartesian_vector)}") + + def find_field(field_options, available_fields): + fields = [field_options] if isinstance(field_options, str) else field_options + return next((f for f in fields if f in available_fields), None) + + x_field_found = find_field(x_field, cartesian_vector.fields) + y_field_found = find_field(y_field, cartesian_vector.fields) + if x_field_found is None or y_field_found is None: + raise ValueError( + "Could not find x/y coordinate fields in Vector. " + f"Available fields: {cartesian_vector.fields}" + ) + + n_scan_y, n_scan_x = cartesian_vector.shape + centers = _standardize_centers(centers, n_scan_y, n_scan_x) + if scan_mask is None: + scan_mask = np.ones((n_scan_y, n_scan_x), dtype=bool) + else: + scan_mask = np.asarray(scan_mask, dtype=bool) + if scan_mask.shape != (n_scan_y, n_scan_x): + raise ValueError(f"scan_mask shape {scan_mask.shape} must match {(n_scan_y, n_scan_x)}") + if sampling_conversion_factor is None: + sampling_conversion_factor = 1.0 + + x_idx = cartesian_vector.fields.index(x_field_found) + y_idx = cartesian_vector.fields.index(y_field_found) + extra_indices = [ + idx for idx in range(len(cartesian_vector.fields)) + if idx not in (x_idx, y_idx) + ] + # ``theta`` is folded when two_fold_rotation_symmetry is set, which maps each + # Friedel pair onto one angle. ``theta_unfolded`` keeps the full 0-2pi angle so + # that information is not lost; peaks are small, so the extra column is cheap. + output_fields = ["r_pixels", "theta", "r_invA", "theta_unfolded"] + [ + cartesian_vector.fields[idx] for idx in extra_indices + ] + output_units = [r_unit, theta_unit, "1/Å", theta_unit] + [ + cartesian_vector.units[idx] for idx in extra_indices + ] + polar_vector = Vector.from_shape( + shape=(n_scan_y, n_scan_x), + fields=output_fields, + units=output_units, + name=cartesian_vector.name + name_suffix, + ) + + theta_period = np.pi if two_fold_rotation_symmetry else 2.0 * np.pi + iterator = tqdm(range(n_scan_y), disable=not use_tqdm, desc="Polar transform peaks") + for i in iterator: + for j in range(n_scan_x): + if not scan_mask[i, j]: + polar_vector[i, j] = np.zeros((0, len(output_fields))) + continue + + cartesian_data = cartesian_vector[i, j].array + if len(cartesian_data) == 0: + polar_vector[i, j] = np.zeros((0, len(output_fields))) + continue + + center_y, center_x = centers[i, j] + dx = cartesian_data[:, x_idx] - center_x + dy = cartesian_data[:, y_idx] - center_y + r_pixels, theta = _cartesian_offsets_to_polar(dx, dy, ellipse_params) + theta_unfolded = np.mod(theta, 2.0 * np.pi) + theta = np.mod(theta, theta_period) + r_invA = r_pixels * sampling_conversion_factor + + polar_data = np.column_stack([r_pixels, theta, r_invA, theta_unfolded]) + if extra_indices: + polar_data = np.column_stack([polar_data, cartesian_data[:, extra_indices]]) + polar_vector[i, j] = polar_data + + return polar_vector + + +def find_origin_angular_grid( + data, + *, + ellipse_params: tuple[float, float, float] | None = None, + num_annular_bins: int = 180, + radial_min: float = 4.0, + radial_max: float | None = None, + radial_step: float = 2.0, + two_fold_rotation_symmetry: bool = False, + device: str = "cpu", + batch_size: int = 16, + local_margin: int = 40, +) -> NDArray: + """Coarse-to-fine angular-variance origin finder. + + This is a surgical port of Karen Ehrhardt's PDF center finder. It first finds + a global center on the mean diffraction pattern, then refines each scan + position by minimizing angular intensity variation in a polar annulus. + """ + array, scan_y, scan_x, n_row, n_col = _as_4d_array(data) + array_t = torch.as_tensor(array, dtype=torch.float32, device=device) + + mean_dp_t = array_t.mean(dim=(0, 1)) + total_intensity = mean_dp_t.clamp(min=0).sum() + 1e-9 + row_grid_t = torch.arange(n_row, dtype=torch.float32, device=device)[:, None] + col_grid_t = torch.arange(n_col, dtype=torch.float32, device=device)[None, :] + com_row = int(round(float(((row_grid_t * mean_dp_t.clamp(min=0)).sum() / total_intensity).item()))) + com_col = int(round(float(((col_grid_t * mean_dp_t.clamp(min=0)).sum() / total_intensity).item()))) + + com_edge_budget = min(com_row, com_col, (n_row - 1) - com_row, (n_col - 1) - com_col) + global_margin = int(min(40, max(2, com_edge_budget // 2))) + safe_radial_max = float( + min( + com_row - global_margin, + (n_row - 1) - (com_row + global_margin), + com_col - global_margin, + (n_col - 1) - (com_col + global_margin), + ) + ) + if radial_max is not None: + safe_radial_max = min(safe_radial_max, float(radial_max)) + if safe_radial_max <= radial_min: + safe_radial_max = radial_min + radial_step + + safe_low = int(np.ceil(safe_radial_max)) + safe_high_row = n_row - 1 - safe_low + safe_high_col = n_col - 1 - safe_low + search_n_phi = max(18, min(int(num_annular_bins), 60)) + local_coarse_step = 5 + + offset_row, offset_col, _, radial_bins = _build_polar_sampling_offsets( + ellipse_params, + search_n_phi, + radial_min, + safe_radial_max, + radial_step, + two_fold_rotation_symmetry, + device, + ) + n_r = radial_bins.numel() + min_r_idx = 0 + max_r_idx = max(1, int(np.ceil(0.9 * n_r))) + col_norm_scale = 2.0 / (n_col - 1) + row_norm_scale = 2.0 / (n_row - 1) + base_col_norm = offset_col * col_norm_scale + base_row_norm = offset_row * row_norm_scale + + mean_dp_batch = mean_dp_t[None, None] + rows, cols, grids = _build_candidate_grids( + base_col_norm, + base_row_norm, + com_row, + com_col, + global_margin, + n_row, + n_col, + col_norm_scale, + row_norm_scale, + device, + step=2, + ) + scores = _angular_std_scores(mean_dp_batch, grids, min_r_idx, max_r_idx) + valid = ( + (rows >= safe_low) & (rows <= safe_high_row) & (cols >= safe_low) & (cols <= safe_high_col) + ) + best = int(scores.masked_fill(~valid, float("inf")).argmin().item()) + coarse_row, coarse_col = int(rows[best].item()), int(cols[best].item()) + + rows, cols, grids = _build_candidate_grids( + base_col_norm, + base_row_norm, + coarse_row, + coarse_col, + 10, + n_row, + n_col, + col_norm_scale, + row_norm_scale, + device, + step=1, + ) + scores = _angular_std_scores(mean_dp_batch, grids, min_r_idx, max_r_idx) + valid = ( + (rows >= safe_low) & (rows <= safe_high_row) & (cols >= safe_low) & (cols <= safe_high_col) + ) + best = int(scores.masked_fill(~valid, float("inf")).argmin().item()) + global_row, global_col = int(rows[best].item()), int(cols[best].item()) + + coarse_rows, coarse_cols, coarse_grids = _build_candidate_grids( + base_col_norm, + base_row_norm, + global_row, + global_col, + int(local_margin), + n_row, + n_col, + col_norm_scale, + row_norm_scale, + device, + step=local_coarse_step, + ) + coarse_valid = ( + (coarse_rows >= safe_low) + & (coarse_rows <= safe_high_row) + & (coarse_cols >= safe_low) + & (coarse_cols <= safe_high_col) + ) + n_coarse = coarse_grids.shape[0] + med_search_range = torch.arange( + -local_coarse_step, local_coarse_step + 1, 1, dtype=torch.long, device=device + ) + med_drow, med_dcol = ( + m.reshape(-1) for m in torch.meshgrid(med_search_range, med_search_range, indexing="ij") + ) + fine_search_range = torch.arange(-2, 3, dtype=torch.long, device=device) + fine_drow, fine_dcol = ( + m.reshape(-1) for m in torch.meshgrid(fine_search_range, fine_search_range, indexing="ij") + ) + flat_dps_t = array_t.reshape(-1, n_row, n_col) + n_pos = flat_dps_t.shape[0] + origin_flat_t = torch.zeros(n_pos, 2, dtype=torch.float32, device=device) + + def refine(dp_batch, current_row, current_col, drow, dcol): + n_cands = drow.numel() + cand_rows = (current_row[:, None] + drow[None, :]).clamp(0, n_row - 1) + cand_cols = (current_col[:, None] + dcol[None, :]).clamp(0, n_col - 1) + g_col = ( + base_col_norm + (cand_cols.reshape(-1).float() * col_norm_scale - 1.0)[:, None, None] + ) + g_row = ( + base_row_norm + (cand_rows.reshape(-1).float() * row_norm_scale - 1.0)[:, None, None] + ) + grids = torch.stack([g_col, g_row], dim=-1) + dps = dp_batch.repeat_interleave(n_cands, dim=0) + polars = F.grid_sample( + dps, grids, mode="bilinear", padding_mode="zeros", align_corners=True + ) + region = polars.view(dp_batch.shape[0], n_cands, *base_col_norm.shape)[ + ..., min_r_idx:max_r_idx + ] + scores = region.std(dim=2).sum(dim=2) / (region.mean(dim=2).sum(dim=2).abs() + 1e-6) + valid = ( + (cand_rows >= safe_low) + & (cand_rows <= safe_high_row) + & (cand_cols >= safe_low) + & (cand_cols <= safe_high_col) + ) + best = scores.masked_fill(~valid, float("inf")).argmin(dim=1) + best_row = cand_rows.gather(1, best[:, None]).squeeze(1) + best_col = cand_cols.gather(1, best[:, None]).squeeze(1) + return best_row, best_col, scores, valid + + n_not_converged = 0 + pbar = tqdm(total=n_pos, desc="Finding origins", disable=n_pos < 8) + for start in range(0, n_pos, batch_size): + end = min(start + batch_size, n_pos) + n_dp = end - start + dp_b = flat_dps_t[start:end].unsqueeze(1) + polars_coarse = F.grid_sample( + dp_b.transpose(0, 1).expand(n_coarse, n_dp, n_row, n_col), + coarse_grids, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + ) + region_coarse = polars_coarse[:, :, :, min_r_idx:max_r_idx] + scores_coarse = region_coarse.std(dim=2).sum(dim=2) / ( + region_coarse.mean(dim=2).sum(dim=2).abs() + 1e-6 + ) + scores_coarse = scores_coarse.masked_fill(~coarse_valid[:, None], float("inf")) + best_coarse = scores_coarse.argmin(dim=0) + current_row, current_col = coarse_rows[best_coarse], coarse_cols[best_coarse] + current_row, current_col, _, _ = refine( + dp_b, current_row, current_col, med_drow, med_dcol + ) + best_row, best_col, fine_scores, fine_valid = refine( + dp_b, current_row, current_col, fine_drow, fine_dcol + ) + + side = fine_search_range.numel() + scores_grid = fine_scores.view(n_dp, side, side) + valid_grid = fine_valid.view(n_dp, side, side) + flat_best = scores_grid.masked_fill(~valid_grid, float("inf")).view(n_dp, -1).argmin(dim=1) + i_star, j_star = flat_best // side, flat_best % side + batch_idx = torch.arange(n_dp, device=device) + ii = torch.stack([i_star - 1, i_star, i_star + 1], dim=1).clamp(0, side - 1) + jj = torch.stack([j_star - 1, j_star, j_star + 1], dim=1).clamp(0, side - 1) + patch = scores_grid[batch_idx[:, None, None], ii[:, :, None], jj[:, None, :]] + on_border = (i_star < 1) | (i_star > side - 2) | (j_star < 1) | (j_star > side - 2) + n_not_converged += int(on_border.sum()) + offset = _quadratic_subpixel_offset(patch).to(torch.float32) + offset = torch.where(on_border[:, None], torch.zeros_like(offset), offset) + origin_flat_t[start:end, 0] = best_row.to(torch.float32) + offset[:, 0] + origin_flat_t[start:end, 1] = best_col.to(torch.float32) + offset[:, 1] + pbar.update(n_dp) + pbar.close() + + if n_not_converged: + warnings.warn( + f"find_origin_angular_grid: {n_not_converged} of {n_pos} scan positions did not " + "bracket a sub-pixel minimum. Integer-pixel origins were used there.", + stacklevel=2, + ) + return origin_flat_t.cpu().numpy().reshape(scan_y, scan_x, 2) + + +def find_origin_angular_descent( + data, + *, + ellipse_params: tuple[float, float, float] | None = None, + radial_min: float = 4.0, + radial_max: float | None = None, + n_phi: int = 120, + radial_step: float = 1.0, + kpow: float = 0.0, + device: str = "cpu", +) -> NDArray: + """COM-anchored local descent origin finder. + + The score is the normalized angular standard deviation in a polar annulus. + Lower scores indicate a more radially uniform transform and therefore a + better center. This method is fast enough to use by default in notebooks. + """ + array, scan_y, scan_x, n_row, n_col = _as_4d_array(data) + if radial_max is None: + radial_max = float(min(n_row, n_col) // 2 - 2) + if radial_max <= radial_min: + radial_max = float(radial_min + max(radial_step, 1.0)) + n_radial = max(4, int(round((radial_max - radial_min) / radial_step)) + 1) + + array_t = torch.as_tensor(array, dtype=torch.float32, device=device) + patterns = array_t.reshape(-1, n_row, n_col) + n_patterns = patterns.shape[0] + image_center = torch.tensor( + [(n_row - 1) / 2.0, (n_col - 1) / 2.0], + dtype=torch.float32, + device=device, + ) + blank_patterns = patterns.clamp(min=0).sum(dim=(1, 2)) <= 0 + if bool(blank_patterns.all().item()): + return ( + image_center[None] + .expand(n_patterns, 2) + .reshape(scan_y, scan_x, 2) + .cpu() + .numpy() + ) + offset_row, offset_col, ring_weights = _local_sampling( + radial_min, radial_max, n_phi, n_radial, kpow, ellipse_params, device + ) + + mean_pattern = array_t.mean(dim=(0, 1)) + global_origin = _descend_batched( + mean_pattern[None], + torch.round(_com_anchor(mean_pattern))[None], + offset_row, + offset_col, + ring_weights, + n_phi, + device, + )[0] + start_centers = torch.round(global_origin)[None].expand(n_patterns, 2).clone() + origins = _descend_batched( + patterns, + start_centers, + offset_row, + offset_col, + ring_weights, + n_phi, + device, + ) + origins = torch.where(blank_patterns[:, None], image_center[None], origins) + return origins.reshape(scan_y, scan_x, 2).cpu().numpy() + + +def _as_4d_array(data) -> tuple[NDArray, int, int, int, int]: + array = np.asarray(data.array if hasattr(data, "array") else data) + if array.ndim == 2: + n_row, n_col = array.shape + array = array[None, None] + return np.ascontiguousarray(array), 1, 1, n_row, n_col + if array.ndim == 4: + scan_y, scan_x, n_row, n_col = array.shape + return np.ascontiguousarray(array), scan_y, scan_x, n_row, n_col + raise ValueError( + f"Expected a 2D diffraction pattern or 4D-STEM array, got shape {array.shape}." + ) + + +def _standardize_centers(centers, scan_y: int, scan_x: int) -> NDArray: + centers = np.asarray(centers, dtype=float) + if centers.shape == (2,): + out = np.empty((scan_y, scan_x, 2), dtype=float) + out[...] = centers + return out + if centers.shape == (scan_y, scan_x, 2): + return centers + if centers.shape == (2, scan_y, scan_x): + return np.moveaxis(centers, 0, -1) + raise ValueError( + f"centers must have shape (2,), {(scan_y, scan_x, 2)}, " + f"or {(2, scan_y, scan_x)}, got {centers.shape}." + ) + + +def _resolve_radial_max( + n_row: int, + n_col: int, + origins: NDArray, + radial_min: float, + radial_max: float | None, + radial_step: float, +) -> float: + if radial_step <= 0: + raise ValueError(f"radial_step must be > 0, got {radial_step}.") + if radial_max is not None: + radial_max_eff = float(radial_max) + else: + origin_rows = origins[..., 0] + origin_cols = origins[..., 1] + radial_limits = np.minimum.reduce( + [ + origin_rows, + (n_row - 1) - origin_rows, + origin_cols, + (n_col - 1) - origin_cols, + ] + ) + radial_max_eff = float(np.nanmin(radial_limits)) + if not np.isfinite(radial_max_eff) or radial_max_eff <= radial_min: + radial_max_eff = float(radial_min + radial_step) + return radial_max_eff + + +def _cartesian_offsets_to_polar( + dx: NDArray, + dy: NDArray, + ellipse_params: tuple[float, float, float] | None, +) -> tuple[NDArray, NDArray]: + dx = np.asarray(dx, dtype=float) + dy = np.asarray(dy, dtype=float) + if ellipse_params is None: + return np.hypot(dx, dy), np.arctan2(dy, dx) + if len(ellipse_params) != 3: + raise ValueError("ellipse_params must be (a, b, theta_deg).") + + a, b, theta_deg = ellipse_params + theta = np.deg2rad(theta_deg) + cos_t = np.cos(theta) + sin_t = np.sin(theta) + u = dx * cos_t + dy * sin_t + v_prime = -dx * sin_t + dy * cos_t + scaled_u = (b / a) * u + r_pixels = np.hypot(scaled_u, v_prime) + phi = np.arctan2(v_prime, scaled_u) + theta + return r_pixels, phi + + +def _polar_to_cartesian_offsets( + phi: torch.Tensor, + r_pix: torch.Tensor, + ellipse_params: tuple[float, float, float] | None, + device: str = "cpu", +) -> tuple[torch.Tensor, torch.Tensor]: + if ellipse_params is None: + offset_col = r_pix * torch.cos(phi) + offset_row = r_pix * torch.sin(phi) + else: + if len(ellipse_params) != 3: + raise ValueError("ellipse_params must be (a, b, theta_deg).") + a, b, theta_deg = ellipse_params + theta = torch.deg2rad(torch.tensor(theta_deg, dtype=torch.float32, device=device)) + alpha = phi - theta + u = (a / b) * r_pix * torch.cos(alpha) + v_prime = r_pix * torch.sin(alpha) + cos_t = torch.cos(theta) + sin_t = torch.sin(theta) + offset_col = u * cos_t - v_prime * sin_t + offset_row = u * sin_t + v_prime * cos_t + return offset_row, offset_col + + +def _build_polar_sampling_offsets( + ellipse_params: tuple[float, float, float] | None, + num_annular_bins: int, + radial_min: float, + radial_max_eff: float, + radial_step: float, + two_fold_rotation_symmetry: bool, + device: str = "cpu", +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + if radial_step <= 0: + raise ValueError(f"radial_step must be > 0, got {radial_step}.") + if num_annular_bins < 1: + raise ValueError("num_annular_bins must be >= 1.") + + radial_bins = torch.arange( + radial_min, radial_max_eff, radial_step, dtype=torch.float32, device=device + ) + if radial_bins.numel() == 0: + radial_bins = torch.tensor([radial_min], dtype=torch.float32, device=device) + phi_range = torch.pi if two_fold_rotation_symmetry else 2.0 * torch.pi + phi_bins = torch.linspace( + 0.0, phi_range, num_annular_bins + 1, dtype=torch.float32, device=device + )[:-1] + phi_grid, r_pix_grid = torch.meshgrid(phi_bins, radial_bins, indexing="ij") + offset_row, offset_col = _polar_to_cartesian_offsets( + phi_grid, r_pix_grid, ellipse_params, device + ) + return offset_row, offset_col, phi_bins, radial_bins + + +def _build_candidate_grids( + base_col_norm: torch.Tensor, + base_row_norm: torch.Tensor, + center_row: int, + center_col: int, + margin: int, + n_row: int, + n_col: int, + col_norm_scale: float, + row_norm_scale: float, + device: str = "cpu", + step: int = 1, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + rows = torch.arange( + max(0, center_row - margin), + min(n_row, center_row + margin + 1), + step, + dtype=torch.long, + device=device, + ) + cols = torch.arange( + max(0, center_col - margin), + min(n_col, center_col + margin + 1), + step, + dtype=torch.long, + device=device, + ) + row_grid, col_grid = torch.meshgrid(rows, cols, indexing="ij") + row_flat, col_flat = row_grid.reshape(-1), col_grid.reshape(-1) + grid_col = ( + base_col_norm.unsqueeze(0) + (col_flat.float() * col_norm_scale - 1.0)[:, None, None] + ) + grid_row = ( + base_row_norm.unsqueeze(0) + (row_flat.float() * row_norm_scale - 1.0)[:, None, None] + ) + grids = torch.stack([grid_col, grid_row], dim=-1) + return row_flat, col_flat, grids + + +def _angular_std_scores( + dp_batch: torch.Tensor, + grids: torch.Tensor, + min_r_idx: int, + max_r_idx: int, +) -> torch.Tensor: + n = grids.shape[0] + polars = F.grid_sample( + dp_batch.expand(n, -1, -1, -1), + grids, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + ) + region = polars.squeeze(1)[:, :, min_r_idx:max_r_idx] + return region.std(dim=1).sum(dim=1) / (region.mean(dim=1).sum(dim=1).abs() + 1e-6) + + +def _quadratic_subpixel_offset(patch: torch.Tensor) -> torch.Tensor: + device = patch.device + scores_flat = patch.reshape(patch.shape[0], 9).to(torch.float64) + grid = torch.tensor([-1.0, 0.0, 1.0], dtype=torch.float64, device=device) + uu, vv = torch.meshgrid(grid, grid, indexing="ij") + u, v = uu.reshape(9), vv.reshape(9) + basis = torch.stack([torch.ones_like(u), u, v, u * u, v * v, u * v], dim=1) + fit_matrix = torch.linalg.pinv(basis) + _, b, c, d, e, f = (scores_flat @ fit_matrix.T).unbind(dim=1) + det = 4.0 * d * e - f * f + valid = (2.0 * d > 0) & (det > 1e-12) & torch.isfinite(scores_flat).all(dim=1) + det_safe = torch.where(valid, det, torch.ones_like(det)) + drow = torch.where(valid, (f * c - 2.0 * e * b) / det_safe, torch.zeros_like(det)) + dcol = torch.where(valid, (f * b - 2.0 * d * c) / det_safe, torch.zeros_like(det)) + return torch.stack([drow.clamp(-1.0, 1.0), dcol.clamp(-1.0, 1.0)], dim=1) + + +_NEIGHBOR_STEPS_8 = [ + [1.0, 0.0], + [-1.0, 0.0], + [0.0, 1.0], + [0.0, -1.0], + [1.0, 1.0], + [1.0, -1.0], + [-1.0, 1.0], + [-1.0, -1.0], +] +_PATCH_OFFSETS_3X3 = [ + [-1.0, -1.0], + [-1.0, 0.0], + [-1.0, 1.0], + [0.0, -1.0], + [0.0, 0.0], + [0.0, 1.0], + [1.0, -1.0], + [1.0, 0.0], + [1.0, 1.0], +] + + +def _com_anchor(pattern: torch.Tensor) -> torch.Tensor: + n_row, n_col = pattern.shape + clipped = pattern.clamp(min=0) + total_raw = clipped.sum() + if float(total_raw.item()) <= 0: + return torch.tensor( + [(n_row - 1) / 2.0, (n_col - 1) / 2.0], + dtype=torch.float32, + device=pattern.device, + ) + total = total_raw + 1e-9 + rows = torch.arange(n_row, device=pattern.device, dtype=torch.float32) + cols = torch.arange(n_col, device=pattern.device, dtype=torch.float32) + center_row = (rows[:, None] * clipped).sum() / total + center_col = (cols[None, :] * clipped).sum() / total + return torch.stack([center_row, center_col]) + + +def _local_sampling(radial_min, radial_max, n_phi, n_radial, kpow, ellipse_params, device): + phi = torch.linspace(0, 2 * np.pi, n_phi + 1, device=device)[:-1] + radii = torch.linspace(radial_min, radial_max, n_radial, device=device) + phi_grid, radius_grid = torch.meshgrid(phi, radii, indexing="ij") + offset_row, offset_col = _polar_to_cartesian_offsets( + phi_grid, radius_grid, ellipse_params, device + ) + ring_weights = radii**kpow + return offset_row, offset_col, ring_weights + + +def _local_polar_score(polar_values, valid_mask, n_phi, ring_weights, min_valid_frac): + n_valid = valid_mask.sum(dim=-2).clamp(min=1) + ring_mean = (polar_values * valid_mask).sum(dim=-2) / n_valid + ring_var = (((polar_values - ring_mean.unsqueeze(-2)) ** 2) * valid_mask).sum(dim=-2) / n_valid + ring_std = ring_var.sqrt() + ring_usable = valid_mask.sum(dim=-2) >= (min_valid_frac * n_phi) + weights = ring_weights * ring_usable + usable_weight = weights.sum(dim=-1) + score = (weights * ring_std).sum(dim=-1) / ((weights * ring_mean.abs()).sum(dim=-1) + 1e-6) + score = torch.where(usable_weight > 0, score, torch.full_like(score, float("inf"))) + return score + + +def _local_score_pairs( + patterns, + pattern_index, + centers, + offset_row, + offset_col, + ring_weights, + n_phi, + device, + min_valid_frac=0.5, + chunk=4096, +): + _, n_row, n_col = patterns.shape + ones_image = torch.ones(1, 1, n_row, n_col, device=device) + scores = torch.empty(centers.shape[0], device=device) + for start in range(0, centers.shape[0], chunk): + index = pattern_index[start : start + chunk] + n_chunk = index.shape[0] + center_row = centers[start : start + chunk, 0][:, None, None] + center_col = centers[start : start + chunk, 1][:, None, None] + sample_grid = torch.stack( + [ + 2.0 * (center_col + offset_col[None]) / (n_col - 1) - 1.0, + 2.0 * (center_row + offset_row[None]) / (n_row - 1) - 1.0, + ], + dim=-1, + ) + polar_values = F.grid_sample( + patterns[index][:, None], + sample_grid, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + )[:, 0] + valid_mask = F.grid_sample( + ones_image.expand(n_chunk, 1, n_row, n_col), + sample_grid, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + )[:, 0] > 0.999 + scores[start : start + n_chunk] = _local_polar_score( + polar_values, valid_mask, n_phi, ring_weights, min_valid_frac + ) + return scores + + +def _descend_batched( + patterns, + anchors, + offset_row, + offset_col, + ring_weights, + n_phi, + device, + schedule=(4.0, 2.0, 1.0), + sweeps=2, +): + n_patterns = patterns.shape[0] + pattern_ids = torch.arange(n_patterns, device=device) + neighbor_steps = torch.tensor(_NEIGHBOR_STEPS_8, device=device) + patch_offsets = torch.tensor(_PATCH_OFFSETS_3X3, dtype=torch.float32, device=device) + pattern_ids_per_neighbor = pattern_ids.repeat_interleave(8) + pattern_ids_per_patch = pattern_ids.repeat_interleave(9) + center = anchors.clone() + + def score_at(pattern_index, centers): + return _local_score_pairs( + patterns, + pattern_index, + centers, + offset_row, + offset_col, + ring_weights, + n_phi, + device, + ) + + best_score = score_at(pattern_ids, center) + for step in schedule: + for _ in range(sweeps): + neighbors = (center[:, None, :] + step * neighbor_steps[None]).reshape(n_patterns * 8, 2) + neighbor_scores = score_at(pattern_ids_per_neighbor, neighbors).reshape(n_patterns, 8) + best_neighbor = neighbor_scores.argmin(dim=1) + best_neighbor_score = neighbor_scores.gather(1, best_neighbor[:, None]).squeeze(1) + improved = best_neighbor_score < best_score - 1e-12 + best_neighbor_center = neighbors.reshape(n_patterns, 8, 2)[pattern_ids, best_neighbor] + center = torch.where(improved[:, None], best_neighbor_center, center) + best_score = torch.where(improved, best_neighbor_score, best_score) + + patch_centers = (center[:, None, :] + patch_offsets[None]).reshape(n_patterns * 9, 2) + patch_scores = score_at(pattern_ids_per_patch, patch_centers).reshape(n_patterns, 9) + center = patch_centers.reshape(n_patterns, 9, 2)[pattern_ids, patch_scores.argmin(dim=1)] + patch_centers = (center[:, None, :] + patch_offsets[None]).reshape(n_patterns * 9, 2) + score_patch = score_at(pattern_ids_per_patch, patch_centers).reshape(n_patterns, 3, 3) + subpixel_offset = _quadratic_subpixel_offset(score_patch).to(torch.float32) + return center + subpixel_offset diff --git a/src/quantem/diffraction/polymer_ice.py b/src/quantem/diffraction/polymer_ice.py new file mode 100644 index 000000000..c69363462 --- /dev/null +++ b/src/quantem/diffraction/polymer_ice.py @@ -0,0 +1,1014 @@ +"""Ice-peak detection for polymer diffraction analyses. + +Crystalline ice contaminating a polymer 4D-STEM scan produces six-fold sets of +reflections. Separating them from the polymer signal is awkward because the +strongest ice ring (d ~ 3.66 A) sits on top of the pi-pi stacking peak, so q +alone cannot do it. ``detect_ice`` therefore tests each peak against four +criteria in turn, cheapest first: + +1. **q window** -- within ``dq_invA`` of ``q_target_invA``. +2. **Sharpness** -- radial/annular FWHM measured from the polar volume against + the ``max_width_*`` ceilings. Ice is annularly sharp, both as compact dots + and as radial streaks; polymer at the same q is an annularly broad arc, so + the annular ceiling is the discriminating one. +3. **Intensity** -- at or above ``intensity_cutoff`` (or a scan-wide percentile). +4. **Six-fold geometry** -- the surviving candidates must align to a lattice of + arms 60 degrees apart, with at least ``min_matches`` arms populated. This is + the only criterion that tests structure rather than appearance, and the one + that separates ice from a sharp polymer reflection. + +Sharpness is applied before the geometry search so broad peaks cannot drag the +lattice orientation around. Several passes can run per pattern +(``max_crystallites``) for scans holding crystallites at unrelated orientations. + +Two properties of the input matter throughout. ``process_polar(two_fold_symmetry +=True)`` folds theta to [0, 180), collapsing each Friedel pair onto one angle -- +so only three of the six arms are distinguishable and ``min_matches`` cannot +exceed 3. The unfolded angle survives in the ``theta_unfolded`` field, which +``require_friedel_pair`` uses to demand genuinely opposed spots. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + +import numpy as np +from matplotlib.colors import LogNorm +from scipy.ndimage import gaussian_filter1d +from matplotlib.patches import Rectangle +from numpy.typing import NDArray + +from quantem.core.datastructures import Vector + + +@dataclass(frozen=True) +class IceFlaggerParams: + q_target_invA: float = 1.61 + dq_invA: float = 0.05 + dtheta_deg: float = 6.0 + min_matches: int = 2 + intensity_field: str = "intensities" + intensity_percentile_global: float = 99.0 + intensity_cutoff: float | None = None + intensity_cutoff_mode: Literal["absolute", "percentile"] = "absolute" + conservative: bool = True + + # --- Sharpness gate --------------------------------------------------- + # Ice reflections are annularly sharp -- both the compact dots and the radial + # streaks, which are narrow in theta and extended in q. Polymer signal at the + # same q is an annularly broad arc. Widths are full-width-at-half-maximum + # measured on the polar intensity volume at each candidate's (r, theta). + # + # max_width_theta_deg is the discriminating one: it separates ice from polymer + # for dots and streaks alike. max_width_r_invA is available but rarely useful -- + # ice dots and polymer arcs have similar radial widths, so it separates streaks + # from dots (a distinction within ice) rather than ice from polymer. Ceilings + # are ANDed; None on either leaves that axis ungated, which is the default. + max_width_r_invA: float | None = None + max_width_theta_deg: float | None = None + # Half-width of the search window used to measure each FWHM. A candidate + # whose profile never falls to half maximum inside the window is reported + # as wider than the window, i.e. broad, and is rejected. + sharpness_window_r_invA: float = 0.06 + # Must comfortably exceed the broadest feature you want to measure: a window that the + # feature fills makes the baseline below sit partway up the feature, which drives the + # half-maximum level up and reports the width as far too small. 90 degrees covers the + # whole folded annulus, so nothing saturates. + sharpness_window_theta_deg: float = 90.0 + # Local background level, as a quantile of the windowed profile. The half + # maximum is taken above this, so a peak riding on the amorphous ring is + # measured against the ring rather than against zero. Keep it low: a high + # quantile on a window containing a broad feature reads that feature as + # background and under-reports its width. + sharpness_baseline_quantile: float = 0.05 + # Bins of local argmax refinement, to absorb the sub-bin offset between a + # detected peak position and the polar volume's sampling grid. + sharpness_refine_bins: int = 2 + # Noise robustness. A half-maximum walk that stops at the FIRST sample below + # half is fooled by a weak, diffuse arc: one downward fluctuation ends it a few + # bins out, so a broad noisy feature reports as narrow as a sharp one. The + # profile is Gaussian-smoothed with this sigma (in bins), and the crossing must + # stay below half for `sharpness_crossing_persistence` samples to count. The + # kernel width is then removed in quadrature -- exact for a Gaussian convolved + # with a Gaussian, which is why this is a Gaussian and not a boxcar. Set the + # sigma to 0 and the persistence to 1 for the raw first-crossing behaviour. + sharpness_smooth_sigma_bins: float = 1.0 + sharpness_crossing_persistence: int = 3 + + # --- Multiple crystallites ------------------------------------------- + # One pattern can contain several ice crystallites at unrelated orientations, + # giving overlaid six-fold lattices. Each pass claims the best-supported + # lattice and peels its peaks away before looking again. 1 keeps the previous + # single-lattice behaviour. + max_crystallites: int = 1 + # How far apart two lattices' phi must be, in degrees on the 0-60 wedge, to + # count as separate crystallites. None uses dtheta_deg, i.e. lattices that the + # matcher could not tell apart anyway are not treated as distinct. + min_phi_separation_deg: float | None = None + + def __post_init__(self): + """Reject parameter values that could only ever produce nonsense. + + These are cheap to check here and expensive to diagnose downstream, where + a bad value shows up as "nothing was flagged" rather than as an error. + """ + positive = { + "dq_invA": self.dq_invA, + "dtheta_deg": self.dtheta_deg, + "sharpness_window_r_invA": self.sharpness_window_r_invA, + "sharpness_window_theta_deg": self.sharpness_window_theta_deg, + } + for name, value in positive.items(): + if not value > 0: + raise ValueError(f"{name} must be positive; got {value!r}") + at_least_one = { + "min_matches": self.min_matches, + "min_peaks_per_arm": self.min_peaks_per_arm, + "max_crystallites": self.max_crystallites, + "sharpness_crossing_persistence": self.sharpness_crossing_persistence, + } + for name, value in at_least_one.items(): + if value < 1: + raise ValueError(f"{name} must be at least 1; got {value!r}") + optional_positive = { + "max_width_r_invA": self.max_width_r_invA, + "max_width_theta_deg": self.max_width_theta_deg, + "min_phi_separation_deg": self.min_phi_separation_deg, + "theta_period_deg": self.theta_period_deg, + } + for name, value in optional_positive.items(): + if value is not None and not value > 0: + raise ValueError(f"{name} must be positive when set; got {value!r}") + if not 0.0 <= self.sharpness_baseline_quantile < 1.0: + raise ValueError( + "sharpness_baseline_quantile must be in [0, 1); got " + f"{self.sharpness_baseline_quantile!r}" + ) + if self.sharpness_smooth_sigma_bins < 0: + raise ValueError( + "sharpness_smooth_sigma_bins must be non-negative; got " + f"{self.sharpness_smooth_sigma_bins!r}" + ) + if self.intensity_cutoff_mode not in ("absolute", "percentile"): + raise ValueError( + "intensity_cutoff_mode must be 'absolute' or 'percentile'; got " + f"{self.intensity_cutoff_mode!r}" + ) + # Angular period of the peak thetas, in degrees. process_polar(two_fold_symmetry=True) + # folds theta to [0, 180), collapsing every Friedel pair onto one angle, so only three + # of the six lattice arms are distinguishable and min_matches cannot exceed 3. None lets + # detect_ice read it off the BraggPeaksPolymer; set 360.0 or 180.0 to force it. + theta_period_deg: float | None = None + # Peaks an arm must carry to count towards min_matches. On a folded theta axis a + # Friedel pair (theta and theta+180) lands on one arm as two peaks, while an + # isolated reflection lands as one -- so 2 demands a pair and rejects lone peaks. + # Caveat: two peaks that merely fall within dtheta_deg/dq_invA of each other also + # satisfy it; folding makes them indistinguishable from a true opposed pair. + min_peaks_per_arm: int = 1 + # Demand a genuine opposed pair on each arm: two peaks whose UNFOLDED angles differ + # by 180 +/- dtheta_deg. Unlike min_peaks_per_arm this cannot be satisfied by two + # peaks that merely sit close together, but it needs the "theta_unfolded" field that + # polar_transform_peaks records -- re-run it if polar_peaks predates that field. + require_friedel_pair: bool = False + + +@dataclass(frozen=True) +class IceFlaggerDebug: + n_peaks_total: int + n_candidates_q: int + n_candidates_q_int: int + intensity_threshold_used: float + best_phi_deg: float | None + matched_bins: list[int] + matched_peak_indices: list[int] + n_candidates_sharp: int | None = None + # One entry per crystallite found, strongest first. ``best_phi_deg`` is the + # first of these, retained so single-lattice call sites keep working. + phi_deg: list[float] | None = None + + +@dataclass(frozen=True) +class IceDetectionResult: + """Ice mask and diagnostics produced for an entire scan.""" + + mask: Vector + flagged_peaks_count_map: NDArray[np.integer] + matched_bins_count_map: NDArray[np.integer] + intensity_threshold: float + debug_records: dict[tuple[int, int], IceFlaggerDebug] | None = None + + @property + def ice_mask(self) -> Vector: + """Alias retained for discoverability.""" + return self.mask + + @property + def threshold(self) -> float: + return self.intensity_threshold + + @property + def flagged_count_map(self) -> NDArray[np.integer]: + return self.flagged_peaks_count_map + + @property + def matched_count_map(self) -> NDArray[np.integer]: + return self.matched_bins_count_map + + def filter(self, vector: Vector, *, invert: bool = False) -> Vector: + """Return a filtered copy; the source vector is never mutated.""" + + if vector.shape != self.mask.shape: + raise ValueError( + f"shape mismatch: vector.shape={vector.shape} vs mask.shape={self.mask.shape}" + ) + out = vector.copy() + for iy in range(vector.shape[0]): + for ix in range(vector.shape[1]): + source = vector[iy, ix].array + mask_cell = self.mask[iy, ix].array + if source is None or len(source) == 0 or mask_cell is None: + continue + flags = np.asarray(mask_cell)[:, 0].astype(bool, copy=False) + if len(flags) != len(source): + raise ValueError( + f"Row count mismatch at ({iy},{ix}): vector has {len(source)} " + f"rows but ice mask has {len(flags)}." + ) + out[iy, ix] = source[flags if invert else ~flags] + return out + + +def _smooth_profile(profile: NDArray[np.floating], sigma_bins: float) -> NDArray[np.floating]: + """Gaussian smoothing along a windowed profile, with edge-clamped padding.""" + + if sigma_bins <= 0 or len(profile) < 3: + return profile + return gaussian_filter1d(profile, float(sigma_bins), mode="nearest", truncate=3.0) + + +def _smoothing_fwhm_bins(sigma_bins: float) -> float: + """FWHM of the smoothing kernel, in bins. + + Subtracted in quadrature from the measured width. A Gaussian convolved with a + Gaussian is exactly Gaussian with FWHM = sqrt(w^2 + k^2), so the correction is + exact rather than approximate -- unlike a boxcar, which leaves a residual bias. + """ + + return 0.0 if sigma_bins <= 0 else 2.3548 * float(sigma_bins) + + +def _half_width_bins( + profile: NDArray[np.floating], + center: int, + direction: int, + half: float, + persistence: int = 1, +) -> float: + """Bins from ``center`` to where ``profile`` drops below ``half`` and stays there. + + Requiring the crossing to persist for ``persistence`` samples is what stops a + single noise dip on a broad, weak feature from ending the walk early. Returns + the window half-length when no crossing is found, so a profile that never comes + back down reads as at least as broad as the window. + """ + + n = len(profile) + need = max(1, int(persistence)) + previous = float(profile[center]) + for step in range(1, n): + index = center + direction * step + if index < 0 or index >= n: + return float(step - 1) + value = float(profile[index]) + if not np.isfinite(value) or value <= half: + run = True + for ahead in range(1, need): + probe = center + direction * (step + ahead) + if probe < 0 or probe >= n: + break # window edge: treat the run as sustained + nxt = float(profile[probe]) + if np.isfinite(nxt) and nxt > half: + run = False + break + if run: + span = previous - value + fraction = (previous - half) / span if span > 0 else 0.0 + return (step - 1) + float(np.clip(fraction, 0.0, 1.0)) + previous = value + return float(n) + + +def _profile_fwhm( + profile: NDArray[np.floating], + center: int, + *, + baseline_quantile: float, + refine_bins: int, + smooth_sigma_bins: float = 0.0, + persistence: int = 1, +) -> tuple[float, int]: + """FWHM of ``profile`` in bins about ``center``, plus the refined peak bin. + + The profile is a window already cut out of the polar volume, so running off + its end means "wider than the window" rather than "edge of the detector". + """ + + finite = profile[np.isfinite(profile)] + if not len(finite): + return float("inf"), center + # Smooth first: the argmax refinement below must not latch onto a noise spike, + # which would raise the half-maximum level and end the walk prematurely. + profile = _smooth_profile(profile, smooth_sigma_bins) + finite = profile[np.isfinite(profile)] + if not len(finite): + return float("inf"), center + if refine_bins > 0: + low = max(0, center - refine_bins) + high = min(len(profile), center + refine_bins + 1) + center = low + int(np.nanargmax(profile[low:high])) + peak = float(profile[center]) + baseline = float(np.quantile(finite, baseline_quantile)) + if not np.isfinite(peak) or peak <= baseline: + return float("inf"), center + half = baseline + 0.5 * (peak - baseline) + left = _half_width_bins(profile, center, -1, half, persistence) + right = _half_width_bins(profile, center, +1, half, persistence) + # Remove the smoothing kernel in quadrature so sharp peaks stay unbiased. + measured = left + right + kernel = _smoothing_fwhm_bins(smooth_sigma_bins) + deconvolved = np.sqrt(max(measured**2 - kernel**2, 0.0)) + return float(deconvolved), center + + +def measure_peak_widths( + r_invA, + theta_rad, + polar_intensity: NDArray[np.floating], + r_axis: NDArray[np.floating], + theta_axis: NDArray[np.floating], + *, + params: IceFlaggerParams = IceFlaggerParams(), +) -> tuple[NDArray[np.floating], NDArray[np.floating]]: + """Radial (1/Å) and annular (degrees) FWHM for each peak of one pattern. + + ``polar_intensity`` is that pattern's polar transform, indexed + ``[radial_bin, annular_bin]``; ``r_axis`` and ``theta_axis`` are its + coordinate axes (1/Å and radians). The annular axis is treated as periodic, + the radial axis is not. Peaks that fall outside the sampled radial range + get ``inf``, so they never pass a sharpness ceiling. + """ + + radius = np.asarray(r_invA, dtype=float) + theta = np.asarray(theta_rad, dtype=float) + width_r = np.full(radius.shape, np.inf) + width_theta = np.full(radius.shape, np.inf) + if not radius.size or polar_intensity.size == 0 or len(r_axis) < 2 or len(theta_axis) < 2: + return width_r, width_theta + + r_step = float(r_axis[1] - r_axis[0]) + theta_step_deg = float(np.rad2deg(theta_axis[1] - theta_axis[0])) + n_r, n_theta = polar_intensity.shape + # Window half-widths in bins; at least 2 so a FWHM is measurable at all. + window_r = max(2, int(np.ceil(params.sharpness_window_r_invA / max(r_step, 1e-12)))) + window_theta = max(2, int(np.ceil(params.sharpness_window_theta_deg / max(theta_step_deg, 1e-12)))) + # The annular axis wraps, so a window wider than the circle would repeat bins and let + # the outward walk run back into the peak it started from. + window_theta = min(window_theta, max(1, (n_theta - 1) // 2)) + theta_period = float(theta_axis[-1] - theta_axis[0]) + (theta_axis[1] - theta_axis[0]) + + for index in range(radius.size): + if not (np.isfinite(radius[index]) and np.isfinite(theta[index])): + continue + r_bin = int(np.round((radius[index] - r_axis[0]) / r_step)) + if not 0 <= r_bin < n_r: + continue + theta_bin = int(np.round(np.mod(theta[index], theta_period) / (theta_period / n_theta))) % n_theta + + # Annular cut first: it is periodic, so the window is always full length + # and the refined bin it returns anchors the radial cut. + theta_indices = np.mod(np.arange(theta_bin - window_theta, theta_bin + window_theta + 1), n_theta) + annular = polar_intensity[r_bin, theta_indices] + fwhm_theta, refined = _profile_fwhm( + annular, + window_theta, + baseline_quantile=params.sharpness_baseline_quantile, + refine_bins=params.sharpness_refine_bins, + smooth_sigma_bins=params.sharpness_smooth_sigma_bins, + persistence=params.sharpness_crossing_persistence, + ) + theta_bin = int(theta_indices[min(refined, len(theta_indices) - 1)]) + + low = max(0, r_bin - window_r) + radial = polar_intensity[low : min(n_r, r_bin + window_r + 1), theta_bin] + fwhm_r, _ = _profile_fwhm( + radial, + r_bin - low, + baseline_quantile=params.sharpness_baseline_quantile, + refine_bins=params.sharpness_refine_bins, + smooth_sigma_bins=params.sharpness_smooth_sigma_bins, + persistence=params.sharpness_crossing_persistence, + ) + width_r[index] = fwhm_r * r_step + width_theta[index] = fwhm_theta * theta_step_deg + return width_r, width_theta + + +def _resolve_theta_period(params: IceFlaggerParams, fallback: float | None) -> float: + """Angular period of the peak thetas: explicit params win, then the caller's value.""" + + if params.theta_period_deg is not None: + return float(params.theta_period_deg) + return 360.0 if fallback is None else float(fallback) + + +def _sharpness_enabled(params: IceFlaggerParams) -> bool: + return params.max_width_r_invA is not None or params.max_width_theta_deg is not None + + +def sharpness_mask( + width_r: NDArray[np.floating], width_theta: NDArray[np.floating], params: IceFlaggerParams +) -> NDArray[np.bool_]: + """Which peaks pass the configured width ceilings. + + The ceilings are ANDed; an axis with no ceiling passes everything, so setting + only ``max_width_theta_deg`` gates on annular sharpness alone. Public so a + tuning preview can apply exactly the gate the flagger applies, rather than + reimplementing it. Non-finite widths fail any ceiling that is set. + """ + + radial_ok = ( + np.ones(width_r.shape, dtype=bool) + if params.max_width_r_invA is None + else width_r <= params.max_width_r_invA + ) + annular_ok = ( + np.ones(width_theta.shape, dtype=bool) + if params.max_width_theta_deg is None + else width_theta <= params.max_width_theta_deg + ) + return radial_ok & annular_ok + + +def _angle_distance( + angles: NDArray[np.floating], target: float, period: float = 360.0 +) -> NDArray[np.floating]: + """Separation on a circle of circumference ``period`` degrees. + + ``period`` is 180 when the polar transform folded theta with two-fold + symmetry, which maps every Friedel pair onto a single angle. + """ + + delta = np.abs(np.mod(angles, period) - np.mod(target, period)) + return np.minimum(delta, period - delta) + + +def _has_friedel_pair(unfolded_deg: NDArray[np.floating], tolerance_deg: float) -> bool: + """True when two of these peaks lie 180 degrees apart on the unfolded circle.""" + + finite = unfolded_deg[np.isfinite(unfolded_deg)] + if len(finite) < 2: + return False + # Separation of every ordered pair on the full circle; a Friedel pair is 180 apart. + delta = np.abs(np.mod(finite[:, None], 360.0) - np.mod(finite[None, :], 360.0)) + delta = np.minimum(delta, 360.0 - delta) + return bool(np.any(np.abs(delta - 180.0) <= tolerance_deg)) + + +def _phi_distance(first: float, second: float) -> float: + """Separation of two six-fold orientations, which live on a 0-60 degree wedge.""" + + delta = abs(np.mod(first, 60.0) - np.mod(second, 60.0)) + return float(min(delta, 60.0 - delta)) + + +def _global_threshold( + intensities: Vector, field: str, percentile: float, scan_mask: NDArray[np.bool_] +) -> float: + index = intensities.fields.index(field) + values = [] + for iy, ix in np.argwhere(scan_mask): + cell = intensities[int(iy), int(ix)].array + if cell is not None and len(cell): + finite = np.asarray(cell)[:, index] + finite = finite[np.isfinite(finite)] + if len(finite): + values.append(finite) + return float(np.percentile(np.concatenate(values), percentile)) if values else float("inf") + + +def compute_global_intensity_threshold( + peak_intensities: Vector, + intensity_field: str = "intensities", + percentile: float = 99.0, + scan_mask=None, +) -> float: + """Compute a scan-wide intensity percentile for ice candidate selection.""" + + if intensity_field not in peak_intensities.fields: + raise KeyError( + f"Intensity field {intensity_field!r} is absent from peak_intensities." + ) + selected = ( + np.ones(peak_intensities.shape, dtype=bool) + if scan_mask is None + else np.asarray(scan_mask, dtype=bool) + ) + if selected.shape != peak_intensities.shape: + raise ValueError( + f"scan_mask shape {selected.shape} must match {peak_intensities.shape}." + ) + return _global_threshold( + peak_intensities, intensity_field, percentile, selected + ) + + +def flag_ice_peaks_in_pattern( + r_invA, + theta_rad, + intensities, + *, + params: IceFlaggerParams, + intensity_threshold_global: float, + return_debug: bool = True, + polar_intensity: NDArray[np.floating] | None = None, + r_axis: NDArray[np.floating] | None = None, + theta_axis: NDArray[np.floating] | None = None, + theta_period_deg: float | None = None, + theta_unfolded_rad=None, +): + """Flag peaks belonging to an aligned, possibly incomplete six-fold ice pattern. + + ``polar_intensity`` / ``r_axis`` / ``theta_axis`` are this pattern's polar + transform and its coordinate axes. They are required only when ``params`` + sets a sharpness ceiling, which is measured from that volume. + """ + + radius = np.asarray(r_invA, dtype=float) + theta = np.asarray(theta_rad, dtype=float) + intensity = np.asarray(intensities, dtype=float) + if radius.shape != theta.shape or radius.shape != intensity.shape: + raise ValueError("r_invA, theta_rad, and intensities must have the same shape.") + + q_candidates = np.isfinite(radius) & ( + np.abs(radius - params.q_target_invA) <= params.dq_invA + ) + + # Sharpness gate. Applied to the q band before the six-fold search, so the + # broad polymer peaks neither get flagged nor drag the phi estimate around. + n_sharp = None + if _sharpness_enabled(params): + if polar_intensity is None or r_axis is None or theta_axis is None: + raise ValueError( + "A sharpness ceiling (max_width_r_invA / max_width_theta_deg) requires the " + "polar intensity volume; pass polar_data through detect_ice()." + ) + width_r, width_theta = measure_peak_widths( + radius, theta, polar_intensity, r_axis, theta_axis, params=params + ) + q_candidates &= sharpness_mask(width_r, width_theta, params) + n_sharp = int(np.count_nonzero(q_candidates)) + if params.intensity_cutoff is None: + threshold = float(intensity_threshold_global) + elif params.intensity_cutoff_mode == "absolute": + threshold = float(params.intensity_cutoff) + elif params.intensity_cutoff_mode == "percentile": + finite = intensity[np.isfinite(intensity)] + threshold = ( + float(np.percentile(finite, params.intensity_cutoff)) + if len(finite) + else float("inf") + ) + else: + raise ValueError("intensity_cutoff_mode must be 'absolute' or 'percentile'.") + + candidate_indices = np.flatnonzero( + q_candidates & np.isfinite(intensity) & (intensity >= threshold) + ) + result = np.zeros(radius.shape, dtype=bool) + bins: list[int] = [] + matched: list[int] = [] + phis: list[float] = [] + + # Greedy peel: fit the best-supported six-fold lattice, claim its peaks, remove + # them, and look again in what is left. A pattern can contain several ice + # crystallites at unrelated orientations, and one pass only ever sees the + # strongest. max_crystallites=1 reproduces the single-lattice behaviour. + separation = ( + params.dtheta_deg + if params.min_phi_separation_deg is None + else params.min_phi_separation_deg + ) + # A folded theta axis (period 180) makes only three of the six arms distinguishable, + # because each arm and its Friedel partner share one angle. + period = _resolve_theta_period(params, theta_period_deg) + n_arms = max(1, int(round(period / 60.0))) + unfolded_deg = ( + None if theta_unfolded_rad is None + else np.mod(np.rad2deg(np.asarray(theta_unfolded_rad, dtype=float)), 360.0) + ) + remaining = candidate_indices + for _ in range(max(1, params.max_crystallites)): + if not len(remaining): + break + angles = np.mod(np.rad2deg(theta[remaining]), period) + modulo = np.mod(angles, 60.0) + supports = [ + _angle_distance(modulo, float(center)) <= params.dtheta_deg + for center in modulo + ] + inliers = supports[int(np.argmax([np.count_nonzero(x) for x in supports]))] + radians = np.deg2rad(modulo[inliers]) + phi = float( + np.mod(np.rad2deg(np.arctan2(np.mean(np.sin(radians)), np.mean(np.cos(radians)))), 60) + ) + # A lattice indistinguishable from one already claimed means the leftovers + # are stragglers of it, not a new crystallite. Stop rather than double-count. + if any(_phi_distance(phi, previous) < separation for previous in phis): + break + + expected = phi + 60.0 * np.arange(n_arms) + errors = np.stack( + [_angle_distance(angles, value, period) for value in expected], axis=1 + ) + closest = np.argmin(errors, axis=1) + aligned = errors[np.arange(len(angles)), closest] <= params.dtheta_deg + # Keep only arms carrying enough peaks, then re-restrict the matched set to them. + arm_counts = np.bincount(closest[aligned].astype(int), minlength=n_arms) + good_arms = np.flatnonzero(arm_counts >= params.min_peaks_per_arm) + if len(good_arms) < params.min_matches: + break + aligned &= np.isin(closest, good_arms) + if params.require_friedel_pair: + # Keep only arms holding two peaks genuinely 180 degrees apart. Folding + # cannot tell that from two nearby peaks; the unfolded angle can. + if unfolded_deg is None: + raise ValueError( + "require_friedel_pair needs the 'theta_unfolded' field on polar_peaks. " + "Re-run polar_transform_peaks (or process_polar) to record it." + ) + paired = [ + arm for arm in good_arms + if _has_friedel_pair( + unfolded_deg[remaining[aligned & (closest == arm)]], params.dtheta_deg + ) + ] + if len(paired) < params.min_matches: + break + good_arms = np.asarray(paired, dtype=int) + aligned &= np.isin(closest, good_arms) + lattice_bins = sorted(good_arms.tolist()) + + result[remaining[aligned]] = True + if not params.conservative: + # Sweep in sub-threshold peaks of the q band that sit on this lattice. + q_indices = np.flatnonzero(q_candidates & ~result) + if len(q_indices): + q_angles = np.mod(np.rad2deg(theta[q_indices]), period) + q_errors = np.stack( + [_angle_distance(q_angles, value, period) for value in expected[good_arms]], + axis=1, + ) + result[q_indices[np.min(q_errors, axis=1) <= params.dtheta_deg]] = True + phis.append(phi) + bins.extend(lattice_bins) + remaining = remaining[~aligned] + matched = np.flatnonzero(result).astype(int).tolist() + + debug = IceFlaggerDebug( + n_peaks_total=int(radius.size), + n_candidates_q=int(np.count_nonzero(q_candidates)), + n_candidates_q_int=int(len(candidate_indices)), + intensity_threshold_used=threshold, + best_phi_deg=phis[0] if phis else None, + matched_bins=bins, + matched_peak_indices=matched, + n_candidates_sharp=n_sharp, + phi_deg=phis, + ) + return result, debug if return_debug else None + + +def detect_ice( + polar_peaks: Vector, + peak_intensities: Vector, + *, + params: IceFlaggerParams = IceFlaggerParams(), + scan_mask=None, + intensity_threshold_global: float | None = None, + return_debug: bool = False, + polar_data: dict | None = None, + theta_period_deg: float | None = None, +) -> IceDetectionResult: + """Detect ice peaks across aligned ragged peak and intensity vectors. + + ``polar_data`` is the polar transform dict produced by ``process_polar`` + (keys ``intensity``, ``r_invA``, ``theta``). It is required only when + ``params`` sets a sharpness ceiling. + """ + + if polar_peaks.shape != peak_intensities.shape: + raise ValueError("polar_peaks and peak_intensities must have matching shapes.") + for field in ("r_invA", "theta"): + if field not in polar_peaks.fields: + raise KeyError(f"Required field {field!r} is absent from polar_peaks.") + if params.intensity_field not in peak_intensities.fields: + raise KeyError( + f"Intensity field {params.intensity_field!r} is absent from peak_intensities." + ) + shape = polar_peaks.shape + selected = np.ones(shape, dtype=bool) if scan_mask is None else np.asarray(scan_mask, bool) + if selected.shape != shape: + raise ValueError(f"scan_mask shape {selected.shape} must match {shape}.") + if params.intensity_cutoff is None: + threshold = ( + compute_global_intensity_threshold( + peak_intensities, + intensity_field=params.intensity_field, + percentile=params.intensity_percentile_global, + scan_mask=selected, + ) + if intensity_threshold_global is None + else float(intensity_threshold_global) + ) + elif params.intensity_cutoff_mode == "absolute": + threshold = float(params.intensity_cutoff) + else: + threshold = float("nan") + + # A folded theta axis collapses Friedel pairs, so only period/60 arms are + # distinguishable. Catch an unsatisfiable min_matches here rather than letting + # every pattern silently fail to match. + period = _resolve_theta_period(params, theta_period_deg) + n_arms = max(1, int(round(period / 60.0))) + if params.min_matches > n_arms: + raise ValueError( + f"min_matches={params.min_matches} can never be reached: theta has period " + f"{period:g} degrees, which leaves only {n_arms} distinguishable six-fold arms. " + "process_polar(two_fold_symmetry=True) folds theta to [0, 180), mapping each " + f"Friedel pair onto one angle. Use min_matches <= {n_arms}." + ) + + polar_intensity_stack = r_axis = theta_axis = None + if _sharpness_enabled(params): + if polar_data is None: + raise ValueError( + "A sharpness ceiling (max_width_r_invA / max_width_theta_deg) requires " + "polar_data; run process_polar() first, or clear the ceilings." + ) + polar_intensity_stack = np.asarray(polar_data["intensity"]) + if polar_intensity_stack.shape[:2] != shape: + raise ValueError( + f"polar_data intensity has scan shape {polar_intensity_stack.shape[:2]}, " + f"which must match {shape}." + ) + # process_polar stores the coordinate grids as [radial_bin, annular_bin] meshes. + r_axis = np.asarray(polar_data["r_invA"])[:, 0] + theta_axis = np.asarray(polar_data["theta"])[0, :] + + mask = Vector.from_shape(shape=shape, fields=["is_ice"], units=["bool"], name="ice_peak_mask") + flagged = np.zeros(shape, dtype=int) + matched_bins = np.zeros(shape, dtype=int) + records = {} if return_debug else None + r_index = polar_peaks.fields.index("r_invA") + theta_index = polar_peaks.fields.index("theta") + # Optional: recorded by polar_transform_peaks so folding does not lose the half-circle. + unfolded_index = ( + polar_peaks.fields.index("theta_unfolded") + if "theta_unfolded" in polar_peaks.fields + else None + ) + if params.require_friedel_pair and unfolded_index is None: + raise ValueError( + "require_friedel_pair needs the 'theta_unfolded' field on polar_peaks, which " + "this vector predates. Re-run bp.polar_transform_peaks(...) (cheap) or " + "process_polar(...) to record it, or use min_peaks_per_arm instead." + ) + intensity_index = peak_intensities.fields.index(params.intensity_field) + for iy, ix in np.argwhere(selected): + iy, ix = int(iy), int(ix) + polar_cell = polar_peaks[iy, ix].array + intensity_cell = peak_intensities[iy, ix].array + if polar_cell is None or intensity_cell is None: + continue + if len(polar_cell) != len(intensity_cell): + raise ValueError( + f"Row count mismatch at ({iy},{ix}): polar peaks have {len(polar_cell)} " + f"rows and intensities have {len(intensity_cell)}." + ) + flags, debug = flag_ice_peaks_in_pattern( + np.asarray(polar_cell)[:, r_index], + np.asarray(polar_cell)[:, theta_index], + np.asarray(intensity_cell)[:, intensity_index], + params=params, + intensity_threshold_global=threshold, + return_debug=return_debug, + polar_intensity=None if polar_intensity_stack is None else polar_intensity_stack[iy, ix], + r_axis=r_axis, + theta_axis=theta_axis, + theta_period_deg=period, + theta_unfolded_rad=( + None if unfolded_index is None + else np.asarray(polar_cell)[:, unfolded_index] + ), + ) + if len(flags): + mask[iy, ix] = flags[:, None] + flagged[iy, ix] = np.count_nonzero(flags) + if debug is not None: + matched_bins[iy, ix] = len(debug.matched_bins) + records[(iy, ix)] = debug + return IceDetectionResult(mask, flagged, matched_bins, threshold, records) + + +def plot_q_intensity_density( + polar_peaks: Vector, + peak_intensities: Vector, + *, + q_field="r_invA", + intensity_field="intensities", + q_bins=250, + i_bins=200, + q_range=None, + q_max=0.5, + cutoff=None, + cutoff_mode="absolute", + cutoff_color="cyan", + q_value=None, + q_window=None, + q_value_color="cyan", + q_window_color="cyan", + q_window_alpha=0.18, + q_value_lw=2.0, + q_window_lw=1.5, +): + """Plot q versus intensity density for aligned ragged peak vectors.""" + + import matplotlib.pyplot as plt + + q_index = polar_peaks.fields.index(q_field) + intensity_index = peak_intensities.fields.index(intensity_field) + qs, values = [], [] + for iy in range(polar_peaks.shape[0]): + for ix in range(polar_peaks.shape[1]): + q_cell = polar_peaks[iy, ix].array + i_cell = peak_intensities[iy, ix].array + if q_cell is None or i_cell is None: + continue + if len(q_cell) != len(i_cell): + raise ValueError(f"Row count mismatch at ({iy},{ix}).") + q = np.asarray(q_cell)[:, q_index] + intensity = np.asarray(i_cell)[:, intensity_index] + valid = np.isfinite(q) & np.isfinite(intensity) & (intensity >= 0) & (intensity <= 1) + qs.extend(q[valid]) + values.extend(intensity[valid]) + if not qs: + raise ValueError("No valid (q, intensity) pairs found.") + qs = np.asarray(qs) + values = np.asarray(values) + fig, ax = plt.subplots(figsize=(8, 4)) + histogram = ax.hist2d( + qs, + values, + bins=(q_bins, i_bins), + range=((0, q_max) if q_range is None else q_range, (0, 1)), + norm=LogNorm(), + cmap="magma", + ) + ax.set(xlabel="q (1/Å)", ylabel=intensity_field, ylim=(0, 1)) + fig.colorbar(histogram[3], ax=ax, label="count (log colormap)") + if q_window is not None and q_value is None: + raise ValueError("q_window requires q_value.") + + # Resolve the intensity floor first: it is the bottom edge of the shaded region. + level = label = None + if cutoff is not None: + if cutoff_mode == "absolute": + level, label = float(cutoff), f"cutoff={float(cutoff):.3g}" + elif cutoff_mode == "percentile": + level = float(np.percentile(values, cutoff)) + label = f"p{float(cutoff):g}={level:.3g}" + else: + raise ValueError("cutoff_mode must be 'absolute' or 'percentile'.") + + if q_value is not None: + if q_window is not None: + # Shade the candidate region itself -- the q window, above the intensity + # floor -- rather than drawing bare edge lines. Cyan reads cleanly on magma. + bottom = 0.0 if level is None else level + top = ax.get_ylim()[1] + ax.add_patch( + Rectangle( + (q_value - q_window, bottom), + 2.0 * q_window, + top - bottom, + facecolor=q_window_color, + alpha=q_window_alpha, + edgecolor=q_window_color, + lw=q_window_lw, + zorder=2, + ) + ) + ax.axvline(q_value, color=q_value_color, lw=q_value_lw, ls=":", zorder=3) + + if level is not None: + ax.axhline(level, color=cutoff_color, lw=2, ls="--", zorder=3) + ax.text(ax.get_xlim()[0], level, " " + label, color=cutoff_color, va="bottom", zorder=3) + fig.tight_layout() + return fig, ax + + +def collect_peak_widths( + polar_peaks: Vector, + peak_intensities: Vector, + polar_data: dict, + *, + params: IceFlaggerParams = IceFlaggerParams(), + scan_mask=None, + q_band_only: bool = True, +) -> dict[str, NDArray]: + """Measure every peak's radial/annular width, flattened across the scan. + + This is the tuning counterpart to the sharpness ceilings: histogram + ``width_r_invA`` against ``width_theta_deg`` to see where the sharp ice + population separates from the broad polymer one, then set + ``max_width_r_invA`` / ``max_width_theta_deg`` between them. + + With ``q_band_only`` the measurement is restricted to the flagger's q window, + which is both far cheaper and the only population the gate ever sees. + Returns flat arrays keyed ``iy``, ``ix``, ``q_invA``, ``theta_deg``, + ``intensity``, ``width_r_invA``, ``width_theta_deg``. + """ + + shape = polar_peaks.shape + selected = np.ones(shape, dtype=bool) if scan_mask is None else np.asarray(scan_mask, bool) + if selected.shape != shape: + raise ValueError(f"scan_mask shape {selected.shape} must match {shape}.") + intensity_stack = np.asarray(polar_data["intensity"]) + r_axis = np.asarray(polar_data["r_invA"])[:, 0] + theta_axis = np.asarray(polar_data["theta"])[0, :] + r_index = polar_peaks.fields.index("r_invA") + theta_index = polar_peaks.fields.index("theta") + intensity_index = peak_intensities.fields.index(params.intensity_field) + + columns: dict[str, list] = {key: [] for key in + ("iy", "ix", "q_invA", "theta_deg", "intensity", + "width_r_invA", "width_theta_deg")} + for iy, ix in np.argwhere(selected): + iy, ix = int(iy), int(ix) + polar_cell = polar_peaks[iy, ix].array + intensity_cell = peak_intensities[iy, ix].array + if polar_cell is None or intensity_cell is None or not len(polar_cell): + continue + radius = np.asarray(polar_cell)[:, r_index] + theta = np.asarray(polar_cell)[:, theta_index] + values = np.asarray(intensity_cell)[:, intensity_index] + keep = ( + np.isfinite(radius) & (np.abs(radius - params.q_target_invA) <= params.dq_invA) + if q_band_only + else np.isfinite(radius) + ) + if not keep.any(): + continue + radius, theta, values = radius[keep], theta[keep], values[keep] + width_r, width_theta = measure_peak_widths( + radius, theta, intensity_stack[iy, ix], r_axis, theta_axis, params=params + ) + columns["iy"].extend([iy] * len(radius)) + columns["ix"].extend([ix] * len(radius)) + columns["q_invA"].extend(radius) + columns["theta_deg"].extend(np.rad2deg(theta)) + columns["intensity"].extend(values) + columns["width_r_invA"].extend(width_r) + columns["width_theta_deg"].extend(width_theta) + return {key: np.asarray(value) for key, value in columns.items()} + + +# Compatibility names used by existing analyses. +flag_ice_peaks_in_dataset = detect_ice + + +def apply_ice_mask_to_vector(vector: Vector, ice_mask_vector: Vector, *, invert=False) -> Vector: + result = IceDetectionResult( + ice_mask_vector, + np.zeros(vector.shape, dtype=int), + np.zeros(vector.shape, dtype=int), + float("nan"), + ) + return result.filter(vector, invert=invert) + + +__all__ = [ + "IceDetectionResult", + "IceFlaggerDebug", + "IceFlaggerParams", + "apply_ice_mask_to_vector", + "collect_peak_widths", + "compute_global_intensity_threshold", + "detect_ice", + "flag_ice_peaks_in_dataset", + "flag_ice_peaks_in_pattern", + "measure_peak_widths", + "sharpness_mask", + "plot_q_intensity_density", +] diff --git a/src/quantem/diffraction/polymer_ice_tuning.py b/src/quantem/diffraction/polymer_ice_tuning.py new file mode 100644 index 000000000..b47122c4b --- /dev/null +++ b/src/quantem/diffraction/polymer_ice_tuning.py @@ -0,0 +1,383 @@ +"""Interactive tuning and diagnostics for the polymer ice flagger. + +Plotting and inspection helpers for choosing IceFlaggerParams against a real +dataset. They live here rather than in a notebook so that any notebook, anywhere, +can import them without carrying a copy: + + from quantem.diffraction.polymer_ice_tuning import ( + selection_box, # tune the sharpness ceilings, returns updated params + orientation_histograms, # did the flagger take ice and leave your peaks? + ice_split_widget, # inspect kept / removed peaks interactively + probe_peaks, # per-peak FWHM at one scan position, with the cuts + ) + +Every function takes the BraggPeaksPolymer and an IceFlaggerParams explicitly and +returns its results; none mutate ``bp``. They are free functions rather than more +methods on BraggPeaksPolymer, which is already large -- the algorithms live in +``polymer_ice``, and this module is only the interactive layer over them. + +``ice_split_widget`` needs the optional ``quantem.widget`` package; its import is +deferred into the function so this module stays importable without it. +""" + +from __future__ import annotations + +import dataclasses + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.colors import LogNorm + +from quantem.diffraction.polymer_ice import measure_peak_widths, sharpness_mask + +__all__ = [ + "IcePeakView", + "ice_split_widget", + "orientation_histograms", + "peak_widths", + "probe_peaks", + "selection_box", +] + +_WIDTH_CACHE: dict[tuple, dict] = {} + + +def peak_widths(bp, params, *, use_cache=True): + """Radial/annular FWHM for every peak in the ice q band, across the scan. + + Cached on the q band and intensity field, which are the only inputs that + change the measurement, so sweeping the ceilings costs nothing. + """ + key = (id(bp), params.q_target_invA, params.dq_invA, params.intensity_field, + params.sharpness_window_theta_deg, params.sharpness_smooth_sigma_bins) + if use_cache and key in _WIDTH_CACHE: + return _WIDTH_CACHE[key] + print("measuring peak widths over the q band (cached until the q band changes)...") + widths = bp.measure_ice_peak_widths(params=params) + _WIDTH_CACHE[key] = widths + return widths + + +def selection_box(bp, params, *, max_width_r_invA=None, max_width_theta_deg=None, + fig_path=None, use_cache=True): + """Preview the sharpness gate as a box in (width, intensity), and return the params. + + Assign the result back to your params object; nothing is mutated in place:: + + params_ice_flagger = selection_box(bp, params_ice_flagger, + max_width_theta_deg=8.0) + + The ceilings are ANDed, and an unset axis is not gated. Ice is annularly + sharp -- compact dots and radial streaks alike -- while polymer at the same q + is an annularly broad arc, so ``max_width_theta_deg`` is normally the only one + you need; a radial ceiling mostly rejects the streaks. + """ + tuned = dataclasses.replace( + params, + max_width_r_invA=max_width_r_invA, + max_width_theta_deg=max_width_theta_deg, + ) + widths = peak_widths(bp, tuned, use_cache=use_cache) + width_r = widths["width_r_invA"] + width_theta = widths["width_theta_deg"] + intensity = widths["intensity"] + + # A percentile cutoff is resolved per scan inside detect_ice, so preview it as + # "no floor" rather than guessing the value it will take. + floor = tuned.intensity_cutoff if tuned.intensity_cutoff is not None else -np.inf + floor_label = "none" if not np.isfinite(floor) else f"{floor:.4g}" + selected = sharpness_mask(width_r, width_theta, tuned) & (intensity >= floor) + total = len(width_r) + print(f"box selects {selected.sum()} of {total} q-band peaks " + f"({100 * selected.sum() / max(total, 1):.1f}%)") + + # Where each population sits. If a width's two rows look alike, that axis does + # not separate the populations and its ceiling buys nothing. + bright, dim = intensity >= floor, intensity < floor + quantiles = [0.05, 0.25, 0.5, 0.75, 0.95] + for label, values, digits in (("radial FWHM (1/Å)", width_r, 4), + ("annular FWHM (deg)", width_theta, 2)): + print(f"{label} quantiles {quantiles}") + for group_label, group in ((f"intensity >= {floor_label}", bright), + (f"intensity < {floor_label}", dim)): + finite = group & np.isfinite(values) + print(f" {group_label:<22} (n={finite.sum():>7}):", + np.round(np.quantile(values[finite], quantiles), digits) + if finite.any() else "none") + + fig, ((ax_r, ax_t), (ax_rt, ax_map)) = plt.subplots(2, 2, figsize=(12.5, 9)) + + def width_panel(ax, values, ceiling, xlabel): + finite = np.isfinite(values) + hist = ax.hist2d(values[finite], intensity[finite], bins=(80, 80), + norm=LogNorm(), cmap="magma") + fig.colorbar(hist[3], ax=ax, label="count (log)") + x_lo, x_hi = ax.get_xlim() + y_lo, y_hi = ax.get_ylim() + bottom = y_lo if not np.isfinite(floor) else floor + if np.isfinite(floor): + ax.axhline(floor, color="cyan", ls="--", lw=2) + edge = x_hi if ceiling is None else ceiling + if ceiling is not None: + ax.axvline(ceiling, color="cyan", ls="--", lw=2) + ax.add_patch(plt.Rectangle((x_lo, bottom), edge - x_lo, y_hi - bottom, + facecolor="cyan", alpha=0.15, edgecolor="none")) + ax.set(xlabel=xlabel, ylabel=tuned.intensity_field, + title=f"{xlabel} vs intensity" + + ("" if ceiling is not None else " (no ceiling set)")) + + width_panel(ax_r, width_r, tuned.max_width_r_invA, "radial FWHM (1/Å)") + width_panel(ax_t, width_theta, tuned.max_width_theta_deg, "annular FWHM (deg)") + + # The plane the gate cuts in. Ceilings are ANDed, so the keep-region is the + # corner under both; with no radial ceiling it is the band below the annular one. + plane = np.isfinite(width_r) & np.isfinite(width_theta) & bright + hist = ax_rt.hist2d(width_r[plane], width_theta[plane], bins=(80, 80), + norm=LogNorm(), cmap="magma") + fig.colorbar(hist[3], ax=ax_rt, label="count (log)") + x_lo, x_hi = ax_rt.get_xlim() + y_lo, y_hi = ax_rt.get_ylim() + r_edge = x_hi if tuned.max_width_r_invA is None else tuned.max_width_r_invA + t_edge = y_hi if tuned.max_width_theta_deg is None else tuned.max_width_theta_deg + if tuned.max_width_r_invA is not None: + ax_rt.axvline(tuned.max_width_r_invA, color="cyan", ls="--", lw=2) + if tuned.max_width_theta_deg is not None: + ax_rt.axhline(tuned.max_width_theta_deg, color="cyan", ls="--", lw=2) + ax_rt.add_patch(plt.Rectangle((x_lo, y_lo), r_edge - x_lo, t_edge - y_lo, + facecolor="cyan", alpha=0.15, edgecolor="none")) + ax_rt.set(xlabel="radial FWHM (1/Å)", ylabel="annular FWHM (deg)", + title=f"radial vs annular FWHM (intensity >= {floor_label})") + + # Ice should be compact blobs on the scan, not sprinkled everywhere. + selected_map = np.zeros(bp.polar_peaks.shape, dtype=int) + np.add.at(selected_map, (widths["iy"][selected], widths["ix"][selected]), 1) + image = ax_map.imshow(selected_map, cmap="inferno", interpolation="nearest") + fig.colorbar(image, ax=ax_map, label="selected peaks per position") + ax_map.set_title("selected peaks across the scan") + ax_map.axis("off") + + fig.suptitle(f"selection box — {selected.sum()} of {total} q-band peaks", fontsize=12) + fig.tight_layout(rect=(0, 0, 1, 0.96)) + if fig_path is not None: + fig.savefig(fig_path, format="pdf", bbox_inches="tight") + plt.show() + + print(f"\nreturned params: max_width_r_invA={tuned.max_width_r_invA}, " + f"max_width_theta_deg={tuned.max_width_theta_deg}, " + f"intensity_cutoff={tuned.intensity_cutoff}") + print(f"detect_ice can flag at most these {selected.sum()} peaks; the six-fold " + "alignment then drops those not sitting in an aligned pattern.") + return tuned + + +def orientation_histograms(bp, ice_result, params, *, theta_step_deg=2.0, + orientation_offset_degrees=0.0, fig_path=None): + """Ice-band orientation histograms for all / kept / removed peaks. + + Run BEFORE the filter cell: the removed set needs the unfiltered peaks. + Crystalline ice is a sharp compact blob in the scan image and a narrow spike + in the angular histogram; polymer signal is diffuse and broad. Returns the + three histograms as a dict, keyed 'all peaks' / 'kept (ice removed)' / + 'removed (ice)'. + """ + q_window = (params.q_target_invA - params.dq_invA, + params.q_target_invA + params.dq_invA) + + def histogram(polar_peaks, peak_intensities): + # make_orientation_histogram reads bp.polar_peaks / bp.peak_intensities, + # so swap the subset in, measure, and restore. upsample_factor=1 keeps + # this cheap and at scan resolution. + saved = (bp.polar_peaks, bp.peak_intensities) + bp.polar_peaks, bp.peak_intensities = polar_peaks, peak_intensities + try: + return bp.make_orientation_histogram( + radial_ranges=np.array([q_window]), + upsample_factor=1, theta_step_deg=theta_step_deg, + sigma_x=0.0, sigma_y=0.0, sigma_theta=3.0, + orientation_offset_degrees=orientation_offset_degrees, + normalize_intensity_image=False, normalize_intensity_stack=False, + progress_bar=False)[0] + finally: + bp.polar_peaks, bp.peak_intensities = saved + + hists = { + "all peaks": histogram(bp.polar_peaks, bp.peak_intensities), + "kept (ice removed)": histogram(ice_result.filter(bp.polar_peaks), + ice_result.filter(bp.peak_intensities)), + "removed (ice)": histogram(ice_result.filter(bp.polar_peaks, invert=True), + ice_result.filter(bp.peak_intensities, invert=True)), + } + + theta = np.arange(0, 180, theta_step_deg) + vmax = max(float(h.max()) for h in hists.values()) or 1.0 + ymax = max(float(h.sum(axis=(0, 1)).max()) for h in hists.values()) or 1.0 + fig, axes = plt.subplots(len(hists), 2, figsize=(9, 3.2 * len(hists)), + gridspec_kw={"width_ratios": [1, 1.3]}) + for (label, hist), (ax_map, ax_hist) in zip(hists.items(), np.atleast_2d(axes)): + image = ax_map.imshow(hist.max(axis=2), cmap="inferno", vmin=0, vmax=vmax, + interpolation="nearest") + ax_map.set_title(f"{label} — max over theta", fontsize=9) + ax_map.axis("off") + fig.colorbar(image, ax=ax_map, fraction=0.046) + ax_hist.plot(theta, hist.sum(axis=(0, 1)), lw=1.2) + ax_hist.set(xlim=(0, 180), ylim=(0, 1.05 * ymax), xlabel="theta (deg)", + ylabel="summed intensity") + ax_hist.set_title(f"{label} — angular histogram", fontsize=9) + fig.suptitle(f"Ice band q = {q_window[0]:.3f}–{q_window[1]:.3f} 1/Å " + f"(d = {1 / q_window[1]:.2f}–{1 / q_window[0]:.2f} Å)", fontsize=10) + fig.tight_layout() + if fig_path is not None: + fig.savefig(fig_path, format="pdf", bbox_inches="tight") + plt.show() + return hists + + +class IcePeakView: + """Stand-in for ``bp`` exposing one side of the ice split. + + ``show_polymer_4DSTEM`` is duck-typed and reads the peak vectors live on every + cursor move, so the subset must stay visible for the widget's lifetime and + cannot be restored after construction. Delegating instead of assigning onto + ``bp`` keeps ``bp`` itself pristine. + """ + + def __init__(self, base, cartesian, intensities, polar): + object.__setattr__(self, "_base", base) + object.__setattr__(self, "peak_coordinates_cartesian", cartesian) + object.__setattr__(self, "peak_intensities", intensities) + object.__setattr__(self, "polar_peaks", polar) + + def __getattr__(self, name): + # Reached only for attributes not set above, i.e. everything but the peaks. + return getattr(self._base, name) + + def __setattr__(self, name, value): + # The widget's "Save settings" writes to its source object; keep those + # writes on the view so they never land on bp. + object.__setattr__(self, name, value) + + +def ice_split_widget(bp, ice_result, params, *, view="removed", ice_hists=None, + map_view="match", **widget_kwargs): + """Open the interactive viewer on one side of the ice split. ``bp`` is untouched. + + ``view`` is 'removed' | 'kept' | 'all'. ``map_view`` follows ``view`` by + default; pin it to one of the same names to hold a fixed backdrop while + toggling the overlay. Pass ``ice_hists`` from :func:`orientation_histograms` + to use the ice-band orientation image as the context map. + """ + from quantem.widget import show_polymer_4DSTEM + + peaks = (bp.peak_coordinates_cartesian, bp.peak_intensities, bp.polar_peaks) + if view == "all": + subset = peaks + elif view in ("kept", "removed"): + subset = tuple(ice_result.filter(v, invert=view == "removed") for v in peaks) + else: + raise ValueError(f"view must be 'removed', 'kept' or 'all'; got {view!r}") + + map_key = {"all": "all peaks", "kept": "kept (ice removed)", "removed": "removed (ice)"} + resolved = view if map_view == "match" else map_view + if resolved not in map_key: + raise ValueError(f"map_view must be 'match', 'removed', 'kept' or 'all'; " + f"got {map_view!r}") + if ice_hists is not None: + context_map = ice_hists[map_key[resolved]].max(axis=2) + map_label = f"{resolved} orientation" + else: + context_map = ice_result.flagged_peaks_count_map.astype(float) + map_label = "flagged count — pass ice_hists for the orientation map" + + def count(vector): + total = 0 + for iy in range(vector.shape[0]): + for ix in range(vector.shape[1]): + rows = vector[iy, ix].array + if rows is not None: + total += len(rows) + return total + + print(f"showing '{view}' peaks: {count(subset[0])} of {count(peaks[0])} total " + "(bp itself is untouched)") + # Pass the view, not bp: bp.show_widget() would bind bp as self. + return show_polymer_4DSTEM( + IcePeakView(bp, *subset), + intensity_map=context_map, + title=f"ice split — {view} peaks (map: {map_label})", + show_inset=True, + sharpness_params=params, + **widget_kwargs, + ) + + +def probe_peaks(bp, params, ry, rx, *, n_show=4, q_band_only=True): + """Measure every peak at one scan position and plot the cuts behind each FWHM. + + Use it to check the automated width against a peak you can see: + ``probe_peaks(bp, params_ice_flagger, ice_widget.pos_ry, ice_widget.pos_rx)``. + Returns the per-peak arrays as a dict. + """ + polar_data = bp.polar_data + image = np.asarray(polar_data["intensity"])[ry, rx] + r_axis = np.asarray(polar_data["r_invA"])[:, 0] + theta_axis = np.asarray(polar_data["theta"])[0, :] + + peaks = np.asarray(bp.polar_peaks[ry, rx].array) + intensities = np.asarray(bp.peak_intensities[ry, rx].array) + q = peaks[:, bp.polar_peaks.fields.index("r_invA")] + theta = peaks[:, bp.polar_peaks.fields.index("theta")] + values = intensities[:, bp.peak_intensities.fields.index(params.intensity_field)] + + if q_band_only: + keep = np.abs(q - params.q_target_invA) <= params.dq_invA + q, theta, values = q[keep], theta[keep], values[keep] + if not len(q): + raise ValueError(f"no peaks at ({ry},{rx})" + + (" in the q band" if q_band_only else "")) + + width_r, width_theta = measure_peak_widths(q, theta, image, r_axis, theta_axis, + params=params) + gate_r, gate_t = params.max_width_r_invA, params.max_width_theta_deg + print(f"position ({ry}, {rx}) -- {len(q)} peaks" + f"{' in the q band' if q_band_only else ''}\n") + print(f"{'#':>3} {'q 1/Å':>8} {'d Å':>7} {'theta°':>8} {'intens':>8} " + f"{'radFWHM':>9} {'annFWHM':>9} gate") + for k in np.argsort(-values): + pass_r = gate_r is None or width_r[k] <= gate_r + pass_t = gate_t is None or width_theta[k] <= gate_t + verdict = ("sharp" if (pass_r and pass_t) else + "broad-r" if not pass_r and pass_t else + "broad-t" if pass_r else "broad-rt") + print(f"{k:>3} {q[k]:>8.4f} {1 / max(q[k], 1e-9):>7.2f} " + f"{np.rad2deg(theta[k]):>8.1f} {values[k]:>8.4f} " + f"{width_r[k]:>9.4f} {width_theta[k]:>9.1f} {verdict}") + + order = np.argsort(-values)[:n_show] + theta_step = float(np.rad2deg(theta_axis[1] - theta_axis[0])) + r_step = float(r_axis[1] - r_axis[0]) + fig, axes = plt.subplots(len(order), 2, figsize=(11, 2.6 * len(order)), squeeze=False) + for row, k in enumerate(order): + r_bin = int(np.clip(round((q[k] - r_axis[0]) / r_step), 0, len(r_axis) - 1)) + theta_bin = int(round(np.mod(np.rad2deg(theta[k]), 180.0) / theta_step)) + theta_bin %= len(theta_axis) + + ax = axes[row][0] + ax.plot(r_axis, image[:, theta_bin], lw=1) + ax.axvline(q[k], color="tab:red", ls=":") + ax.axvspan(q[k] - width_r[k] / 2, q[k] + width_r[k] / 2, color="tab:red", alpha=0.15) + ax.set(xlim=(q[k] - 6 * max(width_r[k], r_step), q[k] + 6 * max(width_r[k], r_step)), + xlabel="q (1/Å)", ylabel="intensity", + title=f"peak {k}: radial cut, FWHM={width_r[k]:.4f} 1/Å") + + ax = axes[row][1] + center = np.mod(np.rad2deg(theta[k]), 180.0) + ax.plot(np.rad2deg(theta_axis), image[r_bin, :], lw=1) + ax.axvline(center, color="tab:red", ls=":") + ax.axvspan(center - width_theta[k] / 2, center + width_theta[k] / 2, + color="tab:red", alpha=0.15) + ax.set(xlim=(0, 180), xlabel="theta (deg)", ylabel="intensity", + title=f"peak {k}: annular cut, FWHM={width_theta[k]:.1f}°") + fig.tight_layout() + plt.show() + return {"q_invA": q, "theta_rad": theta, "intensity": values, + "width_r_invA": width_r, "width_theta_deg": width_theta} diff --git a/src/quantem/diffraction/polymer_models.py b/src/quantem/diffraction/polymer_models.py new file mode 100644 index 000000000..ccc9420ce --- /dev/null +++ b/src/quantem/diffraction/polymer_models.py @@ -0,0 +1,507 @@ +from __future__ import annotations + +from dataclasses import dataclass +import hashlib +import json +import os +from pathlib import Path +import tempfile +from typing import TYPE_CHECKING, Any, Callable, Mapping +from urllib.error import HTTPError, URLError +from urllib.request import urlopen + +from quantem.core import config +from math import floor + +from quantem.core.ml.activation_functions import get_activation_function +from quantem.core.ml.blocks import Conv2dBlock, Upsample2dBlock, complex_pool, passfunc + +if TYPE_CHECKING: + import torch + import torch.nn as nn +else: + if config.get("has_torch"): + import torch + import torch.nn as nn + + +class CNN2d(nn.Module): + """ """ + + def __init__( + self, + in_channels: int, # input channels (C_in, H, W) + out_channels: int | None = None, # output channels (C_out, H, W) + start_filters: int = 16, + num_layers: int = 3, # num_layers + num_per_layer: int = 2, # number conv per layer + use_skip_connections: bool = False, + dtype: torch.dtype = torch.float32, + dropout: float = 0, + activation: str | Callable = "relu", + final_activation: str | Callable = nn.Identity(), + use_batchnorm: bool = True, + conv_kernel_size: int = 3, + ): + super().__init__() + self.in_channels = int(in_channels) + self.out_channels = int(out_channels) if out_channels is not None else int(in_channels) + self.start_filters = start_filters + self.num_layers = num_layers + self._num_per_layer = num_per_layer + if use_skip_connections and num_per_layer < 2: + raise ValueError( + "If using skip connections, num_per_layer must be at least 2 to allow for " + + "channel concatenation." + ) + self.use_skip_connections = use_skip_connections + self.dtype = dtype + self.dropout = dropout + self._use_batchnorm = use_batchnorm + + if self.dtype.is_complex: + self.pool = complex_pool + else: + self.pool = passfunc + self._pooler = nn.MaxPool2d(kernel_size=2, stride=2) + + self.concat = torch.cat + self.flatten = nn.Flatten() + + if callable(activation): + self._activation = activation + else: + self._activation = get_activation_function(activation, self.dtype) + if callable(final_activation): + self._final_activation = final_activation + else: + self._final_activation = get_activation_function(final_activation, self.dtype) + if conv_kernel_size <=0: + raise ValueError(f"Convolutional kernel size must be greater than 0. Got value {conv_kernel_size}") + if conv_kernel_size % 2 == 0: + raise ValueError(f"Convolutional kernel size must be an odd number. Got value {conv_kernel_size}") + self._conv_kernel_size = int(conv_kernel_size) + + self._build() + + @property + def activation(self) -> Callable: + return self._activation + + @property + def final_activation(self) -> Callable: + return self._final_activation + + @property + def conv_kernel_size(self) -> int: + return self._conv_kernel_size + + def _build(self): + self.down_conv_blocks = nn.ModuleList() + self.up_conv_blocks = nn.ModuleList() + self.upsample_blocks = nn.ModuleList() + + in_channels = self.in_channels + out_channels = self.start_filters + for a0 in range(self.num_layers): + if a0 != 0: + out_channels = in_channels * 2 + self.down_conv_blocks.append( + Conv2dBlock( + nb_layers=self._num_per_layer, + input_channels=in_channels, + output_channels=out_channels, + use_batchnorm=self._use_batchnorm, + dropout=0, + # dropout=self.dropout, + dtype=self.dtype, + activation=self.activation, + kernel_size=self.conv_kernel_size, + padding=int(floor(self.conv_kernel_size/2)), + ) + ) + in_channels = out_channels + + out_channels = in_channels * 2 + self.bottleneck = Conv2dBlock( + nb_layers=self._num_per_layer, + input_channels=in_channels, + output_channels=out_channels, + use_batchnorm=self._use_batchnorm, + dropout=self.dropout, + dtype=self.dtype, + activation=self.activation, + kernel_size=self.conv_kernel_size, + padding=int(floor(self.conv_kernel_size/2)), + ) + in_channels = out_channels + + for a0 in range(self.num_layers): + out_channels = self.start_filters if a0 == self.num_layers - 1 else in_channels // 2 + + in_channels2 = in_channels if self.use_skip_connections else out_channels + + self.upsample_blocks.append( + Upsample2dBlock( + in_channels, out_channels, use_batchnorm=self._use_batchnorm, dtype=self.dtype + ) + ) + + self.up_conv_blocks.append( + Conv2dBlock( + nb_layers=self._num_per_layer, + input_channels=in_channels2, + output_channels=out_channels, + use_batchnorm=self._use_batchnorm, + dropout=0, + # dropout=self.dropout, + dtype=self.dtype, + activation=self.activation, + kernel_size=self.conv_kernel_size, + padding=int(floor(self.conv_kernel_size/2)), + ) + ) + + in_channels = out_channels + + self.final_conv = Conv2dBlock( + nb_layers=1, + input_channels=self.start_filters, + output_channels=self.out_channels, + use_batchnorm=False, + dropout=0, + # dropout=self.dropout, + dtype=self.dtype, + activation=self.final_activation, + ) + return + + def forward(self, x: torch.Tensor) -> torch.Tensor: + skips = [] + for down_block in self.down_conv_blocks: + x = down_block(x) + if self.use_skip_connections: + skips.append(x) + x = self.pool(x, self._pooler) + + x = self.bottleneck(x) + for upsample_block, up_conv_block in zip(self.upsample_blocks, self.up_conv_blocks): + x = upsample_block(x) + if self.use_skip_connections: + skip = skips.pop() + x = torch.cat((x, skip), dim=1) + x = up_conv_block(x) + + y = self.final_conv(x) + + return y + + def reset_weights(self): + """ + Reset all weights. + """ + + def _reset(m: nn.Module) -> None: + reset_parameters = getattr(m, "reset_parameters", None) + if callable(reset_parameters): + reset_parameters() + + self.apply(_reset) + + +class MultiChannelCNN2d(CNN2d): + def __init__( + self, + in_channels=1, + out_channels: int = 2, + final_activations: list | tuple | None = None, + **kwargs + ): + # Always use identity activation in base CNN, handle activations here + super().__init__(in_channels=in_channels, out_channels=out_channels, final_activation="identity", **kwargs) + self.final_activations = ( + ["sigmoid"] * out_channels if final_activations is None else final_activations + ) + + @property + def final_activations(self): + return self._final_activations + + @final_activations.setter + def final_activations(self, value): + if not isinstance(value, (list, tuple)) or len(value) != self.out_channels: + raise ValueError(f"final_activations must be a list of length {self.out_channels}") + self._final_activations = [get_activation_function(act, self.dtype) for act in value] + + def forward(self, x): + out = super().forward(x) # B,C,H,W + # Apply per-channel activation + outs = [] + for i, fn in enumerate(self.final_activations): + outs.append(fn(out[:, i:i+1])) + return torch.cat(outs, dim=1) + + +# These aliases make the intentional separation from quantem.core.ml.cnn.CNN2d +# explicit while preserving the state-dict key layout of the paper checkpoint. +PolymerCNN2d = CNN2d +PolymerMultiChannelCNN2d = MultiChannelCNN2d + +PAPER_MODEL_ID = "reference-v2" +PAPER_MODEL_VERSION = "2026-06-29" +PAPER_MODEL_SHA256 = "c2a4ed76cccd9313b4821d629b48767bbce6714b8466ff915357f5607b31c1a7" + +_PAPER_SPEC: dict[str, Any] = { + "schema_version": 1, + "model_id": PAPER_MODEL_ID, + "version": PAPER_MODEL_VERSION, + "description": "Pinned paper polymer diffraction-peak detector.", + "architecture": { + "start_filters": 32, + "num_layers": 4, + "num_per_layer": 3, + "kernel_size": 3, + "input_channels": 1, + "output_channels": 2, + "dropout": 0.0, + }, + "normalization": { + "mode": "v1_global_percentile", + "p_lower": 0.0418, + "p_upper": 3.394, + }, + "experimental_normalization": { + "mode": "per_image_minmax_percentile", + "lower_percentile": 1.0, + "upper_percentile": 99.0, + }, + "weights": {"filename": "best.pth", "sha256": PAPER_MODEL_SHA256}, +} + +# The DOI-backed URL is deliberately unset until the public, immutable Zenodo +# record exists. A local directory override remains available for private review. +DEFAULT_MODEL_REGISTRY: dict[str, dict[str, dict[str, Any]]] = { + PAPER_MODEL_ID: { + PAPER_MODEL_VERSION: { + "specification": _PAPER_SPEC, + "weights_url": None, + } + } +} + + +@dataclass(frozen=True) +class PolymerModelResolution: + """A verified, immutable polymer model artifact.""" + + model_id: str + version: str + weights_path: Path + specification: Mapping[str, Any] + checksum: str + + +class PolymerModelError(RuntimeError): + """Raised when a named polymer model cannot be resolved safely.""" + + +def build_polymer_model(specification: Mapping[str, Any]) -> MultiChannelCNN2d: + """Build the checkpoint-compatible network described by a model specification.""" + + architecture = specification.get("architecture", {}) + return MultiChannelCNN2d( + in_channels=int(architecture["input_channels"]), + out_channels=int(architecture["output_channels"]), + start_filters=int(architecture["start_filters"]), + num_layers=int(architecture["num_layers"]), + num_per_layer=int(architecture["num_per_layer"]), + use_skip_connections=True, + dropout=float(architecture.get("dropout", 0.0)), + final_activations=["sigmoid"] * int(architecture["output_channels"]), + conv_kernel_size=int(architecture["kernel_size"]), + ) + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for block in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def _validate_specification( + specification: Mapping[str, Any], *, model_id: str, version: str, checksum: str +) -> None: + if specification.get("schema_version") != 1: + raise PolymerModelError( + f"Model {model_id!r} version {version!r} uses an unsupported specification schema." + ) + if specification.get("model_id") != model_id: + raise PolymerModelError( + f"Model specification identifies {specification.get('model_id')!r}, expected {model_id!r}." + ) + architecture = specification.get("architecture", {}) + required = { + "input_channels": 1, + "output_channels": 2, + "kernel_size": 3, + } + incompatible = { + key: (architecture.get(key), expected) + for key, expected in required.items() + if architecture.get(key) != expected + } + if incompatible: + raise PolymerModelError( + f"Model {model_id!r} version {version!r} is incompatible with " + f"BraggPeaksPolymer: {incompatible}." + ) + declared = specification.get("weights", {}).get("sha256") + if declared != checksum: + raise PolymerModelError( + f"Model specification checksum {declared!r} does not match registry checksum {checksum!r}." + ) + + +def _local_candidates(root: Path, model_id: str, version: str, filename: str): + yield root / model_id / version / filename + yield root / model_id / filename + yield root / filename + + +def resolve_polymer_model( + model_id: str = PAPER_MODEL_ID, + version: str | None = None, + *, + latest: bool = False, + local_model_dir: str | os.PathLike[str] | None = None, + cache_dir: str | os.PathLike[str] | None = None, + registry: Mapping[str, Mapping[str, Mapping[str, Any]]] | None = None, + downloader: Callable[..., Any] = urlopen, +) -> PolymerModelResolution: + """Resolve and verify a named polymer model. + + The paper model and version are pinned by default. ``latest=True`` is the + only way to select a newer registered version. Private development can use + ``local_model_dir`` or ``QUANTEM_POLYMER_MODEL_DIR`` without network access. + """ + + if latest and version is not None: + raise ValueError("version and latest=True are mutually exclusive") + selected_registry = DEFAULT_MODEL_REGISTRY if registry is None else registry + versions = selected_registry.get(model_id) + if not versions: + raise PolymerModelError(f"Unknown polymer model {model_id!r}.") + if latest: + version = sorted(versions)[-1] + elif version is None: + version = PAPER_MODEL_VERSION if model_id == PAPER_MODEL_ID else sorted(versions)[0] + entry = versions.get(version) + if entry is None: + raise PolymerModelError( + f"Unknown version {version!r} for polymer model {model_id!r}." + ) + + specification = dict(entry["specification"]) + weights = specification.get("weights", {}) + checksum = str(weights.get("sha256", "")).lower() + filename = str(weights.get("filename", "best.pth")) + if len(checksum) != 64: + raise PolymerModelError(f"Model {model_id!r} version {version!r} has no valid SHA-256.") + _validate_specification( + specification, model_id=model_id, version=version, checksum=checksum + ) + + local_root_value = local_model_dir or os.environ.get("QUANTEM_POLYMER_MODEL_DIR") + if local_root_value: + local_root = Path(local_root_value).expanduser() + for candidate in _local_candidates(local_root, model_id, version, filename): + if candidate.is_file(): + actual = _sha256(candidate) + if actual != checksum: + raise PolymerModelError( + f"Checksum failure for local model {candidate}: expected {checksum}, got {actual}." + ) + return PolymerModelResolution( + model_id, version, candidate.resolve(), specification, actual + ) + raise PolymerModelError( + f"Model {model_id!r} version {version!r} was not found below local model " + f"directory {local_root}. Expected {filename}." + ) + + cache_root = ( + Path(cache_dir).expanduser() + if cache_dir is not None + else Path(os.environ.get("XDG_CACHE_HOME", Path.home() / ".cache")) + / "quantem" + / "polymer_models" + ) + target = cache_root / model_id / version / checksum / filename + if target.is_file(): + actual = _sha256(target) + if actual == checksum: + return PolymerModelResolution(model_id, version, target, specification, actual) + raise PolymerModelError( + f"Checksum failure for cached model {target}: expected {checksum}, got {actual}. " + "Remove that file and resolve the model again." + ) + + weights_url = entry.get("weights_url") + if not weights_url: + raise PolymerModelError( + f"Model {model_id!r} version {version!r} is not public yet and is absent from " + "the local cache. Set QUANTEM_POLYMER_MODEL_DIR to the private model archive " + "or pass local_model_dir explicitly." + ) + + target.parent.mkdir(parents=True, exist_ok=True) + temporary: Path | None = None + try: + with tempfile.NamedTemporaryFile(dir=target.parent, prefix=".download-", delete=False) as out: + temporary = Path(out.name) + try: + response = downloader(str(weights_url)) + context = response if hasattr(response, "__enter__") else None + source = context.__enter__() if context is not None else response + try: + while True: + block = source.read(1024 * 1024) + if not block: + break + out.write(block) + finally: + if context is not None: + context.__exit__(None, None, None) + except (HTTPError, URLError, OSError) as exc: + raise PolymerModelError( + f"Could not download model {model_id!r} version {version!r} from " + f"{weights_url}: {exc}. An offline cache or local_model_dir may be used." + ) from exc + actual = _sha256(temporary) + if actual != checksum: + raise PolymerModelError( + f"Checksum failure after downloading {model_id!r} version {version!r}: " + f"expected {checksum}, got {actual}." + ) + os.replace(temporary, target) + temporary = None + finally: + if temporary is not None: + temporary.unlink(missing_ok=True) + + return PolymerModelResolution(model_id, version, target, specification, checksum) + + +__all__ = [ + "DEFAULT_MODEL_REGISTRY", + "PAPER_MODEL_ID", + "PAPER_MODEL_VERSION", + "PolymerCNN2d", + "PolymerModelError", + "PolymerModelResolution", + "PolymerMultiChannelCNN2d", + "build_polymer_model", + "resolve_polymer_model", +] diff --git a/src/quantem/diffraction/polymer_normalization.py b/src/quantem/diffraction/polymer_normalization.py new file mode 100644 index 000000000..762ee0b3f --- /dev/null +++ b/src/quantem/diffraction/polymer_normalization.py @@ -0,0 +1,215 @@ +"""Inference normalization strategies for polymer peak-detection models.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable, Mapping, Protocol, runtime_checkable + +import numpy as np + + +@runtime_checkable +class NormalizationStrategy(Protocol): + """A fittable, reusable normalization operation.""" + + def fit(self, sample_batch: Any) -> Any: + """Fit parameters from a representative batch.""" + + def transform(self, batch: Any, parameters: Any) -> Any: + """Normalize a batch using previously fitted parameters.""" + + +def _is_torch(value: Any) -> bool: + try: + import torch + except ImportError: + return False + return isinstance(value, torch.Tensor) + + +def _percentiles(batch: Any, lower: float, upper: float) -> tuple[float, float]: + if _is_torch(batch): + import torch + + flat = batch.detach().flatten().float() + if flat.numel() == 0: + raise ValueError("Cannot fit normalization from an empty batch.") + values = torch.quantile( + flat, + torch.tensor( + [lower / 100.0, upper / 100.0], + device=flat.device, + dtype=flat.dtype, + ), + ) + return float(values[0].item()), float(values[1].item()) + + array = np.asarray(batch) + if array.size == 0: + raise ValueError("Cannot fit normalization from an empty batch.") + values = np.percentile(array, [lower, upper]) + return float(values[0]), float(values[1]) + + +def _percentile_transform(batch: Any, parameters: tuple[float, float]) -> Any: + lower, upper = parameters + if _is_torch(batch): + import torch + + if not (batch.dtype.is_floating_point or batch.dtype.is_complex): + batch = batch.float() + lo = torch.as_tensor(lower, device=batch.device, dtype=batch.dtype) + hi = torch.as_tensor(upper, device=batch.device, dtype=batch.dtype) + return (torch.clamp(batch, lo, hi) - lo) / (hi - lo + 1e-8) + array = np.asarray(batch) + return (np.clip(array, lower, upper) - lower) / (upper - lower + 1e-8) + + +def _per_image_minmax(batch: Any) -> Any: + """Min-max each image over its final two dimensions.""" + + if getattr(batch, "ndim", None) is None or batch.ndim < 2: + raise ValueError("Normalization expects an image or a batch of images.") + axes = (-2, -1) + if _is_torch(batch): + import torch + + if not (batch.dtype.is_floating_point or batch.dtype.is_complex): + batch = batch.float() + minimum = torch.amin(batch, dim=axes, keepdim=True) + maximum = torch.amax(batch, dim=axes, keepdim=True) + span = maximum - minimum + return torch.where(span > 0, (batch - minimum) / span, torch.zeros_like(batch)) + array = np.asarray(batch) + minimum = np.min(array, axis=axes, keepdims=True) + maximum = np.max(array, axis=axes, keepdims=True) + span = maximum - minimum + return np.divide( + array - minimum, + span, + out=np.zeros_like(array, dtype=np.result_type(array.dtype, np.float32)), + where=span > 0, + ) + + +@dataclass(frozen=True) +class GlobalPercentileNormalization: + """Clip and scale using percentiles fitted across the entire sample batch.""" + + lower_percentile: float = 1.0 + upper_percentile: float = 99.0 + + def __post_init__(self) -> None: + if not 0 <= self.lower_percentile < self.upper_percentile <= 100: + raise ValueError("Percentiles must satisfy 0 <= lower < upper <= 100.") + + def fit(self, sample_batch: Any) -> tuple[float, float]: + return _percentiles( + sample_batch, self.lower_percentile, self.upper_percentile + ) + + def transform(self, batch: Any, parameters: Any) -> Any: + return _percentile_transform(batch, parameters) + + +@dataclass(frozen=True) +class PerImageMinMaxPercentileNormalization: + """Min-max each image, then clip and scale by fitted global percentiles.""" + + lower_percentile: float = 1.0 + upper_percentile: float = 99.0 + + def __post_init__(self) -> None: + if not 0 <= self.lower_percentile < self.upper_percentile <= 100: + raise ValueError("Percentiles must satisfy 0 <= lower < upper <= 100.") + + def fit(self, sample_batch: Any) -> tuple[float, float]: + normalized = _per_image_minmax(sample_batch) + return _percentiles( + normalized, self.lower_percentile, self.upper_percentile + ) + + def transform(self, batch: Any, parameters: Any) -> Any: + return _percentile_transform(_per_image_minmax(batch), parameters) + + +@dataclass(frozen=True) +class LegacyNormalizationAdapter: + """Adapt the historical compute/normalize callback pair to the strategy API.""" + + compute_parameters: Callable[..., Any] + normalize_data: Callable[..., Any] + lower_percentile: float = 1.0 + upper_percentile: float = 99.0 + + def fit(self, sample_batch: Any) -> Any: + return self.compute_parameters( + sample_batch, + lower_percentile=self.lower_percentile, + upper_percentile=self.upper_percentile, + ) + + def transform(self, batch: Any, parameters: Any) -> Any: + if isinstance(parameters, tuple): + return self.normalize_data(batch, *parameters) + return self.normalize_data(batch, parameters) + + +_STRATEGIES: dict[str, type[NormalizationStrategy]] = { + "global_percentile": GlobalPercentileNormalization, + "v1_global_percentile": GlobalPercentileNormalization, + "per_scan_percentile": GlobalPercentileNormalization, + "per_image_minmax_percentile": PerImageMinMaxPercentileNormalization, + "v2_per_image_minmax_percentile": PerImageMinMaxPercentileNormalization, +} + + +def resolve_normalization_strategy( + specification: str | Mapping[str, Any] | NormalizationStrategy, +) -> NormalizationStrategy: + """Resolve a registered strategy name/configuration or return a strategy instance.""" + + if isinstance(specification, str): + mode, parameters = specification, {} + elif isinstance(specification, Mapping): + config = dict(specification) + try: + mode = str(config.pop("mode")) + except KeyError as exc: + raise ValueError("Normalization configuration requires a 'mode'.") from exc + parameters = config + elif isinstance(specification, NormalizationStrategy): + return specification + else: + raise TypeError( + "normalization_strategy must be a registered name, configuration mapping, " + "or object implementing fit() and transform()." + ) + + strategy_type = _STRATEGIES.get(mode) + if strategy_type is None: + raise ValueError( + f"Unknown normalization strategy {mode!r}; registered strategies are " + f"{sorted(_STRATEGIES)}." + ) + if "p_lower" in parameters: + parameters.setdefault("lower_percentile", parameters.pop("p_lower")) + if "p_upper" in parameters: + parameters.setdefault("upper_percentile", parameters.pop("p_upper")) + return strategy_type(**parameters) + + +# Concise aliases for callers that prefer strategy-oriented names. +GlobalPercentileStrategy = GlobalPercentileNormalization +PerImageMinMaxPercentileStrategy = PerImageMinMaxPercentileNormalization + + +__all__ = [ + "GlobalPercentileNormalization", + "GlobalPercentileStrategy", + "LegacyNormalizationAdapter", + "NormalizationStrategy", + "PerImageMinMaxPercentileNormalization", + "PerImageMinMaxPercentileStrategy", + "resolve_normalization_strategy", +] diff --git a/src/quantem/diffraction/polymer_utils.py b/src/quantem/diffraction/polymer_utils.py new file mode 100644 index 000000000..992c0730c --- /dev/null +++ b/src/quantem/diffraction/polymer_utils.py @@ -0,0 +1,79 @@ +"""Small numerical helpers specific to polymer diffraction inference.""" + +from __future__ import annotations + +import warnings + +import numpy as np +from scipy.ndimage import uniform_filter + + +def parse_reciprocal_units(unit_string: str) -> tuple[str, float]: + """Return a canonical reciprocal unit and its multiplier to inverse angstroms.""" + + normalized = ( + str(unit_string) + .strip() + .lower() + .replace(" ", "") + .replace("angstrom", "å") + .replace("ang", "å") + ) + nanometer = {"1/nm", "/nm", "nm^-1", "nm-1", "nm⁻¹", "inv_nm", "per_nm"} + angstrom = { + "1/a", + "/a", + "a^-1", + "a-1", + "a⁻¹", + "1/å", + "/å", + "å^-1", + "å-1", + "å⁻¹", + "inv_a", + "inv_å", + "per_a", + "per_å", + } + if normalized in nanometer: + return "1/nm", 0.1 + if normalized in angstrom: + return "1/A", 1.0 + warnings.warn( + f"Unrecognized reciprocal unit {unit_string!r}; assuming 1/Å.", + UserWarning, + stacklevel=2, + ) + return "unknown", 1.0 + + +def sample_average_from_image( + image: np.ndarray, + coordinates: np.ndarray, + radius_dim1: int = 2, + radius_dim2: int = 2, +) -> np.ndarray: + """Sample local means from a polar image, wrapping only its angular axis.""" + + image = np.asarray(image) + coordinates = np.asarray(coordinates) + if image.ndim != 2 or coordinates.ndim != 2 or coordinates.shape[1] != 2: + raise ValueError("image must be 2D and coordinates must have shape (n, 2)") + height, width = image.shape + padded = np.pad(image, ((radius_dim1, radius_dim1), (radius_dim2, radius_dim2)), mode="wrap") + if radius_dim1: + padded[:radius_dim1] = 0 + padded[-radius_dim1:] = 0 + means = uniform_filter( + padded, + size=(2 * radius_dim1 + 1, 2 * radius_dim2 + 1), + mode="constant", + cval=0, + ) + rows = np.clip(coordinates[:, 0].astype(int) + radius_dim1, 0, height + 2 * radius_dim1 - 1) + cols = coordinates[:, 1].astype(int) % width + radius_dim2 + return means[rows, cols] + + +__all__ = ["parse_reciprocal_units", "sample_average_from_image"] diff --git a/tests/diffraction/test_ellipse_ring_fit.py b/tests/diffraction/test_ellipse_ring_fit.py new file mode 100644 index 000000000..d2769d334 --- /dev/null +++ b/tests/diffraction/test_ellipse_ring_fit.py @@ -0,0 +1,170 @@ +import numpy as np +import pytest + +from quantem.diffraction import BraggPeaksPolymer + + +def _elliptical_ring( + shape=(96, 96), + *, + radius=27.0, + ratio_b_over_a=0.9, + theta_deg=35.0, + sigma=1.8, + center_offset=(0.0, 0.0), +): + yy, xx = np.indices(shape, dtype=float) + cy, cx = (np.asarray(shape) - 1) / 2 + cy += center_offset[0] + cx += center_offset[1] + theta = np.deg2rad(theta_deg) + dx, dy = xx - cx, yy - cy + major = dx * np.cos(theta) + dy * np.sin(theta) + minor = -dx * np.sin(theta) + dy * np.cos(theta) + elliptical_radius = np.sqrt( + (major * ratio_b_over_a) ** 2 + minor**2 + ) + return ( + 2.0 + * np.exp(-0.5 * ((elliptical_radius - radius) / sigma) ** 2) + + 0.02 + ) + + +def _fit(pattern, **kwargs): + detector = object.__new__(BraggPeaksPolymer) + result = detector._fit_ellipse_from_ring( + pattern, + ((pattern.shape[0] - 1) / 2, (pattern.shape[1] - 1) / 2), + n_ratio=7, + n_theta=12, + max_ring_candidates=3, + **kwargs, + ) + return detector, result + + +def _fit_ridge(pattern, **kwargs): + detector = object.__new__(BraggPeaksPolymer) + center = ((pattern.shape[0] - 1) / 2, (pattern.shape[1] - 1) / 2) + result = detector._fit_ellipse_from_ridge( + pattern, + center, + num_annular_bins=120, + **kwargs, + ) + return detector, result + + +def test_ring_fit_recovers_synthetic_ellipse(): + detector, (a_axis, b_axis, theta, band) = _fit( + _elliptical_ring(ratio_b_over_a=0.9, theta_deg=35.0) + ) + + assert detector.ellipse_fit_diagnostics["accepted"] is True + assert a_axis / b_axis == pytest.approx(1 / 0.9, abs=0.025) + assert theta == pytest.approx(35.0, abs=3.0) + assert band[0] < 27 < band[1] + + +def test_sparse_outer_bragg_spots_do_not_select_the_calibration_band(): + pattern = _elliptical_ring( + radius=26.0, ratio_b_over_a=0.94, theta_deg=118.0 + ) + cy, cx = (np.asarray(pattern.shape) - 1) / 2 + for angle in np.deg2rad([5, 42, 91, 147, 221, 305]): + row = int(round(cy + 39 * np.sin(angle))) + column = int(round(cx + 39 * np.cos(angle))) + pattern[row - 1 : row + 2, column - 1 : column + 2] += 50.0 + + detector, (_, _, _, band) = _fit(pattern) + + assert detector.ellipse_fit_diagnostics["accepted"] is True + assert detector.ellipse_fit_diagnostics["selected"]["r0"] < 32 + assert band[0] < 26 < band[1] + + +def test_boundary_solution_is_rejected_and_refinement_is_clipped(): + pattern = _elliptical_ring(ratio_b_over_a=0.7) + + with pytest.warns(RuntimeWarning, match="ratio search boundary"): + detector, (a_axis, b_axis, theta, _) = _fit(pattern) + + selected = detector.ellipse_fit_diagnostics["selected"] + assert detector.ellipse_fit_diagnostics["accepted"] is False + assert selected["boundary_limited"] is True + assert 0.85 <= selected["ratio_b_over_a"] <= 1.18 + assert (a_axis / b_axis, theta) == pytest.approx((1.0, 0.0)) + + +def test_low_information_pattern_falls_back_to_circle(): + with pytest.warns(RuntimeWarning, match="using a circular correction"): + detector, (a_axis, b_axis, theta, _) = _fit(np.ones((96, 96))) + + assert detector.ellipse_fit_diagnostics["accepted"] is False + assert (a_axis / b_axis, theta) == pytest.approx((1.0, 0.0)) + + +def test_ridge_fit_jointly_recovers_center_and_ellipse(): + offset = (1.2, -0.8) + detector, (a_axis, b_axis, theta, band) = _fit_ridge( + _elliptical_ring( + ratio_b_over_a=0.9, + theta_deg=35.0, + center_offset=offset, + ) + ) + + expected_center = np.asarray((47.5, 47.5)) + offset + assert detector.ellipse_fit_diagnostics["accepted"] is True + assert detector.ellipse_fit_diagnostics["center_refined"] == pytest.approx( + expected_center, abs=0.25 + ) + assert a_axis / b_axis == pytest.approx(1 / 0.9, abs=0.03) + assert theta == pytest.approx(35.0, abs=3.0) + assert band[0] < 27 < band[1] + + +def test_ridge_fit_rejects_out_of_range_ellipse(): + with pytest.warns(RuntimeWarning, match="ratio search boundary"): + detector, (a_axis, b_axis, theta, _) = _fit_ridge( + _elliptical_ring(ratio_b_over_a=0.7) + ) + + assert detector.ellipse_fit_diagnostics["accepted"] is False + assert (a_axis / b_axis, theta) == pytest.approx((1.0, 0.0)) + + +def test_ridge_fit_ignores_sparse_outer_bragg_spots(): + pattern = _elliptical_ring( + radius=26.0, ratio_b_over_a=0.94, theta_deg=118.0 + ) + cy, cx = (np.asarray(pattern.shape) - 1) / 2 + for angle in np.deg2rad([5, 42, 91, 147, 221, 305]): + row = int(round(cy + 39 * np.sin(angle))) + column = int(round(cx + 39 * np.cos(angle))) + pattern[row - 1 : row + 2, column - 1 : column + 2] += 50.0 + + detector, (_, _, _, band) = _fit_ridge(pattern) + + assert detector.ellipse_fit_diagnostics["accepted"] is True + assert detector.ellipse_fit_diagnostics["selected"]["r0"] < 32 + assert band[0] < 26 < band[1] + + +def test_low_information_ridge_falls_back_to_circle(): + with pytest.warns(RuntimeWarning, match="using a circular correction"): + detector, (a_axis, b_axis, theta, _) = _fit_ridge( + np.ones((96, 96), dtype=float) + ) + + assert detector.ellipse_fit_diagnostics["method"] == "ridge" + assert detector.ellipse_fit_diagnostics["accepted"] is False + assert (a_axis / b_axis, theta) == pytest.approx((1.0, 0.0)) + + +def test_preprocess_rejects_unknown_ellipse_fit_method_before_data_access(): + detector = object.__new__(BraggPeaksPolymer) + + with pytest.raises(ValueError, match="ellipse_fit_method"): + detector.preprocess(ellipse_fit_method="not-a-fit-method") diff --git a/tests/diffraction/test_orientation_correlation.py b/tests/diffraction/test_orientation_correlation.py new file mode 100644 index 000000000..109957a8b --- /dev/null +++ b/tests/diffraction/test_orientation_correlation.py @@ -0,0 +1,237 @@ +import numpy as np +import pytest + +from quantem.diffraction import BraggPeaksPolymer +from quantem.diffraction.orientation_correlation import ( + calculate_orientation_correlation, +) + + +def _direct_correlation_reference(orient_hist, radius_max): + """Small full-volume implementation used only as a correctness oracle.""" + num_radii, size_x, size_y, num_theta = orient_hist.shape + padded_x = max(2 * size_x, 2 * radius_max) + padded_y = max(2 * size_y, 2 * radius_max) + + x = np.mod(np.arange(padded_x) + padded_x / 2, padded_x) - padded_x / 2 + y = np.mod(np.arange(padded_y) + padded_y / 2, padded_y) - padded_y / 2 + yy, xx = np.meshgrid(y, x) + radius = np.sqrt(xx**2 + yy**2) + lower_mask = radius <= radius_max + upper_mask = radius <= radius_max - 1 + lower_floor = np.floor(radius[lower_mask]).astype(int) + upper_floor = np.floor(radius[upper_mask]).astype(int) + bins = np.concatenate((lower_floor, upper_floor + 1)) + weights = np.concatenate( + ( + 1 - (radius[lower_mask] - lower_floor), + radius[upper_mask] - upper_floor, + ) + ) + + spectrum = np.fft.fftn( + orient_hist, + s=(padded_x, padded_y, num_theta), + axes=(1, 2, 3), + ) + pairs = [ + (first, second) + for first in range(num_radii) + for second in range(first, num_radii) + ] + output = [] + for first, second in pairs: + spatial_angular = np.fft.ifftn( + spectrum[first] * np.conj(spectrum[second]), + axes=(0, 1, 2), + ).real + radial = np.stack( + [ + np.bincount( + bins, + weights=weights + * np.concatenate( + ( + spatial_angular[:, :, theta][lower_mask], + spatial_angular[:, :, theta][upper_mask], + ) + ), + minlength=radius_max + 1, + )[: radius_max + 1] + for theta in range(num_theta) + ] + ) + denominator = radial.sum(axis=0) / num_theta + output.append( + radial[: num_theta // 2 + 1] / denominator[None, :] + ) + return np.stack(output), np.asarray(pairs) + + +@pytest.mark.parametrize("backend", ["numpy", "torch"]) +def test_streamed_correlation_matches_full_volume_reference(backend): + histogram = np.random.default_rng(7).random( + (3, 7, 6, 12), dtype=np.float32 + ) + expected, expected_pairs = _direct_correlation_reference( + histogram, radius_max=4 + ) + + actual, actual_pairs = calculate_orientation_correlation( + histogram, + radius_max=4, + backend=backend, + device="cpu", + mode_batch_size=3, + pair_batch_size=2, + progress_bar=False, + ) + + np.testing.assert_array_equal(actual_pairs, expected_pairs) + np.testing.assert_allclose(actual, expected, rtol=5e-6, atol=5e-6) + + +def test_three_dimensional_input_and_autocorrelation_pairs(): + histogram = np.random.default_rng(8).random( + (2, 5, 4, 9), dtype=np.float32 + ) + + single, single_pairs = calculate_orientation_correlation( + histogram[0], + backend="numpy", + progress_bar=False, + ) + diagonal, diagonal_pairs = calculate_orientation_correlation( + histogram, + pairs="autocorrelation", + backend="numpy", + progress_bar=False, + ) + + assert single.shape == (1, 5, 3) + np.testing.assert_array_equal(single_pairs, [[0, 0]]) + assert diagonal.shape == (2, 5, 3) + np.testing.assert_array_equal(diagonal_pairs, [[0, 0], [1, 1]]) + + +@pytest.mark.parametrize( + ("zero_policy", "expected"), + [("nan", "nan"), ("zero", "zero")], +) +def test_empty_histogram_zero_policy(zero_policy, expected): + output, _ = calculate_orientation_correlation( + np.zeros((1, 4, 5, 8), dtype=np.float32), + backend="numpy", + zero_policy=zero_policy, + progress_bar=False, + ) + + if expected == "nan": + assert np.isnan(output).all() + else: + np.testing.assert_array_equal(output, 0) + + +def test_empty_histogram_raise_policy(): + with pytest.raises(ZeroDivisionError): + calculate_orientation_correlation( + np.zeros((1, 4, 5, 8), dtype=np.float32), + backend="numpy", + zero_policy="raise", + progress_bar=False, + ) + + +def test_bragg_peaks_polymer_native_correlation_plot(): + detector = object.__new__(BraggPeaksPolymer) + detector.orient_corr = np.ones((3, 5, 4), dtype=np.float32) + detector.orient_corr_pairs = np.array([[0, 0], [0, 1], [1, 1]]) + + figure, axes, metrics = detector.plot_orientation_correlation( + pixel_size=0.25, + pixel_units="scan pixels", + return_metrics=True, + ) + + assert axes.shape == (1, 3) + assert axes[0, 1].get_title() == "Correlation of Rings 0 and 1" + assert metrics[0]["title"] == "Autocorrelation of Ring 0" + figure.canvas.draw() + + +def test_correlation_plot_reports_half_probability_intercepts(): + detector = object.__new__(BraggPeaksPolymer) + distances = np.arange(11, dtype=float) + angles = np.linspace(0, 180, 37) + boundary = 20 - 0.5 * distances + panel = 1 + 9 * np.exp(-distances[None, :] / 4) * ( + boundary[None, :] - angles[:, None] + ) / 20 + detector.orient_corr = panel[None] + detector.orient_corr_pairs = np.array([[0, 0]]) + + figure, _, metrics = detector.plot_orientation_correlation( + pixel_size=1.0, + pixel_units="nm", + return_metrics=True, + ) + + assert np.isfinite(metrics[0]["radial_distance"]) + assert np.isfinite(metrics[0]["annular_distance_degrees"]) + assert metrics[0]["slope_degrees_per_unit"] == pytest.approx(-0.5, abs=0.05) + assert metrics[0]["slope_fit_r_squared"] == pytest.approx(1.0) + assert metrics[0]["slope_fit_point_count"] >= 2 + assert metrics[0]["slope_contour_probability"] == 1.0 + # Only the two intercept markers remain; neither the 50% contour nor the + # correlation=1 boundary is drawn. The single line is the signed fit. + assert len(figure.axes[0].collections) == 2 + assert len(figure.axes[0].lines) == 1 + assert len(figure.axes[0].texts) == 1 + + +def test_correlation_plot_resolves_below_baseline_feature(): + detector = object.__new__(BraggPeaksPolymer) + distances = np.arange(11, dtype=float) + angles = np.linspace(0, 180, 37) + boundary = 20 - 0.5 * distances + panel = 1 - 0.5 * np.exp(-distances[None, :] / 4) * ( + boundary[None, :] - angles[:, None] + ) / 20 + detector.orient_corr = panel[None] + detector.orient_corr_pairs = np.array([[0, 1]]) + + _, axes, metrics = detector.plot_orientation_correlation( + pixel_size=1.0, + pixel_units="nm", + return_metrics=True, + ) + + assert np.isfinite(metrics[0]["radial_distance"]) + assert np.isfinite(metrics[0]["annular_distance_degrees"]) + assert metrics[0]["slope_degrees_per_unit"] == pytest.approx(-0.5, abs=0.05) + assert axes[0, 0].get_legend() is not None + + +def test_correlation_slope_stops_before_connected_boundary_turns_back(): + detector = object.__new__(BraggPeaksPolymer) + distances = np.arange(101, dtype=float) + angles = np.linspace(0, 180, 91) + boundary = np.where( + distances <= 30, + 40 + 0.5 * distances, + 55 - 0.8 * (distances - 30), + ) + panel = 1 + (boundary[None, :] - angles[:, None]) / 40 + detector.orient_corr = panel[None] + detector.orient_corr_pairs = np.array([[0, 0]]) + + _, _, metrics = detector.plot_orientation_correlation( + pixel_size=1.0, + pixel_units="nm", + show_metrics=False, + return_metrics=True, + ) + + assert metrics[0]["slope_degrees_per_unit"] == pytest.approx(0.5, abs=0.05) + assert metrics[0]["slope_fit_r_squared"] > 0.95 + assert metrics[0]["slope_fit_point_count"] < len(distances) // 2 diff --git a/tests/diffraction/test_origin_finding.py b/tests/diffraction/test_origin_finding.py new file mode 100644 index 000000000..07f5e19e3 --- /dev/null +++ b/tests/diffraction/test_origin_finding.py @@ -0,0 +1,450 @@ +import numpy as np +import pytest +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +from quantem.core.datastructures import Vector +from quantem.core.datastructures.dataset4dstem import Dataset4dstem +import quantem.diffraction.bragg_peaks as bragg_peaks_module +from quantem.diffraction.bragg_peaks import ( + BraggPeaksPolymer, + _central_peak_index, + _display_center, + _intensity_display_limits, + _mean_intensity_map, + _normalized_dp, + _polar_peak_bins, + _resolve_intensity_map, + _zoom_peak_overlay, +) +from quantem.diffraction.polar_transform import ( + find_origin, + find_origin_angular_descent, + find_origin_angular_grid, +) + + +def _ring_pattern(ny, nx, cy, cx, radii=(10, 20, 30), beam_sigma=2.5): + y, x = np.ogrid[:ny, :nx] + r = np.sqrt((y - cy) ** 2 + (x - cx) ** 2) + pattern = np.zeros((ny, nx), dtype=np.float32) + for radius in radii: + pattern += 80.0 * np.exp(-((r - radius) ** 2) / (2 * 1.5**2)) + pattern += 500.0 * np.exp(-(r**2) / (2 * beam_sigma**2)) + return pattern.astype(np.float32) + + +def _dataset(array): + return Dataset4dstem.from_array( + array=np.asarray(array, dtype=np.float32), + name="origin_finding_test", + origin=(0, 0, 0, 0), + sampling=(1.0, 1.0, 1.0, 1.0), + units=["pixels", "pixels", "pixels", "pixels"], + signal_units="counts", + ) + + +def _bragg_for_plotting(): + arr = np.arange(2 * 2 * 8 * 8, dtype=np.float32).reshape(2, 2, 8, 8) + ds = _dataset(arr) + bp = BraggPeaksPolymer.from_data( + ds, + device="cpu", + compute_parameters=lambda x, **kwargs: (0.0, 1.0), + normalize_data=lambda x, lo, hi: x, + ) + + peaks = Vector.from_shape( + shape=(2, 2), + fields=["y_pixels", "x_pixels"], + units=["pixels", "pixels"], + name="cartesian_peaks", + ) + polar = Vector.from_shape( + shape=(2, 2), + fields=["r_invA", "theta"], + units=["1/Å", "radians"], + name="polar_peaks", + ) + intensities = Vector.from_shape( + shape=(2, 2), + fields=["intensities"], + units=["counts"], + name="peak_intensities", + ) + for i in range(2): + for j in range(2): + peaks[i, j] = np.array( + [[4.0, 4.0], [2.0, 6.0], [6.0, 2.0]], dtype=float + ) + polar[i, j] = np.array( + [[0.1, 0.0], [0.4, np.pi / 2], [0.7, np.pi]], dtype=float + ) + intensities[i, j] = np.array([[10.0], [20.0], [30.0]], dtype=float) + + bp.peak_coordinates_cartesian = peaks + bp.polar_peaks = polar + bp.peak_intensities = intensities + bp.polar_data = {"intensity": np.ones((2, 2, 5, 6), dtype=np.float32)} + bp.image_centers = np.zeros((2, 2, 2), dtype=float) + bp.image_centers[:, :, :] = np.array([4.0, 4.0])[:, None, None] + bp.max_radius_invA = 1.0 + bp.num_radial_bins = 5 + bp.num_annular_bins = 3 + bp.two_fold_symmetry = True + return bp + + +def test_descent_recovers_subpixel_centers(): + ny = nx = 96 + true_centers = [(47.3, 48.7), (48.4, 47.6), (46.8, 49.1), (49.0, 46.9)] + arr = np.stack([_ring_pattern(ny, nx, cy, cx) for cy, cx in true_centers]).reshape( + 2, 2, ny, nx + ) + + origins = find_origin_angular_descent( + _dataset(arr), + radial_min=4, + radial_max=36, + n_phi=96, + device="cpu", + ) + + assert origins.shape == (2, 2, 2) + for idx, (cy, cx) in enumerate(true_centers): + row, col = origins[idx // 2, idx % 2] + assert np.hypot(row - cy, col - cx) < 0.35 + + +def test_grid_recovers_center_on_small_detector(): + ny = nx = 64 + cy, cx = 30.4, 31.6 + arr = _ring_pattern(ny, nx, cy, cx, radii=(8, 16, 24))[None, None] + + origins = find_origin_angular_grid( + _dataset(arr), + radial_min=3, + radial_max=26, + num_annular_bins=72, + device="cpu", + ) + + assert origins.shape == (1, 1, 2) + assert np.hypot(origins[0, 0, 0] - cy, origins[0, 0, 1] - cx) < 0.75 + + +def test_dispatch_accepts_2d_arrays_and_rejects_bad_method(): + pattern = _ring_pattern(64, 64, 31.5, 31.5) + + origins = find_origin(pattern, method="descent", radial_min=4, radial_max=26, device="cpu") + + assert origins.shape == (1, 1, 2) + assert np.hypot(origins[0, 0, 0] - 31.5, origins[0, 0, 1] - 31.5) < 0.35 + with pytest.raises(ValueError, match="method"): + find_origin(pattern, method="peaks") + + +def test_descent_blank_pattern_returns_image_center(): + origins = find_origin_angular_descent( + np.zeros((32, 32), dtype=np.float32), + radial_min=4, + radial_max=12, + device="cpu", + ) + + assert origins.shape == (1, 1, 2) + assert np.allclose(origins[0, 0], [(32 - 1) / 2.0, (32 - 1) / 2.0], atol=1.0) + + +def test_bragg_peak_polar_transform_matches_image_polar_convention(): + ds = _dataset(np.zeros((1, 1, 16, 16), dtype=np.float32)) + peaks = Vector.from_shape( + shape=(1, 1), + fields=["y_pixels", "x_pixels", "y_invA", "x_invA"], + units=["pixels", "pixels", "1/Å", "1/Å"], + name="peaks", + ) + peaks[0, 0] = np.array( + [ + [8.0, 11.0, 0.0, 0.0], # +x axis -> theta 0 + [12.0, 8.0, 0.0, 0.0], # +y axis -> theta pi/2 + [4.0, 8.0, 0.0, 0.0], # -y axis -> theta 3pi/2, folded to pi/2 + [8.0, 5.0, 0.0, 0.0], # -x axis -> theta pi, folded to 0 + ], + dtype=float, + ) + bp = BraggPeaksPolymer.from_data( + ds, + device="cpu", + compute_parameters=lambda x, **kwargs: (0.0, 1.0), + normalize_data=lambda x, lo, hi: x, + ) + + polar = bp.polar_transform_peaks( + peaks, + centers=np.array([[[8.0]], [[8.0]]]), + two_fold_symmetry=True, + use_tqdm=False, + ) + + got = polar[0, 0].array + assert np.allclose(got[:, 0], [3.0, 4.0, 4.0, 3.0]) + assert np.allclose(got[:, 1], [0.0, np.pi / 2.0, np.pi / 2.0, 0.0]) + + +def test_bragg_polar_transform_two_fold_sums_opposite_angles(): + arr = np.zeros((1, 1, 7, 9), dtype=np.float32) + arr[0, 0, 5, 4] = 0.2 # 90 degrees at r=2 + arr[0, 0, 1, 4] = 0.3 # 270 degrees at r=2 + ds = _dataset(arr) + bp = BraggPeaksPolymer.from_data( + ds, + device="cpu", + compute_parameters=lambda x, **kwargs: (0.0, 1.0), + normalize_data=lambda x, lo, hi: x, + ) + centers = np.array([[[3.0]], [[4.0]]]) + + full = bp.polar_transform_4d( + ds, + centers=centers, + num_r=5, + num_theta=4, + two_fold_symmetry=False, + use_tqdm=False, + ) + folded = bp.polar_transform_4d( + ds, + centers=centers, + num_r=5, + num_theta=4, + two_fold_symmetry=True, + use_tqdm=False, + ) + + assert full["intensity"].shape == (1, 1, 5, 4) + assert folded["intensity"].shape == (1, 1, 5, 2) + assert np.allclose( + folded["intensity"][0, 0], + full["intensity"][0, 0, :, :2] + full["intensity"][0, 0, :, 2:], + ) + assert folded["intensity"][0, 0, 2, 1] == pytest.approx(0.5, abs=1e-6) + + +def test_bragg_peak_polar_transform_inverts_ellipse_mapping(): + ds = _dataset(np.zeros((1, 1, 17, 17), dtype=np.float32)) + peaks = Vector.from_shape( + shape=(1, 1), + fields=["y_pixels", "x_pixels"], + units=["pixels", "pixels"], + name="peaks", + ) + peaks[0, 0] = np.array([[8.0, 14.0]], dtype=float) + bp = BraggPeaksPolymer.from_data( + ds, + device="cpu", + compute_parameters=lambda x, **kwargs: (0.0, 1.0), + normalize_data=lambda x, lo, hi: x, + ) + + polar = bp.polar_transform_peaks( + peaks, + centers=np.array([[[8.0]], [[8.0]]]), + two_fold_symmetry=False, + ellipse_params=(2.0, 1.0, 0.0), + use_tqdm=False, + ) + + got = polar[0, 0].array + # r_pixels, theta, r_invA, theta_unfolded + assert polar.fields == ["r_pixels", "theta", "r_invA", "theta_unfolded"] + assert got.shape == (1, 4) + assert got[0, 0] == pytest.approx(3.0) + assert got[0, 1] == pytest.approx(0.0) + # Unfolded theta agrees with theta here because two_fold_symmetry=False. + assert got[0, 3] == pytest.approx(0.0) + + +def _bragg_orientation_histogram(theta_values, theta_step_deg=90, normalize_stack=False): + ds = _dataset(np.zeros((1, 1, 8, 8), dtype=np.float32)) + bp = BraggPeaksPolymer.from_data( + ds, + device="cpu", + compute_parameters=lambda x, **kwargs: (0.0, 1.0), + normalize_data=lambda x, lo, hi: x, + ) + theta_values = np.asarray(theta_values, dtype=float) + + polar = Vector.from_shape( + shape=(1, 1), + fields=["r_invA", "theta"], + units=["1/Å", "radians"], + name="polar_peaks", + ) + polar[0, 0] = np.column_stack([np.ones_like(theta_values), theta_values]) + intensities = Vector.from_shape( + shape=(1, 1), + fields=["intensities"], + units=["counts"], + name="peak_intensities", + ) + intensities[0, 0] = np.ones((theta_values.size, 1), dtype=float) + + bp.polar_peaks = polar + bp.peak_intensities = intensities + return bp.make_orientation_histogram( + radial_ranges=np.array([0.5, 1.5]), + upsample_factor=2, + theta_step_deg=theta_step_deg, + sigma_x=None, + sigma_y=None, + sigma_theta=None, + normalize_intensity_image=False, + normalize_intensity_stack=normalize_stack, + progress_bar=False, + ) + + +def test_bragg_orientation_histogram_preserves_karen_angles(): + hist = _bragg_orientation_histogram([0.0, np.pi / 2]) + + assert hist[0, 0, 0, 0] == pytest.approx(1.0) + assert hist[0, 0, 0, 1] == pytest.approx(1.0) + + +def test_bragg_orientation_histogram_folds_unwrapped_angles(): + hist = _bragg_orientation_histogram([0.0, np.pi, np.pi / 2, 3 * np.pi / 2]) + + assert hist[0, 0, 0, 0] == pytest.approx(2.0) + assert hist[0, 0, 0, 1] == pytest.approx(2.0) + + +def test_bragg_orientation_histogram_empty_input_remains_finite(): + hist = _bragg_orientation_histogram([], normalize_stack=True) + assert np.isfinite(hist).all() + assert not np.any(hist) + + +def test_bragg_private_helpers_characterize_shared_plotting_behavior(): + ds = _dataset(np.arange(2 * 2 * 4 * 4, dtype=np.float32).reshape(2, 2, 4, 4)) + + mean_map = _mean_intensity_map(ds, (2, 2)) + assert mean_map.shape == (2, 2) + assert mean_map[0, 0] == pytest.approx(np.mean(ds[0, 0].array)) + + resolved, upsample = _resolve_intensity_map(ds, None, (2, 2)) + assert upsample == 1 + assert np.allclose(resolved, mean_map) + custom = np.zeros((4, 4), dtype=float) + resolved, upsample = _resolve_intensity_map(ds, custom, (2, 2), validate=True) + assert resolved is custom + assert upsample == 2 + with pytest.raises(ValueError, match="integer multiple"): + _resolve_intensity_map(ds, np.zeros((5, 4)), (2, 2), validate=True) + + is_rgb, vmin, vmax = _intensity_display_limits(np.dstack([custom, custom, custom])) + assert is_rgb is True + assert vmin is None and vmax is None + is_rgb, vmin, vmax = _intensity_display_limits(np.array([[0.0, 1.0], [2.0, 3.0]])) + assert is_rgb is False + assert vmin == pytest.approx(0.03) + assert vmax == pytest.approx(2.97) + + normalized = _normalized_dp( + ds, + 0, + 0, + norm_upper_quantile=0.5, + norm_power=2.0, + ) + clipped = np.clip(ds[0, 0].array, 0, np.quantile(ds[0, 0].array, 0.5)) + expected = (clipped / np.nanmax(clipped)) ** 2.0 * np.nanmax(clipped) + assert np.allclose(normalized, expected) + + assert _display_center(None, 0, 0, (4, 6)) == (2.0, 3.0) + centers = np.zeros((2, 2, 2), dtype=float) + centers[:, 1, 1] = [1.5, 2.5] + assert _display_center(centers, 1, 1, (4, 6)) == pytest.approx((1.5, 2.5)) + + peaks_x = np.array([3.0, 5.0]) + peaks_y = np.array([3.0, 1.0]) + peaks_r = np.array([1.0, 0.5]) + assert _central_peak_index(peaks_x, peaks_y, peaks_r, (3.0, 3.0)) == 0 + assert _central_peak_index(peaks_x, peaks_y, None, (3.0, 3.0)) == 0 + + cropped, zx, zy, zr, zi, zcentral, display_center = _zoom_peak_overlay( + np.zeros((6, 6)), + peaks_x, + peaks_y, + peaks_r, + np.array([10.0, 20.0]), + 0, + 2, + (3.0, 3.0), + ) + assert cropped.shape == (3, 3) + assert np.allclose(zx, [1.0]) + assert np.allclose(zy, [1.0]) + assert np.allclose(zr, [1.0]) + assert np.allclose(zi, [10.0]) + assert zcentral == 0 + assert display_center == pytest.approx((1.0, 1.0)) + + cropped, zx, zy, zr, zi, zcentral, display_center = _zoom_peak_overlay( + np.zeros((8, 8)), + np.array([1.0, 4.0]), + np.array([1.0, 4.0]), + np.array([0.1, 2.0]), + np.array([10.0, 20.0]), + 0, + 2, + (4.0, 4.0), + ) + assert cropped.shape == (4, 4) + assert np.allclose(zx, [2.0]) + assert np.allclose(zy, [2.0]) + assert np.allclose(zr, [2.0]) + assert np.allclose(zi, [20.0]) + assert zcentral is None + assert display_center == pytest.approx((2.0, 2.0)) + + r_bins, theta_bins = _polar_peak_bins( + np.array([1.0, 2.0]), + np.array([np.pi / 2, np.pi]), + max_radius_invA=2.0, + num_radial_bins=10, + num_annular_bins=180, + two_fold_symmetry=True, + ) + assert np.allclose(r_bins, [5.0, 10.0]) + assert np.allclose(theta_bins, [90.0, 180.0]) + + +def test_bragg_plotting_and_save_smoke(monkeypatch, tmp_path): + bp = _bragg_for_plotting() + + def fake_interactive_output(fn, controls): + fn(**{name: widget.value for name, widget in controls.items()}) + return bragg_peaks_module.widgets.Output() + + monkeypatch.setattr(bragg_peaks_module, "interactive_output", fake_interactive_output) + monkeypatch.setattr(bragg_peaks_module, "display", lambda *args, **kwargs: None, raising=False) + monkeypatch.setattr(bragg_peaks_module, "clear_output", lambda *args, **kwargs: None) + + bp.plot_interactive_image_map(ry=0, rx=0, show_polar=False) + bp.plot_interactive_peak_map(ry=0, rx=0, show_polar=True) + + bp.save_diffraction_figures(0, 0, save_dir=tmp_path / "diff", show_polar=True) + assert (tmp_path / "diff" / "diffraction_ry0_rx0_intensity_map.pdf").exists() + assert (tmp_path / "diff" / "diffraction_ry0_rx0_diffraction.pdf").exists() + assert (tmp_path / "diff" / "diffraction_ry0_rx0_polar.pdf").exists() + assert (tmp_path / "diff" / "diffraction_ry0_rx0_combined.pdf").exists() + + bp.save_peak_figures(0, 0, save_dir=tmp_path / "peaks", show_polar=True) + assert (tmp_path / "peaks" / "peaks_ry0_rx0_intensity_map.pdf").exists() + assert (tmp_path / "peaks" / "peaks_ry0_rx0_diffraction.pdf").exists() + assert (tmp_path / "peaks" / "peaks_ry0_rx0_polar.pdf").exists() + plt.close("all") diff --git a/tests/diffraction/test_polymer_ice.py b/tests/diffraction/test_polymer_ice.py new file mode 100644 index 000000000..ee9289277 --- /dev/null +++ b/tests/diffraction/test_polymer_ice.py @@ -0,0 +1,563 @@ +from __future__ import annotations + +import numpy as np +import pytest + +from quantem.core.datastructures import Vector +from quantem.diffraction import (IceFlaggerParams, detect_ice, + measure_peak_widths, sharpness_mask) + + +def _vectors(shape=(1, 2)): + polar = Vector.from_shape( + shape=shape, fields=["r_invA", "theta"], units=["1/A", "rad"] + ) + intensity = Vector.from_shape( + shape=shape, fields=["intensities"], units=["normalized"] + ) + return polar, intensity + + +def test_detect_ice_and_filter_does_not_mutate_source(): + polar, intensity = _vectors() + polar[0, 0] = np.column_stack( + [np.full(3, 1.61), np.deg2rad([5, 65, 140])] + ) + intensity[0, 0] = np.array([[0.9], [0.8], [0.7]]) + polar[0, 1] = np.empty((0, 2)) + intensity[0, 1] = np.empty((0, 1)) + original = polar[0, 0].array.copy() + + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams( + intensity_cutoff=0.5, min_matches=2, dtheta_deg=6 + ), + return_debug=True, + ) + assert result.threshold == 0.5 + assert result.flagged_peaks_count_map.tolist() == [[2, 0]] + filtered = result.filter(polar) + np.testing.assert_array_equal(polar[0, 0].array, original) + assert len(filtered[0, 0].array) == 1 + assert (0, 1) in result.debug_records + + +def test_masked_cells_are_not_analyzed(): + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack( + [np.full(2, 1.61), np.deg2rad([0, 60])] + ) + intensity[0, 0] = np.ones((2, 1)) + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.0), + scan_mask=np.zeros((1, 1), dtype=bool), + ) + assert result.flagged_peaks_count_map[0, 0] == 0 + + +def test_misaligned_ragged_vectors_fail_clearly(): + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.zeros((2, 2)) + intensity[0, 0] = np.zeros((1, 1)) + with pytest.raises(ValueError, match="Row count mismatch"): + detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.0), + ) + + +def _polar_volume(peaks, *, shape=(1, 1), n_r=600, n_theta=180, r_max=3.0): + """Polar volume with a Gaussian blob per (q, theta_deg, width_q, width_deg).""" + r_axis = np.linspace(0.0, r_max, n_r) + theta_axis = np.linspace(0.0, np.pi, n_theta, endpoint=False) + r_grid, theta_grid = np.meshgrid(r_axis, theta_axis, indexing="ij") + intensity = np.zeros(shape + (n_r, n_theta)) + for q, theta_deg, width_q, width_deg in peaks: + # width_* are FWHM; convert to the Gaussian sigma that produces them. + sigma_q = width_q / (2 * np.sqrt(2 * np.log(2))) + sigma_theta = np.deg2rad(width_deg) / (2 * np.sqrt(2 * np.log(2))) + delta = np.abs(theta_grid - np.deg2rad(theta_deg)) + delta = np.minimum(delta, np.pi - delta) + intensity += np.exp( + -0.5 * (((r_grid - q) / sigma_q) ** 2 + (delta / sigma_theta) ** 2) + ) + return {"intensity": intensity, "r_invA": r_grid, "theta": theta_grid} + + +def test_measured_widths_recover_the_input_blob_widths(): + from quantem.diffraction.polymer_ice import measure_peak_widths + + polar_data = _polar_volume([(1.61, 30.0, 0.10, 8.0)]) + width_r, width_theta = measure_peak_widths( + [1.61], + [np.deg2rad(30.0)], + polar_data["intensity"][0, 0], + polar_data["r_invA"][:, 0], + polar_data["theta"][0, :], + params=IceFlaggerParams(sharpness_baseline_quantile=0.0), + ) + assert width_r[0] == pytest.approx(0.10, abs=0.03) + assert width_theta[0] == pytest.approx(8.0, abs=2.0) + + +def test_sharpness_gate_keeps_sharp_ice_and_spares_broad_peaks(): + """Two aligned six-fold peaks: one sharp (ice), one broad (polymer).""" + sharp = (1.61, 5.0, 0.04, 5.0) + broad = (1.61, 65.0, 0.30, 40.0) + polar_data = _polar_volume([sharp, broad]) + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack([[sharp[0], broad[0]], np.deg2rad([sharp[1], broad[1]])]) + intensity[0, 0] = np.array([[0.9], [0.9]]) + + base = dict(intensity_cutoff=0.5, min_matches=2, dtheta_deg=6.0, q_target_invA=1.61) + # Without the gate both peaks are aligned six-fold candidates and both go. + ungated = detect_ice(polar, intensity, params=IceFlaggerParams(**base)) + assert ungated.flagged_peaks_count_map[0, 0] == 2 + + # With the gate only the sharp one survives as a candidate, and a single + # candidate no longer reaches min_matches=2, so nothing is flagged. + gated = detect_ice( + polar, + intensity, + params=IceFlaggerParams( + **base, max_width_r_invA=0.10, max_width_theta_deg=15.0 + ), + polar_data=polar_data, + return_debug=True, + ) + assert gated.debug_records[(0, 0)].n_candidates_sharp == 1 + assert gated.flagged_peaks_count_map[0, 0] == 0 + + +def test_annular_only_gate_keeps_dots_and_radial_streaks(): + """Ice is annularly sharp whether it is a dot or a radial streak. + + A radial streak is narrow in theta and extended in q, so a radial ceiling + would reject it; gating on the annular width alone keeps both ice shapes and + still rejects the annularly broad polymer arc at the same q. + """ + dot = (1.61, 5.0, 0.02, 5.0) + streak = (1.61, 5.0, 0.12, 5.0) + arc = (1.61, 5.0, 0.02, 30.0) + params = IceFlaggerParams(max_width_theta_deg=12.0, sharpness_window_r_invA=0.30) + for label, (q, theta_deg, width_q, width_deg), expected in ( + ("dot", dot, True), ("streak", streak, True), ("arc", arc, False) + ): + polar_data = _polar_volume([(q, theta_deg, width_q, width_deg)]) + width_r, width_theta = measure_peak_widths( + [q], [np.deg2rad(theta_deg)], polar_data["intensity"][0, 0], + polar_data["r_invA"][:, 0], polar_data["theta"][0, :], params=params, + ) + assert bool(sharpness_mask(width_r, width_theta, params)[0]) is expected, label + + +def test_sharpness_ceiling_without_polar_data_fails_clearly(): + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack([[1.61], [0.0]]) + intensity[0, 0] = np.array([[0.9]]) + with pytest.raises(ValueError, match="requires polar_data"): + detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, max_width_r_invA=0.1), + ) + + +def _two_lattice_pattern(phi_a=3.0, phi_b=31.0, q=1.61): + """Two six-fold lattices at unrelated orientations, overlaid in one pattern.""" + angles = [phi_a + 60.0 * k for k in range(6)] + [phi_b + 60.0 * k for k in range(6)] + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack([np.full(len(angles), q), np.deg2rad(angles)]) + intensity[0, 0] = np.full((len(angles), 1), 0.9) + return polar, intensity + + +def test_single_crystallite_default_finds_only_the_strongest_lattice(): + polar, intensity = _two_lattice_pattern() + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, min_matches=3, dtheta_deg=6.0), + return_debug=True, + ) + # 12 peaks present, only one lattice's 6 claimed. + assert result.flagged_peaks_count_map[0, 0] == 6 + assert len(result.debug_records[(0, 0)].phi_deg) == 1 + + +def test_max_crystallites_claims_both_lattices(): + polar, intensity = _two_lattice_pattern(phi_a=3.0, phi_b=31.0) + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams( + intensity_cutoff=0.5, min_matches=3, dtheta_deg=6.0, max_crystallites=3 + ), + return_debug=True, + ) + assert result.flagged_peaks_count_map[0, 0] == 12 + found = sorted(result.debug_records[(0, 0)].phi_deg) + assert len(found) == 2 + assert found[0] == pytest.approx(3.0, abs=0.5) + assert found[1] == pytest.approx(31.0, abs=0.5) + + +def test_peel_stops_instead_of_splitting_one_lattice_in_two(): + """A single lattice must not be re-found as a near-duplicate crystallite.""" + polar, intensity = _vectors((1, 1)) + angles = [5.0 + 60.0 * k for k in range(6)] + polar[0, 0] = np.column_stack([np.full(6, 1.61), np.deg2rad(angles)]) + intensity[0, 0] = np.full((6, 1), 0.9) + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams( + intensity_cutoff=0.5, min_matches=3, dtheta_deg=6.0, max_crystallites=5 + ), + return_debug=True, + ) + assert result.flagged_peaks_count_map[0, 0] == 6 + assert len(result.debug_records[(0, 0)].phi_deg) == 1 + + +def test_min_phi_separation_rejects_a_too_close_second_lattice(): + # 8 degrees apart: separable at the default (dtheta_deg=6), not at 15. + polar, intensity = _two_lattice_pattern(phi_a=3.0, phi_b=11.0) + common = dict(intensity_cutoff=0.5, min_matches=3, dtheta_deg=3.0, max_crystallites=3) + both = detect_ice(polar, intensity, params=IceFlaggerParams(**common), return_debug=True) + merged = detect_ice( + polar, + intensity, + params=IceFlaggerParams(**common, min_phi_separation_deg=15.0), + return_debug=True, + ) + assert len(both.debug_records[(0, 0)].phi_deg) == 2 + assert len(merged.debug_records[(0, 0)].phi_deg) == 1 + + +def test_matched_bins_accumulate_across_crystallites(): + polar, intensity = _two_lattice_pattern() + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams( + intensity_cutoff=0.5, min_matches=3, dtheta_deg=6.0, max_crystallites=3 + ), + return_debug=True, + ) + # Six bins per lattice, two lattices. + assert result.matched_bins_count_map[0, 0] == 12 + + +def test_folded_theta_matches_a_friedel_pair_as_one_lattice(): + """With two-fold folding, theta and theta+180 are the same angle. + + The pair must still be flagged; before the period was honoured it counted as a + single bin and was rejected by min_matches=2. + """ + polar, intensity = _vectors((1, 1)) + # As process_polar(two_fold_symmetry=True) would deliver them: both folded to 5 deg. + polar[0, 0] = np.column_stack([np.full(2, 1.61), np.deg2rad([5.0, 5.0])]) + intensity[0, 0] = np.full((2, 1), 0.9) + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, min_matches=1, dtheta_deg=6.0), + theta_period_deg=180.0, + ) + assert result.flagged_peaks_count_map[0, 0] == 2 + + +def test_folded_theta_reaches_all_three_arms(): + """Three arms 60 deg apart are all reachable on a 180 deg period.""" + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack([np.full(3, 1.61), np.deg2rad([5.0, 65.0, 125.0])]) + intensity[0, 0] = np.full((3, 1), 0.9) + params = IceFlaggerParams(intensity_cutoff=0.5, min_matches=3, dtheta_deg=6.0) + folded = detect_ice(polar, intensity, params=params, theta_period_deg=180.0, + return_debug=True) + assert folded.flagged_peaks_count_map[0, 0] == 3 + assert len(folded.debug_records[(0, 0)].matched_bins) == 3 + + +def test_unsatisfiable_min_matches_on_folded_theta_is_rejected(): + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack([np.full(2, 1.61), np.deg2rad([5.0, 65.0])]) + intensity[0, 0] = np.full((2, 1), 0.9) + with pytest.raises(ValueError, match="can never be reached"): + detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, min_matches=5), + theta_period_deg=180.0, + ) + + +def test_params_theta_period_overrides_the_caller(): + polar, intensity = _vectors((1, 1)) + polar[0, 0] = np.column_stack([np.full(2, 1.61), np.deg2rad([5.0, 5.0])]) + intensity[0, 0] = np.full((2, 1), 0.9) + # Caller says folded, params insist on the full circle: params win, so min_matches=5 + # becomes reachable in principle and no error is raised. + result = detect_ice( + polar, + intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, min_matches=5, theta_period_deg=360.0), + theta_period_deg=180.0, + ) + assert result.flagged_peaks_count_map[0, 0] == 0 + + +def _folded(angles_deg, intensity=0.9): + polar, intensity_vec = _vectors((1, 1)) + polar[0, 0] = np.column_stack( + [np.full(len(angles_deg), 1.61), np.deg2rad(angles_deg)] + ) + intensity_vec[0, 0] = np.full((len(angles_deg), 1), intensity) + return polar, intensity_vec + + +def test_min_peaks_per_arm_rejects_a_lone_peak_but_keeps_a_friedel_pair(): + """On a folded axis a Friedel pair is two peaks on one arm; a lone peak is one.""" + common = dict(intensity_cutoff=0.5, dtheta_deg=6.0, min_matches=1, min_peaks_per_arm=2) + + lone_polar, lone_int = _folded([5.0]) + lone = detect_ice(lone_polar, lone_int, params=IceFlaggerParams(**common), + theta_period_deg=180.0) + assert lone.flagged_peaks_count_map[0, 0] == 0 + + # theta and theta+180 both fold to 5 deg. + pair_polar, pair_int = _folded([5.0, 5.0]) + pair = detect_ice(pair_polar, pair_int, params=IceFlaggerParams(**common), + theta_period_deg=180.0) + assert pair.flagged_peaks_count_map[0, 0] == 2 + + +def test_min_peaks_per_arm_default_is_unchanged_behaviour(): + polar, intensity = _folded([5.0]) + result = detect_ice( + polar, intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, dtheta_deg=6.0, min_matches=1), + theta_period_deg=180.0) + assert result.flagged_peaks_count_map[0, 0] == 1 + + +def test_arms_below_the_peak_floor_are_dropped_not_just_uncounted(): + """An under-populated arm must not contribute its peaks to the flagged set.""" + # Arm A (5 deg) has a pair, arm B (65 deg) has a single peak. + polar, intensity = _folded([5.0, 5.0, 65.0]) + result = detect_ice( + polar, intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, dtheta_deg=6.0, + min_matches=1, min_peaks_per_arm=2), + theta_period_deg=180.0, return_debug=True) + assert result.flagged_peaks_count_map[0, 0] == 2 # the pair only + assert result.debug_records[(0, 0)].matched_bins == [0] # arm B dropped + + +def test_min_peaks_per_arm_combines_with_min_matches(): + # Two arms, each a Friedel pair -> 2 arms of 2 peaks. + polar, intensity = _folded([5.0, 5.0, 65.0, 65.0]) + ok = detect_ice( + polar, intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, dtheta_deg=6.0, + min_matches=2, min_peaks_per_arm=2), + theta_period_deg=180.0) + assert ok.flagged_peaks_count_map[0, 0] == 4 + # Same peaks, but demanding three populated arms: nothing qualifies. + strict = detect_ice( + polar, intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, dtheta_deg=6.0, + min_matches=3, min_peaks_per_arm=2), + theta_period_deg=180.0) + assert strict.flagged_peaks_count_map[0, 0] == 0 + + +def _folded_with_unfolded(pairs): + """pairs: list of (folded_deg, unfolded_deg) as polar_transform_peaks records them.""" + polar = Vector.from_shape( + shape=(1, 1), + fields=["r_invA", "theta", "theta_unfolded"], + units=["1/A", "rad", "rad"], + ) + intensity = Vector.from_shape(shape=(1, 1), fields=["intensities"], units=["normalized"]) + polar[0, 0] = np.column_stack([ + np.full(len(pairs), 1.61), + np.deg2rad([f for f, _ in pairs]), + np.deg2rad([u for _, u in pairs]), + ]) + intensity[0, 0] = np.full((len(pairs), 1), 0.9) + return polar, intensity + + +def test_require_friedel_pair_distinguishes_a_true_pair_from_two_near_peaks(): + """min_peaks_per_arm cannot tell these apart; the unfolded angle can.""" + common = dict(intensity_cutoff=0.5, dtheta_deg=6.0, min_matches=1, + min_peaks_per_arm=2, require_friedel_pair=True) + + # Genuinely opposed: 5 and 185 deg, both folding to 5. + true_pair, ints = _folded_with_unfolded([(5.0, 5.0), (5.0, 185.0)]) + assert detect_ice(true_pair, ints, params=IceFlaggerParams(**common), + theta_period_deg=180.0).flagged_peaks_count_map[0, 0] == 2 + + # Two peaks on the same side, 3 deg apart: same arm, same folded angle, not a pair. + near, ints2 = _folded_with_unfolded([(5.0, 5.0), (8.0, 8.0)]) + assert detect_ice(near, ints2, params=IceFlaggerParams(**common), + theta_period_deg=180.0).flagged_peaks_count_map[0, 0] == 0 + + # Without the strict test, min_peaks_per_arm=2 accepts the near pair. + loose = dict(common, require_friedel_pair=False) + assert detect_ice(near, ints2, params=IceFlaggerParams(**loose), + theta_period_deg=180.0).flagged_peaks_count_map[0, 0] == 2 + + +def test_require_friedel_pair_without_the_field_fails_clearly(): + polar, intensity = _folded([5.0, 5.0]) # no theta_unfolded column + with pytest.raises(ValueError, match="theta_unfolded"): + detect_ice( + polar, intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, require_friedel_pair=True), + theta_period_deg=180.0) + + +def test_unfolded_field_is_optional_when_not_required(): + """Vectors predating theta_unfolded still work for everything else.""" + polar, intensity = _folded([5.0, 5.0]) + result = detect_ice( + polar, intensity, + params=IceFlaggerParams(intensity_cutoff=0.5, dtheta_deg=6.0, + min_matches=1, min_peaks_per_arm=2), + theta_period_deg=180.0) + assert result.flagged_peaks_count_map[0, 0] == 2 + + +def _annular_arc(fwhm_deg, q=0.27, theta_deg=45.0, n_theta=90, n_r=200, seed=0): + theta_axis = np.deg2rad(np.arange(0, 180, 180 / n_theta)) + r_axis = np.linspace(0.0, 1.0, n_r) + grid_r, grid_t = np.meshgrid(r_axis, theta_axis, indexing="ij") + sigma_q = 0.02 / (2 * np.sqrt(2 * np.log(2))) + sigma_t = np.deg2rad(fwhm_deg) / (2 * np.sqrt(2 * np.log(2))) + delta = np.abs(grid_t - np.deg2rad(theta_deg)) + delta = np.minimum(delta, np.pi - delta) + image = np.exp(-0.5 * (((grid_r - q) / sigma_q) ** 2 + (delta / sigma_t) ** 2)) + image += np.random.default_rng(seed).normal(0, 0.02, image.shape) + return image, r_axis, theta_axis, q, theta_deg + + +@pytest.mark.parametrize("true_fwhm", [6.0, 30.0, 60.0, 100.0]) +def test_annular_width_tracks_broad_arcs_not_just_sharp_ones(true_fwhm): + """The default window must not saturate: a broad arc has to read broad. + + With a 40 degree window and a 0.25 baseline quantile, a 60 degree arc measured + ~37 degrees and a 100 degree arc ~39 -- both look sharper than a real threshold. + """ + from quantem.diffraction.polymer_ice import measure_peak_widths + + image, r_axis, theta_axis, q, theta_deg = _annular_arc(true_fwhm) + _, width_theta = measure_peak_widths( + [q], [np.deg2rad(theta_deg)], image, r_axis, theta_axis, + params=IceFlaggerParams(), + ) + assert width_theta[0] == pytest.approx(true_fwhm, rel=0.15) + + +def test_annular_window_cannot_wrap_onto_itself(): + """A window wider than the annulus must be clamped, not allowed to repeat bins.""" + from quantem.diffraction.polymer_ice import measure_peak_widths + + image, r_axis, theta_axis, q, theta_deg = _annular_arc(6.0) + _, width_theta = measure_peak_widths( + [q], [np.deg2rad(theta_deg)], image, r_axis, theta_axis, + params=IceFlaggerParams(sharpness_window_theta_deg=10_000.0), + ) + assert width_theta[0] == pytest.approx(6.0, rel=0.2) + + +@pytest.mark.parametrize("amplitude,noise,true_fwhm", [ + (1.0, 0.01, 5.0), # sharp and clean + (1.0, 0.01, 40.0), # broad and clean + (0.3, 0.05, 40.0), # broad and weak -- the case that used to read as sharp +]) +def test_annular_width_survives_noise(amplitude, noise, true_fwhm): + """A weak, diffuse arc must not measure as narrow as a sharp peak. + + A half-maximum walk that stops at the first sample below half is ended early by + one downward fluctuation, so before smoothing + persistent crossings a true 40 + degree arc at low SNR measured ~10-30 degrees. + """ + from quantem.diffraction.polymer_ice import measure_peak_widths + + n_theta = 90 + theta_axis = np.deg2rad(np.arange(0, 180, 2.0)) + r_axis = np.linspace(0.0, 1.0, 300) + grid_r, grid_t = np.meshgrid(r_axis, theta_axis, indexing="ij") + delta = np.abs(grid_t - np.deg2rad(45.0)) + delta = np.minimum(delta, np.pi - delta) + sigma_t = np.deg2rad(true_fwhm) / 2.3548 + image = amplitude * np.exp( + -0.5 * (((grid_r - 0.27) / (0.02 / 2.3548)) ** 2 + (delta / sigma_t) ** 2) + ) + image += 0.3 * np.exp(-0.5 * ((grid_r - 0.27) / (0.05 / 2.3548)) ** 2) # amorphous ring + + measured = [] + for seed in range(8): + noisy = image + np.random.default_rng(seed).normal(0, noise, image.shape) + _, width_theta = measure_peak_widths( + [0.27], [np.deg2rad(45.0)], noisy, r_axis, theta_axis, + params=IceFlaggerParams(), + ) + measured.append(width_theta[0]) + assert np.median(measured) == pytest.approx(true_fwhm, rel=0.25) + + +def test_smoothing_is_deconvolved_so_sharp_peaks_stay_unbiased(): + """The Gaussian kernel is removed in quadrature, exactly for Gaussian peaks.""" + from quantem.diffraction.polymer_ice import measure_peak_widths + + theta_axis = np.deg2rad(np.arange(0, 180, 2.0)) + r_axis = np.linspace(0.0, 1.0, 300) + grid_r, grid_t = np.meshgrid(r_axis, theta_axis, indexing="ij") + delta = np.abs(grid_t - np.deg2rad(45.0)) + delta = np.minimum(delta, np.pi - delta) + image = np.exp(-0.5 * (((grid_r - 0.27) / (0.02 / 2.3548)) ** 2 + + (delta / (np.deg2rad(6.0) / 2.3548)) ** 2)) + smoothed, raw = ( + measure_peak_widths([0.27], [np.deg2rad(45.0)], image, r_axis, theta_axis, + params=p)[1][0] + for p in (IceFlaggerParams(), + IceFlaggerParams(sharpness_smooth_sigma_bins=0.0, + sharpness_crossing_persistence=1)) + ) + # Deconvolution keeps the smoothed estimate within a bin of the raw one. + assert abs(smoothed - raw) < 2.0 + + +@pytest.mark.parametrize("kwargs,match", [ + ({"dq_invA": 0}, "dq_invA must be positive"), + ({"dtheta_deg": -1}, "dtheta_deg must be positive"), + ({"min_matches": 0}, "min_matches must be at least 1"), + ({"min_peaks_per_arm": 0}, "min_peaks_per_arm must be at least 1"), + ({"max_crystallites": 0}, "max_crystallites must be at least 1"), + ({"sharpness_crossing_persistence": 0}, "sharpness_crossing_persistence"), + ({"max_width_theta_deg": 0}, "max_width_theta_deg must be positive when set"), + ({"sharpness_baseline_quantile": 1.0}, "must be in \\[0, 1\\)"), + ({"sharpness_smooth_sigma_bins": -1}, "must be non-negative"), + ({"intensity_cutoff_mode": "nope"}, "must be 'absolute' or 'percentile'"), +]) +def test_invalid_params_are_rejected_at_construction(kwargs, match): + """A bad value here would otherwise surface as 'nothing was flagged'.""" + with pytest.raises(ValueError, match=match): + IceFlaggerParams(**kwargs) + + +def test_valid_edge_values_are_accepted(): + IceFlaggerParams( + max_width_r_invA=None, max_width_theta_deg=None, + sharpness_baseline_quantile=0.0, sharpness_smooth_sigma_bins=0.0, + min_matches=1, min_peaks_per_arm=1, max_crystallites=1, + ) diff --git a/tests/diffraction/test_polymer_models.py b/tests/diffraction/test_polymer_models.py new file mode 100644 index 000000000..a03bfb2d9 --- /dev/null +++ b/tests/diffraction/test_polymer_models.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import hashlib +import io + +import pytest + +from quantem.diffraction.polymer_models import PolymerModelError, resolve_polymer_model + + +def _registry(payload: bytes, *, architecture=None): + checksum = hashlib.sha256(payload).hexdigest() + spec = { + "schema_version": 1, + "model_id": "test-model", + "architecture": architecture + or {"input_channels": 1, "output_channels": 2, "kernel_size": 3}, + "weights": {"filename": "weights.pth", "sha256": checksum}, + } + return {"test-model": {"v1": {"specification": spec, "weights_url": "mock://weights"}}} + + +def test_download_then_offline_cache_hit(tmp_path): + payload = b"immutable model weights" + calls = [] + + def download(url): + calls.append(url) + return io.BytesIO(payload) + + first = resolve_polymer_model( + "test-model", version="v1", cache_dir=tmp_path, registry=_registry(payload), downloader=download + ) + assert first.weights_path.read_bytes() == payload + assert first.weights_path.parent.name == first.checksum + second = resolve_polymer_model( + "test-model", version="v1", cache_dir=tmp_path, registry=_registry(payload), + downloader=lambda _: pytest.fail("offline cache hit attempted a download"), + ) + assert second == first + assert calls == ["mock://weights"] + + +def test_local_override_uses_verified_weights(tmp_path): + payload = b"private weights" + model_dir = tmp_path / "test-model" / "v1" + model_dir.mkdir(parents=True) + (model_dir / "weights.pth").write_bytes(payload) + result = resolve_polymer_model( + "test-model", version="v1", local_model_dir=tmp_path, registry=_registry(payload) + ) + assert result.weights_path == (model_dir / "weights.pth").resolve() + + +def test_checksum_failure_leaves_no_partial_file(tmp_path): + with pytest.raises(PolymerModelError, match="Checksum failure after downloading"): + resolve_polymer_model( + "test-model", version="v1", cache_dir=tmp_path, registry=_registry(b"expected"), + downloader=lambda _: io.BytesIO(b"corrupt"), + ) + assert not list(tmp_path.rglob("weights.pth")) + assert not list(tmp_path.rglob(".download-*")) + + +def test_interrupted_download_leaves_no_partial_file(tmp_path): + class Interrupted(io.BytesIO): + def read(self, size=-1): + raise OSError("connection interrupted") + + with pytest.raises(PolymerModelError, match="Could not download"): + resolve_polymer_model( + "test-model", version="v1", cache_dir=tmp_path, registry=_registry(b"expected"), + downloader=lambda _: Interrupted(b"partial"), + ) + assert not list(tmp_path.rglob(".download-*")) + + +def test_incompatible_specification_is_rejected(tmp_path): + with pytest.raises(PolymerModelError, match="incompatible"): + resolve_polymer_model( + "test-model", version="v1", cache_dir=tmp_path, + registry=_registry(b"weights", architecture={"input_channels": 3, "output_channels": 2, "kernel_size": 3}), + ) + + +def test_latest_is_explicit_and_exclusive(tmp_path): + payload = b"weights" + versions = _registry(payload)["test-model"] + versions["v2"] = versions["v1"] + result = resolve_polymer_model( + "test-model", latest=True, cache_dir=tmp_path, registry={"test-model": versions}, + downloader=lambda _: io.BytesIO(payload), + ) + assert result.version == "v2" + with pytest.raises(ValueError, match="mutually exclusive"): + resolve_polymer_model("test-model", version="v1", latest=True, registry={"test-model": versions}) diff --git a/tests/diffraction/test_polymer_normalization.py b/tests/diffraction/test_polymer_normalization.py new file mode 100644 index 000000000..77614d51b --- /dev/null +++ b/tests/diffraction/test_polymer_normalization.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +import quantem.diffraction.bragg_peaks as bragg_peaks_module +from quantem.diffraction import ( + BraggPeaksPolymer, + GlobalPercentileNormalization, + LegacyNormalizationAdapter, + PerImageMinMaxPercentileNormalization, + resolve_normalization_strategy, +) + + +@pytest.mark.parametrize("as_torch", [False, True]) +def test_global_percentile_matches_reference(as_torch): + array = np.arange(48, dtype=np.float32).reshape(3, 4, 4) + batch = torch.from_numpy(array) if as_torch else array + strategy = GlobalPercentileNormalization(10, 90) + parameters = strategy.fit(batch) + lower, upper = np.percentile(array, [10, 90]) + expected = (np.clip(array, lower, upper) - lower) / (upper - lower + 1e-8) + actual = strategy.transform(batch, parameters) + if as_torch: + actual = actual.numpy() + np.testing.assert_allclose(actual, expected, rtol=1e-6, atol=1e-7) + + +@pytest.mark.parametrize("as_torch", [False, True]) +def test_per_image_minmax_percentile_matches_reference(as_torch): + array = np.array( + [ + [[1, 2], [3, 5]], + [[10, 20], [30, 50]], + [[7, 7], [7, 7]], + ], + dtype=np.float32, + ) + batch = torch.from_numpy(array) if as_torch else array + per_image = np.stack( + [ + (image - image.min()) / (image.max() - image.min()) + if image.max() > image.min() + else np.zeros_like(image) + for image in array + ] + ) + lower, upper = np.percentile(per_image, [5, 95]) + expected = ( + np.clip(per_image, lower, upper) - lower + ) / (upper - lower + 1e-8) + strategy = PerImageMinMaxPercentileNormalization(5, 95) + actual = strategy.transform(batch, strategy.fit(batch)) + if as_torch: + actual = actual.numpy() + np.testing.assert_allclose(actual, expected, rtol=1e-6, atol=1e-7) + + +def test_registered_resolution_and_legacy_adapter(): + strategy = resolve_normalization_strategy( + { + "mode": "per_image_minmax_percentile", + "lower_percentile": 2, + "upper_percentile": 98, + } + ) + assert isinstance(strategy, PerImageMinMaxPercentileNormalization) + + def compute(batch, lower_percentile, upper_percentile): + return {"low": float(batch.min()), "high": float(batch.max())} + + def transform(batch, parameters): + return (batch - parameters["low"]) / ( + parameters["high"] - parameters["low"] + ) + + adapter = LegacyNormalizationAdapter(compute, transform, 1, 99) + batch = np.array([[1.0, 3.0]]) + np.testing.assert_allclose( + adapter.transform(batch, adapter.fit(batch)), [[0.0, 1.0]] + ) + + +def test_strategy_and_model_changes_invalidate_inference_caches(): + detector = object.__new__(BraggPeaksPolymer) + detector._normalization_strategy = GlobalPercentileNormalization() + detector._normalization_is_explicit = False + detector._normalization_parameters = object() + detector._norm_median = 1 + detector._norm_iqr = 2 + detector._bn_adapted = True + detector._live_chunk_cache = object() + + detector.normalization_strategy = PerImageMinMaxPercentileNormalization() + assert detector._normalization_parameters is None + assert detector._bn_adapted is False + assert detector._live_chunk_cache is None + + detector._normalization_parameters = object() + detector._bn_adapted = True + detector.model = object() + assert detector._normalization_parameters is None + assert detector._bn_adapted is False + + +def test_missing_custom_checkpoint_normalization_fails_clearly(): + detector = object.__new__(BraggPeaksPolymer) + detector._normalization_strategy = None + with pytest.raises(RuntimeError, match="custom checkpoint"): + detector._require_normalization_strategy() + + +def test_named_model_selects_metadata_strategy_and_explicit_override_wins( + monkeypatch, +): + class FakeModel: + def load_state_dict(self, state): + self.state = state + + def to(self, device): + return self + + resolution = SimpleNamespace( + model_id="example", + weights_path="/unused/weights.pth", + specification={ + "experimental_normalization": { + "mode": "per_image_minmax_percentile", + "lower_percentile": 3, + "upper_percentile": 97, + } + }, + ) + monkeypatch.setattr(bragg_peaks_module, "resolve_polymer_model", lambda **_: resolution) + monkeypatch.setattr(bragg_peaks_module, "build_polymer_model", lambda _: FakeModel()) + monkeypatch.setattr(torch, "load", lambda *_, **__: {}) + + detector = object.__new__(BraggPeaksPolymer) + detector._device = "cpu" + detector._model = FakeModel() + detector._normalization_strategy = None + detector._normalization_is_explicit = False + detector._invalidate_inference_caches() + detector.set_model_weights() + assert isinstance( + detector.normalization_strategy, PerImageMinMaxPercentileNormalization + ) + assert detector.normalization_strategy.lower_percentile == 3 + + override = GlobalPercentileNormalization(4, 96) + detector.normalization_strategy = override + detector.set_model_weights() + assert detector.normalization_strategy is override diff --git a/tests/diffraction/test_polymer_numerics.py b/tests/diffraction/test_polymer_numerics.py new file mode 100644 index 000000000..5c0e01350 --- /dev/null +++ b/tests/diffraction/test_polymer_numerics.py @@ -0,0 +1,29 @@ +import numpy as np + +from quantem.diffraction.peak_detection import detect_blobs +from quantem.diffraction.polar_transform import polar_transform +from quantem.diffraction.polymer_utils import parse_reciprocal_units + + +def test_reciprocal_unit_conversion_is_explicit(): + assert parse_reciprocal_units("nm^-1") == ("1/nm", 0.1) + assert parse_reciprocal_units("Å⁻¹") == ("1/A", 1.0) + + +def test_peak_coordinates_remain_row_column_order(): + yy, xx = np.mgrid[:17, :19] + image = np.exp(-((yy - 6.25) ** 2 + (xx - 11.4) ** 2) / 2.0) + peaks, _, success = detect_blobs(image, sigma=0.5, threshold=0.2) + assert success.tolist() == [True] + np.testing.assert_allclose(peaks[0], [6.25, 11.4], atol=0.15) + + +def test_polar_transform_orientation_and_shape(): + data = np.zeros((1, 1, 15, 15), dtype=np.float32) + data[0, 0, 7, 11] = 1.0 + polar = polar_transform( + data, origin_array=np.array([7.0, 7.0]), num_annular_bins=8, + radial_min=0, radial_max=7, radial_step=1, device="cpu", show_progress=False, + ) + assert polar.array.shape == (1, 1, 8, 7) + assert np.unravel_index(np.argmax(polar.array[0, 0]), (8, 7)) == (0, 4) diff --git a/tests/diffraction/test_scan_mask_editor.py b/tests/diffraction/test_scan_mask_editor.py new file mode 100644 index 000000000..3df0e594a --- /dev/null +++ b/tests/diffraction/test_scan_mask_editor.py @@ -0,0 +1,208 @@ +from types import SimpleNamespace + +import matplotlib.pyplot as plt +import numpy as np +import pytest + +from quantem.diffraction import BraggPeaksPolymer, ScanMaskEditor + + +def _analysis(scan_shape=(7, 9)): + rows, columns = scan_shape + analysis = object.__new__(BraggPeaksPolymer) + analysis.dataset_cartesian = SimpleNamespace( + shape=(rows, columns, 3, 4), + array=np.arange(rows * columns * 12, dtype=float).reshape( + rows, columns, 3, 4 + ), + sampling=(0.5, 0.5, 1.0, 1.0), + units=("nm", "nm", "1/Å", "1/Å"), + virtual_images={}, + ) + analysis._scan_mask = None + return analysis + + +def test_controls_follow_probe_directions_and_apply_is_explicit(): + analysis = _analysis() + editor = analysis.edit_scan_mask( + initial_x=4, + initial_y=3, + initial_radius=2, + display_widget=False, + ) + try: + assert isinstance(editor, ScanMaskEditor) + initially_applied = analysis.scan_mask.copy() + assert editor.widget.layout.width == "500px" + assert editor.output.layout.width == "438px" + assert tuple(editor.figure.get_size_inches()) == pytest.approx((4.0, 3.25)) + assert editor.circle_artist.get_linewidth() == pytest.approx(1.05) + assert editor.circle_artist.get_linestyle() != "-" + assert editor.figure.number not in plt.get_fignums() + # Radius is the row immediately above the image/Y-slider row. + assert editor.radius_slider in editor.widget.children[3].children + assert editor.output in editor.widget.children[5].children + + editor.x_slider.value += 1 + assert editor.x == 5 + assert editor.circle_artist.center == (5, 3) + + # Increasing the vertical slider moves its thumb and the probe upward. + editor.y_slider.value += 1 + assert editor.y == 2 + assert editor.circle_artist.center == (5, 2) + + assert not np.array_equal(editor.mask, initially_applied) + np.testing.assert_array_equal(analysis.scan_mask, initially_applied) + # This is the exact compatibility path used by older notebook cells. + mask_arr = editor + assert mask_arr.sum() == initially_applied.sum() + np.testing.assert_array_equal(np.asarray(mask_arr), initially_applied) + editor.apply() + np.testing.assert_array_equal(analysis.scan_mask, editor.mask) + assert mask_arr.sum() == editor.mask.sum() + finally: + editor.close() + + +def test_saved_mask_is_loaded_and_applied(tmp_path): + path = tmp_path / "scan_mask.npz" + first = _analysis() + editor = first.edit_scan_mask( + initial_x=2, + initial_y=3, + initial_radius=2, + state_path=path, + display_widget=False, + ) + try: + editor.set_mask( + x=6, + y=1, + geometry="rectangle", + size_x=2, + size_y=1, + ) + expected = editor.mask + assert editor.save() == path + finally: + editor.close() + + second = _analysis() + loaded = second.edit_scan_mask( + initial_x=0, + initial_y=0, + initial_radius=1, + state_path=path, + display_widget=False, + ) + try: + assert (loaded.x, loaded.y) == (6, 1) + assert (loaded.geometry, loaded.size_x, loaded.size_y) == ( + "rectangle", + 2, + 1, + ) + np.testing.assert_array_equal(loaded.mask, expected) + np.testing.assert_array_equal(second.scan_mask, expected) + with np.load(path, allow_pickle=False) as state: + assert int(state["schema_version"]) == ScanMaskEditor.SCHEMA_VERSION + assert tuple(state["scan_shape"]) == (7, 9) + finally: + loaded.close() + + +def test_saved_mask_rejects_a_different_scan_shape(tmp_path): + path = tmp_path / "scan_mask.npz" + editor = _analysis().edit_scan_mask( + state_path=path, display_widget=False + ) + try: + editor.save() + finally: + editor.close() + + with pytest.raises(ValueError, match="does not match current scan shape"): + _analysis((8, 9)).edit_scan_mask( + state_path=path, display_widget=False + ) + plt.close("all") + + +def test_version_one_circle_state_remains_loadable(tmp_path): + path = tmp_path / "scan_mask_v1.npz" + yy, xx = np.ogrid[:7, :9] + mask = (yy - 3) ** 2 + (xx - 4) ** 2 <= 2**2 + np.savez_compressed( + path, + schema_version=np.asarray(1), + scan_shape=np.asarray((7, 9)), + mask=mask, + center_row=np.asarray(3), + center_column=np.asarray(4), + radius=np.asarray(2), + ) + editor = _analysis().edit_scan_mask( + state_path=path, display_widget=False + ) + try: + assert (editor.geometry, editor.size_x, editor.size_y) == ( + "circle", + 2, + 2, + ) + np.testing.assert_array_equal(editor.mask, mask) + finally: + editor.close() + + +def test_legacy_wrapper_preserves_historical_row_column_arguments(): + analysis = _analysis() + editor = analysis.create_interactive_circular_mask( + initial_x0=2, + initial_y0=6, + initial_r=3, + display_widget=False, + ) + try: + assert (editor.x, editor.y, editor.radius) == (6, 2, 3) + assert editor["x0"] == 2 + assert editor["y0"] == 6 + finally: + editor.close() + + +@pytest.mark.parametrize( + ("geometry", "size_x", "size_y", "expected_count"), + [ + ("circle", 2, 2, 13), + ("ellipse", 3, 1, 9), + ("square", 2, 2, 25), + ("rectangle", 2, 1, 15), + ], +) +def test_geometry_selector_builds_expected_masks( + geometry, size_x, size_y, expected_count +): + editor = _analysis().edit_scan_mask( + initial_x=4, + initial_y=3, + display_widget=False, + ) + try: + editor.set_mask( + geometry=geometry, + size_x=size_x, + size_y=size_y, + ) + assert editor.geometry == geometry + assert editor.mask.sum() == expected_count + assert editor.circle_artist.get_linewidth() == pytest.approx(1.05) + if geometry in {"circle", "square"}: + assert editor.size_y == editor.size_x + assert editor.size_y_row.layout.display == "none" + else: + assert editor.size_y_row.layout.display != "none" + finally: + editor.close()