diff --git a/pyproject.toml b/pyproject.toml index c75263123..e463bb581 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -48,6 +48,8 @@ dependencies = [ "optuna>=4.5.0", "hdf5plugin>=6.0.0", "torchinfo>=1.8.0", + "ase>=3.23", + "spglib>=2.5", "em-database>=0.5", ] diff --git a/src/quantem/__init__.py b/src/quantem/__init__.py index 7872457be..b802333a7 100644 --- a/src/quantem/__init__.py +++ b/src/quantem/__init__.py @@ -9,6 +9,7 @@ from quantem.core import visualization as visualization from quantem import imaging as imaging +from quantem import diffraction as diffraction from quantem import spectroscopy as spectroscopy from quantem import diffractive_imaging as diffractive_imaging diff --git a/src/quantem/core/datastructures/__init__.py b/src/quantem/core/datastructures/__init__.py index 60cddd992..e8ecb4643 100644 --- a/src/quantem/core/datastructures/__init__.py +++ b/src/quantem/core/datastructures/__init__.py @@ -2,6 +2,7 @@ from quantem.core.datastructures.vector import Vector as Vector from quantem.core.datastructures.dataset4dstem import Dataset4dstem as Dataset4dstem +from quantem.core.datastructures.polar4dstem import Polar4dstem as Polar4dstem 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 diff --git a/src/quantem/core/datastructures/dataset4dstem.py b/src/quantem/core/datastructures/dataset4dstem.py index 004db4278..32ef023c0 100644 --- a/src/quantem/core/datastructures/dataset4dstem.py +++ b/src/quantem/core/datastructures/dataset4dstem.py @@ -1,3 +1,4 @@ +from os import PathLike from typing import Any, Self import matplotlib.pyplot as plt @@ -8,6 +9,10 @@ from quantem.core.datastructures.dataset2d import Dataset2d from quantem.core.datastructures.dataset4d import Dataset4d +from quantem.core.datastructures.polar4dstem import dataset4dstem_polar_transform +from quantem.core.utils.diffractive_imaging_utils import ( + fit_probe_circle as _fit_probe_circle, +) from quantem.core.utils.validators import ensure_valid_array from quantem.core.visualization import show_2d from quantem.core.visualization.visualization_utils import ScalebarConfig @@ -72,12 +77,22 @@ def __init__( signal_units : str, optional Units for the array values, by default "arb. units" metadata : dict - "r_to_q_rotation_cw_deg": rotation r to q clockwise in degrees - "ellipticity": 3 parameters (a, b, theta (degrees)) + Missing keys below are set to None. + + "q_to_r_rotation_ccw_deg" : float + Rotation in degrees that maps detector (q) vectors onto the + scan (r) frame. A detector vector (dr, dc) is first swapped to + (dc, dr) if "q_transpose" is True, then rotated as + dr' = cos(t) dr - sin(t) dc, dc' = sin(t) dr + cos(t) dc. + "q_transpose" : bool + If True, swap the detector row and column axes before the + rotation. + "ellipticity" : tuple + 3 parameters (a, b, theta in degrees). _token : object | None, optional Token to prevent direct instantiation, by default None """ - mdata_keys_4dstem = ["r_to_q_rotation_cw_deg", "ellipticity"] + mdata_keys_4dstem = ["q_to_r_rotation_ccw_deg", "q_transpose", "ellipticity"] for k in mdata_keys_4dstem: if k not in metadata.keys(): metadata[k] = None @@ -97,15 +112,15 @@ def __init__( self._virtual_detectors = {} # Store detector information for regeneration @classmethod - def from_file(cls, file_path: str, file_type: str) -> "Dataset4dstem": + def from_file(cls, file_path: str | PathLike, file_type: str | None = None) -> "Dataset4dstem": """ Create a new Dataset4dstem from a file. Parameters ---------- - file_path : str + file_path : str | PathLike Path to the data file - file_type : str + file_type : str | None The type of file reader needed. See rosettasciio for supported formats https://hyperspy.org/rosettasciio/supported_formats/index.html @@ -274,6 +289,39 @@ def get_dp_mean(self, attach: bool = True) -> Dataset2d: return dp_mean_dataset + def fit_probe_circle( + self, + array: NDArray | Dataset2d | None = None, + threshold: float | None = None, + show: bool = True, + ) -> tuple[float, float, float]: + """Fit a circle to the probe in a diffraction pattern. + + Parameters + ---------- + array : NDArray | Dataset2d | None, optional + 2D diffraction pattern to fit, either as an array or a Dataset2d. If + None, uses this dataset's mean diffraction pattern, computing it if needed. + threshold : float | None, optional + Threshold for binarizing the diffraction pattern. If None, Otsu's method + is used. + show : bool, optional + Whether to display the fitted circle, by default True. + + Returns + ------- + tuple[float, float, float] + Probe center in diffraction-pattern row and column coordinates + (probe_qy0, probe_qx0), followed by the fitted radius. + """ + if array is None: + dp_mean = ( + self._dp_mean if hasattr(self, "_dp_mean") else self.get_dp_mean(attach=False) + ) + array = dp_mean.array + + return _fit_probe_circle(array, threshold=threshold, show=show) + @property def dp_max(self) -> Dataset2d: """ @@ -566,7 +614,15 @@ def _create_annular_mask( return (distance >= r_inner) & (distance <= r_outer) - def show_virtual_images(self, figsize: tuple[int, int] | None = None, **kwargs) -> tuple: + def show_virtual_images( + self, + figsize: tuple[int, int] | None = None, + *, + positions: list[tuple[int, int]] | None = None, + position_color: str = "red", + position_size: float = 60.0, + **kwargs, + ) -> tuple: """ Display all virtual images stored in the dataset using show_2d. @@ -574,6 +630,13 @@ def show_virtual_images(self, figsize: tuple[int, int] | None = None, **kwargs) ---------- figsize : tuple[int, int] | None, optional Figure size in inches. If None, automatically calculated based on number of images + positions : list of tuple of int, optional + ``(row, col)`` scan positions to mark on every image, e.g. the + positions used to tune Bragg disk detection. + position_color : str, default="red" + Color of the position markers. + position_size : float, default=60.0 + Area of the position markers in points squared. **kwargs Additional keyword arguments passed to show_2d (e.g., cmap, norm, cbar, etc.) @@ -623,8 +686,120 @@ def show_virtual_images(self, figsize: tuple[int, int] | None = None, **kwargs) kwargs.setdefault("scalebar", [scalebar] + [False] * (len(arrays) - 1)) fig, axs = show_2d(arrays_organized, title=titles_organized, figsize=figsize, **kwargs) + if positions is not None and len(positions) > 0: + pos = np.asarray(positions, dtype=float).reshape(-1, 2) + for ax in np.atleast_1d(np.asarray(axs, dtype=object)).ravel(): + ax.scatter( + pos[:, 1], + pos[:, 0], + s=position_size, + facecolors="none", + edgecolors=position_color, + linewidths=1.5, + ) + return fig, axs + def show_virtual_detectors( + self, + names: str | list[str] | None = None, + *, + colors: list[str] | None = None, + alpha: float = 0.2, + linewidth: float = 1.5, + legend: bool = True, + **kwargs, + ) -> tuple: + """Show the mean diffraction pattern with the virtual detectors drawn on it. + + Every detector attached by :meth:`get_virtual_image` is drawn on a single + mean pattern, so their placement relative to the direct beam and the + diffracted rings can be checked at a glance. + + Parameters + ---------- + names : str or list of str, optional + Detector(s) to draw. ``None`` (default) draws all attached detectors. + colors : list of str, optional + One color per detector; defaults to the matplotlib color cycle. + alpha : float, default=0.2 + Opacity of the filled detector area. Set to 0 for outlines only. + linewidth : float, default=1.5 + Width of the detector outline. + legend : bool, default=True + If ``True``, label the detectors in a legend. + **kwargs + Passed to :func:`~quantem.core.visualization.show_2d`; ``norm``, + ``scalebar`` and ``title`` have diffraction-pattern defaults. + + Returns + ------- + tuple + ``(fig, ax)`` from :func:`~quantem.core.visualization.show_2d`. + """ + if not self._virtual_detectors: + raise ValueError("No virtual detectors attached. Create one with get_virtual_image().") + if names is None: + names = list(self._virtual_detectors) + elif isinstance(names, str): + names = [names] + missing = [n for n in names if n not in self._virtual_detectors] + if missing: + raise ValueError( + f"Virtual detector(s) {missing} not found. " + f"Available detectors: {list(self._virtual_detectors)}" + ) + + if colors is None: + cycle = plt.rcParams["axes.prop_cycle"].by_key().get("color", ["red"]) + colors = [cycle[i % len(cycle)] for i in range(len(names))] + + dp_mean = self.dp_mean + kwargs.setdefault("norm", {"power": 0.4, "upper_quantile": 0.999}) + kwargs.setdefault("title", "mean DP with virtual detectors") + kwargs.setdefault( + "scalebar", ScalebarConfig(sampling=self.sampling[2], units=self.units[2]) + ) + fig, ax = show_2d(dp_mean.array, **kwargs) + + for name, color in zip(names, colors): + det = self._virtual_detectors[name] + mode, geometry, mask = det["mode"], det["geometry"], det["mask"] + if mask is not None: + ax.contour(mask, levels=[0.5], colors=[color], linewidths=linewidth) + ax.plot([], [], color=color, linewidth=linewidth, label=name) + continue + (cy, cx) = geometry[0] + if mode == "circle": + radius = geometry[1] + ax.add_patch( + Circle((cx, cy), radius, color=color, fill=True, alpha=alpha, label=name) + ) + ax.add_patch( + Circle((cx, cy), radius, color=color, fill=False, linewidth=linewidth) + ) + elif mode == "annular": + r_inner, r_outer = geometry[1] + ax.add_patch( + Wedge( + (cx, cy), + r_outer, + 0, + 360, + width=r_outer - r_inner, + color=color, + fill=True, + alpha=alpha, + label=name, + ) + ) + for r in (r_inner, r_outer): + ax.add_patch(Circle((cx, cy), r, color=color, fill=False, linewidth=linewidth)) + + if legend: + ax.legend(loc="upper right", framealpha=0.8) + return fig, ax + def regenerate_virtual_images(self) -> None: """ Regenerate virtual images from stored detector information. @@ -798,3 +973,5 @@ def median_filter_masked_pixels(self, mask: np.ndarray, kernel_width: int = 3): self.array[:, :, index_x, index_y] = np.median( self.array[:, :, x_min:x_max, y_min:y_max], axis=(2, 3) ) + + polar_transform = dataset4dstem_polar_transform diff --git a/src/quantem/core/datastructures/polar4dstem.py b/src/quantem/core/datastructures/polar4dstem.py new file mode 100644 index 000000000..81ae5d037 --- /dev/null +++ b/src/quantem/core/datastructures/polar4dstem.py @@ -0,0 +1,309 @@ +from typing import TYPE_CHECKING, Any + +import numpy as np +from numpy.typing import NDArray +from scipy.ndimage import map_coordinates + +from quantem.core.datastructures.dataset4d import Dataset4d + +if TYPE_CHECKING: + from .dataset4dstem import Dataset4dstem + + +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": + """ + Create a Polar4dstem from a 4D array shaped (scan_y, scan_x, phi, r). + + Parameters + ---------- + array : NDArray + Polar data, shape (scan_y, scan_x, n_phi, n_r). + name : str, optional + Dataset name, by default "Polar 4D-STEM dataset". + origin : NDArray | tuple | list | float | int, optional + Origin of each axis in calibrated units, by default zeros. + sampling : NDArray | tuple | list | float | int, optional + Sampling of each axis, by default ones. + units : list[str] | tuple | list, optional + Units of each axis, by default ["pixels", "pixels", "deg", "pixels"]. + signal_units : str, optional + Units of the array values, by default "arb. units". + metadata : dict, optional + Metadata. Missing "polar_*" keys are set to None. + + Returns + ------- + 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: + """Number of azimuthal bins (axis 2).""" + return int(self.shape[2]) + + @property + def n_r(self) -> int: + """Number of radial bins (axis 3).""" + return int(self.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 dataset4dstem_polar_transform( + self: "Dataset4dstem", + origin_row: float, + origin_col: float, + 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, +) -> Polar4dstem: + """ + Resample every diffraction pattern onto a polar (phi, r) grid. + + Bound to `Dataset4dstem.polar_transform`. Uses bilinear interpolation + (`scipy.ndimage.map_coordinates`, order 1); samples outside the detector + are set to 0. + + Parameters + ---------- + origin_row, origin_col : float + Center of the polar grid on the detector, in detector pixels. The same + center is used for all scan positions. + ellipse_params : tuple[float, float, float], optional + Elliptical distortion (a, b, theta_deg): the radius along the + direction theta_deg (degrees) is scaled by a / b. None for circular + sampling. + num_annular_bins : int, optional + Number of azimuthal bins, by default 180. + radial_min : float, optional + First radial bin in detector pixels, by default 0. + radial_max : float, optional + Radial upper limit (exclusive) in detector pixels. If None, the + distance from the origin to the nearest detector edge. + radial_step : float, optional + Radial bin width in detector pixels, by default 1. + two_fold_rotation_symmetry : bool, optional + If True, phi covers [0, 180) degrees instead of [0, 360). + name : str, optional + Name of the output, by default "_polar". + signal_units : str, optional + Signal units of the output, by default those of this dataset. + + Returns + ------- + Polar4dstem + Array shaped (scan_y, scan_x, n_phi, n_r). Phi is in degrees, 0 along + +column and increasing toward +row. The radial axis uses the sampling + and units of the last detector axis. The polar parameters are stored + in metadata under "polar_*" keys. + + Notes + ----- + Tensor-backed datasets are copied to a CPU numpy array first. + """ + array = self.numpy() + if array.ndim != 4: + raise ValueError("polar_transform requires a 4D-STEM dataset (ndim=4).") + scan_y, scan_x, ny, nx = array.shape + origin_row_f = float(origin_row) + origin_col_f = float(origin_col) + 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(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 = array[iy, ix] + out[iy, ix] = map_coordinates( + dp, + coords, + order=1, + mode="constant", + cval=0.0, + ) + if two_fold_rotation_symmetry: + phi_range = np.pi + else: + phi_range = 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": origin_row_f, + "polar_origin_col": origin_col_f, + "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/core/io/__init__.py b/src/quantem/core/io/__init__.py index de0df5f81..68b2d8eb8 100644 --- a/src/quantem/core/io/__init__.py +++ b/src/quantem/core/io/__init__.py @@ -4,5 +4,6 @@ read_3d_spectroscopy as read_3d_spectroscopy, ) from quantem.core.io.serialize import AutoSerialize as AutoSerialize +from quantem.core.io.serialize import Bundle as Bundle from quantem.core.io.serialize import load as load from quantem.core.io.serialize import print_file as print_file diff --git a/src/quantem/core/io/file_readers.py b/src/quantem/core/io/file_readers.py index 319467ca4..395d828e0 100644 --- a/src/quantem/core/io/file_readers.py +++ b/src/quantem/core/io/file_readers.py @@ -4,6 +4,7 @@ from typing import Any import h5py +import numpy as np from quantem.core.datastructures import Dataset as Dataset from quantem.core.datastructures import Dataset2d as Dataset2d @@ -18,6 +19,69 @@ ) +def _resolve_rsciio_plugin(file_path: str | PathLike, file_type: str | None = None) -> str: + """ + Resolve the RosettaSciIO plugin module used to read a file. + + Parameters + ---------- + file_path : str | PathLike + Path to the file. Its extension is used when ``file_type`` is None. + file_type : str, optional + RosettaSciIO plugin name (e.g. "digitalmicrograph", "quantumdetector") + or a file extension (e.g. "dm4", "mib"). Case-insensitive. + + Returns + ------- + str + Module name of the plugin, e.g. "rsciio.digitalmicrograph". + + Raises + ------ + ValueError + If no plugin matches, or if an extension is listed by more than one + plugin (e.g. ".h5"); pass the plugin name as ``file_type`` in that case. + """ + import rsciio + + key = file_type if file_type is not None else Path(file_path).suffix.lstrip(".") + key = str(key).lower().lstrip(".") + if not key: + raise ValueError( + f"Cannot infer the file type of '{file_path}'; pass file_type= " + "(a RosettaSciIO plugin name such as 'digitalmicrograph')." + ) + + plugins = rsciio.IO_PLUGINS + by_name = sorted({p["api"] for p in plugins if p["api"].lower() == f"rsciio.{key}"}) + if by_name: + return by_name[0] + + by_ext = sorted( + {p["api"] for p in plugins if key in (ext.lower() for ext in p["file_extensions"])} + ) + if len(by_ext) == 1: + return by_ext[0] + if len(by_ext) > 1: + names = ", ".join(f"'{api.removeprefix('rsciio.')}'" for api in by_ext) + raise ValueError( + f"File extension '{key}' is used by several RosettaSciIO plugins ({names}). " + "Pass one of them as file_type=." + ) + raise ValueError(f"No RosettaSciIO reader for file type '{key}'.") + + +def _rsciio_reader(file_path: str | PathLike, file_type: str | None = None): + """ + Return ``(plugin, file_reader)`` for a file, see `_resolve_rsciio_plugin`. + + An ImportError raised here means the plugin exists but one of its optional + dependencies is missing. + """ + plugin = _resolve_rsciio_plugin(file_path, file_type) + return plugin, importlib.import_module(plugin).file_reader + + def _print_available_datasets(data_list): print("Available datasets:") for index, entry in enumerate(data_list): @@ -30,26 +94,47 @@ def read_4dstem( file_type: str | None = None, dataset_index: int | None = None, hot_pixel_filter: bool = False, + scan_length: int | None = None, + scan_axis: int = 0, + transpose_scan_axes: bool = False, **kwargs, ) -> Dataset4dstem: """ - File reader for 4D-STEM data + File reader for 4D-STEM data. Parameters ---------- - file_path: str | PathLike - Path to data - file_type: str - The type of file reader needed. See rosettasciio for supported formats + file_path : str | PathLike + Path to data. + file_type : str, optional + RosettaSciIO plugin name (e.g. "arina", "digitalmicrograph") or file + extension. If None, the extension of `file_path` is used. Extensions + shared by several plugins (e.g. "h5") require the plugin name. See https://hyperspy.org/rosettasciio/supported_formats/index.html - dataset_index: int, optional + dataset_index : int, optional Index of the dataset to load if file contains multiple datasets. If None, automatically selects the first 4D dataset found. - hot_pixel_filter: bool, optional + If no 4D dataset is found but a 3D stack exists, a 3D dataset can be + interpreted as 4D if `scan_length` is provided. + hot_pixel_filter : bool, default False If True, detect and replace hot detector pixels immediately after loading using `quantem.core.utils.filter.filter_hot_pixels` with its default parameters. For custom thresholds, call `filter_hot_pixels` directly on the array. + scan_length : int, optional + For 3D datasets shaped (n_frames, ny, nx) (after possibly moving the + scan axis to the front), interpret the data as a raster scan with shape + (scan_y, scan_x, ny, nx), where scan_y = n_frames // scan_length and + scan_x = scan_length. Required if you want to treat a 3D stack as 4D. + scan_axis : int, default 0 + Which axis of a 3D dataset is the scan/time axis before reshaping. + Must be 0 or 1. The specified axis is moved to axis 0 before the + (scan_y, scan_x) reshape. + transpose_scan_axes : bool, default False + Only used when interpreting a 3D dataset as 4D via `scan_length`. + If True, transpose the scan axes after reshaping so that + (scan_y, scan_x) -> (scan_x, scan_y). This effectively swaps the + interpretation of scan rows and columns in the final 4D array. **kwargs: dict Additional keyword arguments to pass to the file reader. @@ -67,7 +152,7 @@ def read_4dstem( Units for the array values, by default "arb. units" Returns - -------- + ------- Dataset4dstem Examples @@ -90,39 +175,180 @@ def read_4dstem( ... hot_pixel_filter=True, ... ) """ - if file_type is None: - file_type = Path(file_path).suffix.lower().lstrip(".") + + def _reshape_3d_to_4d( + imported_data: dict, + *, + dataset_index_local: int, + scan_length_local: int, + scan_axis_local: int, + transpose_scan_axes_local: bool, + ) -> dict: + data = imported_data["data"] + if data.ndim != 3: + raise ValueError( + f"Expected 3D data to reshape, got ndim={data.ndim} with shape {data.shape}" + ) + + # Move scan axis to front so it becomes the frame axis + if scan_axis_local != 0: + data = np.moveaxis(data, scan_axis_local, 0) + + n_frames, ny, nx = data.shape + + if scan_length_local <= 0: + raise ValueError(f"scan_length must be positive, got {scan_length_local}") + if n_frames % scan_length_local != 0: + raise ValueError( + f"scan_length={scan_length_local} is not compatible with n_frames={n_frames}; " + f"n_frames % scan_length = {n_frames % scan_length_local}" + ) + + scan_y = n_frames // scan_length_local + scan_x = scan_length_local + + data_4d = data.reshape(scan_y, scan_x, ny, nx) + + if transpose_scan_axes_local: + data_4d = np.transpose(data_4d, (1, 0, 2, 3)) + scan_y, scan_x = scan_x, scan_y + + old_axes = imported_data.get("axes", None) + if old_axes is None or len(old_axes) != 3: + raise ValueError( + f"Expected 3 axes for 3D data when reshaping to 4D; got axes={old_axes}" + ) + + ax_scan_y = { + "scale": 1.0, + "offset": 0.0, + "units": "pixels", + "name": "scan_y", + } + ax_scan_x = { + "scale": 1.0, + "offset": 0.0, + "units": "pixels", + "name": "scan_x", + } + + # Detector calibrations come from the two axes that are not the scan axis. + ax_qy, ax_qx = (dict(ax) for i, ax in enumerate(old_axes) if i != scan_axis_local) + + imported_data_4d = imported_data.copy() + imported_data_4d["data"] = data_4d + imported_data_4d["axes"] = [ax_scan_y, ax_scan_x, ax_qy, ax_qx] + + original_shape = imported_data["data"].shape + new_shape = data_4d.shape + print( + f"Using 3D dataset {dataset_index_local} with shape {original_shape} " + f"interpreted as 4D with shape={new_shape} " + f"(scan_axis={scan_axis_local}, scan_length={scan_length_local}, " + f"transpose_scan_axes={transpose_scan_axes_local})." + ) + + return imported_data_4d + + if scan_axis not in (0, 1): + raise ValueError(f"scan_axis must be 0 or 1, got {scan_axis}") sampling_override = kwargs.pop("sampling", None) origin_override = kwargs.pop("origin", None) units_override = kwargs.pop("units", None) name_override = kwargs.pop("name", None) - file_reader = importlib.import_module(f"rsciio.{file_type}").file_reader + plugin, file_reader = _rsciio_reader(file_path, file_type) data_list = file_reader(file_path, **kwargs) - # If specific index provided, use it + if not data_list: + raise ValueError(f"No datasets returned by {plugin} for '{file_path}'") + + # Case 1: dataset_index specified explicitly if dataset_index is not None: imported_data = data_list[dataset_index] - if imported_data["data"].ndim != 4: + ndim = imported_data["data"].ndim + + if ndim == 4: + # Use 4D as-is + pass + elif ndim == 3: + if scan_length is None: + raise ValueError( + f"Dataset at index {dataset_index} is 3D (shape={imported_data['data'].shape}). " + "To interpret it as 4D-STEM, please provide scan_length." + ) + imported_data = _reshape_3d_to_4d( + imported_data, + dataset_index_local=dataset_index, + scan_length_local=scan_length, + scan_axis_local=scan_axis, + transpose_scan_axes_local=transpose_scan_axes, + ) + else: raise ValueError( - f"Dataset at index {dataset_index} has {imported_data['data'].ndim} dimensions, " - f"expected 4D. Shape: {imported_data['data'].shape}" + f"Dataset at index {dataset_index} has ndim={ndim}, " + f"expected 4D or 3D. Shape: {imported_data['data'].shape}" ) + else: - # Automatically find first 4D dataset + # Case 2: auto-select dataset four_d_datasets = [(i, d) for i, d in enumerate(data_list) if d["data"].ndim == 4] _print_available_datasets(data_list) - if len(four_d_datasets) == 0: - print(f"No 4D datasets found in {file_path}.") - raise ValueError("No 4D dataset found in file") - - dataset_index, imported_data = four_d_datasets[0] - - print( - f"Using first 4D dataset at index {dataset_index} with shape {imported_data['data'].shape}" - ) + if four_d_datasets: + dataset_index, imported_data = four_d_datasets[0] + if len(data_list) > 1: + print( + f"File contains {len(data_list)} dataset(s). Using 4D dataset " + f"{dataset_index} with shape {imported_data['data'].shape}" + ) + else: + three_d_datasets = [(i, d) for i, d in enumerate(data_list) if d["data"].ndim == 3] + + if not three_d_datasets: + print(f"No 4D datasets found in {file_path}.") + raise ValueError("No 4D or 3D dataset found in file") + + if scan_length is None: + print(f"No 4D datasets found in {file_path}.") + raise ValueError( + "File contains only 3D datasets. To interpret one as 4D-STEM, " + "please specify scan_length so that n_frames % scan_length == 0." + ) + + # Choose first 3D dataset compatible with scan_length along scan_axis + candidates: list[tuple[int, dict]] = [] + for i, d in three_d_datasets: + shape = d["data"].shape + n_frames_axis = shape[scan_axis] + if n_frames_axis % scan_length == 0: + candidates.append((i, d)) + + if not candidates: + print(f"3D datasets in {file_path}:") + for i, d in three_d_datasets: + print(f" Dataset {i}: shape {d['data'].shape}") + raise ValueError( + f"No 3D dataset has length along scan_axis={scan_axis} " + f"divisible by scan_length={scan_length}." + ) + + dataset_index, imported_data = candidates[0] + if len(candidates) > 1: + print( + f"Multiple 3D datasets compatible with scan_length={scan_length} " + f"along scan_axis={scan_axis}. Using dataset {dataset_index} " + f"with shape {imported_data['data'].shape}" + ) + + imported_data = _reshape_3d_to_4d( + imported_data, + dataset_index_local=dataset_index, + scan_length_local=scan_length, + scan_axis_local=scan_axis, + transpose_scan_axes_local=transpose_scan_axes, + ) imported_axes = imported_data["axes"] @@ -180,7 +406,7 @@ def read_3d_spectroscopy( """ data_type_normalized = str(data_type).upper() - file_reader = importlib.import_module(f"rsciio.{file_type}").file_reader # type: ignore + plugin, file_reader = _rsciio_reader(file_path, file_type) data_list = file_reader(file_path) # If specific index provided, use it @@ -209,13 +435,10 @@ def read_3d_spectroscopy( ) imported_axes = imported_data["axes"] - # axis_order = (0, 1, 2) if file_type == "digitalmicrograph" else (2, 0, 1) - axis_order = (1, 2, 0) if file_type == "digitalmicrograph" else (0, 1, 2) - array = ( - imported_data["data"].transpose(axis_order) - if file_type == "digitalmicrograph" - else imported_data["data"] - ) + # DigitalMicrograph spectrum images are reordered so that axis 0 moves last. + is_dm = plugin == "rsciio.digitalmicrograph" + axis_order = (1, 2, 0) if is_dm else (0, 1, 2) + array = imported_data["data"].transpose(axis_order) if is_dm else imported_data["data"] ordered_axes = [imported_axes[idx] for idx in axis_order] sampling = [ax.get("scale", 1) for ax in ordered_axes] origin = [ax.get("offset", 0) for ax in ordered_axes] @@ -266,10 +489,7 @@ def read_2d( -------- Dataset """ - if file_type is None: - file_type = Path(file_path).suffix.lower().lstrip(".") - - file_reader = importlib.import_module(f"rsciio.{file_type}").file_reader + _, file_reader = _rsciio_reader(file_path, file_type) imported_data = file_reader(file_path)[0] dataset = Dataset2d.from_array( diff --git a/src/quantem/core/io/serialize.py b/src/quantem/core/io/serialize.py index 8c5523bf7..d1f959a86 100644 --- a/src/quantem/core/io/serialize.py +++ b/src/quantem/core/io/serialize.py @@ -1,5 +1,6 @@ import gzip import io +import json import os import shutil import tempfile @@ -171,6 +172,42 @@ def _convert_string_to_path_if_needed(val: Any, group: zarr.Group, key: str) -> return val return val + @staticmethod + def _convert_string_to_device_if_needed(val: Any, group: zarr.Group, key: str) -> Any: + """Convert string back to torch.device if it was originally a device.""" + if isinstance(val, str) and group.attrs.get(f"{key}.is_torch_device", False): + try: + return torch.device(val) + except (ValueError, RuntimeError): + return val + return val + + @staticmethod + def _read_ase_atoms(group: zarr.Group) -> Any: + """Rebuild an ase.Atoms written by `_serialize_value`.""" + from ase import Atoms + + if "arrays" in group.group_keys(): + arrays_group = AutoSerialize._get_group(group, "arrays") + arrays = { + k: AutoSerialize._read_array_np(arrays_group, k) for k in arrays_group.array_keys() + } + arrays.update( + {k: np.asarray(v) for k, v in dict(group.attrs.get("text_arrays", {})).items()} + ) + else: # files written before per-atom arrays were stored + arrays = {k: AutoSerialize._read_array_np(group, k) for k in ("numbers", "positions")} + atoms = Atoms( + numbers=arrays.pop("numbers"), + positions=arrays.pop("positions"), + cell=AutoSerialize._read_array_np(group, "cell"), + pbc=AutoSerialize._read_array_np(group, "pbc"), + ) + for key, arr in arrays.items(): + atoms.set_array(key, arr) + atoms.info.update(dict(group.attrs.get("info", {}))) + return atoms + @staticmethod def _is_autoserialize_instance(value: Any) -> bool: """Return True if value behaves like an AutoSerialize instance, even across autoreloads.""" @@ -406,6 +443,12 @@ def _serialize_value( group.attrs[name] = str(value) group.attrs[f"{name}.is_path"] = True + elif isinstance(value, torch.device): + # A device belongs to the machine, not to the data: store the string + # so the object reloads on a host that does not have that device. + group.attrs[name] = str(value) + group.attrs[f"{name}.is_torch_device"] = True + elif self._is_autoserialize_instance(value): # Nested AutoSerialize subtree subgroup = group.require_group(name) @@ -445,6 +488,27 @@ def _serialize_value( subgroup.attrs["_rng_type"] = "torch.Generator" # Don't try to save the state - it's not essential for core functionality + elif type(value).__module__.startswith("ase.") and type(value).__name__ == "Atoms": + # Stored as plain arrays so the file stays readable without pickling + # an ase version in: cell, pbc, every per-atom array in atoms.arrays + # (numbers, positions, occupancy, tags, masses, ...) and the + # JSON-serializable entries of atoms.info. Constraints and attached + # calculators are not saved. + subgroup = group.require_group(name) + subgroup.attrs["_ase_atoms"] = True + self._write_ndarray(subgroup, "cell", np.asarray(value.get_cell()), compressors) + self._write_ndarray(subgroup, "pbc", np.asarray(value.get_pbc()), compressors) + arrays_group = subgroup.require_group("arrays") + text_arrays = {} + for key, arr in value.arrays.items(): + arr = np.asarray(arr) + if arr.dtype.kind in "biufc": + self._write_ndarray(arrays_group, key, arr, compressors) + else: + text_arrays[key] = arr.tolist() + subgroup.attrs["text_arrays"] = _json_entries(text_arrays, f"{name}.arrays") + subgroup.attrs["info"] = _json_entries(dict(value.info), f"{name}.info") + else: # Fallback: dill-serialize + gzip-compress print(f"falling back in serialize for {name} of type {type(value)}") @@ -522,6 +586,7 @@ def _recursive_load( name == "_autoserialize" or name.endswith(".torch_save") or name.endswith(".is_path") + or name.endswith(".is_torch_device") ): continue # Skip metadata/flags if name in skip_names: @@ -531,6 +596,7 @@ def _recursive_load( # Convert string paths back to pathlib.Path objects if needed val = cls._convert_string_to_path_if_needed(val, group, name) + val = cls._convert_string_to_device_if_needed(val, group, name) setattr(obj, name, val) set_attrs.add(name) @@ -558,8 +624,16 @@ def _recursive_load( continue subgrp = AutoSerialize._get_group(group, name) + # ase.Atoms group + if subgrp.attrs.get("_ase_atoms"): + atoms = AutoSerialize._read_ase_atoms(subgrp) + if type(atoms) in skip_types: + continue + setattr(obj, name, atoms) + set_attrs.add(name) + # torch tensor group - if subgrp.attrs.get("_torch_tensor"): + elif subgrp.attrs.get("_torch_tensor"): data = AutoSerialize._read_array_np(subgrp, "tensor").tobytes() buf = io.BytesIO(data) tensor = torch.load(buf, map_location="cpu", weights_only=False) @@ -879,6 +953,7 @@ def maybe_tensor(group, key): val = group.attrs[key] # Convert string paths back to Path objects if needed val = cls._convert_string_to_path_if_needed(val, group, key) + val = cls._convert_string_to_device_if_needed(val, group, key) items.append(val) elif key in group.array_keys(): items.append(maybe_tensor(group, key)) @@ -994,6 +1069,7 @@ def maybe_tensor(group, key): val = group.attrs[key] # Convert string paths back to Path objects if needed val = cls._convert_string_to_path_if_needed(val, group, key) + val = cls._convert_string_to_device_if_needed(val, group, key) items.append(val) elif key in group.array_keys(): items.append(maybe_tensor(group, key)) @@ -1083,11 +1159,13 @@ def maybe_tensor(group, key): key == "_container_type" or key.endswith(".torch_save") or key.endswith(".is_path") + or key.endswith(".is_torch_device") ): continue val = group.attrs[key] # Convert string paths back to Path objects if needed val = cls._convert_string_to_path_if_needed(val, group, key) + val = cls._convert_string_to_device_if_needed(val, group, key) result[key] = val # Restore arrays (including torch tensors) for key in group.array_keys(): @@ -1504,3 +1582,57 @@ def _recurse(obj: Any, prefix: str = "", current_depth: int = 0, is_last: bool = ) _recurse(root) + + +def _json_entries(entries: dict, label: str) -> dict: + """Return the JSON-serializable entries of a dict, warning about the rest.""" + kept, dropped = {}, [] + for key, val in entries.items(): + try: + json.dumps({str(key): val}) + except (TypeError, ValueError): + dropped.append(str(key)) + else: + kept[str(key)] = val + if dropped: + print(f"Not saving non-JSON-serializable entries of {label}: {dropped}") + return kept + + +class Bundle(AutoSerialize): + """A named collection of serializable objects, saved as one file. + + Groups any AutoSerialize objects (Datasets, Vectors, ...) plus plain + metadata values under attribute names:: + + bundle = Bundle(adf=dataset2d, peaks=vector, note="IM689") + bundle.save("data.zip") + b = load("data.zip"); b.adf, b.peaks + + Parameters + ---------- + **objects + Objects to store, keyed by attribute name. Names must not shadow an + existing attribute or method of the class (e.g. ``save``). + + Raises + ------ + ValueError + If a name shadows a class attribute or method. + """ + + def __init__(self, **objects): + reserved = sorted(name for name in objects if hasattr(type(self), name)) + if reserved: + raise ValueError( + f"Bundle names {reserved} shadow Bundle/AutoSerialize attributes; " + "choose different names." + ) + for name, obj in objects.items(): + setattr(self, name, obj) + + def __repr__(self) -> str: + items = ", ".join( + f"{k}: {type(v).__name__}" for k, v in vars(self).items() if not k.startswith("_") + ) + return f"Bundle({items})" diff --git a/src/quantem/core/utils/clustering.py b/src/quantem/core/utils/clustering.py new file mode 100644 index 000000000..ff3d0abb3 --- /dev/null +++ b/src/quantem/core/utils/clustering.py @@ -0,0 +1,261 @@ +"""Density-based clustering (DBSCAN) in pure torch, and Vector helpers. + +No external clustering dependency: neighbors are found with blockwise +distance computations pruned by a sliding sorted window along the widest +dimension. Cluster connectivity among core points is resolved in a single +pass with scipy's sparse connected-components graph routine, rather than +iterative label propagation. Runs on CPU or any torch device. +""" + +from __future__ import annotations + +import numpy as np +import scipy.sparse as sp +import torch +from scipy.sparse.csgraph import connected_components + + +def dbscan( + points, + eps: float, + min_samples: int, + device: str | torch.device = "cpu", + block: int = 2048, + sort_by_size: bool = True, +) -> np.ndarray: + """DBSCAN cluster labels for a point set. + + Same semantics as the standard algorithm: points with at least + `min_samples` neighbors within `eps` (self included) are core points; + core points within `eps` of each other share a cluster; non-core points + within `eps` of a core point join that core's cluster (ties resolved to + the NEAREST core, deterministically); everything else is noise (-1). + + Parameters + ---------- + points : array-like (N, D) + Point coordinates. Scale the columns beforehand to weight + dimensions (the metric is plain Euclidean). + eps : float + Neighborhood radius. + min_samples : int + Neighbors (including the point itself) required for a core point. + device : str | torch.device, default="cpu" + Torch device for the distance computations. + block : int, default=2048 + Rows per distance block; lower to reduce memory. + sort_by_size : bool, default=True + Renumber clusters largest-first (label 0 is the biggest cluster). + + Returns + ------- + np.ndarray + (N,) integer labels, -1 for noise. + """ + pts = torch.as_tensor(np.asarray(points), dtype=torch.float32, device=device) + N, D = pts.shape + if N == 0: + return np.zeros(0, dtype=int) + + # sort along the widest dimension so neighbor candidates live in a + # contiguous window of the sorted order + spans = pts.max(dim=0).values - pts.min(dim=0).values + d0 = int(spans.argmax()) + order = torch.argsort(pts[:, d0]) + ps = pts[order] + key = ps[:, d0].contiguous() + + def block_candidates(i0: int, i1: int) -> tuple[int, int]: + lo = float(key[i0]) - eps + hi = float(key[i1 - 1]) + eps + j0 = int(torch.searchsorted(key, torch.tensor(lo, device=device))) + j1 = int(torch.searchsorted(key, torch.tensor(hi, device=device), right=True)) + return j0, j1 + + # pass 1: neighbor counts -> core points + counts = torch.zeros(N, dtype=torch.long, device=device) + for i0 in range(0, N, block): + i1 = min(i0 + block, N) + j0, j1 = block_candidates(i0, i1) + d = torch.cdist(ps[i0:i1], ps[j0:j1]) + counts[i0:i1] = (d <= eps).sum(dim=1) + core = counts >= min_samples + + # pass 2: connect core points within eps of each other, once, then + # resolve clusters as connected components (single pass, no iterative + # relabeling rounds) + edge_rows: list[torch.Tensor] = [] + edge_cols: list[torch.Tensor] = [] + for i0 in range(0, N, block): + i1 = min(i0 + block, N) + if not bool(core[i0:i1].any()): + continue + j0, j1 = block_candidates(i0, i1) + d = torch.cdist(ps[i0:i1], ps[j0:j1]) + adj = (d <= eps) & core[i0:i1, None] & core[None, j0:j1] + ii, jj = torch.nonzero(adj, as_tuple=True) + edge_rows.append(ii + i0) + edge_cols.append(jj + j0) + + labels = torch.full((N,), -1, dtype=torch.long, device=device) + if edge_rows: + rows = torch.cat(edge_rows).cpu().numpy() + cols = torch.cat(edge_cols).cpu().numpy() + graph = sp.coo_matrix((np.ones(rows.shape[0], dtype=np.int8), (rows, cols)), shape=(N, N)) + _, comp = connected_components(graph, directed=False) + core_np = core.cpu().numpy() + labels = torch.as_tensor(np.where(core_np, comp, -1), dtype=torch.long, device=device) + + # pass 3: border points join the nearest core cluster within eps + for i0 in range(0, N, block): + i1 = min(i0 + block, N) + bmask = ~core[i0:i1] + if not bool(bmask.any()): + continue + j0, j1 = block_candidates(i0, i1) + d = torch.cdist(ps[i0:i1], ps[j0:j1]) + d = torch.where(core[None, j0:j1], d, torch.full_like(d, torch.inf)) + d_min, j_min = d.min(dim=1) + near = bmask & (d_min <= eps) + if bool(near.any()): + lab = labels[i0:i1] + lab[near] = labels[j0 + j_min[near]] + labels[i0:i1] = lab + + # back to input order, compact label ids + out = torch.full((N,), -1, dtype=torch.long, device=device) + out[order] = labels + out_np = out.cpu().numpy() + uniq, inv = np.unique(out_np[out_np >= 0], return_inverse=True) + if uniq.size: + if sort_by_size: + sizes = np.bincount(inv) + rank = np.empty_like(sizes) + rank[np.argsort(sizes)[::-1]] = np.arange(sizes.size) + out_np[out_np >= 0] = rank[inv] + else: + out_np[out_np >= 0] = inv + return out_np + + +def cluster_vector( + vector, + fields, + eps: float, + min_samples: int, + field_scales=None, + scan_scales=None, + device: str | torch.device = "cpu", + label_field: str = "cluster", +): + """DBSCAN over the rows of a Vector, in any combination of field and + scan coordinates. + + Builds the clustering space from the named fields (optionally scaled per + field) plus, when `scan_scales` is given, the scan-grid indices of each + row (row, col of the cell it belongs to, scaled). Digital dark field + clustering is the special case fields=(qx, qy) with scan_scales set. + + Parameters + ---------- + vector : Vector + Ragged vector over a scan grid. + fields : sequence of str + Field names contributing dimensions. + eps : float + DBSCAN neighborhood radius, in the scaled clustering space + (Euclidean metric). + min_samples : int + Neighbors (including the point itself) required for a core point. + field_scales : sequence of float | None + Multiplier per field; default 1. + scan_scales : (float, float) | None + If given, append (row * s0, col * s1) of each row's scan cell. + device : str | torch.device, default="cpu" + Torch device for the DBSCAN distance computations. + label_field : str, default="cluster" + Name of the label field on the returned Vector. + + Returns + ------- + labeled : Vector + Copy of `vector` with the integer labels appended as a new field + (-1 = noise). + labels : np.ndarray + The flat label array, aligned with vector.numpy().astype(np.float64). + """ + flat = vector.select_fields(*fields).numpy().astype(float) + if field_scales is not None: + flat = flat * np.asarray(field_scales, dtype=float)[None, :] + dims = [flat] + if scan_scales is not None: + counts = np.asarray(vector.row_counts(), dtype=int) + shape = vector.shape[:2] + cell_r, cell_c = np.divmod(np.arange(counts.size), shape[1]) + rr = np.repeat(cell_r, counts) * float(scan_scales[0]) + cc = np.repeat(cell_c, counts) * float(scan_scales[1]) + dims.append(np.stack([rr, cc], axis=1)) + space = np.concatenate(dims, axis=1) + + labels = dbscan(space, eps=eps, min_samples=min_samples, device=device) + + labeled = vector.copy() + labeled.add_fields([label_field], units=["index"]) + full = labeled.numpy().astype(np.float64) + full[:, -1] = labels + labeled.set_flattened(full) + return labeled, labels + + +def filter_rows(vector, mask): + """Copy of a ragged Vector keeping only the flattened rows where mask. + + The scan-grid shape is unchanged; rows are dropped from their cells. + + Parameters + ---------- + vector : Vector + Ragged vector over a 2D scan grid. + mask : array-like of bool + (vector.total_rows,) keep flag per row, in the flattened order of + vector.numpy(). + + Returns + ------- + Vector + New Vector with the same fields, units, name, dtype and metadata, + holding only the kept rows. + + Raises + ------ + ValueError + If len(mask) differs from vector.total_rows. + """ + mask = np.asarray(mask, dtype=bool).ravel() + if mask.size != vector.total_rows: + raise ValueError( + f"mask has {mask.size} entries but the Vector has {vector.total_rows} rows." + ) + counts = np.asarray(vector.row_counts(), dtype=int) + flat = vector.numpy().astype(np.float64) + starts = np.concatenate([[0], np.cumsum(counts)]) + shape = vector.shape[:2] + nested = [] + for r in range(shape[0]): + row = [] + for c in range(shape[1]): + k = r * shape[1] + c + sel = mask[starts[k] : starts[k + 1]] + row.append(flat[starts[k] : starts[k + 1]][sel]) + nested.append(row) + from quantem.core.datastructures.vector import Vector + + out = Vector.from_data( + nested, + fields=list(vector.fields), + units=list(vector.units), + name=vector.name, + dtype=vector.dtype, + ) + out.metadata.update(vector.metadata) + return out diff --git a/src/quantem/core/utils/diffractive_imaging_utils.py b/src/quantem/core/utils/diffractive_imaging_utils.py index 0959c1b6c..ffb6d75fb 100644 --- a/src/quantem/core/utils/diffractive_imaging_utils.py +++ b/src/quantem/core/utils/diffractive_imaging_utils.py @@ -1,4 +1,6 @@ -from typing import Optional, Tuple +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional, Tuple import matplotlib.pyplot as plt import numpy as np @@ -8,21 +10,29 @@ from quantem.core.utils.filter import otsu_threshold +if TYPE_CHECKING: + from quantem.core.datastructures.dataset2d import Dataset2d + def fit_probe_circle( - img: np.ndarray, threshold: Optional[float] = None, show: bool = True + img: np.ndarray | Dataset2d, threshold: Optional[float] = None, show: bool = True ) -> Tuple[float, float, float]: """ Fit a circle to the probe shape in an image. Args: - img (np.ndarray): Input image containing the probe. + img (np.ndarray | Dataset2d): Input diffraction pattern containing the probe. threshold (Optional[float]): Threshold for binarization. If None, Otsu's method is used. show (bool): Whether to display the fitted circle. Default is True. Returns: Tuple[float, float, float]: Center coordinates (xc, yc) and radius R of the fitted circle. """ + if not isinstance(img, np.ndarray) and hasattr(img, "array"): + img = img.array + if img.ndim != 2: + raise ValueError(f"Expected a 2D diffraction pattern, got shape {img.shape}.") + if threshold is None: threshold = otsu_threshold(img) binary = ndi.binary_closing(img > threshold, iterations=2) diff --git a/src/quantem/core/visualization/visualization_utils.py b/src/quantem/core/visualization/visualization_utils.py index 126cb0962..5bac8029e 100644 --- a/src/quantem/core/visualization/visualization_utils.py +++ b/src/quantem/core/visualization/visualization_utils.py @@ -335,6 +335,11 @@ def _normalize_length_units(length_units: float, units: str) -> tuple[float, str return length_units, units +# Minimum AnchoredSizeBar padding (fraction of the font size) when a box is +# drawn behind the scale bar, so the box does not clip the label. +_SCALEBAR_BOX_MIN_PAD = 0.35 + + def add_scalebar_to_ax( ax: Axes, array_size: float, @@ -347,6 +352,9 @@ def add_scalebar_to_ax( loc: str | int, fontsize: int = 12, bold: bool = True, + box: bool = False, + box_color: str = "black", + box_alpha: float = 0.5, ) -> None: """Add a scale bar to a matplotlib axis. @@ -375,6 +383,18 @@ def add_scalebar_to_ax( Font size of the scale bar label in points. bold : bool Whether to render the scale bar label in bold. + box : bool, default=False + Draw a translucent box behind the bar and label so it stays + readable on any image (e.g. white bar on a black box, or black bar + on a white box). The padding is raised to at least + ``_SCALEBAR_BOX_MIN_PAD`` (fraction of the font size) so the box clears the + label. ``box``, ``box_color`` and ``box_alpha`` are only available when + calling this function directly; ScalebarConfig and show_2d do not + pass them. + box_color : str, default="black" + Fill color of the box. + box_alpha : float, default=0.5 + Opacity of the box. """ from matplotlib.font_manager import FontProperties @@ -403,14 +423,16 @@ def add_scalebar_to_ax( length_px, label, loc, - pad=pad_px, + pad=pad_px if not box else max(pad_px, _SCALEBAR_BOX_MIN_PAD), color=color, - frameon=False, + frameon=box, label_top=label_top, size_vertical=int(width_px), fontproperties=fontprops, sep=2 if label_top else int(round(0.3 * fontsize)), ) + if box: + bar.patch.set(facecolor=box_color, edgecolor="none", alpha=box_alpha) ax.add_artist(bar) @@ -452,7 +474,9 @@ def add_cbar_to_ax( formatter = ticker.ScalarFormatter(useMathText=True) formatter.set_scientific(True) - formatter.set_powerlimits((-1, 1)) + # only fall back to a shared exponent for genuinely extreme ranges: a + # correlation running 0 to 0.8 should read 0.0, 0.2, ... not 0, 2 x 10^-1 + formatter.set_powerlimits((-3, 4)) sm = cm.ScalarMappable(norm=norm, cmap=cmap) cb = fig.colorbar(sm, cax=cax, ticks=ticks, format=formatter) diff --git a/src/quantem/diffraction/__init__.py b/src/quantem/diffraction/__init__.py index e69de29bb..f167408e5 100644 --- a/src/quantem/diffraction/__init__.py +++ b/src/quantem/diffraction/__init__.py @@ -0,0 +1,12 @@ +from quantem.diffraction.bragg_vectors import BraggVectors as BraggVectors +from quantem.diffraction.crystal import Crystal as Crystal +from quantem.diffraction.crystal_map import CrystalMap as CrystalMap +from quantem.diffraction.orientation import OrientationMap as OrientationMap +from quantem.diffraction.phase import PhaseMap as PhaseMap +from quantem.diffraction.reverse_monte_carlo import ReverseMonteCarlo as ReverseMonteCarlo +from quantem.diffraction.strain import StrainMap as StrainMap +from quantem.diffraction import bloch as bloch +from quantem.diffraction import calibration as calibration +from quantem.diffraction import rotations as rotations +from quantem.diffraction import digital_dark_field as digital_dark_field +from quantem.diffraction import illumination as illumination diff --git a/src/quantem/diffraction/bloch.py b/src/quantem/diffraction/bloch.py new file mode 100644 index 000000000..1b469ef6d --- /dev/null +++ b/src/quantem/diffraction/bloch.py @@ -0,0 +1,4282 @@ +"""Dynamical (Bloch wave) electron diffraction. + +Simulation and refinement with the Bloch wave formulation of De Graef +(2003), ch. 5. The module has five families of functions: + +- Spot patterns: dynamical_pattern() gives the Bloch intensities of one + orientation at every thickness; refine_thickness() refits thickness and + phase of a fitted PhaseMap with them. +- Convergent beam patterns: calculate_cbed() and calculate_cbed_library() + (disk patterns), calculate_lacbed() (one reflection's rocking surface) + and calculate_kossel() (wide-angle Kossel patterns). +- Kossel reference patterns: calculate_kossel_reference() computes the + bright field over all beam directions once; kossel_from_reference() and + kossel_polar_from_reference() look patterns up from it. The line model + (kossel_lines(), render_kossel_lines(), kossel_line_segments()) describes + the same patterns as one profile per systematic row, with + kossel_reference_residual() adding the many-beam correction near zone + axes. +- Bragg-vector refinement: refine_dynamical() refines orientation, + thickness, in-plane deformation and phase per position against the + measured peak intensities; dynamical_maps(), plot_dynamical_maps(), + strain_crystal_frame() and plot_strain_crystal_frame() present the result. +- Image refinement: fit_disk_shape() and refine_dynamical_image() refine + against the diffraction pattern pixels. + +The structure matrix uses U_g = gamma_rel * F_g / pi with F_g the +kinematical structure factors (scattering amplitude per volume, +1/Angstrom^2), or the absorptive Weickenmeier-Kohl factors when the crystal +carries them (Crystal.calculate_dynamical_structure_factors), off-diagonals +U_(g-h) and diagonal 2 k0 s_g. One eigendecomposition per incident +direction gives the intensities at every thickness: + + psi(t) = C exp(2 pi i gamma t) C^-1 psi_0, A C = 2 k0 gamma C +""" + +from __future__ import annotations + +import os +import threading +import warnings +from concurrent.futures import ThreadPoolExecutor + +import numpy as np +import torch +from tqdm import tqdm + +from quantem.core.utils.utils import electron_wavelength_angstrom +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.defaults import ( + MIN_NUMBER_PEAKS, + MIN_SIM_INTENSITY_REL, + PAIR_DISTANCE, + POWER_INTENSITY, + SG_MAX, + resolve, +) +from quantem.diffraction.rotations import qrotate, sample_zone_axes + + +def relativistic_gamma(energy_ev: float) -> float: + """Relativistic mass factor 1 + eV / (m0 c^2) at beam energy energy_ev (eV).""" + return 1.0 + float(energy_ev) / 510998.95 + + +def _coupling_matrix( + crystal: Crystal, hkl_beams: torch.Tensor, gamma_rel: float +) -> tuple[torch.Tensor, float, bool]: + """Off-diagonal Bloch coupling matrix U_(g-h) for a beam list. + + Prefers the absorptive Weickenmeier-Kohl factors when the crystal has + them (calculate_dynamical_structure_factors); they carry the + relativistic and 1/pi factors already. Falls back to the kinematical + (Lobato) factors, purely elastic. + + Returns + ------- + U : torch.Tensor + (nb, nb) complex coupling matrix with zero diagonal. + u0_imag : float + Imaginary part of U_000 (mean absorption), 0 without absorption. + absorptive : bool + Whether absorptive factors were used. + """ + absorptive = getattr(crystal, "U_dyn", None) is not None + if absorptive: + hkl_all = crystal.hkl_dyn + U_all = crystal.U_dyn + else: + hkl_all = crystal.hkl + U_all = crystal.struct_factors * (gamma_rel / np.pi) + nb = hkl_beams.shape[0] + diff = hkl_beams[:, None, :] - hkl_beams[None, :, :] # (nb, nb, 3) + # integer key that is injective over the union of the stored indices + # and the queried differences, so a difference outside the stored box + # can never alias a stored factor (a scalar key over the stored box + # alone did: (2,0,0) unstored returned the factor of (-1,1,0)) + m = max(int(hkl_all.abs().max()), int(diff.abs().max())) + span = 2 * m + 1 + key_mult = torch.tensor([1, span, span**2], dtype=torch.long) + + def keys(h): + return ((h + m) * key_mult).sum(dim=-1) + + # vectorized lookup: binary search of the queried keys in the sorted + # stored keys + keys_all = keys(hkl_all) + order = torch.argsort(keys_all) + keys_sorted = keys_all[order] + + def lookup(k): + pos = torch.searchsorted(keys_sorted, k).clamp(max=keys_sorted.shape[0] - 1) + return torch.where(keys_sorted[pos] == k, order[pos], -1) + + U = torch.zeros((nb, nb), dtype=torch.complex128) + idx = lookup(keys(diff).reshape(-1)).reshape(nb, nb) + has = idx >= 0 + U[has] = U_all[idx[has]] + U.fill_diagonal_(0) + + u0_imag = 0.0 + if absorptive: + i0 = int(lookup(keys(torch.zeros((1, 3), dtype=torch.long)))[0]) + if i0 >= 0: + u0_imag = float(U_all[i0].imag) + return U, u0_imag, absorptive + + +# guards the per-crystal record of issued coverage warnings, which the +# threads of refine_dynamical() check concurrently +_coverage_lock = threading.Lock() + + +def _beam_universe(crystal: Crystal) -> tuple[torch.Tensor, torch.Tensor]: + """Reciprocal lattice points that can carry dynamical intensity. + + With absorptive factors present, every index of the stored factor set + (calculate_dynamical_structure_factors), including reflections whose + structure factor is zero: those fill by multiple scattering through + intermediate beams and would never enter the Bloch state if the beams + were taken from the kinematical list, which drops them. Without + absorptive factors, the kinematical list. Returns (hkl (N, 3) long, + g crystal-frame (N, 3)) without the 000 beam. + """ + if getattr(crystal, "U_dyn", None) is not None: + hkl = crystal.hkl_dyn + g = hkl.to(torch.float64) @ crystal.lat_recip + keep = torch.linalg.norm(g, dim=1) > 1e-9 + # points of the conventional index box that are not reciprocal + # lattice points of the primitive cell (centering absences: odd + # h+k+l in bcc, mixed parity in fcc) are never excited and are + # dropped; glide and screw absences such as Si 200 and 222 are + # lattice points and stay + keep &= _primitive_lattice_mask(crystal, g) + return hkl[keep], g[keep] + return crystal.hkl, crystal.g_vec + + +def _primitive_lattice_mask(crystal: Crystal, g: torch.Tensor) -> torch.Tensor: + """True where the Cartesian reciprocal vectors g are points of the + primitive reciprocal lattice of the crystal (integer coordinates + g . a_i for the primitive real-space vectors a_i).""" + lat_p = getattr(crystal, "_primitive_lattice", None) + if lat_p is None: + import spglib + + from quantem.diffraction.crystal import _spglib_raises + + cell = ( + crystal.lat_real.numpy(), + crystal.positions_frac.numpy(), + crystal.numbers.numpy(), + ) + try: + with _spglib_raises(): + prim = spglib.standardize_cell(cell, to_primitive=True, no_idealize=True) + except Exception: + # no primitive cell found: keep every point of the stored box + prim = None + lat_p = np.asarray(crystal.lat_real.numpy() if prim is None else prim[0], dtype=float) + crystal._primitive_lattice = lat_p + m = g.to(torch.float64) @ torch.as_tensor(lat_p, dtype=torch.float64).T + return (m - torch.round(m)).abs().amax(dim=1) < 1e-6 + + +def _check_dynamical_factors(crystal: Crystal, energy_ev: float, g_max_beams: float) -> None: + """Warn when the absorptive factors are missing, were computed at + another energy, or stop short of 1.5 times the largest beam (the + couplings g - h reach twice it, but the factors beyond 1.5 times are + negligible). + + A warning is issued once per crystal and factor set: the record is + kept on the crystal, keyed by the energy and extent of its dynamical + factors and the energy of the calculation, so recomputing the factors + or running at another energy is checked again. + """ + key = ( + getattr(crystal, "dyn_energy_ev", None), + getattr(crystal, "dyn_k_max", None), + round(float(energy_ev)), + ) + warned = getattr(crystal, "_bloch_coverage_warned", None) + if warned is not None and key in warned: + return + msgs = [] + if getattr(crystal, "U_dyn", None) is None: + k_kin = getattr(crystal, "k_max", None) + msg = ( + "no absorptive structure factors (calculate_dynamical_structure_factors), " + "so the Bloch calculation uses the elastic kinematical factors" + ) + if k_kin is not None and 1.5 * g_max_beams > k_kin + 1e-9: + msg += ( + f", which stop at {k_kin:.2f} 1/A, short of the " + f"{1.5 * g_max_beams:.2f} 1/A the couplings of this beam list need" + ) + msgs.append(msg) + else: + e_dyn = getattr(crystal, "dyn_energy_ev", None) + k_dyn = getattr(crystal, "dyn_k_max", None) + if e_dyn is not None and abs(e_dyn - energy_ev) > 1.0: + msgs.append( + f"dynamical structure factors were computed at {e_dyn:.0f} eV, the " + f"calculation runs at {energy_ev:.0f} eV" + ) + # couplings g - h reach twice the beam radius, but the factors fall + # off fast: 1.5 times it keeps every coupling that matters (to 5%, + # since the fitted in-plane strain stretches the beams a little past + # k_max) + if k_dyn is not None and 1.5 * g_max_beams > 1.05 * k_dyn: + msgs.append( + f"dynamical structure factors extend to {k_dyn:.2f} 1/A but the beam " + f"list reaches {g_max_beams:.2f} 1/A, so couplings beyond " + f"{k_dyn:.2f} 1/A are missing (treated as zero); recompute with " + f"k_max >= {1.5 * g_max_beams:.2f}" + ) + if not msgs: + return + with _coverage_lock: + warned = getattr(crystal, "_bloch_coverage_warned", None) + if warned is None: + warned = set() + crystal._bloch_coverage_warned = warned + if key in warned: + return + warned.add(key) + warnings.warn(f"{crystal.name}: " + "; ".join(msgs), stacklevel=3) + + +def select_dynamical_beams( + crystal: Crystal, + orientation: torch.Tensor, + energy_ev: float, + alpha_max_rad: float = 0.0, + sg_max: float = SG_MAX, + k_max: float | None = None, + deform: torch.Tensor | None = None, +) -> torch.Tensor: + """Beam list for a Bloch calculation over a range of incident directions. + + Selects the reflections with |s_g| < sg_max + alpha_max_rad |g| at the + given orientation, which covers every incident direction within + alpha_max_rad of the optic axis (a tilt t shifts s_g by at most + |t| |g| / k0 to leading order). One list computed with the largest tilt + of a refinement search keeps every stage of that search in the same + truncated system. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + energy_ev : float + Beam energy in eV. + alpha_max_rad : float, default=0.0 + Largest incident tilt from the optic axis, in radians. + sg_max : float, default=SG_MAX + Excitation error cutoff in 1/Angstroms at zero tilt. + k_max : float | None + Largest |g| in 1/Angstroms; None keeps every candidate reflection. + deform : torch.Tensor | None + (3, 3) deformation applied to the lab-frame reciprocal vectors. + + Returns + ------- + torch.Tensor + (nb, 3) Miller indices, with the 000 beam first. + """ + lam = electron_wavelength_angstrom(energy_ev) + hkl_u, g_u = _beam_universe(crystal) + g_lab = qrotate(orientation, g_u) + if deform is not None: + g_lab = g_lab @ deform.to(torch.float64).T + g_len = torch.linalg.norm(g_lab, dim=1) + gz, g2 = g_lab[:, 2], (g_lab**2).sum(dim=1) + s0 = (2 * gz - lam * g2) / (2 - 2 * lam * gz) + sel = torch.abs(s0) < sg_max + alpha_max_rad * g_len + if k_max is not None: + sel &= g_len <= k_max + return torch.cat([torch.zeros((1, 3), dtype=torch.long), hkl_u[sel]]) + + +def dynamical_pattern( + crystal: Crystal, + orientation: torch.Tensor, + thicknesses_A: torch.Tensor | np.ndarray | float, + energy_ev: float = 300e3, + sg_max: float = SG_MAX, + k_max: float | None = None, +) -> dict[str, torch.Tensor]: + """Bloch-wave diffraction intensities for one orientation, all thicknesses. + + The beams are the reflections with |s_g| < sg_max (and |g| <= k_max); + one eigendecomposition gives the intensities at every thickness. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. Preferably also with + calculate_dynamical_structure_factors (absorptive factors, at this + energy, covering at least 1.5 times k_max so the couplings g - h + that matter have a factor); a warning is issued otherwise. + orientation : torch.Tensor + Unit quaternion (4,) rotating crystal vectors into the lab frame. + thicknesses_A : array-like or float + Specimen thicknesses in Angstroms. + energy_ev : float, default=300e3 + Beam energy in eV. + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms) for including a beam. + k_max : float | None + Largest |g| (1/Angstroms) of an included beam; None keeps every + reflection within sg_max. + + Returns + ------- + dict + 'qx', 'qy' (N,) lab-frame positions (1/Angstroms), 'hkl' (N, 3), + 's_g' (N,) excitation errors, 'intensity' (T, N) diffracted + intensities per thickness, 'intensity_000' (T,) the direct beam, + 'thicknesses' (T,) in Angstroms. + """ + if crystal.g_vec is None: + raise RuntimeError("Run crystal.calculate_structure_factors() first.") + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + gamma_rel = relativistic_gamma(energy_ev) + + t = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + + # beam selection in the lab frame, from every lattice point that can + # carry dynamical intensity + hkl_u, g_u = _beam_universe(crystal) + g_lab = qrotate(orientation, g_u) + gz, g2 = g_lab[:, 2], (g_lab**2).sum(dim=1) + s_g = (2 * gz - lam * g2) / (2 - 2 * lam * gz) + sel = torch.abs(s_g) < sg_max + if k_max is not None: + sel &= torch.linalg.norm(g_lab, dim=1) <= k_max + hkl_sel = hkl_u[sel] + g_sel = g_lab[sel] + s_sel = s_g[sel] + _check_dynamical_factors( + crystal, energy_ev, float(torch.linalg.norm(g_sel, dim=1).max()) if g_sel.shape[0] else 0.0 + ) + + # beams list includes the (000) beam at index 0 + hkl_beams = torch.cat([torch.zeros((1, 3), dtype=torch.long), hkl_sel]) + s_beams = torch.cat([torch.zeros(1, dtype=torch.float64), s_sel]) + + U, u0_imag, absorptive = _coupling_matrix(crystal, hkl_beams, gamma_rel) + A = U.clone() + diag = (2 * k0 * s_beams).to(torch.complex128) + if absorptive: + # mean absorption: imaginary part of U_000 damps every beam + diag = diag + 1j * u0_imag + A += torch.diag(diag) + + if absorptive: + # non-Hermitian: general eigendecomposition, complex gamma damps + evals, C = torch.linalg.eig(A) + gam = evals / (2 * k0) + else: + evals, C = torch.linalg.eigh(A) + gam = (evals.real / (2 * k0)).to(torch.complex128) + psi0 = torch.linalg.inv(C)[:, 0] # C^-1 @ e_0 + phase = torch.exp(2j * np.pi * gam[None, :] * t.to(torch.complex128)[:, None]) + psi = torch.einsum("ij,tj,j->ti", C, phase, psi0) # (T, nb) + intensity = torch.abs(psi[:, 1:]) ** 2 # drop the (000) beam + + return { + "qx": g_sel[:, 0], + "qy": g_sel[:, 1], + "hkl": hkl_sel, + "s_g": s_sel, + "intensity": intensity, + "intensity_000": torch.abs(psi[:, 0]) ** 2, + "thicknesses": t, + } + + +def refine_thickness( + phase_map, + thicknesses_A: np.ndarray | None = None, + pair_distance: float | None = None, + power_intensity: float | None = None, + sg_max: float = SG_MAX, + k_max: float | None = None, + min_number_peaks: int | None = None, + progress_bar: bool = True, +): + """Thickness and phase refinement with dynamical intensities. + + For every probe position, the winning candidates of a fitted PhaseMap are + re-simulated with Bloch waves over a thickness grid at their matched + orientations. The peak pairing is fixed (positions are kinematic); the + intensity cost is evaluated for all thicknesses from a single + eigendecomposition per candidate, and the best (thickness, candidate) + combination updates the phase decision. refine_dynamical() also refines + the orientation and the in-plane deformation. + + Parameters left as None inherit the phase fit's values (see + PhaseMap.fit); the resolved values are recorded in + phase_map.metadata['thickness']. + + Parameters + ---------- + phase_map : PhaseMap + A fitted PhaseMap (fit() has been run). + thicknesses_A : np.ndarray | None + Thickness grid in Angstroms; default 50 to 1000 in 25 A steps. + pair_distance : float | None + Largest distance (1/Angstroms) at which a simulated and a measured + peak are paired. + power_intensity : float | None + Intensities are compared as I ** power_intensity. + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms) of the Bloch beam list. + k_max : float | None + Largest |g| (1/Angstroms) of a beam; None keeps every reflection + within sg_max. + min_number_peaks : int | None + Positions with fewer measured peaks, direct beam included, are + skipped; None inherits the phase fit's minimum. At least 3. + progress_bar : bool, default=True + Show a progress bar over positions. + + Returns + ------- + dict + 'thickness' (R, C) best-fit thickness of the winning candidate, + 'cost' (R, C, F) dynamical cost per candidate at its best + thickness, 'phase_index' (R, C) updated phase assignment (-1 where + no candidate was refined), 'thickness_per_candidate' (R, C, F). + NaN where a candidate was not refined. + """ + if thicknesses_A is None: + thicknesses_A = np.arange(50.0, 1000.0 + 1e-6, 25.0) + t_grid = torch.as_tensor(thicknesses_A, dtype=torch.float64) + + oms = phase_map.orientation_maps + fit_md = phase_map.metadata.get("fit") if hasattr(phase_map, "metadata") else None + pair_distance = resolve(pair_distance, "pair_distance", fit_md, default=PAIR_DISTANCE) + power_intensity = resolve(power_intensity, "power_intensity", fit_md, default=POWER_INTENSITY) + min_number_peaks = int( + resolve(min_number_peaks, "min_number_peaks", fit_md, default=MIN_NUMBER_PEAKS) + ) + if min_number_peaks < 3: + raise ValueError( + f"min_number_peaks={min_number_peaks}: a dynamical fit needs at least the " + "direct beam and two non-collinear reflections" + ) + if hasattr(phase_map, "metadata"): + phase_map.metadata["thickness"] = dict( + thicknesses_A=np.asarray(thicknesses_A, dtype=float).tolist(), + pair_distance=float(pair_distance), + power_intensity=float(power_intensity), + sg_max=float(sg_max), + k_max=k_max, + min_number_peaks=int(min_number_peaks), + ) + cands = phase_map.candidates + peaks = oms[0].peaks + R, C = peaks.shape[0], peaks.shape[1] + F = len(cands) + delta = pair_distance + + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + + cost_out = torch.full((R, C, F), torch.nan, dtype=torch.float64) + thick_out = torch.full((R, C, F), torch.nan, dtype=torch.float64) + + iterator = list(np.ndindex(R, C)) + if progress_bar: + iterator = tqdm(iterator, desc="dynamical refinement") + for rx, ry in iterator: + data = peaks[rx, ry].numpy().astype(np.float64) + if data.shape[0] < min_number_peaks: + continue + qxy = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + im = torch.as_tensor(data[:, ix[2]], dtype=torch.float64).clamp_min(0) + im = im**power_intensity + int_total = float(im.sum()) + + for f, (i_om, m) in enumerate(cands): + om = oms[i_om] + if om.corr[rx, ry, m] <= 0: + continue + # only refine candidates that won weight in the first pass + if ( + phase_map.phase_weights is not None + and float(phase_map.phase_weights[rx, ry, f]) <= 0 + ): + continue + sim = dynamical_pattern( + om.crystal, + om.quats[rx, ry, m], + t_grid, + energy_ev=om.energy_ev, + sg_max=sg_max, + k_max=k_max, + ) + sq = torch.stack((sim["qx"], sim["qy"]), dim=1) + if sq.shape[0] == 0: + continue + si = sim["intensity"] ** power_intensity # (T, N) + d = torch.cdist(sq, qxy) + d_min, j_min = d.min(dim=1) + pair = d_min < delta + frac = (d_min[pair] / delta).clamp(0, 1) + + a = si[:, pair] * (1 - frac)[None, :] # (T, P) + b = im[j_min[pair]][None, :] + w = (a * b).sum(dim=1) / (a * a).sum(dim=1).clamp_min(1e-12) # (T,) + w = w.clamp_min(0) + + c_paired = ( + (b - w[:, None] * a).abs() * (1 - frac)[None, :] + w[:, None] * a * frac[None, :] + ).sum(dim=1) + c_unpaired_sim = 0.5 * w * si[:, ~pair].sum(dim=1) + matched = torch.zeros(im.shape[0], dtype=torch.bool) + matched[j_min[pair]] = True + c_unpaired_exp = 0.5 * float(im[~matched].sum()) + cost_t = (c_paired + c_unpaired_sim + c_unpaired_exp) / (int_total + 1e-12) + + t_best = int(cost_t.argmin()) + cost_out[rx, ry, f] = cost_t[t_best] + thick_out[rx, ry, f] = t_grid[t_best] + + # updated per-crystal phase decision from the dynamical costs + n_maps = len(oms) + cost_phase = torch.full((R, C, n_maps), torch.inf, dtype=torch.float64) + for f, (i_om, _) in enumerate(cands): + c = torch.nan_to_num(cost_out[..., f], nan=torch.inf) + cost_phase[..., i_om] = torch.minimum(cost_phase[..., i_om], c) + done = torch.isfinite(cost_out).any(dim=-1) + phase_index = torch.where(done, cost_phase.argmin(dim=-1), -1) + + f_best = torch.nan_to_num(cost_out, nan=torch.inf).argmin(dim=-1) + thickness = torch.gather(thick_out, 2, f_best[..., None]).squeeze(-1) + + return { + "thickness": thickness, + "cost": cost_out, + "phase_index": phase_index, + "thickness_per_candidate": thick_out, + } + + +def _cbed_amplitudes( + crystal: Crystal, + orientation: torch.Tensor, + tilts: torch.Tensor, + thicknesses_A: torch.Tensor, + energy_ev: float, + sg_max: float, + k_max: float | None, + tilt_batch: int = 64, + progress_bar: bool = False, + fast_absorption: bool = False, + deform: torch.Tensor | None = None, + beams: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Bloch intensities of every beam at every incident tilt. + + The coupling matrix is built once; only the diagonal (excitation + errors) changes with tilt, and the eigendecompositions are batched + over tilt chunks. + + Parameters + ---------- + tilts : torch.Tensor + (M, 2) in-plane incident wavevector components (1/Angstroms). + deform : torch.Tensor | None + (3, 3) deformation applied to the lab-frame reciprocal lattice + (g' = deform @ g), for a strained cell; the structure factors are + those of the ideal cell. + beams : torch.Tensor | None + Explicit beam list (nb, 3) hkl with 000 first, e.g. from + select_dynamical_beams(); when None the list is selected here from + the tilts given. + + Returns + ------- + intensity : torch.Tensor + (M, T, nb) beam intensities per tilt and thickness; beam 0 is the + direct (000) beam. + g_xy : torch.Tensor + (nb, 2) in-plane reciprocal vectors of the beams (000 first). + hkl_beams : torch.Tensor + (nb, 3) Miller indices of the beams. + """ + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + gamma_rel = relativistic_gamma(energy_ev) + t_thick = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + + # beam selection: near the Ewald sphere for ANY tilt in the aperture + if beams is None: + alpha_max = float(torch.linalg.norm(tilts, dim=1).max()) / k0 + hkl_beams = select_dynamical_beams( + crystal, orientation, energy_ev, alpha_max, sg_max, k_max, deform + ) + else: + hkl_beams = beams + g_beams = qrotate(orientation, hkl_beams[1:].to(torch.float64) @ crystal.lat_recip) + if deform is not None: + g_beams = g_beams @ deform.to(torch.float64).T + g_beams = torch.cat([torch.zeros((1, 3), dtype=torch.float64), g_beams]) + nb = hkl_beams.shape[0] + _check_dynamical_factors( + crystal, energy_ev, float(torch.linalg.norm(g_beams, dim=1).max()) if nb > 1 else 0.0 + ) + + U, u0_imag, absorptive = _coupling_matrix(crystal, hkl_beams, gamma_rel) + + gx, gy, gzb = g_beams[:, 0], g_beams[:, 1], g_beams[:, 2] + g2b = (g_beams**2).sum(dim=1) + M = tilts.shape[0] + out = torch.zeros((M, t_thick.shape[0], nb), dtype=torch.float64) + chunks = range(0, M, tilt_batch) + if progress_bar: + chunks = tqdm(chunks, desc="Bloch tilts") + for m0 in chunks: + m1 = min(m0 + tilt_batch, M) + tt = tilts[m0:m1] # (B, 2) + kz = torch.sqrt(k0**2 - (tt**2).sum(dim=1)) # (B,) + # s_g for incident k = (tx, ty, -kz), surface normal along z + num = ( + 2 * kz[:, None] * gzb[None, :] + - 2 * (tt[:, 0, None] * gx[None, :] + tt[:, 1, None] * gy[None, :]) + - g2b[None, :] + ) + den = 2 * (kz[:, None] - gzb[None, :]) + s_t = num / den # (B, nb) + + out[m0:m1] = _bloch_solve( + U, u0_imag, absorptive, s_t, k0, t_thick, fast_absorption=fast_absorption + ) + return out, g_beams[:, :2], hkl_beams + + +def _bloch_solve( + U: torch.Tensor, + u0_imag: float, + absorptive: bool, + s_t: torch.Tensor, + k0: float, + t_thick: torch.Tensor, + fast_absorption: bool = False, +) -> torch.Tensor: + """Batched Bloch solve: intensities (B, T, nb) for excitation errors s_t + (B, nb) with a shared coupling matrix U (nb, nb). + + With fast_absorption=True the Hermitian part is diagonalized (eigh, much + faster and better batched than the general complex eig) and the weak + absorption enters first order: gamma_imag = diag(C^dagger U'' C)/(2 k0). + Standard for reference (master) pattern computations; the absorptive + parts of U are a few percent of the elastic parts, so the first-order + error is small. + """ + nb = U.shape[0] + if absorptive and fast_absorption: + H = 0.5 * (U + U.conj().T) + W = (U - H) / 1j # Hermitian absorptive part (off-diagonal) + A = H[None].expand(s_t.shape[0], nb, nb).clone() + A += torch.diag_embed((2 * k0 * s_t).to(torch.complex128)) + evals, C = torch.linalg.eigh(A) + gam_r = evals / (2 * k0) # (B, nb) real + CW = torch.einsum("bji,jk,bki->bi", C.conj(), W, C).real + gam_i = (CW + u0_imag) / (2 * k0) # (B, nb) + gam = gam_r.to(torch.complex128) + 1j * gam_i + psi0 = C.conj().transpose(1, 2)[:, :, 0] # unitary: C^-1 = C^dagger + else: + diag = (2 * k0 * s_t).to(torch.complex128) + if absorptive: + diag = diag + 1j * u0_imag + A = U[None].expand(s_t.shape[0], nb, nb).clone() + A += torch.diag_embed(diag) + if absorptive: + evals, C = torch.linalg.eig(A) + gam = evals / (2 * k0) + else: + evals, C = torch.linalg.eigh(A) + gam = (evals.real / (2 * k0)).to(torch.complex128) + psi0 = torch.linalg.inv(C)[:, :, 0] # (B, nb) + phase = torch.exp( + 2j * np.pi * gam[:, None, :] * t_thick.to(torch.complex128)[None, :, None] + ) # (B, T, nb) + psi = torch.einsum("bij,btj,bj->bti", C, phase, psi0) + return torch.abs(psi) ** 2 + + +def tilt_grid(semiconv_mrad: float, energy_ev: float, n_rings: int = 8): + """Concentric-ring sampling of the illumination aperture. + + Parameters + ---------- + semiconv_mrad : float + Convergence semiangle in mrad. + energy_ev : float + Beam energy in eV. + n_rings : int, default=8 + Rings outside the center point; ring r has ceil(2 pi r) points. + + Returns + ------- + torch.Tensor + (M, 2) in-plane incident wavevectors (1/Angstroms) covering the + disk with approximately uniform density, the center first. + """ + lam = electron_wavelength_angstrom(energy_ev) + alpha_k = semiconv_mrad * 1e-3 / lam + pts = [(0.0, 0.0)] + for r in range(1, n_rings + 1): + rad = alpha_k * r / n_rings + n_az = int(np.ceil(2 * np.pi * r)) + th = 2 * np.pi * (np.arange(n_az) + 0.5 * (r % 2)) / n_az + pts += [(rad * np.cos(a), rad * np.sin(a)) for a in th] + return torch.tensor(pts, dtype=torch.float64) + + +def calculate_cbed( + crystal: Crystal, + orientation: torch.Tensor, + thicknesses_A, + energy_ev: float = 300e3, + semiconv_mrad: float = 3.0, + n_rings: int = 8, + sg_max: float = SG_MAX, + k_max: float | None = None, + pixel_size: float | None = None, + q_max_plot: float | None = None, + tilt_batch: int = 64, +) -> dict: + """Simulate a CBED pattern with Bloch waves. + + Every incident direction inside the aperture is an independent plane + wave (incoherent illumination): its Bloch intensities are placed at + g + t in the detector plane, filling each diffraction disk with the + rocking-curve intensity variation. Uses the absorptive + Weickenmeier-Kohl structure factors when the crystal carries them. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated, and preferably + calculate_dynamical_structure_factors at this energy (absorption). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + thicknesses_A : float | array-like + One or more specimen thicknesses in Angstroms. + energy_ev : float, default=300e3 + Beam energy in eV. + semiconv_mrad : float, default=3.0 + Convergence semiangle in mrad. Disks overlap when it exceeds half + the smallest g spacing times the wavelength. + n_rings : int, default=8 + Radial sampling rings across the aperture (~200 tilts at 8). + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms) for beam selection, widened + automatically by the aperture tilt range. + k_max : float | None + Largest |g| (1/Angstroms) of an included beam. + pixel_size : float | None + Detector sampling (1/Angstroms per pixel); default disk radius / 12. + q_max_plot : float | None + Half-width of the detector (1/Angstroms); default covers all beams + plus a disk. Disk samples beyond it are dropped. + tilt_batch : int, default=64 + Incident tilts per batched eigendecomposition (memory versus + speed). + + Returns + ------- + dict + 'pattern' ((T, H, H), squeezed to (H, H) for one thickness; rows + are qx, columns qy, the direct beam at the center pixel; the + intensity is averaged over the incident tilts, so it sums to the + transmitted fraction when every disk is on the detector), + 'sampling' (1/Angstroms per pixel), 'disk_radius' (1/Angstroms), + 'thicknesses', 'hkl' (nb, 3) and 'g_xy' (nb, 2) of the beams, 000 + first. + """ + lam = electron_wavelength_angstrom(energy_ev) + alpha_k = semiconv_mrad * 1e-3 / lam + tilts = tilt_grid(semiconv_mrad, energy_ev, n_rings=n_rings) + t_thick = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + + inten, g_xy, hkl_beams = _cbed_amplitudes( + crystal, orientation, tilts, t_thick, energy_ev, sg_max, k_max, tilt_batch + ) + + if pixel_size is None: + pixel_size = alpha_k / 12 + if q_max_plot is None: + q_max_plot = float(torch.linalg.norm(g_xy, dim=1).max()) + 2 * alpha_k + half = int(np.ceil(q_max_plot / pixel_size)) + H = 2 * half + 1 + + # deposit every (beam, tilt) sample with bilinear weights + qx = (g_xy[:, 0][None, :] + tilts[:, 0][:, None]).numpy() # (M, nb) + qy = (g_xy[:, 1][None, :] + tilts[:, 1][:, None]).numpy() + fx = qx / pixel_size + half + fy = qy / pixel_size + half + ix0 = np.floor(fx).astype(int) + iy0 = np.floor(fy).astype(int) + wx = fx - ix0 + wy = fy - iy0 + + T = t_thick.shape[0] + pattern = np.zeros((T, H, H)) + inten_np = inten.numpy() # (M, T, nb) + for dx in (0, 1): + for dy in (0, 1): + w = (wx if dx else 1 - wx) * (wy if dy else 1 - wy) + jx = ix0 + dx + jy = iy0 + dy + # samples beyond the detector are dropped, not piled on its edge + ok = (jx >= 0) & (jx < H) & (jy >= 0) & (jy < H) + for ti in range(T): + np.add.at(pattern[ti], (jx[ok], jy[ok]), (w * inten_np[:, ti, :])[ok]) + pattern /= tilts.shape[0] + + return { + "pattern": pattern[0] if T == 1 else pattern, + "sampling": float(pixel_size), + "disk_radius": float(alpha_k), + "thicknesses": t_thick.numpy(), + "hkl": hkl_beams.numpy(), + "g_xy": g_xy.numpy(), + } + + +def calculate_lacbed( + crystal: Crystal, + orientation: torch.Tensor, + thicknesses_A, + hkl, + energy_ev: float = 300e3, + semiconv_mrad: float = 10.0, + n_pixels: int = 48, + sg_max: float = SG_MAX, + k_max: float | None = None, + tilt_batch: int = 64, +) -> dict: + """Large-angle CBED: one reflection's rocking surface over the aperture. + + The intensity of the chosen reflection is mapped over the incident-tilt + disk on a square grid (parallax / LACBED view of a single disk, without + the geometric overlap of neighboring disks). + + Parameters + ---------- + crystal : Crystal + With structure factors calculated, and preferably + calculate_dynamical_structure_factors at this energy (absorption). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + thicknesses_A : float | array-like + One or more specimen thicknesses in Angstroms. + hkl : sequence of int + The reflection to map; (0, 0, 0) gives the bright field disk. + energy_ev : float, default=300e3 + Beam energy in eV. + semiconv_mrad : float, default=10.0 + Convergence semiangle in mrad. + n_pixels : int, default=48 + Pixels across the disk (the incident-tilt sampling). + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms), widened automatically by + the aperture tilt range. + k_max : float | None + Largest |g| (1/Angstroms) of an included beam. + tilt_batch : int, default=64 + Incident tilts per batched eigendecomposition. + + Returns + ------- + dict + 'disk' ((T, n, n), squeezed for one thickness; rows are the y + tilt, columns the x tilt; NaN outside the aperture), 'tilt_max' + (aperture radius, 1/Angstroms), 'thicknesses'. + + Raises + ------ + ValueError + If the reflection is not among the excited beams. + """ + lam = electron_wavelength_angstrom(energy_ev) + alpha_k = semiconv_mrad * 1e-3 / lam + ax = torch.linspace(-alpha_k, alpha_k, n_pixels, dtype=torch.float64) + ty, tx = torch.meshgrid(ax, ax, indexing="ij") + inside = (tx**2 + ty**2) <= alpha_k**2 + tilts = torch.stack([tx[inside], ty[inside]], dim=1) + t_thick = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + + inten, _, hkl_beams = _cbed_amplitudes( + crystal, orientation, tilts, t_thick, energy_ev, sg_max, k_max, tilt_batch + ) + match = (hkl_beams == torch.as_tensor(hkl, dtype=torch.long)[None, :]).all(dim=1) + if not bool(match.any()): + raise ValueError(f"reflection {tuple(hkl)} is not among the excited beams") + b = int(match.nonzero()[0]) + + T = t_thick.shape[0] + disk = np.full((T, n_pixels, n_pixels), np.nan) + m = inside.numpy() + for ti in range(T): + plane = np.full((n_pixels, n_pixels), np.nan) + plane[m] = inten[:, ti, b].numpy() + disk[ti] = plane + return { + "disk": disk[0] if T == 1 else disk, + "tilt_max": float(alpha_k), + "thicknesses": t_thick.numpy(), + } + + +def calculate_cbed_library( + crystal: Crystal, + orientations: torch.Tensor, + thickness_A: float, + energy_ev: float = 300e3, + semiconv_mrad: float = 3.0, + k_max: float | None = None, + q_max_plot: float | None = None, + pixel_size: float | None = None, + progress_bar: bool = True, + **kwargs, +) -> dict: + """A stack of simulated CBED patterns on one common detector grid. + + The starting point for CBED orientation matching: all patterns share + the same sampling and extent, ready for polar transformation and + correlation. One entry per orientation, all at one thickness. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. + orientations : torch.Tensor + (N, 4) unit quaternions, crystal to lab. + thickness_A : float + Specimen thickness in Angstroms. + energy_ev : float, default=300e3 + Beam energy in eV. + semiconv_mrad : float, default=3.0 + Convergence semiangle in mrad. + k_max : float | None + Largest |g| (1/Angstroms) of an included beam. + q_max_plot : float | None + Half-width of the detector (1/Angstroms); default k_max (or half + the crystal's structure factor range) plus two disk radii. + pixel_size : float | None + Detector sampling (1/Angstroms per pixel); default disk radius / 12. + progress_bar : bool, default=True + Show a progress bar over orientations. + **kwargs + Passed to calculate_cbed() (n_rings, sg_max, tilt_batch). + + Returns + ------- + dict + 'patterns' (N, H, H), 'quats' (N, 4), 'sampling' (1/Angstroms per + pixel), 'disk_radius' (1/Angstroms), 'thickness_A'. + """ + lam = electron_wavelength_angstrom(energy_ev) + alpha_k = semiconv_mrad * 1e-3 / lam + if pixel_size is None: + pixel_size = alpha_k / 12 + if q_max_plot is None: + base = k_max if k_max is not None else float(crystal.k_max) / 2 + q_max_plot = base + 2 * alpha_k + + quats = torch.atleast_2d(torch.as_tensor(orientations, dtype=torch.float64)) + pats = [] + it = range(quats.shape[0]) + if progress_bar: + it = tqdm(it, desc="CBED library") + for i in it: + res = calculate_cbed( + crystal, + quats[i], + thickness_A, + energy_ev=energy_ev, + semiconv_mrad=semiconv_mrad, + k_max=k_max, + pixel_size=pixel_size, + q_max_plot=q_max_plot, + **kwargs, + ) + pats.append(res["pattern"]) + return { + "patterns": np.stack(pats), + "quats": quats.numpy(), + "sampling": float(pixel_size), + "disk_radius": float(alpha_k), + "thickness_A": float(thickness_A), + } + + +def calculate_kossel( + crystal: Crystal, + orientation: torch.Tensor, + thicknesses_A, + energy_ev: float = 300e3, + semiconv_mrad: float = 40.0, + n_pixels: int = 192, + sg_max: float = 0.05, + k_max: float | None = None, + tilt_batch: int = 64, + fast_absorption: bool = False, + progress_bar: bool = True, +) -> dict: + """Wide-angle convergent beam (Kossel) pattern with Bloch waves. + + At convergence angles far beyond the Bragg angles the diffraction disks + overlap completely and the pattern becomes a continuous map of + deficiency and excess lines (the Kossel regime of CBED; the bright + field disk alone is the LACBED view). One Bloch computation over the + incident-tilt grid yields both: + + - 'bright_field': the (000) beam intensity at each incident tilt, the + deficiency (dark) line system, every line at a Bragg condition. + - 'pattern': the full detector intensity, the incoherent sum of every + diffracted cone shifted by its g: deficiency lines from the direct + beam plus the excess (bright) lines of the diffracted beams. + + Line positions are exact; line profiles carry the many-beam dynamical + structure, with the deficiency/excess asymmetry from the absorptive + structure factors when the crystal has them. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated, and preferably + calculate_dynamical_structure_factors at this energy (absorption). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + thicknesses_A : float | array-like + One or more thicknesses in Angstroms. + energy_ev : float, default=300e3 + Beam energy in eV. + semiconv_mrad : float, default=40.0 + Convergence semiangle in mrad; the pattern covers this angular + radius. + n_pixels : int, default=192 + Detector pixels across the pattern (also the tilt sampling; the + 1-2 mrad dynamical line widths need ~0.5 mrad per pixel). + sg_max : float, default=0.05 + Excitation error cutoff (1/Angstroms). Smaller than the SG_MAX of + the spot pattern functions: the beam list is widened by the + aperture (alpha |g|, already 0.04 1/A for |g| = 1 at 40 mrad), so + the base cutoff can be tighter without losing lines, and the + eigensolves over tens of thousands of tilts stay affordable. + k_max : float | None + Largest |g| (1/Angstroms) of an included reflection. + tilt_batch : int, default=64 + Incident tilts per batched eigendecomposition. + fast_absorption : bool, default=False + First-order absorption (Hermitian eigensolver, faster); see + _bloch_solve. + progress_bar : bool, default=True + Show a progress bar over tilt batches. + + Returns + ------- + dict + 'bright_field' and 'pattern' ((T, n, n), squeezed for one + thickness; rows are theta_y, columns theta_x, as in + render_kossel_lines; NaN / 0 outside the aperture), 'sampling' + (1/Angstroms per pixel), 'mrad_per_pixel', 'thicknesses', 'hkl'. + """ + lam = electron_wavelength_angstrom(energy_ev) + alpha_k = semiconv_mrad * 1e-3 / lam + ax = torch.linspace(-alpha_k, alpha_k, n_pixels, dtype=torch.float64) + px = float(ax[1] - ax[0]) + ty, tx = torch.meshgrid(ax, ax, indexing="ij") + inside = (tx**2 + ty**2) <= alpha_k**2 + tilts = torch.stack([tx[inside], ty[inside]], dim=1) + t_thick = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + T = t_thick.shape[0] + + inten, g_xy, hkl_beams = _cbed_amplitudes( + crystal, + orientation, + tilts, + t_thick, + energy_ev, + sg_max, + k_max, + tilt_batch, + progress_bar=progress_bar, + fast_absorption=fast_absorption, + ) + inten_np = inten.numpy() # (M, T, nb) + m = inside.numpy() + + # bright field: beam 0 on the tilt grid directly + bright = np.full((T, n_pixels, n_pixels), np.nan) + for ti in range(T): + plane = np.full((n_pixels, n_pixels), np.nan) + plane[m] = inten_np[:, ti, 0] + bright[ti] = plane + + # full pattern: every diffracted cone shifted by its g, bilinear deposit + pattern = np.zeros((T, n_pixels, n_pixels)) + tx_in = tilts[:, 0].numpy() + ty_in = tilts[:, 1].numpy() + g_np = g_xy.numpy() + for b in range(g_np.shape[0]): + fx = (tx_in + g_np[b, 0] + alpha_k) / px + fy = (ty_in + g_np[b, 1] + alpha_k) / px + ix0 = np.floor(fx).astype(int) + iy0 = np.floor(fy).astype(int) + wx = fx - ix0 + wy = fy - iy0 + for dx in (0, 1): + for dy in (0, 1): + jx = ix0 + dx + jy = iy0 + dy + ok = (jx >= 0) & (jx < n_pixels) & (jy >= 0) & (jy < n_pixels) + w = (wx if dx else 1 - wx) * (wy if dy else 1 - wy) + # (row, col) = (theta_y, theta_x), as the bright field + for ti in range(T): + np.add.at(pattern[ti], (jy[ok], jx[ok]), (w * inten_np[:, ti, b])[ok]) + pattern[:, ~m] = 0.0 + + return { + "bright_field": bright[0] if T == 1 else bright, + "pattern": pattern[0] if T == 1 else pattern, + "sampling": px, + "mrad_per_pixel": px * lam * 1e3, + "thicknesses": t_thick.numpy(), + "hkl": hkl_beams.numpy(), + } + + +def _lambert_raster( + crystal: Crystal, dirs: torch.Tensor, values: torch.Tensor, step: float +) -> np.ndarray: + """Expand wedge samples by the crystal symmetry (plus inversion) and + splat them bilinearly onto a Lambert equal-area grid of the upper + hemisphere. values is (N, T); returns (T, n, n) with NaN where unhit + (raster holes inside the disk are filled from their neighbors).""" + from quantem.diffraction.rotations import quat_to_matrix + + T = values.shape[1] + Rs = quat_to_matrix(crystal.sym_quats_matching) + d_all = torch.einsum("sij,nj->sni", Rs, dirs).reshape(-1, 3) + I_all = values[None, :, :].expand(Rs.shape[0], -1, -1).reshape(-1, T) + d_all = torch.cat([d_all, -d_all]) + I_all = torch.cat([I_all, I_all]) + up = d_all[:, 2] >= 0 + d_all, I_all = d_all[up], I_all[up] + + # grid always spans the full hemisphere: symmetry expansion moves wedge + # samples to any polar angle, and clipping them onto a smaller rim + # corrupts the equatorial region + rho_max = float(np.sqrt(2.0)) + half = int(np.ceil(rho_max / step)) + n = 2 * half + 1 + rho = torch.sqrt((2 * (1 - d_all[:, 2])).clamp_min(0)) + dxy = torch.linalg.norm(d_all[:, :2], dim=1).clamp_min(1e-12) + px_x = (d_all[:, 0] / dxy * rho / step + half).numpy() + px_y = (d_all[:, 1] / dxy * rho / step + half).numpy() + + acc = np.zeros((T, n, n)) + wgt = np.zeros((n, n)) + ix0 = np.floor(px_x).astype(int) + iy0 = np.floor(px_y).astype(int) + wx = px_x - ix0 + wy = px_y - iy0 + I_np = I_all.numpy() + for dx in (0, 1): + for dy in (0, 1): + jx = np.clip(ix0 + dx, 0, n - 1) + jy = np.clip(iy0 + dy, 0, n - 1) + w = (wx if dx else 1 - wx) * (wy if dy else 1 - wy) + np.add.at(wgt, (jx, jy), w) + for ti in range(T): + np.add.at(acc[ti], (jx, jy), w * I_np[:, ti]) + lambert = np.where(wgt[None] > 1e-6, acc / np.maximum(wgt[None], 1e-6), np.nan) + + # fill raster holes (unhit pixels between splatted samples) from their + # neighbors so bilinear lookups never touch NaN inside the disk + yy, xx = np.mgrid[0:n, 0:n] + in_disk = ((xx - half) ** 2 + (yy - half) ** 2) <= (rho_max / step) ** 2 + for ti in range(T): + L = lambert[ti] + for _ in range(4): + holes = np.isnan(L) & in_disk + if not holes.any(): + break + Lp = np.pad(L, 1, constant_values=np.nan) + stack = np.stack( + [ + Lp[1 + dy : n + 1 + dy, 1 + dx : n + 1 + dx] + for dy in (-1, 0, 1) + for dx in (-1, 0, 1) + ] + ) + with np.errstate(all="ignore"): + fill = np.nanmean(stack, axis=0) + L[holes] = fill[holes] + lambert[ti] = L + return lambert + + +def calculate_kossel_reference( + crystal: Crystal, + thicknesses_A, + energy_ev: float = 300e3, + angle_step_mrad: float = 1.0, + sg_max: float = 0.05, + k_max: float | None = None, + theta_max_deg: float = 90.0, + chunk: int = 256, + fast_absorption: bool = True, + progress_bar: bool = True, +) -> dict: + """Kossel reference pattern: the dynamical bright field over all directions. + + The bright field intensity depends only on the incident beam direction + in the CRYSTAL frame (each incident plane wave is independent), so one + Bloch computation over the symmetry-reduced direction wedge gives the + pattern for every specimen orientation at once (called a master pattern + in parts of the EBSD literature). Patterns for arbitrary orientations, + convergence angles, and all precomputed thicknesses are then + interpolation lookups via kossel_from_reference(), milliseconds instead + of a fresh dynamical calculation. + + The wedge samples are expanded by the crystal's proper rotations plus + inversion and rasterized onto a Lambert azimuthal equal-area grid of + the upper hemisphere. (The inversion expansion assumes Friedel symmetry + of the bright field; for non-centrosymmetric crystals with absorption + this neglects a small polarity contrast.) + + Parameters + ---------- + crystal : Crystal + With structure factors calculated, and preferably + calculate_dynamical_structure_factors at this energy (absorption). + thicknesses_A : float | array-like + Thickness grid in Angstroms; all thicknesses share the + eigendecompositions, so a thickness axis is nearly free. + energy_ev : float, default=300e3 + Beam energy in eV. + angle_step_mrad : float, default=1.0 + Angular sampling of the wedge, and the pixel size of the Lambert + grid (in Lambert radius units of 1e-3). The dynamical line widths + are 1-2 mrad; 0.5 for production references, 1-2 for quick looks. + sg_max : float, default=0.05 + Excitation error cutoff (1/Angstroms) at the center of each chunk + of directions, widened by the chunk's angular radius; as in + calculate_kossel, tighter than SG_MAX because of that widening. + k_max : float | None + Largest |g| (1/Angstroms) of an included beam; None keeps every + reflection of the factor set, which is slow for large sets. The + cost grows steeply with it. + theta_max_deg : float, default=90.0 + Polar cutoff of the wedge samples. Keep at 90 unless the wedge's + far corners are never observed: cutting the wedge leaves coverage + holes at all their symmetry equivalents. + chunk : int, default=256 + Directions per batch; each batch shares one beam list and one + coupling matrix. + fast_absorption : bool, default=True + First-order absorption (Hermitian eigensolver, several times + faster); see _bloch_solve. + progress_bar : bool, default=True + Show a progress bar over chunks. + + Returns + ------- + dict + 'lambert' (T, n, n) bright field on the equal-area grid of the + upper hemisphere (NaN where unsampled), 'rho_max' (Lambert radius + of the equator, sqrt(2)), 'step' (Lambert grid spacing), + 'thicknesses', 'energy_ev', 'k_max' (as given, possibly None), and + the raw wedge samples 'directions' (N, 3) and 'intensity' (N, T). + """ + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + gamma_rel = relativistic_gamma(energy_ev) + t_thick = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + T = t_thick.shape[0] + + msg = crystal.matching_symmetry_warning() + if msg is not None: + warnings.warn(msg, stacklevel=2) + wedge = crystal.zone_axis_wedge() + step_deg = np.rad2deg(angle_step_mrad * 1e-3) + if wedge is None: + from quantem.diffraction.orientation import fibonacci_hemisphere + + n_dirs = int(np.ceil(2 * np.pi / np.deg2rad(step_deg) ** 2)) + dirs = fibonacci_hemisphere(n_dirs) + else: + dirs, _ = sample_zone_axes(wedge, step_deg) + keep = dirs[:, 2] >= np.cos(np.deg2rad(theta_max_deg)) + dirs = dirs[keep] + N = dirs.shape[0] + + hkl_u, g = _beam_universe(crystal) # crystal frame, orientation is identity + g2 = (g**2).sum(dim=1) + g_len = torch.linalg.norm(g, dim=1) + + out = torch.zeros((N, T), dtype=torch.float64) + chunks = range(0, N, chunk) + if progress_bar: + chunks = tqdm(chunks, desc="Kossel reference") + for c0 in chunks: + c1 = min(c0 + chunk, N) + d = dirs[c0:c1] # (B, 3) beam directions in the crystal frame + d_c = d.mean(dim=0) + d_c = d_c / torch.linalg.norm(d_c) + radius = float(torch.arccos((d @ d_c).clamp(-1, 1)).max()) + + # normal-tracking geometry: each sample is computed with the foil + # normal along the sampled direction (the crystal is conceptually + # re-tilted per sample). The slab problem is then a function of the + # crystal-frame beam direction ALONE, which is what makes the + # symmetry expansion below exact; a pattern lookup only ever probes + # directions within the convergence semiangle of the true normal, + # so the approximation error is O(alpha^2). + u_c = g @ d_c + s_c = (2 * k0 * u_c - g2) / (2 * (k0 - u_c)) + sel = torch.abs(s_c) < sg_max + (radius + 1e-4) * g_len + if k_max is not None: + sel &= g_len <= k_max + hkl_beams = torch.cat([torch.zeros((1, 3), dtype=torch.long), hkl_u[sel]]) + g_b = torch.cat([torch.zeros((1, 3), dtype=torch.float64), g[sel]]) + U, u0_imag, absorptive = _coupling_matrix(crystal, hkl_beams, gamma_rel) + + g2b = (g_b**2).sum(dim=1) + u = torch.einsum("bk,nk->bn", d, g_b) # g . d_hat per sample + s_t = (2 * k0 * u - g2b[None, :]) / (2 * (k0 - u)) + inten_b = _bloch_solve( + U, u0_imag, absorptive, s_t, k0, t_thick, fast_absorption=fast_absorption + ) + out[c0:c1] = inten_b[:, :, 0] + + # symmetry expansion and Lambert raster: with normal-tracking geometry + # the intensity is a function of the crystal-frame beam direction only, + # so proper rotations apply directly; the reversed beam (with reversed + # normal) gives the same bright field by reciprocity. + step = angle_step_mrad * 1e-3 + lambert = _lambert_raster(crystal, dirs, out, step) + rho_max = float(np.sqrt(2.0)) + + return { + "lambert": lambert, + "rho_max": rho_max, + "step": step, + "thicknesses": t_thick.numpy(), + "energy_ev": float(energy_ev), + "k_max": k_max, + "directions": dirs.numpy(), + "intensity": out.numpy(), + } + + +def _lambert_lookup(lambert: np.ndarray, step: float, d_c: torch.Tensor) -> np.ndarray: + """Bilinear lookup of a (T, n, n) Lambert grid at crystal-frame + directions d_c (..., 3); returns (T, ...). Directions are folded to + the upper hemisphere (reciprocity).""" + d_c = torch.where(d_c[..., 2:3] < 0, -d_c, d_c) + half = (lambert.shape[-1] - 1) // 2 + rho = torch.sqrt((2 * (1 - d_c[..., 2])).clamp_min(0)) + dxy = torch.linalg.norm(d_c[..., :2], dim=-1).clamp_min(1e-12) + fx = (d_c[..., 0] / dxy * rho / step + half).numpy() + fy = (d_c[..., 1] / dxy * rho / step + half).numpy() + n_l = lambert.shape[-1] + ix0 = np.clip(np.floor(fx).astype(int), 0, n_l - 2) + iy0 = np.clip(np.floor(fy).astype(int), 0, n_l - 2) + wx = np.clip(fx - ix0, 0, 1) + wy = np.clip(fy - iy0, 0, 1) + out = np.zeros((lambert.shape[0],) + fx.shape) + for ti in range(lambert.shape[0]): + L = lambert[ti] + out[ti] = ( + L[ix0, iy0] * (1 - wx) * (1 - wy) + + L[ix0 + 1, iy0] * wx * (1 - wy) + + L[ix0, iy0 + 1] * (1 - wx) * wy + + L[ix0 + 1, iy0 + 1] * wx * wy + ) + return out + + +def kossel_from_reference( + reference: dict, + orientation: torch.Tensor, + semiconv_mrad: float = 40.0, + n_pixels: int = 192, +) -> dict: + """Extract a bright field Kossel pattern from a reference pattern. + + Interpolation only, no Bloch calculation. The detector tilt grid is + mapped into the crystal frame by the orientation and looked up + bilinearly on the reference's Lambert grid, so the pattern has the + reference's angular resolution (angle_step_mrad), whatever n_pixels. + + Parameters + ---------- + reference : dict + From calculate_kossel_reference(). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + semiconv_mrad : float, default=40.0 + Convergence semiangle in mrad; the pattern covers this radius. + n_pixels : int, default=192 + Pixels across the pattern. + + Returns + ------- + dict + 'bright_field' ((T, n, n), squeezed for one thickness; rows are + theta_y, columns theta_x; NaN outside the aperture), + 'mrad_per_pixel', 'thicknesses'. + """ + lam = electron_wavelength_angstrom(reference["energy_ev"]) + d_c, inside, _, _ = _detector_directions( + lam, orientation, semiconv_mrad, False, n_pixels, 1, 1 + ) + bf = _lambert_lookup(reference["lambert"], reference["step"], d_c) + bf[:, ~inside.numpy()] = np.nan + T = bf.shape[0] + + return { + "bright_field": bf[0] if T == 1 else bf, + "mrad_per_pixel": 2 * semiconv_mrad / (n_pixels - 1), + "thicknesses": reference["thicknesses"], + } + + +def plot_kossel_reference( + reference: dict, + crystal: Crystal, + thickness_index: int = 0, + max_index: int = 2, + theta_max_label_deg: float = 75.0, + min_crossing: float = 1.0, + min_crossing_rim: float = 0.3, + sigma_mrad: float = 10.0, + lines: dict | None = None, + theta_circles=(), + label_color=(0.9, 0.0, 0.0), + label_fontsize: float = 12, + stroke_color=(1.0, 1.0, 1.0, 0.7), + stroke_width: float = 6.0, + upsample: int = 2, + cmap: str = "gray", + axsize: tuple[float, float] = (9.0, 9.0), + filename: str | None = None, + figax=None, +): + """The reference pattern with polar angle circles and low index zone labels. + + Every symmetry copy of the zone axes with direction indices up to + `max_index` is labeled with its own signed indices (4-index for + hexagonal and trigonal crystals). The circles and labels are vector + graphics; saving to a PDF via `filename` keeps them sharp at any zoom, + with the pattern embedded as a smoothly interpolated image. + + Parameters + ---------- + reference : dict + From calculate_kossel_reference(). + crystal : Crystal + The crystal the reference was computed for. + thickness_index : int, default=0 + Which thickness of the reference to show. + max_index : int, default=2 + Largest direction index to label. + theta_max_label_deg : float, default=75.0 + Zones between this polar angle and the equator are left unlabeled + (rim clutter); the equatorial zones themselves are labeled just + outside the disk edge. + min_crossing : float, default=1.0 + Only label a zone whose crossing strength reaches this value. Each + Kossel band is a pair of lines at +-theta_B about the zone plane, + so the rows in a zone (zone law g . [uvw] = 0) form a rosette + around the zone axis rather than lines through it. The crossing + strength sums, over the rows in the zone, the line depth weighted + by exp(-theta_B^2 / 2 sigma^2), and subtracts the strongest row: + a zone on a single band scores zero (such as <221> or <223> in + diamond, which contain only the 220 row), and a rosette of several + strong rows with small Bragg angles scores high. In silicon at + 200 kV the default keeps <001>, <011>, <111>, <112> and <013>, + and drops <113> (0.6, its 422 and 620 rows sit 11-15 mrad out), + <123> (0.65) and <233> (0.9). + min_crossing_rim : float, default=0.3 + The same threshold for the equatorial zones labeled outside the + disk, where there is room for weaker crossings: keeps <120> and + <130> in silicon and drops <230>, which is a single 400 band. + sigma_mrad : float, default=10.0 + Rosette scale of the crossing strength: rows with Bragg angles + beyond this contribute little, since their band edges are too far + from the zone axis to read as a crossing. + lines : dict | None + Line set from kossel_lines() for the crossing strength, on the + reference's thickness grid; computed from the crystal, the chosen + thickness and the reference's k_max (1.2 1/A when it has none) if + omitted. + theta_circles : sequence, default=() + Polar angles (degrees) at which to draw dashed circles; off by + default. + label_color, label_fontsize, stroke_color, stroke_width : + Zone label styling: text color, size, and the translucent outline + drawn behind each label. + upsample : int, default=2 + Bilinear upsampling factor of the displayed pattern. + cmap : str, default="gray" + Colormap of the pattern. + axsize : tuple[float, float], default=(9.0, 9.0) + Figure size in inches when a new figure is made. + filename : str | None + If given, save the figure (PDF recommended). + figax : tuple | None + (fig, ax) to draw into; a new figure if None. + + Returns + ------- + fig, ax + The matplotlib figure and axes. + """ + import matplotlib.pyplot as plt + from matplotlib import patheffects + + from quantem.diffraction.crystal import miller_to_miller_bravais + from quantem.diffraction.rotations import quat_to_matrix + + L = reference["lambert"][thickness_index] + step = reference["step"] + half = (L.shape[-1] - 1) // 2 + if upsample > 1: + from scipy.ndimage import zoom + + L = zoom(np.nan_to_num(L, nan=np.nanmax(L)), upsample, order=1) + scale = upsample if upsample > 1 else 1 + + if figax is None: + fig, ax = plt.subplots(figsize=axsize) + else: + fig, ax = figax + ax.imshow(L, cmap=cmap, interpolation="bilinear") + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_frame_on(False) + + def to_px(v): + return (v / step + half) * scale + (scale - 1) / 2 + + phi = np.linspace(0, 2 * np.pi, 721) + for theta_deg in theta_circles: + r = 2 * np.sin(np.deg2rad(theta_deg) / 2) / step * scale + c = to_px(0.0) + ax.plot(c + r * np.cos(phi), c + r * np.sin(phi), ls="--", color="0.45", lw=0.7) + ax.text( + c, + c - r, + " %d°" % theta_deg, + color="0.35", + fontsize=9, + va="bottom", + ) + + # unique low index zone directions, expanded over the crystal symmetry + hexagonal = crystal.hexagonal_matching + A_T = crystal.lat_real.numpy().T # d_cartesian = A_T @ [u, v, w] + A_T_inv = np.linalg.inv(A_T) + Rs = quat_to_matrix(crystal.sym_quats_matching).numpy() + + # crossing strength from the line set: per row, the deepest line + # weighted by its Bragg angle + if lines is None: + lines = kossel_lines( + crystal, + reference["thicknesses"][thickness_index], + energy_ev=reference["energy_ev"], + # a reference computed without a cutoff stores k_max=None; the + # line set needs a finite one + k_max=reference.get("k_max") or 1.2, + ) + ti = 0 + else: + ti = thickness_index + g_hat = lines["g_hat"].numpy() + n_rows = g_hat.shape[0] + row_of = lines["line_row"].numpy() + line_w = lines["line_depth"].numpy()[:, ti] * np.exp( + -0.5 * (lines["line_u"].numpy() / (sigma_mrad * 1e-3)) ** 2 + ) + row_weight = np.zeros(n_rows) + np.maximum.at(row_weight, row_of, line_w) + + def crossing_strength(dc): + w = row_weight[np.abs(g_hat @ dc) < 1e-4] + return float(w.sum() - w.max()) if w.size else 0.0 + + # one label per distinct crystallographic direction, keyed by its + # canonical index tuple so a direction reached by several symmetry + # operations (common at the equatorial rim) is drawn only once + # integer index-space representation of each symmetry rotation, so the + # zone index of a symmetry copy is computed by exact integer arithmetic + # rather than by rounding a projected direction (which can alias a + # high-index direction onto a low-index label) + M = [np.rint(A_T_inv @ R @ A_T).astype(int) for R in Rs] + + placed: dict[tuple, tuple] = {} + rng = range(-max_index, max_index + 1) + for u in rng: + for v in rng: + for w in rng: + uvw = np.array([u, v, w]) + if not uvw.any() or np.gcd.reduce(np.abs(uvw)) != 1: + continue + d = A_T @ uvw + d = d / np.linalg.norm(d) + for R, Mi in zip(Rs, M): + for sgn in (1, -1): + dc = sgn * (R @ d) + idx = sgn * (Mi @ uvw) + # fold to the upper hemisphere (the reference is + # stored there); flip the index to match + if dc[2] < 0: + dc = -dc + idx = -idx + rim = dc[2] < np.sin(np.deg2rad(1.0)) + if not rim and dc[2] < np.cos(np.deg2rad(theta_max_label_deg)): + continue + if crossing_strength(dc) < (min_crossing_rim if rim else min_crossing): + continue + key = tuple(int(k) for k in idx) + if key in placed: + continue + ks = np.atleast_2d(miller_to_miller_bravais(idx))[0] if hexagonal else idx + txt = ( + "$[" + + "".join((r"\bar{%d\!}" % abs(k)) if k < 0 else str(k) for k in ks) + + "]$" + ) + rho = np.sqrt(max(2 * (1 - dc[2]), 0)) + if rim: + rho = np.sqrt(2.0) * 1.07 # just outside the disk + dxy = max(np.hypot(dc[0], dc[1]), 1e-12) + px = to_px(dc[0] / dxy * rho) + py = to_px(dc[1] / dxy * rho) + placed[key] = (px, py, txt) + + for px, py, txt in placed.values(): + t = ax.text( + py, + px, + txt, + color=label_color, + fontsize=label_fontsize, + ha="center", + va="center", + ) + t.set_path_effects( + [patheffects.withStroke(linewidth=stroke_width, foreground=stroke_color)] + ) + n_px = L.shape[-1] + ax.set_xlim(-0.06 * n_px, 1.06 * n_px) + ax.set_ylim(1.06 * n_px, -0.06 * n_px) + if filename is not None: + fig.savefig(filename, bbox_inches="tight", dpi=300) + return fig, ax + + +def kossel_polar_from_reference( + reference: dict, + orientation: torch.Tensor, + semiconv_mrad: float = 40.0, + n_radial: int = 64, + n_azimuthal: int = 180, +) -> dict: + """A bright field Kossel pattern sampled directly on a polar grid. + + Dictionary matching correlates over the in-plane rotation, which is a + cyclic shift of the azimuthal axis in polar coordinates: sampling the + reference directly at the polar detector positions avoids the intermediate + Cartesian raster and its interpolation. + + Parameters + ---------- + reference : dict + From calculate_kossel_reference(). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + semiconv_mrad : float, default=40.0 + Outer radius of the polar grid in mrad. + n_radial : int, default=64 + Radial samples, at radii semiconv_mrad * (1 ... n_radial) / n_radial. + n_azimuthal : int, default=180 + Azimuthal samples, at 2 pi (0 ... n_azimuthal - 1) / n_azimuthal. + + Returns + ------- + dict + 'polar' ((T, n_azimuthal, n_radial), squeezed for one thickness; + rows are azimuth, columns radius, matching the quantem polar + transform convention), 'radii_mrad', 'azimuth_rad', 'thicknesses'. + """ + lam = electron_wavelength_angstrom(reference["energy_ev"]) + d_c, _, axes, _ = _detector_directions( + lam, orientation, semiconv_mrad, True, 1, n_radial, n_azimuthal + ) + out = _lambert_lookup(reference["lambert"], reference["step"], d_c) + T = out.shape[0] + return { + "polar": out[0] if T == 1 else out, + "radii_mrad": axes["radii_mrad"], + "azimuth_rad": axes["azimuth_rad"], + "thicknesses": reference["thicknesses"], + } + + +def kossel_lines( + crystal: Crystal, + thicknesses_A, + energy_ev: float = 300e3, + k_max: float = 1.2, + u_step_mrad: float = 0.05, + u_tail_mrad: float = 150.0, + min_depth: float = 0.005, + fast_absorption: bool = False, +) -> dict: + """Vector representation of the Kossel lines: one profile per systematic row. + + The bright field depends on the beam direction d (a unit vector, the + anti-propagation direction in the crystal frame) only through the + projections u = d . g_hat onto the row normals. For each systematic row + {n g} the profile is a Bloch calculation over the row beams alone versus + the signed projection u, which places the deficiency line of reflection + +n g at u = +n lambda |g| / 2 and that of -n g at u = -n lambda |g| / 2: + the two lines of a Kossel band, 2 theta_B apart, and their higher + orders, all with their dynamical widths and thickness fringes. Rows + combine multiplicatively as independent attenuation channels; the + many-beam coupling between different rows at the zone axis crossings is + the one approximation. + + A row profile is a smooth function of a continuous variable, so patterns + rendered from the line set (render_kossel_lines) are exact in geometry + and free of raster interpolation at any pixel size, in Cartesian or + polar coordinates. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated, and preferably + calculate_dynamical_structure_factors at this energy (absorption). + thicknesses_A : float | array-like + Thickness grid in Angstroms. + energy_ev : float, default=300e3 + Beam energy in eV. + k_max : float, default=1.2 + Reflections with |g| up to this (1/Angstroms) are included; a row + keeps every order |n| |g| <= k_max. Must be a number. + u_step_mrad : float, default=0.05 + Profile sampling; the line widths are 1-2 mrad. + u_tail_mrad : float, default=150.0 + Profile extent beyond the outermost line of each row. The rocking + curve tails fall off as 1 / (s xi)^2 and are still ~1% at 40 mrad + for the strong reflections, so the window has to be wide for the + far-from-line background to be the true mean-absorption level. + min_depth : float, default=0.005 + Lines (and rows) whose deepest deficit at any thickness is below + this fraction of the background are dropped. + fast_absorption : bool, default=False + First-order absorption in the row calculations; see _bloch_solve. + + Returns + ------- + dict with, per row, 'g_hat' (L, 3) crystal-frame unit normals, 'g_len' + (L,), 'hkl_row' (L, 3), 'log_trans' (L, T, n_u) log transmission + versus 'u' (n_u,) (0 far from the lines), 'background' (T,) the + far-from-line bright field; and per line 'line_row' (K,) row index, + 'line_order' (K,) the order n, 'line_hkl' (K, 3), 'line_u' (K,) the + cone position u = n lambda |g| / 2, 'line_depth' (K, T) the deepest + deficit fraction and 'line_width_mrad' (K, T) the equivalent width + (integrated deficit over depth); also 'energy_ev' and 'thicknesses'. + + Raises + ------ + ValueError + If k_max is None, or no line reaches min_depth. + """ + if crystal.g_vec is None: + raise RuntimeError("Run crystal.calculate_structure_factors() first.") + if k_max is None: + raise ValueError( + "kossel_lines needs a numeric k_max: every reflection up to it gets a " + "row profile, so the full factor set would be very slow" + ) + k_max = float(k_max) + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + gamma_rel = relativistic_gamma(energy_ev) + t_thick = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + + # unique rows: group reflections by ray direction (g and -g together), + # keep the shortest g of each as the row vector + hkl, g_all = _beam_universe(crystal) + g_len = torch.linalg.norm(g_all, dim=1) + sel = (g_len <= k_max) & (g_len > 1e-8) + idx = torch.nonzero(sel).squeeze(1) + idx = idx[torch.argsort(g_len[idx])] + rows: list[int] = [] + dirs: list[torch.Tensor] = [] + for i in idx.tolist(): + d = g_all[i] / g_len[i] + if any(float(torch.abs(d @ e)) > 0.9999 for e in dirs): + continue + rows.append(i) + dirs.append(d) + + u_max = 0.5 * lam * k_max + u_tail_mrad * 1e-3 + du = u_step_mrad * 1e-3 + n_half = int(np.ceil(u_max / du)) + u = torch.arange(-n_half, n_half + 1, dtype=torch.float64) * du + + g_hat_out, g_len_out, hkl_out, lt_out, bg_rows = [], [], [], [], [] + l_row, l_order, l_hkl, l_u, l_depth, l_width = [], [], [], [], [], [] + for i in rows: + g1 = float(g_len[i]) + h1 = hkl[i] + n_ord = int(np.floor(k_max / g1 + 1e-9)) + ns = torch.arange(-n_ord, n_ord + 1, dtype=torch.long) + ns = ns[torch.argsort((ns != 0).to(torch.long), stable=True)] # 000 first + hkl_beams = ns[:, None] * h1[None, :] + U, u0_imag, absorptive = _coupling_matrix(crystal, hkl_beams, gamma_rel) + # projection of each row beam on the beam direction: n |g| u, and + # the same excitation error geometry as the reference pattern + # (foil normal along the beam) + ng = ns.to(torch.float64) * g1 + uu = u[:, None] * ng[None, :] + s_t = (2 * k0 * uu - ng[None, :] ** 2) / (2 * (k0 - uu)) + inten_b = _bloch_solve( + U, u0_imag, absorptive, s_t, k0, t_thick, fast_absorption=fast_absorption + ) + bf = inten_b[:, :, 0].transpose(0, 1) # (T, n_u) + bg = 0.5 * (bf[:, 0] + bf[:, -1]) # (T,) far-from-line level + trans = (bf / bg[:, None]).clamp_min(1e-6) + deficit = 1 - trans + + # per-line depth and width, each order in its own window of half + # the order spacing on either side of its cone + lines_here = [] + for n in range(-n_ord, n_ord + 1): + if n == 0: + continue + u_n = n * lam * g1 / 2 + win = torch.abs(u - u_n) <= lam * g1 / 4 + dep = deficit[:, win].amax(dim=1) # (T,) + if float(dep.max()) < min_depth: + continue + width = deficit[:, win].clamp_min(0).sum(dim=1) * du / dep.clamp_min(1e-9) + lines_here.append((n, u_n, dep, width * 1e3)) + if not lines_here: + continue + row_id = len(g_hat_out) + g_hat_out.append(g_all[i] / g_len[i]) + g_len_out.append(g1) + hkl_out.append(h1) + lt_out.append(torch.log(trans)) + bg_rows.append(bg) + for n, u_n, dep, width in lines_here: + l_row.append(row_id) + l_order.append(n) + l_hkl.append(n * h1) + l_u.append(u_n) + l_depth.append(dep) + l_width.append(width) + if not g_hat_out: + raise ValueError( + f"no Kossel line reaches min_depth={min_depth} with k_max={k_max} 1/A: " + "raise k_max or lower min_depth" + ) + + return { + "g_hat": torch.stack(g_hat_out), + "g_len": torch.tensor(g_len_out, dtype=torch.float64), + "hkl_row": torch.stack(hkl_out), + "u": u, + "log_trans": torch.stack(lt_out), # (L, T, n_u) + "background": torch.stack(bg_rows).mean(dim=0), # (T,) + "line_row": torch.tensor(l_row, dtype=torch.long), + "line_order": torch.tensor(l_order, dtype=torch.long), + "line_hkl": torch.stack(l_hkl), + "line_u": torch.tensor(l_u, dtype=torch.float64), + "line_depth": torch.stack(l_depth), # (K, T) + "line_width_mrad": torch.stack(l_width), # (K, T) + "energy_ev": float(energy_ev), + "thicknesses": t_thick.numpy(), + } + + +def _detector_directions( + lam: float, + orientation: torch.Tensor, + semiconv_mrad: float, + polar: bool, + n_pixels: int, + n_radial: int, + n_azimuthal: int, +): + """Crystal-frame anti-propagation directions of a Cartesian or polar + detector grid, plus the grid axes. Polar grids follow the quantem + convention: rows are azimuth, columns radius.""" + from quantem.diffraction.rotations import quat_to_matrix + + k0 = 1.0 / lam + alpha_k = semiconv_mrad * 1e-3 / lam + if polar: + r = torch.linspace(0, alpha_k, n_radial + 1, dtype=torch.float64)[1:] + phi = torch.arange(n_azimuthal, dtype=torch.float64) * (2 * np.pi / n_azimuthal) + tx = r[None, :] * torch.cos(phi)[:, None] + ty = r[None, :] * torch.sin(phi)[:, None] + inside = torch.ones_like(tx, dtype=torch.bool) + axes = {"radii_mrad": (r * lam * 1e3).numpy(), "azimuth_rad": phi.numpy()} + else: + ax = torch.linspace(-alpha_k, alpha_k, n_pixels, dtype=torch.float64) + ty, tx = torch.meshgrid(ax, ax, indexing="ij") + inside = (tx**2 + ty**2) <= alpha_k**2 + axes = {"mrad_per_pixel": 2 * semiconv_mrad / (n_pixels - 1)} + tz = torch.sqrt((k0**2 - tx**2 - ty**2).clamp_min(0)) + # the beam landing at detector tilt +t propagates along (t, -tz); the + # line set and the reference parameterize the anti-propagation direction + d_lab = torch.stack([-tx, -ty, tz], dim=-1) / k0 + R = quat_to_matrix(torch.atleast_2d(torch.as_tensor(orientation, dtype=torch.float64))[0]).to( + torch.float64 + ) + d_c = torch.einsum("ji,rcj->rci", R, d_lab) # crystal frame, R^T d + return d_c, inside, axes, R + + +def _lines_bright_field(lines: dict, d_c: torch.Tensor) -> torch.Tensor: + """Line-model bright field at crystal-frame directions d_c (..., 3): + product of the row transmissions read at u = d . g_hat; returns + (..., T).""" + u = torch.einsum("...i,li->...l", d_c, lines["g_hat"]) # (.., L) + u_ax = lines["u"] + n_u = u_ax.shape[0] + du = float(u_ax[1] - u_ax[0]) + lt = lines["log_trans"].permute(0, 2, 1) # (L, n_u, T) + f = ((u - float(u_ax[0])) / du).clamp(0, n_u - 1 - 1e-9) + i0 = f.floor().to(torch.long) + w = (f - i0)[..., None] + L_idx = torch.arange(lt.shape[0]).reshape((1,) * (u.dim() - 1) + (-1,)) + v = lt[L_idx, i0] * (1 - w) + lt[L_idx, (i0 + 1).clamp(max=n_u - 1)] * w + return lines["background"] * torch.exp(v.sum(dim=-2)) + + +def kossel_reference_residual(reference: dict, lines: dict, crystal: Crystal) -> dict: + """Add the many-beam residual of the line model to a reference pattern. + + The line model is evaluated at the reference's own wedge samples and + rasterized onto the same Lambert grid, and the difference (reference + minus line model) is stored as reference['residual']. It is zero away + from the zone axes, where the rows are independent, and carries the + many-beam correction of the zone axis rosettes. render_kossel_lines() + adds it by lookup when given the reference. + + Parameters + ---------- + reference : dict + From calculate_kossel_reference(); modified in place. + lines : dict + From kossel_lines(), on the same thickness grid and energy. + crystal : Crystal + The crystal both were computed for. + + Returns + ------- + dict + The reference, with 'residual' (T, n, n) added. + + Raises + ------ + ValueError + If the thickness grids differ. + """ + if not np.allclose(reference["thicknesses"], lines["thicknesses"]): + raise ValueError("reference and line set must share the thickness grid") + dirs = torch.as_tensor(reference["directions"], dtype=torch.float64) + I_lines = _lines_bright_field(lines, dirs) # (N, T) + lambert_lines = _lambert_raster(crystal, dirs, I_lines, reference["step"]) + reference["residual"] = np.nan_to_num(reference["lambert"] - lambert_lines, nan=0.0) + return reference + + +def render_kossel_lines( + lines: dict, + orientation: torch.Tensor, + semiconv_mrad: float = 40.0, + n_pixels: int = 256, + polar: bool = False, + n_radial: int = 64, + n_azimuthal: int = 180, + reference: dict | None = None, +) -> dict: + """Render the bright field from the Kossel line set, all thicknesses. + + Every detector direction is projected on every row normal and the row + profiles are read there: one evaluation per pixel and row, no raster + in between, so the result is smooth at any resolution in Cartesian or + polar coordinates. A polar pattern is sampled directly at the polar + detector positions. + + Parameters + ---------- + lines : dict + From kossel_lines(). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + semiconv_mrad : float, default=40.0 + Convergence semiangle in mrad; the pattern covers this radius. + n_pixels : int, default=256 + Pixels across a Cartesian pattern. + polar : bool, default=False + Sample on a polar grid instead (see kossel_polar_from_reference + for the grid). + n_radial, n_azimuthal : int, default=64, 180 + Polar grid size. + reference : dict | None + A reference pattern carrying the many-beam residual from + kossel_reference_residual(). If given, the residual is added to + the rendered pattern: the line model then also carries the + many-beam intensity of the zone axis rosettes (which the + independent-row product gets too dark), while the lines themselves + keep their exact analytic geometry. + + Returns + ------- + dict with 'bright_field' ((T, n, n), squeezed; rows are theta_y, + columns theta_x; NaN outside the aperture) or, with polar=True, 'polar' ((T, n_azimuthal, n_radial), + squeezed; rows are azimuth, columns radius), plus the grid axes and + 'thicknesses'. + """ + lam = electron_wavelength_angstrom(lines["energy_ev"]) + d_c, inside, axes, _ = _detector_directions( + lam, orientation, semiconv_mrad, polar, n_pixels, n_radial, n_azimuthal + ) + bf = _lines_bright_field(lines, d_c).permute(2, 0, 1).numpy() # (T, ..) + if reference is not None: + if "residual" not in reference: + raise ValueError( + "reference has no many-beam residual: run " + "kossel_reference_residual(reference, lines, crystal) first." + ) + bf = bf + _lambert_lookup(reference["residual"], reference["step"], d_c) + bf[:, ~inside.numpy()] = np.nan + T = bf.shape[0] + + out = {"polar" if polar else "bright_field": bf[0] if T == 1 else bf} + out.update(axes) + out["thicknesses"] = lines["thicknesses"] + return out + + +def kossel_line_segments( + lines: dict, + orientation: torch.Tensor, + semiconv_mrad: float = 40.0, + thickness_index: int = 0, +) -> dict: + """The Kossel lines crossing the aperture as vector segments. + + Each line is the cone d . g_hat = u of its reflection, which within + the aperture is a straight line in the detector tilt plane (the + curvature term is |g_z| alpha^2 / 2, below 0.1 mrad at 40 mrad). The + end points on the aperture edge are computed exactly from the cone. + + Parameters + ---------- + lines : dict + From kossel_lines(). + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + semiconv_mrad : float, default=40.0 + Aperture radius in mrad. + thickness_index : int, default=0 + Thickness of the line set for 'depth' and 'width_mrad'. + + Returns + ------- + dict of arrays over the K visible lines. Cartesian positions are + (row, col) tilt angles in mrad, matching the image axes of + render_kossel_lines: 'start_mrad', 'stop_mrad' (K, 2) the end points + on the aperture edge; 'normal' (K, 2) the unit normal of the line and + 'distance_mrad' (K,) its signed distance from the optic axis, so the + line is the set of points with p . normal = distance. Polar positions + are (azimuth_rad, radius_mrad): 'start_polar', 'stop_polar' (K, 2), the + end points at radius = semiconv_mrad; in between the line follows + radius = distance / cos(azimuth - azimuth_normal). Also 'hkl' (K, 3), + 'depth' (K,) the deficit fraction and 'width_mrad' (K,) the equivalent + width at the chosen thickness. + """ + from quantem.diffraction.rotations import quat_to_matrix + + alpha = semiconv_mrad * 1e-3 + R = ( + quat_to_matrix(torch.atleast_2d(torch.as_tensor(orientation, dtype=torch.float64))[0]) + .to(torch.float64) + .numpy() + ) + g_lab = (R @ lines["g_hat"].numpy().T).T # d_lab . g_lab = d_c . g_c + rows = lines["line_row"].numpy() + u_k = lines["line_u"].numpy() + g = g_lab[rows] # (K, 3) + # cone in tilt angles theta = (theta_x, theta_y), d_lab = (-theta, sqrt(1 - theta^2)): + # -g_x theta_x - g_y theta_y + g_z sqrt(1 - theta^2) = u + gxy = np.hypot(g[:, 0], g[:, 1]) + ok = gxy > 1e-9 + phi_g = np.arctan2(g[:, 1], g[:, 0]) + cz = np.sqrt(1 - alpha**2) + # on the aperture edge theta = alpha (cos phi, sin phi): + # cos(phi - phi_g) = (g_z cz - u) / (|g_xy| alpha) + c = np.where(ok, (g[:, 2] * cz - u_k) / np.maximum(gxy * alpha, 1e-12), 2.0) + ok &= np.abs(c) < 1 + dphi = np.arccos(np.clip(c[ok], -1, 1)) + phi_a = phi_g[ok] + dphi + phi_b = phi_g[ok] - dphi + # small-angle line: (g_x, g_y) . theta = g_z - u + p = (g[ok, 2] - u_k[ok]) / gxy[ok] # signed distance (rad) along -normal + normal = np.stack([g[ok, 1], g[ok, 0]], axis=1) / gxy[ok, None] # (row, col) + + def pt(phi): + # (row, col) = (theta_y, theta_x) in mrad + return np.stack([alpha * np.sin(phi), alpha * np.cos(phi)], axis=1) * 1e3 + + ti = thickness_index + return { + "hkl": lines["line_hkl"].numpy()[ok], + "start_mrad": pt(phi_a), + "stop_mrad": pt(phi_b), + "start_polar": np.stack( + [np.mod(phi_a, 2 * np.pi), np.full(phi_a.shape, semiconv_mrad)], axis=1 + ), + "stop_polar": np.stack( + [np.mod(phi_b, 2 * np.pi), np.full(phi_b.shape, semiconv_mrad)], axis=1 + ), + "normal": normal, + "distance_mrad": p * 1e3, + "depth": lines["line_depth"].numpy()[ok, ti], + "width_mrad": lines["line_width_mrad"].numpy()[ok, ti], + } + + +def overlay_kossel_segments( + ax, + segments: dict, + semiconv_mrad: float, + n_pixels: int | None = None, + polar: bool = False, + n_radial: int | None = None, + n_azimuthal: int | None = None, + color=(0.9, 0.0, 0.0), + width_scale: float = 1.0, + min_depth: float = 0.05, +): + """Draw the vector line segments over a rendered pattern. + + Line width is the equivalent width of each line in pixels (times + width_scale) and the opacity is its depth. On a Cartesian axis the + segments run between their aperture-edge end points; on a polar axis + (rows azimuth, columns radius) each straight line becomes the curve + radius = distance / cos(azimuth - azimuth_normal), drawn from end + point to end point and split at the azimuth wrap. The pixel registration + is that of render_kossel_lines and the reference lookups: Cartesian + pixel i at angle -semiconv + 2 semiconv i / (n_pixels - 1), polar + column j at radius semiconv (j + 1) / n_radial and row i at azimuth + 2 pi i / n_azimuthal. + + Parameters + ---------- + ax : matplotlib.axes.Axes + Axes showing the rendered pattern (imshow pixel coordinates). + segments : dict + From kossel_line_segments(). + semiconv_mrad : float + Aperture radius of the pattern in mrad. + n_pixels : int | None + Pixels across a Cartesian pattern. + polar : bool, default=False + Draw on a polar pattern instead. + n_radial, n_azimuthal : int | None + Polar grid size. + color : color, default=(0.9, 0.0, 0.0) + Line color. + width_scale : float, default=1.0 + Multiplier of the drawn line width. + min_depth : float, default=0.05 + Lines shallower than this deficit fraction are not drawn. + + Returns + ------- + matplotlib.axes.Axes + The axes. + """ + sel = segments["depth"] >= min_depth + n_lines = int(sel.sum()) + p = segments["distance_mrad"][sel] + nrm = segments["normal"][sel] + start = segments["start_mrad"][sel] + stop = segments["stop_mrad"][sel] + dep = segments["depth"][sel] + wid = segments["width_mrad"][sel] + if polar: + # radius r_j = semiconv (j + 1) / n_radial, azimuth phi_i = 2 pi i / n_az + px_r = n_radial / semiconv_mrad + px_phi = n_azimuthal / (2 * np.pi) + tang = np.stack([-nrm[:, 1], nrm[:, 0]], axis=1) + t_edge = np.sqrt(np.maximum(semiconv_mrad**2 - p**2, 0)) + t = np.linspace(-1, 1, 400) + for k in range(n_lines): + pts = p[k] * nrm[k][None, :] + (t * t_edge[k])[:, None] * tang[k][None, :] + r = np.hypot(pts[:, 0], pts[:, 1]) + phi = np.mod(np.arctan2(pts[:, 0], pts[:, 1]), 2 * np.pi) + jumps = np.abs(np.diff(phi)) > np.pi + phi = np.ma.array(phi, mask=np.r_[False, jumps]) + ax.plot( + r * px_r - 1.0, + phi * px_phi, + color=color, + lw=wid[k] * px_r * width_scale, + alpha=float(dep[k]), + solid_capstyle="butt", + ) + else: + # linspace(-semiconv, semiconv, n_pixels): pixel centers at the ends + px = (n_pixels - 1) / (2 * semiconv_mrad) + for k in range(n_lines): + ax.plot( + [ + (start[k, 1] + semiconv_mrad) * px, + (stop[k, 1] + semiconv_mrad) * px, + ], + [ + (start[k, 0] + semiconv_mrad) * px, + (stop[k, 0] + semiconv_mrad) * px, + ], + color=color, + lw=wid[k] * px * width_scale, + alpha=float(dep[k]), + solid_capstyle="butt", + ) + return ax + + +def average_bloch_fourier( + crystal: Crystal, + orientation: torch.Tensor, + trial_tilts: torch.Tensor, + thicknesses_A, + energy_ev: float, + precession_deg: float, + sg_max: float = SG_MAX, + k_max: float | None = None, + deform: torch.Tensor | None = None, + beams: torch.Tensor | None = None, + n_harmonics: int = 48, + n_geometry: int = 128, + n_matrix_harmonics: int | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Precession-averaged Bloch intensities by harmonic propagation. + + On the precession ring the structure matrix is a Fourier series in the + azimuth, A(phi) = sum_m A_m exp(i m phi); for a ring centered on the + optic axis only m = 0, +-1 are nonzero and the coefficients are exact, + + A_0 = U + diag(2 k0 c_g + i U0''), + A_(+1) = diag[-k0 r (g_x - i g_y) / (K - g_z)], A_(-1) = conj., + + with r = k0 sin(theta_p), K = sqrt(k0^2 - r^2), c_g the ring-centered + excitation error. Expanding the wave function in azimuthal modes, + psi(phi, z) = sum_n x_n(z) exp(i n phi), the Bloch equation couples + neighboring modes, dx_n/dz = (i pi / k0) sum_m A_m x_(n-m), starting + from x_0(0) = e_000, and the ring average of the intensity is the + incoherent sum over modes, I_g = sum_n |x_(n,g)|^2, exactly (Parseval). + The mode chain is truncated at |n| <= n_harmonics (48 reproduces a + converged quadrature to 1e-15 for silicon at 600 A and 0.4 degrees; 24 + leaves 1e-7) and a uniformly spaced thickness grid comes from one + action of the matrix exponential of the block-tridiagonal generator + at all its time points. For a ring displaced by a trial tilt + the coefficients are no longer three: the exact excitation errors on + n_geometry azimuths are Fourier transformed and n_matrix_harmonics + of them kept (default n_harmonics // 3); both counts and n_harmonics + must be converged for the result to be exact. + + The absorption is the full complex matrix (there is no first-order + variant here). Cost: a sparse block matrix of size (2 n_harmonics + + 1) x nb per trial tilt and one Krylov exponential action over the + thickness grid. + + Status: an alternative to the azimuthal quadrature of + illumination_nodes(), verified against it in the tests but not used by + the refinement functions of this module, which average batched + eigensolves instead (one eigendecomposition serves every thickness). + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + trial_tilts : torch.Tensor + (M, 2) ring centers as in-plane incident wavevectors (1/Angstroms). + thicknesses_A : float | array-like + Thickness grid in Angstroms; a uniform grid is propagated in one + pass. + energy_ev : float + Beam energy in eV. + precession_deg : float + Precession semi-angle in degrees. + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms) of the beam list. + k_max : float | None + Largest |g| (1/Angstroms) of a beam. + deform : torch.Tensor | None + (3, 3) deformation of the lab-frame reciprocal vectors. + beams : torch.Tensor | None + Explicit beam list (nb, 3), 000 first; selected here if None. + n_harmonics : int, default=48 + Azimuthal modes kept, |n| <= n_harmonics. + n_geometry : int, default=128 + Azimuths sampled for the coefficients of a displaced ring. + n_matrix_harmonics : int | None + Fourier coefficients of the displaced ring kept; default + n_harmonics // 3. + + Returns + ------- + intensities : torch.Tensor + (M, T, nb) ring-averaged intensities, direct beam first. + g_xy : torch.Tensor + (nb, 2) in-plane positions of the beams (1/Angstroms). + """ + from scipy.sparse import csr_matrix, diags, kron + from scipy.sparse.linalg import expm_multiply + + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + gamma_rel = relativistic_gamma(energy_ev) + t_grid = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + r = k0 * np.sin(np.deg2rad(precession_deg)) + K = np.sqrt(k0**2 - r**2) + trial = torch.atleast_2d(torch.as_tensor(trial_tilts, dtype=torch.float64)) + if beams is None: + alpha_max = (float(torch.linalg.norm(trial, dim=1).max()) + r) / k0 + beams = select_dynamical_beams( + crystal, orientation, energy_ev, alpha_max, sg_max, k_max, deform + ) + g_beams = qrotate(orientation, beams[1:].to(torch.float64) @ crystal.lat_recip) + if deform is not None: + g_beams = g_beams @ deform.to(torch.float64).T + g_beams = torch.cat([torch.zeros((1, 3), dtype=torch.float64), g_beams]).numpy() + nb = g_beams.shape[0] + U, u0_imag, _ = _coupling_matrix(crystal, beams, gamma_rel) + U_np = U.numpy() + L = int(n_harmonics) + nm = 2 * L + 1 + H = n_harmonics // 3 if n_matrix_harmonics is None else int(n_matrix_harmonics) + g2 = (g_beams**2).sum(axis=1) + + out = np.zeros((trial.shape[0], t_grid.shape[0], nb)) + for it, t0 in enumerate(trial.numpy()): + if np.hypot(*t0) < 1e-12: + den = K - g_beams[:, 2] + c = (2 * K * g_beams[:, 2] - g2) / (2 * den) + coeff = { + 0: U_np + np.diag(2 * k0 * c + 1j * u0_imag), + 1: np.diag(-k0 * r * (g_beams[:, 0] - 1j * g_beams[:, 1]) / den), + -1: np.diag(-k0 * r * (g_beams[:, 0] + 1j * g_beams[:, 1]) / den), + } + else: + Q = int(n_geometry) + phi = 2 * np.pi * np.arange(Q) / Q + t = t0[None, :] + r * np.stack([np.cos(phi), np.sin(phi)], axis=1) + kz = np.sqrt(k0**2 - (t**2).sum(1))[:, None] + s_phi = (2 * kz * g_beams[:, 2] - 2 * (t @ g_beams[:, :2].T) - g2) / ( + 2 * (kz - g_beams[:, 2]) + ) + d = np.fft.fft(2 * k0 * s_phi, axis=0) / Q + coeff = {m: np.diag(d[m % Q]) for m in range(-H, H + 1)} + coeff[0] = coeff[0] + U_np + 1j * u0_imag * np.eye(nb) + operator = csr_matrix((nm * nb, nm * nb), dtype=complex) + for m, A in coeff.items(): + if abs(m) > 2 * L: + continue + shift = diags(np.ones(nm - abs(m)), -m, shape=(nm, nm), format="csr") + operator = operator + kron(shift, csr_matrix(A), format="csr") + generator = (1j * np.pi / k0) * operator + x0 = np.zeros(nm * nb, dtype=complex) + x0[L * nb] = 1.0 + zs = t_grid.numpy() + uniform = zs.shape[0] > 1 and np.allclose(np.diff(zs), zs[1] - zs[0]) + if uniform: + x = expm_multiply( + generator, + x0, + start=float(zs[0]), + stop=float(zs[-1]), + num=zs.shape[0], + endpoint=True, + ).reshape(zs.shape[0], nm, nb) + out[it] = (np.abs(x) ** 2).sum(axis=1) + else: + for iz, z in enumerate(zs): + x = expm_multiply(z * generator, x0).reshape(nm, nb) + out[it, iz] = (np.abs(x) ** 2).sum(axis=0) + return torch.as_tensor(out), torch.as_tensor(g_beams[:, :2]) + + +def illumination_nodes( + energy_ev: float, + precession_deg: float = 0.0, + n_precession: int = 32, + semiconv_mrad: float = 0.0, + n_disk_radial: int = 4, + n_disk_azimuthal: int = 16, + maped_tilts_deg=None, + maped_weights=None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Incident beam tilts and weights that model one measured pattern. + + Precession is a ring of radius k0 sin(theta_p) sampled uniformly in + azimuth (the Gauss-Chebyshev quadrature of the ring integral); the + convergence disk of radius k0 sin(alpha) is sampled with Gauss-Legendre + nodes in (radius / R)^2 and uniform azimuth, which integrates the + uniform-area measure exactly for polynomials (an equal-weight ring + grid gives the disk a second moment of 0.70 R^2 instead of 0.50 R^2). + MAPED is an explicit tilt list with exposure weights. Ring and disk + combine as a product measure. Zero tilt with unit weight when none + apply. + + Parameters + ---------- + energy_ev : float + Beam energy in eV. + precession_deg : float, default=0.0 + Precession semi-angle in degrees; 0 for none. + n_precession : int, default=32 + Azimuthal samples on the precession ring. + semiconv_mrad : float, default=0.0 + Convergence semiangle in mrad; 0 for a parallel beam. + n_disk_radial, n_disk_azimuthal : int, default=4, 16 + Gauss-Legendre radii and azimuths of the convergence disk. + maped_tilts_deg : array-like | None + (M, 2) explicit beam tilts in degrees (MAPED), replacing the ring. + maped_weights : array-like | None + (M,) exposure weights of the MAPED tilts; equal if None. + + Returns + ------- + tilts : torch.Tensor + (M, 2) in-plane incident wavevectors (1/Angstroms). + weights : torch.Tensor + (M,) weights summing to one. + """ + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + if maped_tilts_deg is not None: + ring = k0 * torch.sin(torch.deg2rad(torch.as_tensor(maped_tilts_deg, dtype=torch.float64))) + if maped_weights is None: + w_ring = torch.full((ring.shape[0],), 1.0 / ring.shape[0], dtype=torch.float64) + else: + w_ring = torch.as_tensor(maped_weights, dtype=torch.float64) + w_ring = w_ring / w_ring.sum() + elif precession_deg > 0: + phi = torch.arange(n_precession, dtype=torch.float64) * (2 * np.pi / n_precession) + r = k0 * np.sin(np.deg2rad(precession_deg)) + ring = torch.stack([r * torch.cos(phi), r * torch.sin(phi)], dim=1) + w_ring = torch.full((n_precession,), 1.0 / n_precession, dtype=torch.float64) + else: + ring = torch.zeros((1, 2), dtype=torch.float64) + w_ring = torch.ones(1, dtype=torch.float64) + if semiconv_mrad > 0: + R = k0 * np.sin(semiconv_mrad * 1e-3) + x, w = np.polynomial.legendre.leggauss(n_disk_radial) + radius = R * np.sqrt((x + 1) / 2) + psi = 2 * np.pi * np.arange(n_disk_azimuthal) / n_disk_azimuthal + disk = torch.as_tensor( + (radius[:, None, None] * np.stack([np.cos(psi), np.sin(psi)], axis=1)[None]).reshape( + -1, 2 + ) + ) + w_disk = torch.as_tensor(np.repeat(w / 2 / n_disk_azimuthal, n_disk_azimuthal)) + else: + disk = torch.zeros((1, 2), dtype=torch.float64) + w_disk = torch.ones(1, dtype=torch.float64) + tilts = (ring[:, None, :] + disk[None, :, :]).reshape(-1, 2) + weights = (w_ring[:, None] * w_disk[None, :]).reshape(-1) + return tilts, weights + + +def _dynamical_cost(inten, sq, qxy, im, delta, power: float, min_sim_rel: float = 0.0): + """Intensity cost (M, T) of simulated raw intensities (M, T, N) at + positions sq (N, 2) against measured peaks (P, 2) with intensities im + (P,) already raised to `power`, with a free scale per (tilt, + thickness). Both sides are compared as I ** power; the power is applied + here, after any illumination averaging, never before it. Pairing is by + position (within delta); the position residuals themselves do not + enter, they belong to the deformation fit. Unpaired simulated beams + weaker than min_sim_rel times the strongest simulated beam (a + visibility mask on the RAW intensities, so it is defined for power = + 0 as well) are ignored: they are the beams a detector would not see, + and with power < 1 they would otherwise dominate the unpaired term.""" + d = torch.cdist(sq, qxy) + d_min, j_min = d.min(dim=1) + pair = d_min < delta + if int(pair.sum()) == 0: + return None, pair, j_min, d_min + visible = inten > min_sim_rel * inten.amax(dim=2, keepdim=True) + si = inten.clamp_min(0) ** power + a = si[:, :, pair] + b = im[j_min[pair]][None, None, :] + w = ((a * b).sum(dim=2) / (a * a).sum(dim=2).clamp_min(1e-12)).clamp_min(0)[:, :, None] + c_paired = (b - w * a).abs().sum(dim=2) + c_unpaired_sim = 0.5 * w[:, :, 0] * (si[:, :, ~pair] * visible[:, :, ~pair]).sum(dim=2) + matched = torch.zeros(im.shape[0], dtype=torch.bool) + matched[j_min[pair]] = True + # the measured direct beam is not a diffracted intensity: leave it out + # of the unexplained-measured term and of the normalization + direct = torch.linalg.norm(qxy, dim=1) < delta + matched |= direct + c_unpaired_exp = 0.5 * float(im[~matched].sum()) + norm = float(im[~direct].sum()) + 1e-12 + cost = (c_paired + c_unpaired_sim + c_unpaired_exp) / norm + return cost, pair, j_min, d_min + + +def _fit_deformation(sq, qxy, w_exp, delta): + """Symmetric in-plane deformation S and in-plane rotation angle wz + (radians) from the paired positions: A = (sum w qm qs^T)(sum w qs qs^T)^-1 + with measured = A ideal, split by polar decomposition A = S Q.""" + d = torch.cdist(sq, qxy) + d_min, j_min = d.min(dim=1) + pair = d_min < delta + if int(pair.sum()) < 3: + return None, 0.0, pair + qs = sq[pair] + qm = qxy[j_min[pair]] + w = w_exp[j_min[pair]] * (1 - d_min[pair] / delta).clamp_min(0) + # the paired positions must span the plane: collinear pairs (one + # systematic row) leave the deformation across the row undetermined + sv = torch.linalg.svdvals(torch.sqrt(w)[:, None] * qs) + if float(sv[-1]) < 0.2 * float(sv[0]): + return None, 0.0, pair + M1 = torch.einsum("p,pi,pj->ij", w, qm, qs) + M2 = torch.einsum("p,pi,pj->ij", w, qs, qs) + A = M1 @ torch.linalg.inv(M2 + 1e-12 * torch.eye(2, dtype=torch.float64)) + U_, _, Vh_ = torch.linalg.svd(A) + Q = U_ @ Vh_ + if torch.linalg.det(Q) < 0: + return None, 0.0, pair + S = A @ Q.T + S = 0.5 * (S + S.T) + wz = float(torch.atan2(Q[1, 0], Q[0, 0])) + return S, wz, pair + + +# matched orientations of two positions closer than this (degrees) belong to +# one grain for the neighbor rescue: above the error of kinematical matching +# (a few tenths of a degree), below typical grain boundary angles +_RESCUE_SAME_GRAIN_DEG = 2.0 + + +def _closest_symmetry_variant(q: torch.Tensor, ref: torch.Tensor, sym_quats) -> torch.Tensor: + """The symmetry equivalent q * s of orientation q closest to ref.""" + if sym_quats is None: + return q + from quantem.diffraction.rotations import qmult + + variants = qmult(q[None], torch.as_tensor(sym_quats, dtype=q.dtype)) # (S, 4) + return variants[int((variants @ ref).abs().argmax())] + + +def _tilt_twist(dq: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Split a lab-frame rotation into dq = tilt * twist: the twist is a + rotation about the beam (z), the tilt one about an in-plane axis. + + Returns the tilt as its rotation vector (wx, wy) in radians, and the + twist quaternion.""" + from quantem.diffraction.rotations import qconj, qmult + + dq = dq / torch.linalg.norm(dq) + n = float(torch.hypot(dq[0], dq[3])) + if n < 1e-12: + twist = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=dq.dtype) + else: + twist = torch.stack([dq[0], torch.zeros_like(dq[0]), torch.zeros_like(dq[0]), dq[3]]) / n + swing = qmult(dq, qconj(twist)) + if swing[0] < 0: + swing = -swing + sin_half = float(torch.linalg.norm(swing[1:3])) + if sin_half < 1e-15: + return torch.zeros(2, dtype=dq.dtype), twist + angle = 2 * np.arctan2(sin_half, float(swing[0])) + return swing[1:3] / sin_half * angle, twist + + +def refine_dynamical( + phase_map, + thicknesses_A: np.ndarray | None = None, + tilt_stages=((0.3, 0.1), (0.1, 0.025), (0.03, 0.006)), + precession_deg: float | None = None, + n_precession: int = 32, + n_precession_search: int | None = None, + semiconv_mrad: float | None = None, + n_disk_radial: int = 4, + n_disk_azimuthal: int = 16, + maped_tilts_deg=None, + refine_deformation: bool = True, + pair_distance: float | None = None, + power_intensity: float | None = None, + min_sim_intensity_rel: float | None = None, + sg_max: float = SG_MAX, + k_max: float | None = None, + min_number_peaks: int | None = None, + mask: np.ndarray | None = None, + fast_absorption: bool = True, + update_orientations: bool = True, + require_phase_weight: bool = True, + warm_start: bool = True, + neighbor_rescue: bool = True, + rescue_thickness_A: float = 100.0, + rescue_tilt_deg: float = 0.05, + rescue_max_starts: int = 2, + num_workers: int | None = None, + progress_bar: bool = True, +) -> dict: + """Dynamical refinement on the Bragg vectors: orientation, thickness, + in-plane deformation and candidate, pixel by pixel. + + Starting from the kinematically matched orientation of each candidate, + the crystal is re-initialized at every trial orientation of a + coarse-to-fine tilt grid and its diffracted intensities computed with + Bloch waves, averaged over the precession ring, the convergence disk or + the MAPED tilt list, for all thicknesses at once (one batched + eigendecomposition per stage). The peak pairing is fixed by the + positions, which the tilt does not move; the intensity cost is + minimized over (tilt, thickness), and the tilt is interpolated + parabolically at the finest stage. At the refined orientation the + symmetric in-plane deformation of the tilted cell is solved in closed + form from the paired positions (weighted least squares), and its + antisymmetric part, an in-plane rotation, is folded into the + orientation. The candidate with the lowest cost decides the phase. + + The intensities are far more tilt-sensitive than the positions: at + 500 A the rocking curve width is ~2e-3 1/A, so a 0.05 degree tilt + error is already visible in the weak beams. The default stages search + +-0.3 degrees at 0.1, +-0.1 at 0.025 and +-0.03 at 0.006 degrees, 251 + trial orientations per candidate, and should start from orientations + refined by refine_orientations(). The compute goes into the Bloch + eigensolves, one per trial orientation, illumination node and + candidate; the search stages use fewer precession nodes and the + first-order absorption, and the reported solution is then evaluated + once with the full node count and the exact absorption, so the stored + cost, thickness and orientation are at full accuracy while the search + costs a fraction of a full-accuracy grid. + + Parameters left as None inherit from the previous stages: the + pairing distance, intensity power, weak-beam cut and peak minimum from + the phase fit, and the precession and convergence angles from the + OrientationMaps (from_vectors). The resolved values are recorded in + phase_map.metadata['dynamical'] and returned under 'metadata'. + + Parameters + ---------- + phase_map : PhaseMap + A fitted PhaseMap (fit() has been run). + thicknesses_A : np.ndarray | None + Thickness grid in Angstroms; default 50 to 2000 in 25 A steps (the + thickness axis is nearly free: all thicknesses come from one + eigendecomposition). + tilt_stages : sequence of (half_range_deg, step_deg) + Successive tilt grids, each centered on the previous optimum. + precession_deg, n_precession : float | None, int + Precession semi-angle (inherited from the OrientationMap) and the + number of azimuthal samples on the ring for the last stage and the + final evaluation. Uniform sampling is the Gauss-Chebyshev + quadrature of the ring integral; the rocking curves oscillate at + pi t rho k0 g along the ring, so 32 or more samples are needed at + 500 A and 0.5 degrees. + n_precession_search : int | None + Ring samples for the search stages before the last one; defaults + to half of n_precession (at least 8). The coarse stages only have + to find the basin, which the reduced sampling does. + semiconv_mrad : float | None + Convergence semiangle (inherited); the intensities are averaged + over the disk with n_disk_radial Gauss-Legendre radii times + n_disk_azimuthal azimuths (64 nodes by default, times the ring). + n_disk_radial, n_disk_azimuthal : int, default=4, 16 + Convergence disk sampling, see semiconv_mrad. + maped_tilts_deg : array-like | None + Explicit (M, 2) beam tilt list (degrees) for MAPED, overriding + precession. + pair_distance : float | None + Largest distance (1/Angstroms) at which a simulated and a measured + peak are paired; inherited from the phase fit. + power_intensity : float | None + Intensities are compared as I ** power_intensity; inherited from + the phase fit. + min_sim_intensity_rel : float | None + Unpaired simulated beams weaker than this fraction of the + strongest simulated beam do not count against a candidate (the + detector would not have seen them). Inherited from the phase fit, + else MIN_SIM_INTENSITY_REL (0.02). + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms) of the Bloch beam list, + widened by the tilt search range and the illumination. + k_max : float | None + Largest |g| (1/Angstroms) of a beam; None keeps every reflection + within the cutoff. Recorded in the metadata, so the image + refinement uses the same beam set. + min_number_peaks : int | None + Positions with fewer measured peaks, direct beam included, are + skipped. None inherits the minimum of the phase fit, itself the + matching's (5 by default). At least 3: the direct beam and two + non-collinear reflections. + refine_deformation : bool, default=True + Solve the symmetric in-plane deformation and the in-plane rotation + from the paired positions before the intensity search; the + deformation is applied to the tilted cell in the Bloch calculation + and the rotation folded into the orientation. + mask : np.ndarray | None + Positions to refine: an (R, C) boolean mask or a list of + (row, col), as for OrientationMap.match_orientations. None + (default) refines every position the orientation maps reached, so + a staged test run on a few positions carries through. + fast_absorption : bool, default=True + First-order treatment of absorption during the search (Hermitian + eigh, ~4x faster, 0.5% rms intensity error); the final evaluation + of the reported solution always uses the exact complex absorption. + The tilt and thickness are never treated perturbatively. + update_orientations : bool, default=True + Write the refined quaternions back into the OrientationMaps. + require_phase_weight : bool, default=True + Refine only candidates that carried weight in the kinematical + phase fit; False refines every matched candidate, so the dynamical + pass can rescue a candidate the kinematical model rejected. + warm_start : bool, default=True + Start each position from the refined orientation of an already + refined neighbor (above or to the left, same candidate) and skip + the coarsest stage. Neighbors within one grain share their + orientation to well inside the fine stages, so this removes about + half of the eigensolves; the in-plane deformation and rotation are + still fit from the position's own peaks, and the final evaluation + is unchanged. The reported tilt and zero-tilt cost still refer to + the position's own matched orientation. + neighbor_rescue : bool, default=True + Second pass: positions whose winning solution differs from a + 4-neighbor in the same grain (same crystal, matched orientations + within 2 degrees; the nearest refined position + within two steps, so a mask of every second position works too) by + more than rescue_thickness_A in thickness or rescue_tilt_deg in + orientation are refined again from that neighbor's solution, and + the lower cost is kept. Repairs isolated wrong basins (thickness + aliases, tilt minima at a grid edge). + rescue_thickness_A : float, default=100.0 + Thickness difference (Angstroms) to a neighbor that triggers a + rescue. + rescue_tilt_deg : float, default=0.05 + Misorientation (degrees) between the refined orientations of a + position and a neighbor that triggers a rescue. Neighbors in one + grain differ by the true orientation gradient, so keep it above + that. + rescue_max_starts : int, default=2 + Neighbor solutions tried per rescued position, lowest cost first, + skipping neighbors whose solution repeats one already tried. + num_workers : int | None + Threads refining positions side by side; None uses every core. + The Bloch eigensolves are too small to spread over cores on their + own, so this is where the speed comes from. Positions are handed + out in contiguous raster-order blocks and warm starts stay inside a + block, so the result depends on the number of blocks, never on + which thread finishes first. + progress_bar : bool, default=True + Show progress bars over positions and rescues. + + Returns + ------- + dict + Per position (R, C) at the winning candidate: 'thickness', + 'thickness_contrast' (range of the final cost over the thickness + grid; a small value means the thickness is not determined) and + 'tilt_deg' (R, C, 2), the tilt about the lab x and y axes (degrees) + from 'quats_base' to 'quats'. Per candidate (R, C, F): 'quats' + (..., 4) refined orientations; 'quats_base' (..., 4) the matched + orientation with the in-plane rotation of the refinement folded in + (a symmetry equivalent of it closest to the solution), so quats = + tilt x quats_base whether or not the position was warm started or + rescued; 'deformation' (..., 2, 2) symmetric in-plane deformation + A of the tilted cell in the calibrated frame (measured reciprocal + positions = A x ideal); 'cost'; 'cost_zero_tilt', the best cost + over thickness at 'quats_base', evaluated like the final cost (its + difference to 'cost' is the gain of the tilt search; a small gain + means the intensities do not constrain the tilt); + 'thickness_per_candidate'; 'warm_started' flags. Also 'rescued' + (R, C) flags, 'phase_index' (R, C) the winning crystal and + 'candidate' (R, C) the winning candidate (both -1 where nothing + was refined), and 'metadata', the resolved parameters. Values are + NaN where a candidate was not refined. + """ + from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qconj, + qmult, + quat_from_axis_angle, + ) + + if thicknesses_A is None: + thicknesses_A = np.arange(50.0, 2000.0 + 1e-6, 25.0) + t_grid = torch.as_tensor(thicknesses_A, dtype=torch.float64) + T = t_grid.shape[0] + + oms = phase_map.orientation_maps + fit_md = phase_map.metadata.get("fit") if hasattr(phase_map, "metadata") else None + om_md = oms[0].metadata if hasattr(oms[0], "metadata") else None + pair_distance = resolve(pair_distance, "pair_distance", fit_md, default=PAIR_DISTANCE) + power_intensity = resolve(power_intensity, "power_intensity", fit_md, default=POWER_INTENSITY) + min_sim_intensity_rel = resolve( + min_sim_intensity_rel, "min_sim_intensity_rel", fit_md, default=MIN_SIM_INTENSITY_REL + ) + min_number_peaks = int( + resolve(min_number_peaks, "min_number_peaks", fit_md, default=MIN_NUMBER_PEAKS) + ) + if min_number_peaks < 3: + raise ValueError( + f"min_number_peaks={min_number_peaks}: a dynamical fit needs at least the " + "direct beam and two non-collinear reflections" + ) + precession_deg = float(resolve(precession_deg, "precession_deg", om_md, default=0.0)) + semiconv_mrad = float(resolve(semiconv_mrad, "semiconv_mrad", om_md, default=0.0)) + if n_precession_search is None: + n_precession_search = max(8, n_precession // 2) + used = dict( + thicknesses_A=np.asarray(thicknesses_A, dtype=float).tolist(), + tilt_stages=[tuple(float(v) for v in st) for st in tilt_stages], + precession_deg=precession_deg, + n_precession=int(n_precession), + n_precession_search=int(n_precession_search), + semiconv_mrad=semiconv_mrad, + n_disk_radial=int(n_disk_radial), + n_disk_azimuthal=int(n_disk_azimuthal), + maped_tilts_deg=maped_tilts_deg, + refine_deformation=bool(refine_deformation), + pair_distance=float(pair_distance), + power_intensity=float(power_intensity), + min_sim_intensity_rel=float(min_sim_intensity_rel), + sg_max=float(sg_max), + k_max=k_max, + min_number_peaks=int(min_number_peaks), + fast_absorption=bool(fast_absorption), + require_phase_weight=bool(require_phase_weight), + warm_start=bool(warm_start), + neighbor_rescue=bool(neighbor_rescue), + rescue_thickness_A=float(rescue_thickness_A), + rescue_tilt_deg=float(rescue_tilt_deg), + rescue_max_starts=int(rescue_max_starts), + num_workers=None if num_workers is None else int(num_workers), + ) + if hasattr(phase_map, "metadata"): + phase_map.metadata["dynamical"] = used + cands = phase_map.candidates + peaks = oms[0].peaks + R, C = peaks.shape[0], peaks.shape[1] + F = len(cands) + delta = pair_distance + energy_ev = oms[0].energy_ev + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + + ring, w_ring = illumination_nodes( + energy_ev, + precession_deg, + n_precession, + semiconv_mrad, + n_disk_radial, + n_disk_azimuthal, + maped_tilts_deg=maped_tilts_deg, + ) # (Mr, 2), (Mr,) + Mr = ring.shape[0] + alpha_ill = float(torch.linalg.norm(ring, dim=1).max()) / k0 + ring_s, w_ring_s = illumination_nodes( + energy_ev, + precession_deg, + n_precession_search, + semiconv_mrad, + n_disk_radial, + n_disk_azimuthal, + maped_tilts_deg=maped_tilts_deg, + ) + Mr_s = ring_s.shape[0] + + def stage_grid(center, half, step): + n = int(round(2 * half / step)) + 1 + tg = torch.linspace(-half, half, n, dtype=torch.float64) + wx_g, wy_g = torch.meshgrid(tg, tg, indexing="ij") + w = torch.stack([wx_g.reshape(-1), wy_g.reshape(-1)], dim=1) + center[None, :] + return w, n, tg + + cost_out = torch.full((R, C, F), torch.nan, dtype=torch.float64) + cost0_out = torch.full((R, C, F), torch.nan, dtype=torch.float64) + thick_out = torch.full((R, C, F), torch.nan, dtype=torch.float64) + tcontrast_out = torch.full((R, C, F), torch.nan, dtype=torch.float64) + tilt_out = torch.zeros((R, C, F, 2), dtype=torch.float64) + quat_out = torch.zeros((R, C, F, 4), dtype=torch.float64) + quat_out[..., 0] = 1.0 + quat_base = quat_out.clone() + deform_out = torch.zeros((R, C, F, 2, 2), dtype=torch.float64) + deform_out[..., 0, 0] = 1.0 + deform_out[..., 1, 1] = 1.0 + warm_out = torch.zeros((R, C, F), dtype=torch.bool) + rescued_out = torch.zeros((R, C), dtype=torch.bool) + + def refine_from(crystal, q_start, q_match, qxy, im, w_exp, stages): + """Search from q_start: in-plane deformation and rotation from the + positions, one beam list, the tilt stages, and the exact final + evaluation. q_match is the kinematically matched orientation of the + position, the reference of the reported tilt and of the zero-tilt + cost (q_start differs from it on a warm start or a rescue). Returns + None or a dict with the solution.""" + q0 = q_start + S = None + deform3 = None + if refine_deformation: + # in-plane deformation and rotation from the positions first: + # the rotation is folded into the orientation and the + # symmetric deformation applied to the tilted cell, so the + # intensity search below sees the strained lattice and a + # pairing free of position residuals + g0 = qrotate(q0, crystal.g_vec) + near = ( + torch.abs((2 * g0[:, 2] - lam * (g0**2).sum(1)) / (2 - 2 * lam * g0[:, 2])) + < sg_max + ) + S, wz, _ = _fit_deformation(g0[near, :2], qxy, w_exp, delta) + if S is not None: + half_z = torch.tensor(wz / 2, dtype=torch.float64) + dqz = torch.stack( + [torch.cos(half_z), torch.zeros(()), torch.zeros(()), torch.sin(half_z)] + ).to(torch.float64) + q0 = qmult(dqz, q0) + deform3 = torch.eye(3, dtype=torch.float64) + deform3[:2, :2] = S + # the matched orientation with the in-plane rotation of q0: the + # base the reported tilt is measured from (equal to q0 on a cold + # start), and its tilt away from q0 + q_m = _closest_symmetry_variant(q_match, q0, crystal.sym_quats) + base_tilt, twist = _tilt_twist(qmult(q0, qconj(q_m))) + q_base = qmult(twist, q_m) + offset = float(torch.linalg.norm(base_tilt)) + # one beam list for the whole search of this candidate: every + # trial center, the base, the illumination and the deformation are + # inside its selection, so all stages compare the same truncated + # system + beam_list = select_dynamical_beams( + crystal, + q0, + energy_ev, + np.deg2rad(stages[0][0]) * np.sqrt(2) + alpha_ill + offset, + sg_max, + k_max, + deform3, + ) + if beam_list.shape[0] < 2: + return None + + def exact_cost(q): + # full illumination and exact absorption at one orientation + inten, g_xy, _ = _cbed_amplitudes( + crystal, + q, + ring, + t_grid, + energy_ev, + sg_max, + k_max, + tilt_batch=max(64, Mr * 8), + progress_bar=False, + fast_absorption=False, + deform=deform3, + beams=beam_list, + ) + inten = (inten * w_ring[:, None, None]).sum(dim=0, keepdim=True) + cost, _, _, _ = _dynamical_cost( + inten[:, :, 1:], g_xy[1:], qxy, im, delta, power_intensity, min_sim_intensity_rel + ) + return cost + + # untilted reference: the best thickness at the matched orientation, + # for the gain the tilt search achieves + cost_base = exact_cost(q_base) + cost0 = float("nan") if cost_base is None else float(cost_base[0].min()) + center = torch.zeros(2, dtype=torch.float64) + best = None + n_stages = len(stages) + for i_stage, (half, step) in enumerate(stages): + half = np.deg2rad(half) + step = np.deg2rad(step) + w_grid, n, tg = stage_grid(center, half, step) + Mt = w_grid.shape[0] + last = i_stage == n_stages - 1 + nodes, w_nodes, M_nodes = (ring, w_ring, Mr) if last else (ring_s, w_ring_s, Mr_s) + # crystal tilt (wx, wy) about the in-plane axes shifts s_g by + # wx g_y - wy g_x; the same excitation errors come from a beam + # tilt k0 (wy, -wx) in the fixed-normal Bloch geometry, so every + # trial orientation is a full re-solve of the Bloch problem + # with the coupling matrix shared + trial = k0 * torch.stack([w_grid[:, 1], -w_grid[:, 0]], dim=1) + tilts = (trial[:, None, :] + nodes[None, :, :]).reshape(-1, 2) + inten, g_xy, _ = _cbed_amplitudes( + crystal, + q0, + tilts, + t_grid, + energy_ev, + sg_max, + k_max, + tilt_batch=max(64, M_nodes * 8), + progress_bar=False, + fast_absorption=fast_absorption, + deform=deform3, + beams=beam_list, + ) + inten = (inten.reshape(Mt, M_nodes, T, -1) * w_nodes[None, :, None, None]).sum(dim=1) + sq = g_xy[1:] + if sq.shape[0] == 0: + break + cost, pair, j_min, d_min = _dynamical_cost( + inten[:, :, 1:], sq, qxy, im, delta, power_intensity, min_sim_intensity_rel + ) + if cost is None: + break + flat = int(cost.argmin()) + m_best, t_best = flat // T, flat % T + i_b, j_b = m_best // n, m_best % n + cost_t = cost[:, t_best].reshape(n, n) + wx, wy = float(w_grid[m_best, 0]), float(w_grid[m_best, 1]) + if 0 < i_b < n - 1: + c0, c1, c2 = cost_t[i_b - 1, j_b], cost_t[i_b, j_b], cost_t[i_b + 1, j_b] + den = float(c0 - 2 * c1 + c2) + if den > 1e-12: + wx += 0.5 * float(c0 - c2) / den * step + if 0 < j_b < n - 1: + c0, c1, c2 = cost_t[i_b, j_b - 1], cost_t[i_b, j_b], cost_t[i_b, j_b + 1] + den = float(c0 - 2 * c1 + c2) + if den > 1e-12: + wy += 0.5 * float(c0 - c2) / den * step + center = torch.tensor([wx, wy], dtype=torch.float64) + best = (float(cost[m_best, t_best]), float(t_grid[t_best]), wx, wy) + if best is None: + return None + c_best, t_fit, wx, wy = best + q = q0 + ang = float(np.hypot(wx, wy)) + if ang > 1e-12: + axis = torch.tensor([wx / ang, wy / ang, 0.0], dtype=torch.float64) + q = qmult(quat_from_axis_angle(axis, torch.tensor(ang, dtype=torch.float64)), q0) + # the reported solution is the crystal rotated by the interpolated + # tilt with the beam along the foil normal: evaluate that model once + # more, with the full illumination and the exact absorption, so the + # stored cost and thickness belong to the stored orientation at full + # accuracy (the beam-offset search geometry is paraxially, not + # exactly, equivalent to it) + cost = exact_cost(q) + t_contrast = float("nan") + if cost is not None: + t_best = int(cost[0].argmin()) + c_best, t_fit = float(cost[0, t_best]), float(t_grid[t_best]) + # how much the cost varies over the thickness grid at this + # orientation: a flat curve means the thickness is not + # determined by these intensities (precession, few beams) + t_contrast = float(cost[0].max() - cost[0].min()) + # the reported tilt and base split the rotation from the matched + # orientation exactly, q = tilt x base; on a cold start they are the + # search's own (wx, wy) and q0, on a warm start the base differs + # from the zero-tilt one above only at second order in the tilt + tilt, twist = _tilt_twist(qmult(q, qconj(q_m))) + q_base = qmult(twist, q_m) + return dict( + cost=c_best, + t=t_fit, + wx=float(tilt[0]), + wy=float(tilt[1]), + q=q, + q_base=q_base, + S=S, + cost0=cost0, + t_contrast=t_contrast, + ) + + def store(rx, ry, f, sol): + cost_out[rx, ry, f] = sol["cost"] + thick_out[rx, ry, f] = sol["t"] + tcontrast_out[rx, ry, f] = sol.get("t_contrast", float("nan")) + tilt_out[rx, ry, f, 0] = sol["wx"] + tilt_out[rx, ry, f, 1] = sol["wy"] + quat_base[rx, ry, f] = sol["q_base"] + quat_out[rx, ry, f] = sol["q"] + if sol["S"] is not None: + deform_out[rx, ry, f] = sol["S"] + if np.isfinite(sol["cost0"]): + cost0_out[rx, ry, f] = sol["cost0"] + + def peaks_at(rx, ry): + data = peaks[rx, ry].numpy().astype(np.float64) + if data.shape[0] < min_number_peaks: + return None + qxy = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + im = torch.as_tensor(data[:, ix[2]], dtype=torch.float64).clamp_min(0) ** power_intensity + return qxy, im, im / im.max().clamp_min(1e-12) + + from quantem.diffraction.orientation import position_mask + + mask_rc = position_mask(mask, (R, C)) + for om in oms: + if om.computed is not None: + mask_rc = mask_rc & om.computed + positions = [(r, c) for r, c in np.ndindex(R, C) if mask_rc[r, c]] + + def refine_position(rx, ry, block): + pk = peaks_at(rx, ry) + if pk is None: + return + qxy, im, w_exp = pk + for f, (i_om, m) in enumerate(cands): + om = oms[i_om] + if om.corr[rx, ry, m] <= 0: + continue + if ( + require_phase_weight + and phase_map.phase_weights is not None + and float(phase_map.phase_weights[rx, ry, f]) <= 0 + ): + continue + q_start = om.quats[rx, ry, m] + quat_out[rx, ry, f] = q_start + stages = tilt_stages + if warm_start and len(tilt_stages) > 1: + # an already refined neighbor of the same candidate (raster + # order: above or to the left, in the same block, one or two + # steps away so a mask of every second position still warm + # starts) is a start inside the fine stages' reach; its + # solution costs one coarse stage less + for nr, nc in ((rx - 1, ry), (rx, ry - 1), (rx - 2, ry), (rx, ry - 2)): + if (nr, nc) not in block: + continue + if not torch.isfinite(cost_out[nr, nc, f]): + continue + # same grain: the kinematically matched orientations of + # the two positions agree within the coarse stage + miso = float( + misorientation_angle_deg( + om.quats[rx, ry, m][None], + om.quats[nr, nc, m][None], + om.crystal.sym_quats, + )[0] + ) + if miso < tilt_stages[0][0]: + q_start = quat_out[nr, nc, f] + stages = tilt_stages[1:] + warm_out[rx, ry, f] = True + break + sol = refine_from(om.crystal, q_start, om.quats[rx, ry, m], qxy, im, w_exp, stages) + if sol is None: + continue + store(rx, ry, f, sol) + + def run_parallel(jobs, work, desc): + """Run work(job) over jobs on num_workers threads. Each Bloch + eigensolve is too small to use more than one core, and torch + releases the GIL inside it, so positions run side by side; every + job writes only its own positions.""" + bar = tqdm(total=sum(len(j) for j in jobs), desc=desc) if progress_bar else None + lock = threading.Lock() + + def run(job): + for item in job: + work(item) + if bar is not None: + with lock: + bar.update(1) + + if n_workers == 1: + for job in jobs: + run(job) + else: + # one intra-op thread per worker: the worker threads already + # fill the cores, and torch's own pool on top of them would + # oversubscribe. The setting is process-global, so the previous + # value is restored afterwards. + n_threads = torch.get_num_threads() + torch.set_num_threads(1) + try: + with ThreadPoolExecutor(n_workers) as pool: + for fut in [pool.submit(run, job) for job in jobs]: + fut.result() + finally: + torch.set_num_threads(n_threads) + if bar is not None: + bar.close() + + n_workers = max(1, int(num_workers if num_workers is not None else os.cpu_count() or 1)) + # fill the lazily cached lattice data once, before any thread reads it + for om in oms: + _beam_universe(om.crystal) + # contiguous runs of positions in raster order, several per worker so + # they balance but long enough that most positions still warm start; + # warm starts stay inside a run, so the result does not depend on which + # thread finished first + n_blocks = 1 if n_workers == 1 else max(1, min(4 * n_workers, len(positions) // 16)) + blocks = [] + for idx in np.array_split(np.arange(len(positions)), max(n_blocks, 1)): + block = {positions[i] for i in idx} + blocks.append([(*positions[i], block) for i in idx]) + run_parallel( + [b for b in blocks if b], lambda item: refine_position(*item), "dynamical refinement" + ) + + if neighbor_rescue and len(tilt_stages) > 1: + cost_f0 = torch.nan_to_num(cost_out, nan=torch.inf) + f_win = cost_f0.argmin(dim=-1) + done = torch.isfinite(cost_out).any(dim=-1) + rescue_list = [] + + def miso(r0, c0, f0, r1, c1, f1): + # misorientation (degrees) of two refined solutions of the same + # crystal: the tilt corrections have different bases (the + # matched orientations of the two positions), the solutions not + return float( + misorientation_angle_deg( + quat_out[r0, c0, f0][None], + quat_out[r1, c1, f1][None], + oms[cands[f0][0]].crystal.sym_quats, + )[0] + ) + + def nearest_done(rx, ry, dr, dc): + # the refined position one step away, or two on a sparse mask + for k in (1, 2): + nr, nc = rx + k * dr, ry + k * dc + if 0 <= nr < R and 0 <= nc < C and done[nr, nc]: + return nr, nc + return None + + for rx, ry in np.ndindex(R, C): + if not done[rx, ry]: + continue + f = int(f_win[rx, ry]) + i_om = cands[f][0] + starts = [] + for dr, dc in ((-1, 0), (1, 0), (0, -1), (0, 1)): + nb = nearest_done(rx, ry, dr, dc) + if nb is None: + continue + nr, nc = nb + fn = int(f_win[nr, nc]) + if cands[fn][0] != i_om: + continue + # same grain only: a start from another grain is no rescue, + # and its tilt from the matched orientation would widen the + # beam list without bound + if ( + misorientation_angle_deg( + oms[i_om].quats[rx, ry, cands[f][1]][None], + oms[i_om].quats[nr, nc, cands[fn][1]][None], + oms[i_om].crystal.sym_quats, + )[0] + >= _RESCUE_SAME_GRAIN_DEG + ): + continue + dt = abs(float(thick_out[nr, nc, fn]) - float(thick_out[rx, ry, f])) + if dt > rescue_thickness_A or miso(nr, nc, fn, rx, ry, f) > rescue_tilt_deg: + starts.append((float(cost_out[nr, nc, fn]), nr, nc, fn)) + if starts: + # lowest-cost neighbors first, one start per distinct + # solution, at most rescue_max_starts (each start is a full + # fine-stage search) + starts.sort(key=lambda x: x[0]) + kept: list = [] + for c_n, nr, nc, fn in starts: + dup = False + for _, kr, kc, kf in kept: + if ( + abs(float(thick_out[nr, nc, fn]) - float(thick_out[kr, kc, kf])) + <= rescue_thickness_A + and miso(nr, nc, fn, kr, kc, kf) <= rescue_tilt_deg + ): + dup = True + break + if not dup: + kept.append((c_n, nr, nc, fn)) + if len(kept) >= max(1, rescue_max_starts): + break + # the neighbors' solutions as they stand now: rescues run + # in parallel and must not start from each other's updates + rescue_list.append( + (rx, ry, f, [quat_out[nr, nc, fn].clone() for _, nr, nc, fn in kept]) + ) + + def rescue(item): + rx, ry, f, starts = item + pk = peaks_at(rx, ry) + if pk is None: + return + qxy, im, w_exp = pk + i_om, m = cands[f] + crystal = oms[i_om].crystal + q_match = oms[i_om].quats[rx, ry, m] + for q_n in starts: + sol = refine_from(crystal, q_n, q_match, qxy, im, w_exp, tilt_stages[1:]) + if sol is not None and sol["cost"] < float(cost_out[rx, ry, f]) - 1e-9: + store(rx, ry, f, sol) + rescued_out[rx, ry] = True + + n_jobs = 1 if n_workers == 1 else 4 * n_workers + run_parallel( + [rescue_list[k::n_jobs] for k in range(n_jobs) if rescue_list[k::n_jobs]], + rescue, + "neighbor rescue", + ) + + n_maps = len(oms) + cost_f = torch.nan_to_num(cost_out, nan=torch.inf) + cost_phase = torch.full((R, C, n_maps), torch.inf, dtype=torch.float64) + for f, (i_om, _) in enumerate(cands): + cost_phase[..., i_om] = torch.minimum(cost_phase[..., i_om], cost_f[..., f]) + done = torch.isfinite(cost_out).any(dim=-1) + phase_index = torch.where(done, cost_phase.argmin(dim=-1), -1) + f_best = cost_f.argmin(dim=-1) + thickness = torch.gather(thick_out, 2, f_best[..., None]).squeeze(-1) + thickness_contrast = torch.gather(tcontrast_out, 2, f_best[..., None]).squeeze(-1) + tilt_deg = torch.rad2deg( + torch.gather(tilt_out, 2, f_best[..., None, None].expand(R, C, 1, 2)).squeeze(2) + ) + tilt_deg[~done] = torch.nan + + if update_orientations: + for f, (i_om, m) in enumerate(cands): + done = torch.isfinite(cost_out[..., f]) + oms[i_om].quats[..., m, :][done] = quat_out[..., f, :][done] + + return { + "thickness": thickness, + "thickness_contrast": thickness_contrast, + "tilt_deg": tilt_deg, + "quats": quat_out, + "deformation": deform_out, + "cost": cost_out, + "cost_zero_tilt": cost0_out, + "quats_base": quat_base, + "warm_started": warm_out, + "rescued": rescued_out, + "phase_index": phase_index, + "candidate": torch.where(done, f_best, -1), + "thickness_per_candidate": thick_out, + "metadata": used, + } + + +def dynamical_maps( + result: dict, + phase_map, + crystal_index: int | None = None, + min_thickness_contrast: float = 0.02, +) -> dict: + """Maps of the winning candidate of a refine_dynamical() result. + + Parameters + ---------- + result : dict + From refine_dynamical(). + phase_map : PhaseMap + The PhaseMap that was refined (for the candidate list). + crystal_index : int | None + Keep only positions won by this crystal; None keeps all. + min_thickness_contrast : float, default=0.02 + The thickness is NaN where the thickness contrast is below this: a + flat cost curve, typical of precessed data with few beams, does not + determine the thickness and the grid minimum there is not a + measurement. + + Returns + ------- + dict + (R, C) maps 'thickness' (Angstroms), 'tilt_deg' (magnitude of the + tilt from the matched orientation, degrees), 'gain' (cost at the + matched orientation minus the final cost), 'cost', + 'thickness_contrast', 'phase_index' (-1 outside the mask), 'mask' + (positions refined, and of the given crystal when crystal_index is + set), 'quats' (R, C, 4), 'deformation' (R, C, 2, 2) and 'strain', + the crystal-frame strain components of strain_crystal_frame(). + Positions outside the mask are NaN. + """ + cand = result["candidate"] + R, C = cand.shape + refined = cand >= 0 + cand = cand.clamp_min(0) + idx4 = cand[..., None, None] + quats = torch.gather(result["quats"], 2, idx4.expand(R, C, 1, 4)).squeeze(2) + deform = torch.gather( + result["deformation"], 2, cand[..., None, None, None].expand(R, C, 1, 2, 2) + ).squeeze(2) + cost = torch.gather(result["cost"], 2, cand[..., None]).squeeze(-1) + cost0 = torch.gather(result["cost_zero_tilt"], 2, cand[..., None]).squeeze(-1) + tcon = result.get("thickness_contrast") + if tcon is None: + tcon = torch.full_like(cost, torch.nan) + mask = refined & torch.isfinite(cost) + if crystal_index is not None: + i_om = torch.tensor([c[0] for c in phase_map.candidates]) + mask &= i_om[cand] == crystal_index + nan = torch.full((R, C), torch.nan, dtype=torch.float64) + strain = strain_crystal_frame(deform, quats) + out = { + "thickness": torch.where( + mask & ~(tcon < min_thickness_contrast), result["thickness"], nan + ), + "tilt_deg": torch.where(mask, torch.linalg.norm(result["tilt_deg"], dim=-1), nan), + "gain": torch.where(mask, cost0 - cost, nan), + "thickness_contrast": torch.where(mask, tcon, nan), + "cost": torch.where(mask, cost, nan), + "phase_index": torch.where(mask, result["phase_index"], -1), + "mask": mask, + "quats": torch.where(mask[..., None], quats, torch.nan), + "deformation": torch.where(mask[..., None, None], deform, torch.nan), + "strain": {k: torch.where(mask, v, nan) for k, v in strain.items() if k != "eps_crystal"}, + } + return out + + +def plot_dynamical_maps( + maps: dict, + scalebar=None, + thickness_range_A: tuple[float, float] = (0.0, 2000.0), + tilt_range_deg: tuple[float, float] = (0.0, 0.3), + gain_range: tuple[float, float] = (0.0, 0.05), + axsize: tuple[float, float] = (4.0, 4.0), +): + """Thickness, tilt correction, gain of the tilt search and final cost + of a dynamical refinement. + + Parameters + ---------- + maps : dict + From dynamical_maps(). + scalebar : dict | None + Passed to show_2d. + thickness_range_A : tuple[float, float], default=(0.0, 2000.0) + Color range of the thickness map (Angstroms). + tilt_range_deg : tuple[float, float], default=(0.0, 0.3) + Color range of the tilt map (degrees). + gain_range : tuple[float, float], default=(0.0, 0.05) + Color range of the gain map. + axsize : tuple[float, float], default=(4.0, 4.0) + Size of each panel in inches. + + Returns + ------- + fig, axs + From show_2d. NaN positions are shown as zero. + """ + from quantem.core.visualization import show_2d + + imgs = [ + [np.nan_to_num(maps["thickness"].numpy()), np.nan_to_num(maps["tilt_deg"].numpy())], + [np.nan_to_num(maps["gain"].numpy()), np.nan_to_num(maps["cost"].numpy())], + ] + cmax = ( + float(np.nanmax(maps["cost"].numpy())) if np.isfinite(maps["cost"].numpy()).any() else 1.0 + ) + return show_2d( + imgs, + title=[["thickness (A)", "tilt correction (deg)"], ["gain of the tilt search", "cost"]], + cmap=[["viridis", "magma"], ["magma", "gray_r"]], + cbar=True, + norm=[ + [ + { + "interval_type": "manual", + "vmin": thickness_range_A[0], + "vmax": thickness_range_A[1], + }, + {"interval_type": "manual", "vmin": tilt_range_deg[0], "vmax": tilt_range_deg[1]}, + ], + [ + {"interval_type": "manual", "vmin": gain_range[0], "vmax": gain_range[1]}, + {"interval_type": "manual", "vmin": 0.0, "vmax": cmax}, + ], + ], + scalebar=scalebar, + axsize=axsize, + ) + + +def strain_crystal_frame(deformation: torch.Tensor, quats: torch.Tensor) -> dict: + """Strain tensor components in the crystal Cartesian frame. + + The measured in-plane reciprocal deformation A (2, 2) of the tilted + cell (measured = A x ideal) is the reciprocal image of the real-space + deformation F = A^-T restricted to the beam-normal plane; only that + in-plane part is observable from one projection, and the components + along the beam are set to zero before the tensor is rotated into the + crystal frame with the 3x3 orientation matrix R (v_lab = R v_crystal): + eps_crystal = R^T eps_lab R. The crystal axes are the Cartesian frame + of the cell (x along a, z along c; for hexagonal cells 'b' is the + in-basal-plane direction perpendicular to a). + + Parameters + ---------- + deformation : torch.Tensor + (..., 2, 2) symmetric in-plane deformation from refine_dynamical + (or the columns of OrientationMap.calculate_strain's A). + quats : torch.Tensor + (..., 4) orientations. + + Returns + ------- + dict of (...) tensors 'aa', 'bb', 'cc', 'ab', 'ac', 'bc' (strain + components) and 'eps_crystal' (..., 3, 3). + """ + from quantem.diffraction.rotations import quat_to_matrix + + A = deformation.to(torch.float64) + Fp = torch.linalg.inv(A).transpose(-1, -2) # real-space in-plane deformation + eps2 = 0.5 * (Fp + Fp.transpose(-1, -2)) - torch.eye(2, dtype=torch.float64) + eps_lab = torch.zeros(A.shape[:-2] + (3, 3), dtype=torch.float64) + eps_lab[..., :2, :2] = eps2 + Rm = quat_to_matrix(quats.to(torch.float64)) + eps_c = torch.einsum("...ji,...jk,...kl->...il", Rm, eps_lab, Rm) + return { + "aa": eps_c[..., 0, 0], + "bb": eps_c[..., 1, 1], + "cc": eps_c[..., 2, 2], + "ab": eps_c[..., 0, 1], + "ac": eps_c[..., 0, 2], + "bc": eps_c[..., 1, 2], + "eps_crystal": eps_c, + } + + +def plot_strain_crystal_frame( + strain: dict, + mask: np.ndarray | None = None, + strain_range_percent: tuple[float, float] = (-2.0, 2.0), + scalebar=None, + axsize: tuple[float, float] = (4.0, 4.0), + cmap: str = "RdBu_r", +): + """Six strain components in the crystal frame as maps. + + Normal strains along the crystal a, b and c axes on the top row and + the ab, ac and bc shears below, in percent, masked where the fit is + not trusted. Components with a c (beam-direction) index are the + rotated in-plane measurement only; see strain_crystal_frame(). + + Parameters + ---------- + strain : dict + (R, C) arrays 'aa', 'bb', 'cc', 'ab', 'ac', 'bc', as returned by + strain_crystal_frame() or dynamical_maps()['strain']. + mask : np.ndarray | None + (R, C) weights (0 hides a position); None shows all. + strain_range_percent : tuple[float, float], default=(-2.0, 2.0) + Color range in percent. + scalebar : dict | None + Passed to show_2d. + axsize : tuple[float, float], default=(4.0, 4.0) + Size of each panel in inches. + cmap : str, default="RdBu_r" + Colormap. + + Returns + ------- + fig, axs + From show_2d. + """ + from quantem.core.visualization import show_2d + + keys = [["aa", "bb", "cc"], ["ab", "ac", "bc"]] + names = [["ε_aa", "ε_bb", "ε_cc"], ["ε_ab", "ε_ac", "ε_bc"]] + m = 1.0 if mask is None else np.asarray(mask, dtype=float) + imgs = [[np.asarray(strain[k]) * 100 * m for k in row] for row in keys] + lo, hi = strain_range_percent + return show_2d( + imgs, + title=[[n + " (%)" for n in row] for row in names], + cmap=cmap, + cbar=True, + norm={"interval_type": "manual", "vmin": lo, "vmax": hi}, + scalebar=scalebar, + axsize=axsize, + ) + + +# ---------------------------------------------------------------------- +# image-based dynamical refinement (the final step, on the pattern pixels) +# ---------------------------------------------------------------------- + + +def _q_to_pixels( + q_xy: torch.Tensor, origin_rc, pixel_size: float, rotation_ccw_deg: float, ellipse +): + """Calibrated (qx, qy) [row, col frame] -> detector pixel (row, col): + undo the scan rotation and the ellipse correction of + calibration.peaks_to_calibrated, then scale and shift to the origin.""" + q = q_xy.to(torch.float64) + if rotation_ccw_deg: + th = np.deg2rad(-rotation_ccw_deg) + rot = torch.tensor( + [[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]], dtype=torch.float64 + ) + q = q @ rot.T + if ellipse is not None: + e11, e12 = float(ellipse[0]), float(ellipse[1]) + A = torch.tensor([[1 + e11, e12], [e12, 1 - e11]], dtype=torch.float64) + q = q @ torch.linalg.inv(A).T + return q / pixel_size + torch.as_tensor(origin_rc, dtype=torch.float64)[None, :] + + +def render_disks( + centers_px: torch.Tensor, + intensities: torch.Tensor, + shape: tuple[int, int], + disk_radius_px: float, + edge_px: float, +) -> torch.Tensor: + """Sum of soft-edged disks on a pixel grid. + + The edge is a logistic of width edge_px (the disk profile of a + defocused or blurred aperture). + + Parameters + ---------- + centers_px : torch.Tensor + (N, 2) disk centers in pixels, (row, col). + intensities : torch.Tensor + (..., N) disk intensities; leading dimensions give a stack. + shape : tuple[int, int] + Image shape (ny, nx). + disk_radius_px : float + Disk radius in pixels (the logistic's half point). + edge_px : float + Edge width in pixels. + + Returns + ------- + torch.Tensor + (..., ny, nx) images. + """ + ny, nx = shape + rows = torch.arange(ny, dtype=torch.float64) + cols = torch.arange(nx, dtype=torch.float64) + d = torch.sqrt( + (rows[None, :, None] - centers_px[:, 0, None, None]) ** 2 + + (cols[None, None, :] - centers_px[:, 1, None, None]) ** 2 + ) # (N, ny, nx) + disks = torch.sigmoid((disk_radius_px - d) / max(edge_px, 1e-3)) + return torch.einsum("...n,nyx->...yx", intensities.to(torch.float64), disks) + + +def render_pattern_image( + crystal: Crystal, + orientation: torch.Tensor, + thicknesses_A, + energy_ev: float, + shape: tuple[int, int], + origin_rc, + pixel_size: float, + rotation_ccw_deg: float = 0.0, + ellipse=None, + deform: torch.Tensor | None = None, + disk_radius_px: float = 3.0, + edge_px: float = 1.0, + tilts: torch.Tensor | None = None, + trial_tilts: torch.Tensor | None = None, + sg_max: float = SG_MAX, + k_max: float | None = None, + fast_absorption: bool = True, + tilt_weights: torch.Tensor | None = None, + beams: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Dynamical diffraction pattern images on the detector grid. + + Bloch intensities, averaged over the precession / convergence tilt set + `tilts`, rendered as disks for every trial orientation tilt and every + thickness. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. + orientation : torch.Tensor + Unit quaternion (4,), crystal to lab. + thicknesses_A : float | array-like + Thicknesses in Angstroms. + energy_ev : float + Beam energy in eV. + shape : tuple[int, int] + Detector shape (ny, nx). + origin_rc : array-like + Direct beam position (row, col) in pixels. + pixel_size : float + Detector sampling (1/Angstroms per pixel). + rotation_ccw_deg : float, default=0.0 + Scan-to-detector rotation of the calibration, undone here. + ellipse : sequence | None + Elliptic distortion (e11, e12) of the calibration, undone here. + deform : torch.Tensor | None + (3, 3) deformation of the lab-frame reciprocal vectors. + disk_radius_px, edge_px : float, default=3.0, 1.0 + Disk shape, see render_disks. + tilts : torch.Tensor | None + (Mr, 2) illumination tilts (1/Angstroms), e.g. from + illumination_nodes(); a single untilted beam if None. + trial_tilts : torch.Tensor | None + (M, 2) crystal tilts about the lab x and y axes (radians); none if + None. + sg_max : float, default=SG_MAX + Excitation error cutoff (1/Angstroms) of the beam list. + k_max : float | None + Largest |g| (1/Angstroms) of a beam. + fast_absorption : bool, default=True + First-order absorption; see _bloch_solve. + tilt_weights : torch.Tensor | None + (Mr,) weights of the illumination tilts; equal if None. + beams : torch.Tensor | None + Explicit beam list (nb, 3), 000 first. + + Returns + ------- + images : torch.Tensor + (M, T, ny, nx), M the number of trial tilts (1 when None). + centers_px : torch.Tensor + (nb, 2) disk centers (row, col), direct beam first. + intensities : torch.Tensor + (M, T, nb) illumination-averaged intensities. + """ + t_grid = torch.atleast_1d(torch.as_tensor(thicknesses_A, dtype=torch.float64)) + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + ring = torch.zeros((1, 2), dtype=torch.float64) if tilts is None else tilts + Mr = ring.shape[0] + w_ring = ( + torch.full((Mr,), 1.0 / Mr, dtype=torch.float64) if tilt_weights is None else tilt_weights + ) + if trial_tilts is None: + trial = torch.zeros((1, 2), dtype=torch.float64) + else: + trial = k0 * torch.stack([trial_tilts[:, 1], -trial_tilts[:, 0]], dim=1) + Mt = trial.shape[0] + all_tilts = (trial[:, None, :] + ring[None, :, :]).reshape(-1, 2) + inten, g_xy, _ = _cbed_amplitudes( + crystal, + orientation, + all_tilts, + t_grid, + energy_ev, + sg_max, + k_max, + tilt_batch=max(64, Mr * 8), + progress_bar=False, + fast_absorption=fast_absorption, + deform=deform, + beams=beams, + ) + inten = (inten.reshape(Mt, Mr, t_grid.shape[0], -1) * w_ring[None, :, None, None]).sum(dim=1) + centers = _q_to_pixels(g_xy, origin_rc, pixel_size, rotation_ccw_deg, ellipse) + images = render_disks(centers, inten, shape, disk_radius_px, edge_px) + return images, centers, inten + + +def _image_cost( + meas: torch.Tensor, + sims: torch.Tensor, + mask: torch.Tensor, + power: float, + background: str = "constant", + radius: torch.Tensor | None = None, +) -> torch.Tensor: + """Normalized residual of the measured image (ny, nx) against each + simulated one (..., ny, nx). The intensity scale and the background + are solved by least squares in the raw domain over the mask; the + residual is then taken between the power-law images so the weak + diffracted disks weigh as they do in the Bragg-vector cost. + + background : {"constant", "radial"} + A constant, or a quadratic in the distance from the direct beam + (b0 + b1 r + b2 r^2, with `radius` in pixels), the smooth + diffuse-scattering floor of a diffraction pattern. + """ + m = mask.to(torch.float64) + lead = sims.shape[:-2] + sel = m.reshape(-1) > 0 + x = sims.reshape(-1, sims.shape[-2] * sims.shape[-1])[:, sel].to(torch.float64) # (K, npix) + y = meas.reshape(-1)[sel].to(torch.float64) + ones = torch.ones_like(y) + if background == "radial" and radius is not None: + r = radius.reshape(-1)[sel].to(torch.float64) + r = r / r.max().clamp_min(1e-12) + basis = torch.stack([ones, r, r * r], dim=1) # (npix, 3) + else: + basis = ones[:, None] + # normal equations for [scale, background coefficients], per image + nb = basis.shape[1] + G_bb = basis.T @ basis + G_xb = x @ basis # (K, nb) + G_xx = (x * x).sum(dim=1) + rhs_x = x @ y + rhs_b = basis.T @ y + K = x.shape[0] + A = torch.zeros((K, nb + 1, nb + 1), dtype=torch.float64) + A[:, 0, 0] = G_xx + A[:, 0, 1:] = G_xb + A[:, 1:, 0] = G_xb + A[:, 1:, 1:] = G_bb[None] + rhs = torch.cat([rhs_x[:, None], rhs_b[None].expand(K, -1)], dim=1) + A = A + 1e-12 * torch.eye(nb + 1, dtype=torch.float64)[None] + coef = torch.linalg.solve(A, rhs[:, :, None])[:, :, 0] + a = coef[:, 0].clamp_min(0) + bg = (basis @ coef[:, 1:].T).T + model = (a[:, None] * x + bg).clamp_min(0) ** power + yp = y.clamp_min(0) ** power + resid = ((yp[None] - model) ** 2).sum(dim=1) + return (resid / (yp * yp).sum().clamp_min(1e-12)).reshape(lead) + + +def _image_mask(shape, origin, r_max_px, exclude_direct_px): + yy, xx = np.mgrid[0 : shape[0], 0 : shape[1]] + r = np.hypot(yy - origin[0], xx - origin[1]) + m = np.ones(shape, dtype=bool) if r_max_px is None else r <= r_max_px + if exclude_direct_px is not None and exclude_direct_px > 0: + m &= r > exclude_direct_px + return torch.as_tensor(m) + + +def fit_disk_shape( + dataset, + phase_map, + result: dict, + origins: np.ndarray, + pixel_size: float, + rotation_ccw_deg: float = 0.0, + ellipse=None, + positions=None, + n_positions: int = 20, + radii_px=None, + edges_px=None, + power_intensity: float | None = None, + r_max_px: float | None = None, + exclude_direct_px: float | None = None, + background: str = "constant", + sg_max: float | None = None, + k_max: float | None = None, + fast_absorption: bool | None = None, + progress_bar: bool = True, +) -> dict: + """Global disk radius and edge width from the best-fit patterns. + + The convergence disk shape is a property of the illumination, not of + the position, so it is fit once: on the `n_positions` positions with + the lowest dynamical cost (or the given `positions`), the rendered + pattern at the refined orientation, thickness and deformation is + compared with the measured image over a grid of (radius, edge), and + the pair minimizing the summed image cost is returned for + refine_dynamical_image() to use. + + The illumination, intensity power and beam set (sg_max, k_max, + fast_absorption) default to those of the refine_dynamical() result. + + Parameters + ---------- + dataset : Dataset4dstem + Measured patterns; anything with `.array` (R, C, ny, nx) and + `.shape`. + phase_map : PhaseMap + The refined PhaseMap. + result : dict + From refine_dynamical(). + origins : np.ndarray + (R, C, 2) direct beam positions (row, col) in pixels. + pixel_size : float + Detector sampling (1/Angstroms per pixel). + rotation_ccw_deg : float, default=0.0 + Scan-to-detector rotation of the calibration. + ellipse : sequence | None + Elliptic distortion (e11, e12) of the calibration. + positions : list of (int, int) | None + Positions to fit; None takes the n_positions with the lowest + dynamical cost. + n_positions : int, default=20 + Number of positions when positions is None. + radii_px : array-like | None + Disk radii to try (pixels); default 1.5 to 6 in steps of 0.5. + edges_px : array-like | None + Edge widths to try (pixels); default 0.5, 0.75, 1, 1.5, 2. + power_intensity : float | None + Images are compared as I ** power_intensity; inherited. + r_max_px : float | None + Ignore pixels farther than this from the origin; None uses all. + exclude_direct_px : float | None + Ignore pixels within this distance of the origin; default 1.5 + times the largest radius tried. + background : {"constant", "radial"}, default="constant" + Background model of the image cost: a constant, or a quadratic in + the distance from the direct beam. + sg_max, k_max : float | None + Beam set; inherited from the result's metadata (SG_MAX if absent). + fast_absorption : bool | None + First-order absorption; inherited (True if absent). + progress_bar : bool, default=True + Show a progress bar over positions. + + Returns + ------- + dict + 'disk_radius_px', 'edge_px' the best pair, 'cost' (n_radii, + n_edges) the summed image cost, 'radii_px', 'edges_px', + 'positions'. + """ + oms = phase_map.orientation_maps + cands = phase_map.candidates + md = result.get("metadata", {}) + energy_ev = oms[0].energy_ev + power_intensity = float( + resolve(power_intensity, "power_intensity", md, default=POWER_INTENSITY) + ) + # the beam set of the Bragg-vector refinement, unless overridden + sg_max = float(resolve(sg_max, "sg_max", md, default=SG_MAX)) + k_max = resolve(k_max, "k_max", md) + fast_absorption = bool(resolve(fast_absorption, "fast_absorption", md, default=True)) + tilts, tilt_w = illumination_nodes( + energy_ev, + md.get("precession_deg", 0.0), + md.get("n_precession", 32), + md.get("semiconv_mrad", 0.0), + md.get("n_disk_radial", 4), + md.get("n_disk_azimuthal", 16), + maped_tilts_deg=md.get("maped_tilts_deg"), + ) + if radii_px is None: + radii_px = np.arange(1.5, 6.01, 0.5) + if edges_px is None: + edges_px = np.array([0.5, 0.75, 1.0, 1.5, 2.0]) + cost = torch.nan_to_num(result["cost"], nan=torch.inf).amin(dim=-1) + if positions is None: + flat = torch.argsort(cost.reshape(-1))[:n_positions] + positions = [ + (int(i) // cost.shape[1], int(i) % cost.shape[1]) + for i in flat + if torch.isfinite(cost.reshape(-1)[i]) + ] + shape = tuple(dataset.shape[-2:]) + if exclude_direct_px is None: + exclude_direct_px = 1.5 * float(np.max(radii_px)) + total = torch.zeros((len(radii_px), len(edges_px)), dtype=torch.float64) + it = tqdm(positions, desc="disk shape") if progress_bar else positions + for rx, ry in it: + f = int(result["candidate"][rx, ry]) + i_om, m = cands[f] + om = oms[i_om] + q = result["quats"][rx, ry, f] + t = float(result["thickness_per_candidate"][rx, ry, f]) + d3 = torch.eye(3, dtype=torch.float64) + d3[:2, :2] = result["deformation"][rx, ry, f] + o = origins[rx, ry] + meas = torch.as_tensor(np.asarray(dataset.array[rx, ry], dtype=float)).clamp_min(0) + mask = _image_mask(shape, o, r_max_px, exclude_direct_px) + radius = torch.as_tensor( + np.hypot( + *(np.mgrid[0 : shape[0], 0 : shape[1]] - np.asarray(o, dtype=float)[:, None, None]) + ) + ) + _, centers, inten = render_pattern_image( + om.crystal, + q, + [t], + energy_ev, + shape, + o, + pixel_size, + rotation_ccw_deg, + ellipse, + d3, + 1.0, + 1.0, + tilts, + None, + sg_max, + k_max, + fast_absorption, + tilt_weights=tilt_w, + ) + for i, r in enumerate(radii_px): + for j, e in enumerate(edges_px): + sim = render_disks(centers, inten[0, 0], shape, float(r), float(e)) + total[i, j] += _image_cost(meas, sim, mask, power_intensity, background, radius) + k = int(total.argmin()) + i, j = k // len(edges_px), k % len(edges_px) + return { + "disk_radius_px": float(radii_px[i]), + "edge_px": float(edges_px[j]), + "cost": total, + "radii_px": np.asarray(radii_px), + "edges_px": np.asarray(edges_px), + "positions": positions, + } + + +def refine_dynamical_image( + dataset, + phase_map, + result: dict, + origins: np.ndarray, + pixel_size: float, + disk_radius_px: float, + edge_px: float, + rotation_ccw_deg: float = 0.0, + ellipse=None, + thickness_half_range_A: float = 100.0, + thickness_step_A: float = 10.0, + tilt_stage=(0.03, 0.01), + power_intensity: float | None = None, + r_max_px: float | None = None, + exclude_direct_px: float | None = None, + background: str = "constant", + mask: np.ndarray | None = None, + sg_max: float | None = None, + k_max: float | None = None, + fast_absorption: bool | None = None, + update_orientations: bool = True, + progress_bar: bool = True, +) -> dict: + """Final dynamical refinement against the diffraction images. + + Starting from the Bragg-vector solution of refine_dynamical (winning + candidate, orientation, thickness, in-plane deformation), every pixel + of the measured pattern is compared with a rendered pattern: Bloch + intensities averaged over the precession / convergence tilt set, + drawn as disks of the global radius and edge width from + fit_disk_shape(), with a free intensity scale and a constant or + radial background. The thickness and the orientation tilt are re-searched + on a local grid (thickness +- thickness_half_range_A, tilt +- the + stage half-range), the deformation and in-plane rotation are kept + from the position fit. The image cost is the residual after the + linear fit, normalized by the image power, so it is comparable + across positions. The direct beam disk is excluded from the cost + (exclude_direct_px, default 1.5 disk radii): it carries most of the + counts, its measured intensity is the least reliable (saturation, + detector response), and a fraction of a percent of model error on it + would outweigh every diffracted disk. Run this when the Bragg-vector + refinement is not accurate enough; it costs one rendered image per + trial (tilt, thickness) on top of the Bloch solves. + + The illumination, intensity power and beam set (sg_max, k_max, + fast_absorption) default to those of the refine_dynamical() result. + + Parameters + ---------- + dataset : Dataset4dstem + Measured patterns; anything with `.array` (R, C, ny, nx) and + `.shape`. + phase_map : PhaseMap + The refined PhaseMap. + result : dict + From refine_dynamical(). + origins : np.ndarray + (R, C, 2) direct beam positions (row, col) in pixels. + pixel_size : float + Detector sampling (1/Angstroms per pixel). + disk_radius_px, edge_px : float + Disk shape, from fit_disk_shape(). + rotation_ccw_deg : float, default=0.0 + Scan-to-detector rotation of the calibration. + ellipse : sequence | None + Elliptic distortion (e11, e12) of the calibration. + thickness_half_range_A : float, default=100.0 + Half-width (Angstroms) of the thickness search around the + Bragg-vector thickness. + thickness_step_A : float, default=10.0 + Thickness step in Angstroms. + tilt_stage : (float, float), default=(0.03, 0.01) + Tilt search half-range and step in degrees. + power_intensity : float | None + Images are compared as I ** power_intensity; inherited. + r_max_px : float | None + Ignore pixels farther than this from the origin; None uses all. + exclude_direct_px : float | None + Ignore pixels within this distance of the origin; default 1.5 + disk radii. + background : {"constant", "radial"}, default="constant" + Background model: a constant, or a quadratic in the distance from + the direct beam (the diffuse scattering floor). + mask : np.ndarray | None + Positions to refine, as in refine_dynamical; None refines all. + sg_max, k_max : float | None + Beam set; inherited from the result's metadata (SG_MAX if absent). + fast_absorption : bool | None + First-order absorption; inherited (True if absent). + update_orientations : bool, default=True + Write the refined quaternions back into the OrientationMaps. + progress_bar : bool, default=True + Show a progress bar over positions. + + Returns + ------- + dict + 'thickness' (R, C), 'tilt_deg' (R, C, 2) the additional tilt over + the Bragg-vector result, 'quats' (R, C, 4), 'cost' (R, C) the + normalized image residual (NaN where not refined), and 'metadata'. + """ + from quantem.diffraction.rotations import qmult, quat_from_axis_angle + + oms = phase_map.orientation_maps + cands = phase_map.candidates + md = result.get("metadata", {}) + energy_ev = oms[0].energy_ev + power_intensity = float( + resolve(power_intensity, "power_intensity", md, default=POWER_INTENSITY) + ) + # the beam set of the Bragg-vector refinement, unless overridden + sg_max = float(resolve(sg_max, "sg_max", md, default=SG_MAX)) + k_max = resolve(k_max, "k_max", md) + fast_absorption = bool(resolve(fast_absorption, "fast_absorption", md, default=True)) + if exclude_direct_px is None: + exclude_direct_px = 1.5 * disk_radius_px + tilts, tilt_w = illumination_nodes( + energy_ev, + md.get("precession_deg", 0.0), + md.get("n_precession", 32), + md.get("semiconv_mrad", 0.0), + md.get("n_disk_radial", 4), + md.get("n_disk_azimuthal", 16), + maped_tilts_deg=md.get("maped_tilts_deg"), + ) + R, C = result["thickness"].shape + shape = tuple(dataset.shape[-2:]) + half, step = (np.deg2rad(v) for v in tilt_stage) + n = int(round(2 * half / step)) + 1 + tg = torch.linspace(-half, half, n, dtype=torch.float64) + wx_g, wy_g = torch.meshgrid(tg, tg, indexing="ij") + w_grid = torch.stack([wx_g.reshape(-1), wy_g.reshape(-1)], dim=1) + + thickness = torch.full((R, C), torch.nan, dtype=torch.float64) + tilt_out = torch.zeros((R, C, 2), dtype=torch.float64) + quat_out = torch.zeros((R, C, 4), dtype=torch.float64) + quat_out[..., 0] = 1.0 + cost_out = torch.full((R, C), torch.nan, dtype=torch.float64) + + from quantem.diffraction.orientation import position_mask + + mask_rc = position_mask(mask, (R, C)) + iterator = [(r, c) for r, c in np.ndindex(R, C) if mask_rc[r, c]] + if progress_bar: + iterator = tqdm(iterator, desc="image refinement") + for rx, ry in iterator: + f = int(result["candidate"][rx, ry]) + if f < 0 or not torch.isfinite(result["cost"][rx, ry, f]): + continue + i_om, m = cands[f] + om = oms[i_om] + q0 = result["quats"][rx, ry, f] + t0 = float(result["thickness_per_candidate"][rx, ry, f]) + d3 = torch.eye(3, dtype=torch.float64) + d3[:2, :2] = result["deformation"][rx, ry, f] + o = origins[rx, ry] + meas = torch.as_tensor(np.asarray(dataset.array[rx, ry], dtype=float)).clamp_min(0) + pmask = _image_mask(shape, o, r_max_px, exclude_direct_px) + radius = torch.as_tensor( + np.hypot( + *(np.mgrid[0 : shape[0], 0 : shape[1]] - np.asarray(o, dtype=float)[:, None, None]) + ) + ) + t_grid = np.arange( + max(thickness_step_A, t0 - thickness_half_range_A), + t0 + thickness_half_range_A + 1e-6, + thickness_step_A, + ) + images, _, _ = render_pattern_image( + om.crystal, + q0, + t_grid, + energy_ev, + shape, + o, + pixel_size, + rotation_ccw_deg, + ellipse, + d3, + disk_radius_px, + edge_px, + tilts, + w_grid, + sg_max, + k_max, + fast_absorption, + tilt_weights=tilt_w, + ) + cost = _image_cost(meas, images, pmask, power_intensity, background, radius) # (M, T) + T = len(t_grid) + flat = int(cost.argmin()) + m_best, t_best = flat // T, flat % T + i_b, j_b = m_best // n, m_best % n + wx, wy = float(w_grid[m_best, 0]), float(w_grid[m_best, 1]) + cost_t = cost[:, t_best].reshape(n, n) + if 0 < i_b < n - 1: + c0, c1, c2 = cost_t[i_b - 1, j_b], cost_t[i_b, j_b], cost_t[i_b + 1, j_b] + den = float(c0 - 2 * c1 + c2) + if den > 1e-12: + wx += 0.5 * float(c0 - c2) / den * step + if 0 < j_b < n - 1: + c0, c1, c2 = cost_t[i_b, j_b - 1], cost_t[i_b, j_b], cost_t[i_b, j_b + 1] + den = float(c0 - 2 * c1 + c2) + if den > 1e-12: + wy += 0.5 * float(c0 - c2) / den * step + q = q0 + ang = float(np.hypot(wx, wy)) + if ang > 1e-12: + axis = torch.tensor([wx / ang, wy / ang, 0.0], dtype=torch.float64) + q = qmult(quat_from_axis_angle(axis, torch.tensor(ang, dtype=torch.float64)), q0) + thickness[rx, ry] = float(t_grid[t_best]) + tilt_out[rx, ry, 0] = wx + tilt_out[rx, ry, 1] = wy + quat_out[rx, ry] = q + cost_out[rx, ry] = cost[m_best, t_best] + if update_orientations: + om.quats[rx, ry, m] = q + + used = dict( + disk_radius_px=float(disk_radius_px), + edge_px=float(edge_px), + thickness_half_range_A=float(thickness_half_range_A), + thickness_step_A=float(thickness_step_A), + tilt_stage=tuple(float(v) for v in tilt_stage), + power_intensity=float(power_intensity), + r_max_px=r_max_px, + exclude_direct_px=float(exclude_direct_px), + background=str(background), + sg_max=float(sg_max), + k_max=k_max, + fast_absorption=bool(fast_absorption), + inherited=dict(md), + ) + if hasattr(phase_map, "metadata"): + phase_map.metadata["dynamical_image"] = used + return { + "thickness": thickness, + "tilt_deg": torch.rad2deg(tilt_out), + "quats": quat_out, + "cost": cost_out, + "metadata": used, + } diff --git a/src/quantem/diffraction/bragg_vectors.py b/src/quantem/diffraction/bragg_vectors.py new file mode 100644 index 000000000..a733d3d20 --- /dev/null +++ b/src/quantem/diffraction/bragg_vectors.py @@ -0,0 +1,1980 @@ +from __future__ import annotations + +import copy as _copy +import warnings +from pathlib import Path +from typing import Any, Literal, Sequence, Union + +import numpy as np +import torch +from numpy.typing import NDArray + +from quantem.core.datastructures.dataset2d import Dataset2d +from quantem.core.datastructures.dataset4dstem import Dataset4dstem +from quantem.core.datastructures.vector import Vector +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.bragg_vectors_visualization import ( + plot_basis_vectors, + plot_bvm, + plot_detection, + plot_diffraction_grid, + plot_lattice_fit, + plot_reference_lattice, + plot_template, +) +from quantem.diffraction.disk_detection import ( + cross_correlation, + detect_disks_batch, + estimate_central_beam, + make_template, + probe_centroid, + synthetic_probe, + template_fourier, +) +from quantem.diffraction.strain import StrainMap + +PEAK_FIELDS = ("q_row", "q_col", "intensity") + + +class BraggVectors(AutoSerialize): + """Correlation-based Bragg disk detection and lattice fitting for 4D-STEM. + + Workflow (each step writes state consumed by the next): + + 1. ``make_template_*`` – build a cross-correlation template, either from a + synthetic soft disk (:meth:`make_template_synthetic`), by averaging data + over an ROI (:meth:`make_template_from_data`), or from an explicit probe + image (:meth:`make_template_from_probe`). + 2. :meth:`detect_disks` – template-match every scan position; detected peaks + are stored in :attr:`peaks` (a :class:`Vector` of ``[q_row, q_col, + intensity]`` in detector pixels) and accumulated into the Bragg vector map + :attr:`bvm`. + 3. :meth:`choose_basis_vectors` – pick the lattice basis ``(origin, g1, g2)`` + from the BVM, automatically or by hand; also stores the numbered candidate + peaks. + 4. :meth:`index_peaks` – quick, lightweight: index just the picked candidate + peaks into a reference lattice (``reference_ab``/``reference_qpos``). + 5. :meth:`fit_lattice` – the heavy step: at every scan position, match the + detections to the reference within ``max_peak_shift``, intensity-weighted + least-squares fit the lattice vectors into ``g1_array``/``g2_array`` of shape + ``(scan_row, scan_col, 2)``, and compute the per-position ``mask_weight``. + 6. :meth:`calculate_strain_map` – hand the lattice vectors (and + ``mask_weight``) to a :class:`~quantem.diffraction.strain.StrainMap`. + + Detection runs in torch on :attr:`device` (CPU or GPU); the ragged peak table + is held in a torch-backed :class:`Vector`. The detector-to-scan rotation is + read from the parent dataset metadata (``q_to_r_rotation_ccw_deg`` and + ``q_transpose``), the same keys used by the DPC/CoM workflow, and applied in + :meth:`calculate_strain_map`. + + Use :meth:`from_dataset` to construct an instance. + + Parameters + ---------- + dataset : Dataset4dstem + The 4D-STEM dataset to analyze. + device : str, default="cpu" + Torch device used for detection (e.g. ``"cpu"`` or ``"cuda"``). + """ + + _token = object() + + # nanobeam lattice vectors are measured in reciprocal space -> sign flip + real_space: bool = False + + def __init__( + self, + dataset: Dataset4dstem, + device: str = "cpu", + _token: object | None = None, + ): + if _token is not self._token: + raise RuntimeError("Use BraggVectors.from_dataset() to instantiate this class.") + super(BraggVectors, self).__init__() + self.dataset = dataset + self.device = device + + self.peaks: Vector | None = None + self.calibration = None + self.bvm: Dataset2d | None = None + + self.origin: np.ndarray | None = None + self.g1: np.ndarray | None = None + self.g2: np.ndarray | None = None + + # candidate peaks picked from the BVM in choose_basis_vectors(), and the + # reference lattice (those candidates indexed) set by index_peaks(). + self.candidates_rc: np.ndarray | None = None + self.candidates_intensity: np.ndarray | None = None + self.reference_ab: np.ndarray | None = None + self.reference_qpos: np.ndarray | None = None + self.reference_intensity: np.ndarray | None = None + + self.g1_array: np.ndarray | None = None + self.g2_array: np.ndarray | None = None + # per-position diagnostics from fit_lattice() + self.mask_weight: np.ndarray | None = None + self.fit_error: np.ndarray | None = None + + self._template: torch.Tensor | None = None + self._template_ft: torch.Tensor | None = None + # full dataset resident on self.device, set by detect_disks(save_to_gpu=True) + self._gpu_cache: torch.Tensor | None = None + self.metadata: dict[str, Any] = {} + + @classmethod + def from_dataset( + cls, dataset: Dataset4dstem, *, device: str = "cpu", name: str | None = None + ) -> "BraggVectors": + """Create a BraggVectors workflow bound to a 4D-STEM dataset. + + Parameters + ---------- + dataset : Dataset4dstem + The 4D-STEM dataset to analyze. + device : str, default="cpu" + Torch device used for detection (e.g. ``"cpu"`` or ``"cuda"``). + name : str, optional + If given, sets ``dataset.name``. + + Returns + ------- + BraggVectors + A new workflow instance bound to ``dataset``. + """ + if not isinstance(dataset, Dataset4dstem): + raise TypeError("BraggVectors.from_dataset expects a Dataset4dstem instance.") + if name is not None: + dataset.name = name + return cls(dataset=dataset, device=device, _token=cls._token) + + def save( + self, + path: str | Path, + mode: Literal["w", "o"] = "w", + store: Literal["auto", "zip", "dir"] = "auto", + skip: Union[str, type, Sequence[Union[str, type]]] = (), + compression_level: int | None = 4, + *, + include_dataset: bool = False, + ) -> None: + """Save the workflow to disk, excluding the raw 4D-STEM dataset by default. + + Overrides :meth:`~quantem.core.io.serialize.AutoSerialize.save` to drop + :attr:`dataset` — the raw 4D-STEM cube, which dominates the file size — from + serialization by default. The detected :attr:`peaks`, lattice fit + (:attr:`g1_array`/:attr:`g2_array`), Bragg vector map and all diagnostics are + kept, so the file holds the *results* of the workflow (orders of magnitude + smaller than the data) rather than the data itself. The device copy of the + dataset made by ``detect_disks(save_to_gpu=True)`` is never saved. + + ``"dataset"`` is recorded in the file's skip metadata, so a reloaded workflow + simply has no ``dataset`` attribute. Re-attach one (``bv.dataset = ds``) before + calling methods that read the raw cube — :meth:`detect_disks`, + :meth:`correlation_map`, :meth:`make_template_from_data`, + :meth:`calculate_strain_map`, etc. Pass ``include_dataset=True`` to keep the + dataset in the file instead. + + Parameters + ---------- + path : str or Path + Target file path. Use a ``.zip`` extension for zip format, otherwise a + directory is written. + mode : {'w', 'o'}, default='w' + ``'w'`` writes only if the path does not exist; ``'o'`` overwrites. + store : {'auto', 'zip', 'dir'}, default='auto' + Storage format; ``'auto'`` infers from the file extension. + skip : str, type, or sequence of (str or type), default=() + Additional attribute names/types to skip during serialization, merged with + the default ``dataset`` exclusion. + compression_level : int or None, default=4 + Zstandard/Blosc compression level (0–9); ``0`` disables compression. + include_dataset : bool, default=False + If ``True``, keep the raw 4D-STEM :attr:`dataset` in the file (large). The + default ``False`` excludes it. + """ + if isinstance(skip, (str, type)): + skip = [skip] + else: + skip = list(skip) + if not include_dataset and "dataset" not in skip: + skip.append("dataset") + if "_gpu_cache" not in skip: + skip.append("_gpu_cache") + # Explicit (two-arg) super() rather than the bare super(): the zero-arg form + # needs a compiler-created __class__ closure cell that is absent when this + # method's source is re-exec'd from a string (Jupyter autoreload), which + # raises "super(): __class__ cell not found". The explicit form is immune. + super(BraggVectors, self).save( + path, + mode=mode, + store=store, + skip=skip, + compression_level=compression_level, + ) + + # ---- main methods ---- + + def make_template_synthetic( + self, + radius: float | None = None, + edge: float = 1.0, + center: tuple[float, float] | None = None, + subtract_mean: bool = False, + ) -> "BraggVectors": + """Build the template from a synthetic soft-edged disk. + + Parameters + ---------- + radius : float, optional + Disk radius in pixels. Defaults to a rough estimate from the mean + diffraction pattern + (:func:`~quantem.diffraction.disk_detection.estimate_central_beam`); + pass it explicitly when several disks share comparable intensity. + edge : float, default=1.0 + Width in pixels of the ``tanh`` edge falloff. + center : tuple of float, optional + ``(row, col)`` disk center; defaults to the detector center + ``(H // 2, W // 2)``. + subtract_mean : bool, default=False + If ``True``, make the template zero-sum. The default keeps the + unit-sum positive template, so correlation values stay positive and + roughly measure the probe-weighted counts under each peak. + + Returns + ------- + BraggVectors + ``self``, for method chaining. + """ + H, W = int(self.dataset.shape[-2]), int(self.dataset.shape[-1]) + if radius is None: + dp_mean = torch.as_tensor( + np.asarray(self.dataset.dp_mean.array), dtype=torch.float, device=self.device + ) + _, radius = estimate_central_beam(dp_mean) + if center is None: + center = (H // 2, W // 2) + probe = synthetic_probe((H, W), float(radius), edge=edge, center=center) + self._set_template(probe, center=center, subtract_mean=subtract_mean) + self.metadata["template"] = { + "kind": "synthetic", + "radius": float(radius), + "edge": float(edge), + "center": (float(center[0]), float(center[1])), + } + return self + + def make_template_from_data( + self, + roi: NDArray | None = None, + subtract_mean: bool = False, + center: tuple[float, float] | None = None, + ) -> "BraggVectors": + """Build the template by averaging diffraction patterns from the data. + + Parameters + ---------- + roi : np.ndarray, optional + ``(scan_row, scan_col)`` mask selecting scan positions to average — + ideally a vacuum / single-disk region so the unscattered probe is + isolated. ``None`` (default) averages the whole scan (the mean + diffraction pattern). + subtract_mean : bool, default=False + If ``True``, make the template zero-sum. The default keeps the + unit-sum positive template, so correlation values stay positive and + roughly measure the probe-weighted counts under each peak. + center : tuple of float, optional + ``(row, col)`` probe center rolled to the origin; defaults to the + probe's intensity centroid. + + Returns + ------- + BraggVectors + ``self``, for method chaining. + """ + gpu_cache = getattr(self, "_gpu_cache", None) + if gpu_cache is not None: + # Dataset is already resident on self.device (e.g. from a prior + # detect_disks(save_to_gpu=True)) -- reuse it instead of transferring again. + if roi is None: + probe = gpu_cache.mean(dim=(0, 1)) + else: + m = torch.as_tensor(np.asarray(roi) > 0, device=self.device) + if not bool(m.any()): + raise ValueError("roi selects no scan positions.") + probe = gpu_cache[m].mean(dim=0) + else: + # Select the (small) ROI on the CPU first so we never have to put the + # full dataset on the device just to average a handful of positions. + array = np.asarray(self.dataset.array) + if roi is None: + probe_np = array.mean(axis=(0, 1)) + else: + m = np.asarray(roi) > 0 + if not m.any(): + raise ValueError("roi selects no scan positions.") + probe_np = array[m].mean(axis=0) + probe = torch.as_tensor(probe_np, dtype=torch.float, device=self.device) + if center is None: + center = probe_centroid(probe) + self._set_template(probe, center=center, subtract_mean=subtract_mean) + self.metadata["template"] = { + "kind": "from_data", + "roi": roi is not None, + "center": (float(center[0]), float(center[1])), + } + return self + + def make_template_from_probe( + self, + probe: NDArray | torch.Tensor, + center: tuple[float, float] | None = None, + subtract_mean: bool = False, + ) -> "BraggVectors": + """Build the template from an explicit probe image (e.g. a measured vacuum probe). + + Parameters + ---------- + probe : np.ndarray or torch.Tensor + Probe image; must match the diffraction-pattern shape. + center : tuple of float, optional + ``(row, col)`` probe center rolled to the origin; defaults to the + probe's intensity centroid. + subtract_mean : bool, default=False + If ``True``, make the template zero-sum. The default keeps the + unit-sum positive template, so correlation values stay positive and + roughly measure the probe-weighted counts under each peak. + + Returns + ------- + BraggVectors + ``self``, for method chaining. + """ + probe_t = torch.as_tensor(probe, dtype=torch.float, device=self.device) + if center is None and tuple(probe_t.shape) == tuple(self.dataset.shape[-2:]): + center = probe_centroid(probe_t) + self._set_template(probe_t, center=center, subtract_mean=subtract_mean) + self.metadata["template"] = { + "kind": "from_probe", + "center": (float(center[0]), float(center[1])), + } + return self + + @property + def template(self) -> np.ndarray | None: + """The correlation template, fftshifted to the image center for display (numpy). + + The ``make_template_*`` methods store the template corner-shifted (center + at the ``[0, 0]`` FFT origin) so correlation peaks land at absolute disk + positions; this property shifts it back to the center for plotting. + + Returns + ------- + np.ndarray or None + ``(H, W)`` center-shifted template, or ``None`` if no template has been + built yet. + """ + if self._template is None: + return None + return torch.fft.fftshift(self._template).detach().cpu().numpy() + + def correlation_map( + self, + row: int, + col: int, + background_sigma: float | str | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, + ) -> np.ndarray: + """Cross-correlation map of one diffraction pattern with the template (numpy). + + Peaks in the returned map sit at absolute disk positions (no fftshift + needed), matching what :meth:`detect_disks` searches. + + Parameters + ---------- + row : int + Scan row of the diffraction pattern to correlate. + col : int + Scan column of the diffraction pattern to correlate. + background_sigma, corr_power, sigma_cc + Correlation options, as for :meth:`detect_disks`; pass the same + values to see the map the detection actually searches. + + Returns + ------- + np.ndarray + ``(H, W)`` real-space correlation map. + """ + if self._template_ft is None: + raise ValueError("Run a make_template_* method before correlation_map().") + dp = torch.as_tensor( + np.asarray(self.dataset.array[row, col]), dtype=torch.float, device=self.device + ) + corr, _ = cross_correlation( + dp, + self._template_ft, + self._resolve_background_sigma(background_sigma), + corr_power, + sigma_cc, + ) + return corr.detach().cpu().numpy() + + def detect_disks( + self, + *, + positions: list[tuple[int, int]] | None = None, + min_abs_intensity: float = 0.0, + min_spacing: float = 0.0, + edge_boundary: int = 1, + subpixel: str = "upsample", + upsample_factor: int = 16, + max_num_peaks: int = 1000, + background_sigma: float | str | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, + batch_size: int | None = None, + progressbar: bool = True, + save_to_gpu: bool = True, + ) -> Vector: + """Detect Bragg disks at every scan position (or a subset for testing). + + Pass ``positions`` to test detection hyperparameters on a handful of + patterns without scanning the full grid; the returned :class:`Vector` then + has shape ``(len(positions),)`` and the workflow state is left untouched. + With ``positions=None`` the full scan is processed, :attr:`peaks` and + :attr:`bvm` are populated, and the same Vector is returned. Patterns are + processed in batches (the FFTs and subpixel refinement run together across + the batch, which is far faster on a GPU); results are identical to detecting + each pattern on its own. + + Parameters + ---------- + positions : list of tuple of int, optional + ``(row, col)`` scan positions to test on. ``None`` (default) processes + the full scan and updates the workflow state. + min_abs_intensity : float, default=0.0 + Drop correlation peaks below this absolute intensity. + min_spacing : float, default=0.0 + Minimum spacing in pixels between kept peaks; closer / dimmer peaks are + suppressed. + edge_boundary : int, default=1 + Width in pixels of the border in which peaks are ignored. + subpixel : {"none", "parabolic", "upsample"}, default="upsample" + Subpixel refinement mode; see + :func:`~quantem.diffraction.disk_detection.detect_disks`. + upsample_factor : int, default=16 + Upsampling factor for the ``"upsample"`` subpixel refinement. + max_num_peaks : int, default=1000 + Maximum number of peaks to keep per pattern. + background_sigma : float | "auto" | None, default=None + Width in pixels of a Fourier high-pass applied to the + cross-correlation before peak finding: a copy of the correlation + map smoothed by a Gaussian of this width is subtracted from it. + The filter is isotropic in the Fourier domain and needs no origin, + so it is not a radial background fit. Its purpose is the negative + moat a zero-sum template leaves around the bright unscattered + beam, which pushes weak disk peaks below zero where the + correlation clamp erases them. Set it a little wider than a disk + so the disk-scale peaks pass untouched; "auto" uses twice the + central beam radius. + corr_power : float, default=1.0 + Correlation type: 1 is the plain cross-correlation, 0 the phase + correlation, and values in between the hybrid correlation. Lowering + it equalizes weak and strong disks, which finds many more weak + reflections on a bright background; the reported intensities are + then compressed, so check the effect before using them as weights. + sigma_cc : float | None + Gaussian smoothing of the correlation map (pixels) before peak + finding; merges the speckle of a noisy disk into one maximum. + batch_size : int, optional + Number of patterns per batch. ``None`` (default) picks a size from the + detector dimensions. + progressbar : bool, default=True + If ``True``, show a tqdm progress bar over the full-scan detection. + save_to_gpu : bool, default=True + If ``True`` and :attr:`device` is not the CPU, copy the whole dataset + to the device once and read the batches from that copy, which is + faster and lowers CPU load. If the copy does not fit in device memory, + batches are read from the dataset instead. The copy is kept for later + calls (including :meth:`make_template_from_data`) and is not saved. + + Returns + ------- + Vector + Detected peaks (``[q_row, q_col, intensity]``): shape ``(scan_row, + scan_col)`` for the full scan, or ``(len(positions),)`` for a test run. + """ + if self._template_ft is None: + raise ValueError("Run a make_template_* method before detect_disks().") + + detect_kwargs = dict( + min_abs_intensity=min_abs_intensity, + min_spacing=min_spacing, + edge_boundary=edge_boundary, + subpixel=subpixel, + upsample_factor=upsample_factor, + max_num_peaks=max_num_peaks, + background_sigma=self._resolve_background_sigma(background_sigma), + corr_power=corr_power, + sigma_cc=sigma_cc, + ) + + if save_to_gpu and str(self.device) != "cpu": + if getattr(self, "_gpu_cache", None) is None: + try: + print(f"Loading dataset to {self.device}...", end=" ", flush=True) + self._gpu_cache = torch.as_tensor( + np.asarray(self.dataset.array), + dtype=torch.float32, + device=self.device, + ) + except (RuntimeError, torch.cuda.OutOfMemoryError): + print("out of memory, reading per batch instead.") + self._gpu_cache = None + + if positions is not None: + if len(positions) == 0: + raise ValueError("positions must contain at least one (row, col) to test on.") + coords = [(int(r), int(c)) for r, c in positions] + results = self._detect_positions(coords, detect_kwargs, batch_size, progressbar=False) + return Vector.from_data(results, fields=PEAK_FIELDS, name="bragg_peaks_test") + + scan_r, scan_c = int(self.dataset.shape[0]), int(self.dataset.shape[1]) + coords = list(np.ndindex(scan_r, scan_c)) + results = self._detect_positions( + coords, detect_kwargs, batch_size, progressbar=progressbar + ) + # Store every cell in one pass. Per-cell assignment (peaks[r, c] = arr) + # appends to the backing buffer on every write, while from_data joins all + # cells with a single concatenation. + nested = [results[r * scan_c : (r + 1) * scan_c] for r in range(scan_r)] + peaks = Vector.from_data(nested, fields=PEAK_FIELDS, name="bragg_peaks") + peaks.metadata.update(self._scan_calibration()) + + self.peaks = peaks + self.metadata["detect"] = detect_kwargs + self.compute_bvm() + return peaks + + def correct_peak_origins( + self, + origins: NDArray, + origin_ref: NDArray | tuple[float, float] | None = None, + *, + inplace: bool = False, + ) -> "BraggVectors": + """Shift the detected peak coordinates so all positions share one origin. + + Subtracts each scan position's measured diffraction origin (e.g. the + plane-fitted center-of-mass of the central beam -- the descan) from its + peaks and adds back a common reference ``origin_ref``, so the peak + coordinates from every position live in a single detector frame. The + diffraction data itself is untouched: this calibrates the measurements, + not the images. The Bragg vector map is recomputed from the corrected + peaks. + + By default a corrected *copy* of the workflow is returned and ``self`` + keeps the raw detections, so calling this repeatedly (e.g. re-running a + notebook cell) never double-applies the shift. + + Parameters + ---------- + origins : np.ndarray + ``(scan_row, scan_col, 2)`` per-position diffraction origins in + detector pixels (row, col). + origin_ref : array-like of float, optional + ``(row, col)`` common origin the corrected peaks are referred to. + Defaults to the scan-mean of ``origins``. + inplace : bool, default=False + If ``True``, correct ``self`` instead of returning a corrected copy. + + Returns + ------- + BraggVectors + The workflow holding the corrected peaks (a new instance unless + ``inplace=True``); its :attr:`bvm` is recomputed. + """ + if self.peaks is None: + raise ValueError("Run detect_disks() before correct_peak_origins().") + scan_shape = tuple(int(v) for v in self.dataset.shape[:2]) + origins = np.asarray(origins, dtype=float) + if origins.shape != scan_shape + (2,): + raise ValueError(f"origins must have shape {scan_shape + (2,)}, got {origins.shape}.") + if origin_ref is None: + origin_ref = origins.mean(axis=(0, 1)) + origin_ref = np.asarray(origin_ref, dtype=float).reshape(2) + + if inplace: + bv = self + else: + bv = type(self)(dataset=self.dataset, device=self.device, _token=type(self)._token) + bv._template = self._template + bv._template_ft = self._template_ft + bv.metadata = _copy.deepcopy(self.metadata) + bv.peaks = self.peaks.copy() + + peaks = bv.peaks + # rowwise transform on the flat peak table: one shift per scan cell, + # repeated per detected peak (cells and shifts share raster order) + flat = peaks.numpy().astype(np.float64) + counts = np.asarray(peaks.row_counts(), dtype=int) + shifts = np.repeat(origin_ref[None, :] - origins.reshape(-1, 2), counts, axis=0) + flat[:, :2] += shifts + peaks.set_flattened(flat) + + bv.metadata["origin_correction"] = { + "origin_ref": (float(origin_ref[0]), float(origin_ref[1])), + } + # the fitted origins travel with the peaks: plots that put a pattern + # behind the peaks need them to line the two up + peaks.metadata["origins"] = origins + peaks.metadata["origin_ref"] = (float(origin_ref[0]), float(origin_ref[1])) + bv.compute_bvm() + return bv + + def compute_bvm(self, sampling: float = 1.0) -> Dataset2d: + """Accumulate all detected peaks into a Bragg vector map (intensity histogram). + + Parameters + ---------- + sampling : float, default=1.0 + Reciprocal-space sampling (per pixel) stored on the returned dataset. + + Returns + ------- + Dataset2d + ``(H, W)`` Bragg vector map, also stored on :attr:`bvm`. + """ + if self.peaks is None: + raise ValueError("Run detect_disks() before compute_bvm().") + H, W = (int(self.dataset.shape[-2]), int(self.dataset.shape[-1])) + flat = self.peaks.select_fields("q_row", "q_col", "intensity").numpy().astype(np.float64) + + bvm = np.zeros((H, W), dtype=float) + if flat.shape[0] > 0: + rows = np.clip(np.round(flat[:, 0]).astype(int), 0, H - 1) + cols = np.clip(np.round(flat[:, 1]).astype(int), 0, W - 1) + np.add.at(bvm, (rows, cols), flat[:, 2]) + + self.bvm = Dataset2d.from_array( + bvm, name="bragg_vector_map", sampling=(sampling, sampling), signal_units="intensity" + ) + return self.bvm + + def choose_basis_vectors( + self, + origin: int | tuple[float, float] | NDArray | None = None, + g1: int | tuple[float, float] | NDArray | None = None, + g2: int | tuple[float, float] | NDArray | None = None, + *, + num_candidates: int = 100, + min_spacing: float = 2.0, + min_abs_intensity: float = 0.0, + plot: bool = True, + returnfig: bool = False, + **show_kwargs, + ): + """Select the lattice basis ``(origin, g1, g2)`` from the Bragg vector map. + + Any of ``origin``/``g1``/``g2`` may be given explicitly; the rest are picked + automatically. The origin is the **brightest** candidate peak (the + unscattered central beam). With ``quality = intensity / distance`` rewarding + short, bright vectors, ``g1`` is then the highest-``quality`` peak and ``g2`` + is the highest ``quality * sin^2(theta)`` peak, where ``theta`` is its angle + to ``g1`` (the ``sin^2`` factor vanishes for peaks collinear with ``g1``). + + Each override accepts **either** form, told apart by shape: + + - a scalar **candidate index** (an ``int``) picks one of the numbered + candidate peaks drawn on the plot; + - a **``(row, col)`` vector** is taken literally — an absolute position for + ``origin``, an offset *from the origin* for ``g1``/``g2``. + + With ``plot=True`` (default) the Bragg vector map is shown with the + candidate peaks numbered and the chosen basis overlaid, so the index to + pass back here can be read straight off the figure. + + Parameters + ---------- + origin : int or tuple of float or np.ndarray, optional + Candidate index, or absolute ``(row, col)`` lattice origin. Picked + automatically if omitted. + g1 : int or tuple of float or np.ndarray, optional + Candidate index (vector taken as ``peak - origin``), or a ``(row, col)`` + offset *from the origin*. Picked automatically if omitted. + g2 : int or tuple of float or np.ndarray, optional + Candidate index (vector taken as ``peak - origin``), or a ``(row, col)`` + offset *from the origin*. Picked automatically if omitted. + num_candidates : int, default=100 + Number of brightest candidate peaks to consider (and to number on the + plot). + min_spacing : float, default=2.0 + Minimum spacing in pixels between candidate peaks. + min_abs_intensity : float, default=0.0 + Drop candidate peaks below this absolute intensity. + plot : bool, default=True + If ``True``, show the Bragg vector map with the numbered candidates and + the chosen basis overlaid via + :func:`~quantem.diffraction.bragg_vectors_visualization.plot_basis_vectors`. + returnfig : bool, default=False + If ``True``, return ``(fig, ax)`` from the overlay plot instead of + ``self`` (implies ``plot=True``). + **show_kwargs + Display-scaling options forwarded to the overlay plot's + :func:`~quantem.core.visualization.show_2d` call (e.g. ``norm``, + ``vmin``, ``vmax``, ``cmap``, ``lower_quantile``, ``upper_quantile``). + + Returns + ------- + BraggVectors or tuple + ``self`` for method chaining; or ``(fig, ax)`` when ``returnfig=True``. + """ + if self.bvm is None: + raise ValueError("Run detect_disks()/compute_bvm() before choosing basis vectors.") + + cand_rc, cand_int = self._bvm_candidates(num_candidates, min_spacing, min_abs_intensity) + if cand_rc.shape[0] == 0: + raise RuntimeError("No candidate peaks found in the Bragg vector map.") + # remember the numbered candidates so index_peaks() reuses this exact set. + self.candidates_rc = cand_rc + self.candidates_intensity = cand_int + + def _candidate(idx: int) -> NDArray: + i = int(idx) + if not -cand_rc.shape[0] <= i < cand_rc.shape[0]: + raise IndexError( + f"candidate index {i} out of range for {cand_rc.shape[0]} candidates " + f"(increase num_candidates or loosen min_spacing/min_abs_intensity)." + ) + return cand_rc[i] + + if origin is None: + # candidates are returned brightest-first, so [0] is the central beam. + origin_rc = cand_rc[0] + elif np.ndim(origin) == 0: + origin_rc = _candidate(origin) + else: + origin_rc = np.asarray(origin, dtype=float).reshape(2) + + rel = cand_rc - origin_rc + dist = np.linalg.norm(rel, axis=1) + valid = dist > max(1e-6, min_spacing) + # "shortest/brightest": reward bright peaks close to the origin. + quality = cand_int / (dist + 1e-12) + + if g1 is None: + if not valid.any(): + raise RuntimeError("Could not find a g1 candidate distinct from the origin.") + g1_rc = rel[int(np.argmax(np.where(valid, quality, -np.inf)))] + elif np.ndim(g1) == 0: + g1_rc = _candidate(g1) - origin_rc + else: + g1_rc = np.asarray(g1, dtype=float).reshape(2) + + if g2 is None: + # g2 = highest quality * sin^2(theta) where theta is the angle to g1; + # sin^2 = 1 - cos^2 vanishes for peaks collinear with g1. + g1n = g1_rc / (np.linalg.norm(g1_rc) + 1e-12) + cos = rel @ g1n / (dist + 1e-12) + sin2 = np.clip(1.0 - cos**2, 0.0, 1.0) + g2_score = np.where(valid, quality * sin2, -np.inf) + if g2_score.max() <= 0.0: + raise RuntimeError("Could not find a g2 candidate non-collinear with g1.") + g2_rc = rel[int(np.argmax(g2_score))] + elif np.ndim(g2) == 0: + g2_rc = _candidate(g2) - origin_rc + else: + g2_rc = np.asarray(g2, dtype=float).reshape(2) + + self.origin = np.asarray(origin_rc, dtype=float).reshape(2) + self.g1 = np.asarray(g1_rc, dtype=float).reshape(2) + self.g2 = np.asarray(g2_rc, dtype=float).reshape(2) + + if plot or returnfig: + fig, ax = plot_basis_vectors( + np.asarray(self.bvm.array), + cand_rc, + cand_int, + self.origin, + self.g1, + self.g2, + **show_kwargs, + ) + if returnfig: + return fig, ax + return self + + def index_peaks( + self, + *, + plot: bool = True, + returnfig: bool = False, + **show_kwargs, + ): + """Index the chosen candidate peaks into a reference lattice. + + Assigns integer Miller indices ``(a, b)`` to the numbered candidate peaks + picked in :meth:`choose_basis_vectors` — not the full per-position + detections; that heavy lifting happens in :meth:`fit_lattice`. Each + candidate gets ``[a, b] = round(B^-1 (q - origin))`` with ``B`` columns + ``g1, g2``; when two candidates round to the same ``(a, b)`` only the + brightest is kept. The result is a compact reference lattice stored on + :attr:`reference_ab` / :attr:`reference_qpos` / :attr:`reference_intensity` + which :meth:`fit_lattice` matches against at every scan position. This step + is quick and lightweight. + + With ``plot=True`` (default) the reference lattice is drawn over the Bragg + vector map, each site ringed and labelled with its ``(a, b)`` index. The ring + color encodes how far the picked candidate sits from its ideal lattice site + ``origin + a*g1 + b*g2``, so a mis-picked or duplicate candidate (which rings + far from zero offset) is easy to spot. + + Parameters + ---------- + plot : bool, default=True + If ``True``, show the reference lattice via + :func:`~quantem.diffraction.bragg_vectors_visualization.plot_reference_lattice`. + returnfig : bool, default=False + If ``True``, return ``(fig, ax)`` instead of ``self`` (implies ``plot``). + **show_kwargs + Display-scaling options forwarded to the plot's + :func:`~quantem.core.visualization.show_2d` call (e.g. ``norm``, + ``vmin``, ``vmax``, ``cmap``). + + Returns + ------- + BraggVectors or tuple + ``self`` for method chaining; or ``(fig, ax)`` when ``returnfig=True``. + """ + if self.candidates_rc is None or self.candidates_intensity is None: + raise ValueError("Run choose_basis_vectors() before index_peaks().") + if self.origin is None or self.g1 is None or self.g2 is None: + raise ValueError("Run choose_basis_vectors() before index_peaks().") + + cand_rc = self.candidates_rc + cand_int = self.candidates_intensity + ab = _index_directions(cand_rc, self.origin, self.g1, self.g2) + + # Dedupe by (a, b): candidates are brightest-first, so the first occurrence + # of each index is the brightest -- keep it, drop the dimmer duplicates. + seen: dict[tuple[int, int], int] = {} + for i, (a, b) in enumerate(ab): + key = (int(a), int(b)) + if key not in seen: + seen[key] = i + keep = np.array(sorted(seen.values()), dtype=int) + + self.reference_ab = ab[keep] + self.reference_qpos = cand_rc[keep] + self.reference_intensity = cand_int[keep] + + if not (plot or returnfig): + return self + + fig, ax = plot_reference_lattice( + np.asarray(self.bvm.array), + self.reference_qpos, + self.reference_ab, + self.origin, + self.g1, + self.g2, + **show_kwargs, + ) + if returnfig: + return fig, ax + return self + + def fit_lattice( + self, + min_num_peaks: int = 5, + max_peak_shift: float | None = None, + *, + progressbar: bool = True, + plot: bool = True, + returnfig: bool = False, + ): + """Per-position weighted least-squares fit of the lattice vectors. + + This is the heavy step. At every scan position the detected peaks are + matched to the *ideal* lattice sites from :meth:`index_peaks` — + ``origin + a*g1 + b*g2`` for each reference ``(a, b)`` — keeping a peak only + when it lands within ``max_peak_shift`` of its nearest ideal site (not the + measured candidate position), then ``q = x0 + a*g1 + b*g2`` is fit by + intensity-weighted least squares over the matched peaks. The fitted ``g1``/``g2`` go into :attr:`g1_array`/ + :attr:`g2_array` (shape ``(scan_row, scan_col, 2)``, row/col components); + positions with fewer than ``min_num_peaks`` matched peaks are left ``nan``. + + Two diagnostics are stored per position. :attr:`fit_error` is the RMS fit + residual over the *matched* peaks, in pixels. :attr:`mask_weight` is a lattice + *order parameter* in ``0``–``1``: every detected peak is snapped to the + nearest site of the just-fitted lattice and the intensity-weighted RMS of + those displacements (the zero beam excluded) is normalized by + ``sqrt(|g1 x g2| / 2*pi)`` — the RMS displacement expected from intensity + scattered at random with no lattice — as ``1 - RMS_all / rms_rand`` clipped to + ``[0, 1]``. Because it weighs *all* detected intensity against the lattice + (not just the matched peaks, as :attr:`fit_error` does), a clean single + crystal approaches ``1`` while positions with strong off-lattice intensity (a + second grain, a mis-index, many spurious peaks) fall toward ``0``; weak false + positives carry little intensity and barely move it. Positions with no valid + fit (vacuum, fewer than ``min_num_peaks``) are ``0``. It is the default + reference weighting handed to :meth:`calculate_strain_map`. + + Parameters + ---------- + min_num_peaks : int, default=5 + Minimum number of matched peaks required to fit a position; positions + with fewer are left ``nan``. + max_peak_shift : float, optional + Inclusion radius in pixels: a detected peak is kept only if it lands + within this distance of its nearest *ideal* lattice site + (``origin + a*g1 + b*g2``), excluding peaks that stray too far from where + the best-fit lattice predicts. Defaults to ``0.5 * min(|g1|, |g2|)`` — + half the shorter lattice spacing. + progressbar : bool, default=True + If ``True``, show a tqdm progress bar over the scan positions. + plot : bool, default=True + If ``True``, show the fit diagnostics (mask weight + RMS error) via + :func:`~quantem.diffraction.bragg_vectors_visualization.plot_lattice_fit`. + returnfig : bool, default=False + If ``True``, return ``(fig, ax)`` instead of ``self`` (implies ``plot``). + + Returns + ------- + BraggVectors or tuple + ``self`` for method chaining; or ``(fig, ax)`` when ``returnfig=True``. + """ + if self.peaks is None: + raise ValueError("Run detect_disks() before fit_lattice().") + if self.reference_ab is None or self.reference_qpos is None: + raise ValueError("Run index_peaks() before fit_lattice().") + + ref_ab = self.reference_ab.astype(float) + + # Ideal lattice sites for the reference (a, b) set: origin + a*g1 + b*g2. + # Per-position detections are matched to these *ideal* points (not the measured + # candidate positions), so a peak is included only when it lands within + # max_peak_shift of where the best-fit lattice predicts it should be. + o = np.asarray(self.origin, dtype=float).reshape(2) + g1 = np.asarray(self.g1, dtype=float).reshape(2) + g2 = np.asarray(self.g2, dtype=float).reshape(2) + ideal_qpos = o[None, :] + ref_ab[:, 0:1] * g1[None, :] + ref_ab[:, 1:2] * g2[None, :] + + if max_peak_shift is None: + max_peak_shift = 0.5 * float(min(np.linalg.norm(g1), np.linalg.norm(g2))) + + # Normalization scale for the mask weight: the intensity-weighted RMS + # peak-to-nearest-site displacement expected from intensity scattered at + # random (uniformly over a unit cell of area |g1 x g2|) -- i.e. with no + # lattice order at all. sqrt(A / 2pi) is the equal-area-disk value; the mask + # weight is then 1 - RMS_all / rms_rand, a lattice order parameter in [0, 1]. + cell_area = abs(float(g1[0] * g2[1] - g1[1] * g2[0])) + rms_rand = float(np.sqrt(cell_area / (2.0 * np.pi))) if cell_area > 0 else 1.0 + + scan_r, scan_c = int(self.dataset.shape[0]), int(self.dataset.shape[1]) + g1_array = np.full((scan_r, scan_c, 2), np.nan, dtype=float) + g2_array = np.full((scan_r, scan_c, 2), np.nan, dtype=float) + mask_weight = np.zeros((scan_r, scan_c), dtype=float) + fit_error = np.full((scan_r, scan_c), np.nan, dtype=float) + + fields = self.peaks.fields + i_qr, i_qc = fields.index("q_row"), fields.index("q_col") + i_int = fields.index("intensity") + + coords: Any = list(np.ndindex(scan_r, scan_c)) + if progressbar: + try: + from tqdm.auto import tqdm + + coords = tqdm(coords, desc="fit_lattice", leave=True) + except Exception: + pass + + for r, c in coords: + cell = self.peaks[r, c].numpy().astype(np.float64) + if cell.shape[0] == 0: + continue + qpos = cell[:, [i_qr, i_qc]] + inten = cell[:, i_int] + + # nearest *ideal* lattice site for each detected peak, kept if close enough + d = np.linalg.norm(qpos[:, None, :] - ideal_qpos[None, :, :], axis=2) + nearest = np.argmin(d, axis=1) + matched = d[np.arange(d.shape[0]), nearest] <= max_peak_shift + + if int(matched.sum()) < min_num_peaks: + continue + + beta, rms = _fit_lattice_vectors( + qpos[matched, 0], + qpos[matched, 1], + ref_ab[nearest[matched], 0], + ref_ab[nearest[matched], 1], + inten[matched], + ) + if beta is None: + continue + g1_array[r, c] = beta[1] + g2_array[r, c] = beta[2] + fit_error[r, c] = rms + + # mask weight = lattice "order parameter": snap EVERY detected peak to the + # nearest site of the just-fitted lattice (beta = [x0, g1, g2]) and take + # the intensity-weighted RMS of those displacements, excluding the zero + # beam. Unlike fit_error (matched peaks only), this sees OFF-lattice + # intensity -- a second grain, a mis-index, or many spurious peaks drive + # it up -- while weak false positives, carrying little intensity, barely + # move it. Normalized by rms_rand: 1 (all intensity on the lattice) down + # to 0 (scattered as if there were no lattice). + x0 = beta[0] + mat = np.stack([beta[1], beta[2]]) # (2, 2): rows g1, g2 + ab = np.rint((qpos - x0[None, :]) @ np.linalg.inv(mat)) + sites = x0[None, :] + ab @ mat + disp = np.linalg.norm(qpos - sites, axis=1) + nonzero = ~np.all(ab == 0, axis=1) # drop the central (zero) beam + w = np.clip(inten[nonzero], 0.0, None) + wsum = float(w.sum()) + if wsum > 0: + rms_all = float(np.sqrt(np.sum(w * disp[nonzero] ** 2) / wsum)) + mask_weight[r, c] = float(np.clip(1.0 - rms_all / rms_rand, 0.0, 1.0)) + + self.g1_array = g1_array + self.g2_array = g2_array + self.mask_weight = mask_weight + self.fit_error = fit_error + self.metadata["fit"] = { + "min_num_peaks": int(min_num_peaks), + "max_peak_shift": float(max_peak_shift), + } + + if not (plot or returnfig): + return self + + fig, ax = plot_lattice_fit(mask_weight, fit_error) + if returnfig: + return fig, ax + return self + + def calculate_strain_map( + self, + g1_ref: np.ndarray | None = None, + g2_ref: np.ndarray | None = None, + mask: np.ndarray | None = None, + q_to_r_rotation_ccw_deg: float | None = None, + q_transpose: bool | None = None, + calculation_metric: str = "median", + ) -> StrainMap: + """Build a :class:`StrainMap` from the fitted per-position lattice vectors. + + Parameters + ---------- + g1_ref : np.ndarray, optional + ``(2,)`` reference for the first lattice vector, in detector + ``(row, col)`` pixels (the frame of :attr:`g1_array`). Defaults to the + ``calculation_metric`` over the scan inside :class:`StrainMap`. + g2_ref : np.ndarray, optional + ``(2,)`` reference for the second lattice vector, as ``g1_ref``. + mask : np.ndarray, optional + ``(scan_row, scan_col)`` per-position weighting used when computing the + reference lattice. Defaults to :attr:`mask_weight` from + :meth:`fit_lattice` (the lattice order parameter: how well all detected + intensity snaps to the fitted lattice), so clean single-crystal positions + dominate the reference and positions with off-lattice intensity are + down-weighted. + q_to_r_rotation_ccw_deg : float, optional + Counter-clockwise rotation in degrees from the detector frame to the + scan frame, applied to the lattice vectors before the strain is + computed, so ``e_rr``/``e_cc`` refer to the scan rows and columns. + ``None`` (default) reads ``q_to_r_rotation_ccw_deg`` from the dataset + metadata, else uses 0. + q_transpose : bool, optional + If ``True``, swap the detector row/col axes before the rotation. + ``None`` (default) reads ``q_transpose`` from the dataset metadata, + else uses ``False``. + calculation_metric : {"median", "mean"}, default="median" + Statistic for the automatic reference lattice (weighted by ``mask``). + + Returns + ------- + StrainMap + A strain map initialized from the fitted lattice vectors. The rotation + and transpose used are also stored in :attr:`metadata`. + """ + if self.g1_array is None or self.g2_array is None: + raise ValueError("Run fit_lattice() before calculate_strain_map().") + + if mask is None: + mask = self.mask_weight + + ds_units = None + ds_sampling = None + if hasattr(self.dataset, "units"): + if isinstance(self.dataset.units, (tuple, list)): + ds_units = str(self.dataset.units[0]) + else: + ds_units = str(self.dataset.units) + if hasattr(self.dataset, "sampling"): + if isinstance(self.dataset.sampling, (tuple, list, np.ndarray)): + ds_sampling = float(self.dataset.sampling[0]) + else: + ds_sampling = float(self.dataset.sampling) + + metadata = getattr(self.dataset, "metadata", None) or {} + parent_rot = metadata.get("q_to_r_rotation_ccw_deg", None) + parent_tr = metadata.get("q_transpose", None) + + used_parent = False + if q_to_r_rotation_ccw_deg is None and parent_rot is not None: + q_to_r_rotation_ccw_deg = parent_rot + used_parent = True + if q_transpose is None and parent_tr is not None: + q_transpose = parent_tr + used_parent = True + + if used_parent: + warnings.warn( + "BraggVectors.calculate_strain_map: using Dataset4dstem metadata " + f"(q_to_r_rotation_ccw_deg={q_to_r_rotation_ccw_deg or 0.0}, " + f"q_transpose={q_transpose or False}).", + UserWarning, + ) + + if q_to_r_rotation_ccw_deg is None or q_transpose is None: + q_to_r_rotation_ccw_deg = ( + 0.0 if q_to_r_rotation_ccw_deg is None else q_to_r_rotation_ccw_deg + ) + q_transpose = False if q_transpose is None else q_transpose + warnings.warn( + "BraggVectors.calculate_strain_map: no detector rotation given or in " + f"the dataset metadata; using q_to_r_rotation_ccw_deg=" + f"{q_to_r_rotation_ccw_deg} and q_transpose={q_transpose}.", + UserWarning, + ) + + self.metadata["q_to_r_rotation_ccw_deg"] = float(q_to_r_rotation_ccw_deg) + self.metadata["q_transpose"] = bool(q_transpose) + + return StrainMap( + g1_array=self.g1_array, + g2_array=self.g2_array, + ds_shape=tuple(self.dataset.shape), + real_space=self.real_space, + g1_ref=g1_ref, + g2_ref=g2_ref, + mask=mask, + ds_sampling=ds_sampling, + ds_units=ds_units, + q_to_r_rotation_ccw_deg=float(q_to_r_rotation_ccw_deg), + q_transpose=bool(q_transpose), + calculation_metric=calculation_metric, + ) + + # ---- visualization ---- + + def show_template( + self, + position: tuple[int, int] = (0, 0), + *, + crop_factor: float | None = None, + returnfig: bool = False, + **kwargs, + ): + """Plot the mean diffraction pattern, the template, and one correlation map. + + Parameters + ---------- + position : tuple of int, default=(0, 0) + ``(row, col)`` scan position (scan pixels) whose correlation map is + shown. + crop_factor : float, optional + If given, zoom to a square window of half-width ``crop_factor * radius`` + about the central-beam center, where ``radius`` is the central-beam + radius (the synthetic template radius if known, else estimated from the + mean diffraction pattern). The mean-diffraction and correlation panels are + centered on the beam; the template panel on its own (fftshifted) center. + For example, ``crop_factor=2.0`` shows two beam radii either side of the + center. The window is clamped to the detector, so a large factor shows the + full image. ``None`` (default) shows the full panels. + returnfig : bool, default=False + If ``True``, return the ``(fig, ax)`` for further customization. + **kwargs + Extra keyword arguments forwarded to + :func:`~quantem.diffraction.bragg_vectors_visualization.plot_template`. + + Returns + ------- + tuple + ``(fig, ax)`` when ``returnfig=True``; otherwise nothing. + """ + if self._template is None: + raise ValueError("Run a make_template_* method before show_template().") + r, c = int(position[0]), int(position[1]) + dp_mean = np.asarray(self.dataset.dp_mean.array) + + crop = None + if crop_factor is not None: + # Center the mean-diffraction and correlation panels on the actual + # central-beam position rather than the geometric center (H//2, W//2): + # the unscattered beam -- and the correlation peak that matches it -- are + # generally offset by a few pixels from the detector center. + center, radius_est = estimate_central_beam(dp_mean) + radius = self.metadata.get("template", {}).get("radius") + if radius is None: + radius = radius_est + crop = (float(center[0]), float(center[1]), float(crop_factor) * float(radius)) + + fig, ax = plot_template( + dp_mean, + self.template, + self.correlation_map(r, c), + (r, c), + crop=crop, + **kwargs, + ) + if returnfig: + return fig, ax + + def show_diffraction( + self, + inds: list[tuple[int, int]], + image: np.ndarray | None = None, + *, + ncols: int = 4, + image_kwargs: dict | None = None, + marker_radius: float | None = None, + linewidth: float = 0.5, + sigma_plot: float | None = None, + returnfig: bool = False, + **show_kwargs, + ): + """Preview the diffraction patterns at ``inds``, optionally beside a navigation image. + + No detection is run; this is for choosing scan positions to tune on. The + patterns are tiled ``ncols`` wide and rendered with + :func:`~quantem.core.visualization.show_2d`; the navigation image is styled + independently (real space and reciprocal space rarely want the same + scaling). + + Parameters + ---------- + inds : list of tuple of int + ``(row, col)`` scan positions to preview. + image : np.ndarray or Dataset2d, optional + Real-space navigation image (e.g. a virtual dark-field image) shown on + the left with the positions marked. ``None`` (default) shows only the + diffraction tiles. + ncols : int, default=4 + Number of diffraction tiles per row. + image_kwargs : dict, optional + Keyword arguments styling the navigation image (passed to + :func:`~quantem.core.visualization.show_2d`). + marker_radius : float, optional + Radius in image pixels of the scan-position marker rings. + linewidth : float, default=0.5 + Stroke width of the scan-position markers. + sigma_plot : float, optional + Gaussian blur (sigma) applied to the *displayed* patterns only. + returnfig : bool, default=False + If ``True``, return the ``(fig, ax)`` for further customization. + **show_kwargs + Extra keyword arguments (e.g. ``norm``, ``cmap``, ``cbar``, ``axsize``) + styling the diffraction tiles via + :func:`~quantem.core.visualization.show_2d`. + + Returns + ------- + tuple + ``(fig, ax)`` when ``returnfig=True``; otherwise nothing. + """ + if image is None: + image_arr = None + image_title = "navigation image" + else: + image_title = getattr(image, "name", None) or "navigation image" + image_arr = np.asarray(image.array if hasattr(image, "array") else image) + dps = [np.asarray(self.dataset.array[r, c], dtype=float) for r, c in inds] + fig, ax = plot_diffraction_grid( + image_arr, + dps, + inds, + ncols=ncols, + image_title=image_title, + image_kwargs=image_kwargs, + marker_radius=marker_radius, + linewidth=linewidth, + sigma_plot=sigma_plot, + **show_kwargs, + ) + if returnfig: + return fig, ax + + def show_detection( + self, + positions: list[tuple[int, int]] | None = None, + *, + min_abs_intensity: float = 0.0, + min_spacing: float = 0.0, + edge_boundary: int = 1, + subpixel: str = "upsample", + upsample_factor: int = 16, + max_num_peaks: int = 1000, + background_sigma: float | str | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, + image: np.ndarray | None = None, + peak_radius: float = 6.0, + marker_radius: float | None = None, + linewidth: float = 1.0, + sigma_plot: float | None = None, + image_kwargs: dict | None = None, + returnfig: bool = False, + **plot_kwargs, + ): + """Detect on a few patterns and overlay the peaks, for tuning hyperparameters. + + The detection keywords match :meth:`detect_disks`. The workflow state + (:attr:`peaks`/:attr:`bvm`) is left untouched, so this is safe to re-run + while tuning; for the raw peaks, call :meth:`detect_disks` with the same + ``positions``. + + Parameters + ---------- + positions : list of tuple of int, optional + ``(row, col)`` scan positions to detect on. ``None`` (default) + auto-samples four positions spread across the scan. + min_abs_intensity : float, default=0.0 + Drop correlation peaks below this absolute intensity. + min_spacing : float, default=0.0 + Minimum spacing in pixels between kept peaks. + edge_boundary : int, default=1 + Width in pixels of the border in which peaks are ignored. + subpixel : {"none", "parabolic", "upsample"}, default="upsample" + Subpixel refinement mode; see :meth:`detect_disks`. + upsample_factor : int, default=16 + Upsampling factor for the ``"upsample"`` subpixel refinement. + max_num_peaks : int, default=1000 + Maximum number of peaks to keep per pattern. + background_sigma : float | "auto" | None, default=None + Width in pixels of a Fourier high-pass applied to the + cross-correlation before peak finding: a copy of the correlation + map smoothed by a Gaussian of this width is subtracted from it. + The filter is isotropic in the Fourier domain and needs no origin, + so it is not a radial background fit. Its purpose is the negative + moat a zero-sum template leaves around the bright unscattered + beam, which pushes weak disk peaks below zero where the + correlation clamp erases them. Set it a little wider than a disk + so the disk-scale peaks pass untouched; "auto" uses twice the + central beam radius. + corr_power : float, default=1.0 + Correlation type: 1 is the plain cross-correlation, 0 the phase + correlation, and values in between the hybrid correlation. Lowering + it equalizes weak and strong disks, which finds many more weak + reflections on a bright background; the reported intensities are + then compressed, so check the effect before using them as weights. + sigma_cc : float | None + Gaussian smoothing of the correlation map (pixels) before peak + finding; merges the speckle of a noisy disk into one maximum. + image : np.ndarray or Dataset2d, optional + Real-space navigation image (e.g. a virtual dark-field image) shown on + the left with the chosen positions marked. ``None`` (default) shows only + the diffraction tiles. + peak_radius : float, default=6.0 + Radius in diffraction pixels of the cyan rings drawn at detected peaks. + marker_radius : float, optional + Radius in image pixels of the scan-position marker rings. + linewidth : float, default=1.0 + Stroke width of the peak and scan-position rings. + sigma_plot : float, optional + Gaussian blur (sigma) applied to the *displayed* patterns only; + detection still uses the raw data. + image_kwargs : dict, optional + Keyword arguments styling the navigation image (passed to + :func:`~quantem.core.visualization.show_2d`). + returnfig : bool, default=False + If ``True``, return the ``(fig, ax)`` for further customization. + **plot_kwargs + Extra keyword arguments (e.g. ``ncols``, ``norm``, ``cmap``, ``cbar``, + ``axsize``) styling the diffraction tiles via + :func:`~quantem.core.visualization.show_2d`. + + Returns + ------- + tuple + ``(fig, ax)`` when ``returnfig=True``; otherwise nothing. + """ + if positions is None: + positions = self._sample_positions() + sub = self.detect_disks( + positions=positions, + min_abs_intensity=min_abs_intensity, + min_spacing=min_spacing, + edge_boundary=edge_boundary, + subpixel=subpixel, + upsample_factor=upsample_factor, + max_num_peaks=max_num_peaks, + background_sigma=background_sigma, + corr_power=corr_power, + sigma_cc=sigma_cc, + progressbar=False, + ) + if image is None: + image_arr = None + image_title = "virtual image" + else: + image_title = getattr(image, "name", None) or "virtual image" + image_arr = np.asarray(image.array if hasattr(image, "array") else image) + dps = [np.asarray(self.dataset.array[r, c], dtype=float) for r, c in positions] + peaks = [sub[i].numpy().astype(np.float64) for i in range(len(positions))] + fig, ax = plot_detection( + image_arr, + dps, + peaks, + positions, + peak_radius=peak_radius, + marker_radius=marker_radius, + linewidth=linewidth, + sigma_plot=sigma_plot, + image_title=image_title, + image_kwargs=image_kwargs, + **plot_kwargs, + ) + if returnfig: + return fig, ax + + def peak_histogram(self, *, returnfig: bool = False, **kwargs): + """Plot the Bragg vector map beside the per-position peak count. + + The Bragg vector map is the 2-D histogram of all detected peak positions + (intensity-weighted) accumulated over the scan; the second panel shows the + number of peaks detected at each scan position. + + Parameters + ---------- + returnfig : bool, default=False + If ``True``, return the ``(fig, ax)`` for further customization. + **kwargs + Extra keyword arguments forwarded to + :func:`~quantem.diffraction.bragg_vectors_visualization.plot_bvm`. + + Returns + ------- + tuple + ``(fig, ax)`` when ``returnfig=True``; otherwise nothing. + """ + if self.peaks is None or self.bvm is None: + raise ValueError("Run detect_disks() before peak_histogram().") + scan_r, scan_c = int(self.dataset.shape[0]), int(self.dataset.shape[1]) + # row_counts() reads cell_lengths directly; far cheaper than building a + # per-cell view (self.peaks[r, c]) just to read its row count. + counts = np.asarray(self.peaks.row_counts(), dtype=int).reshape(scan_r, scan_c) + fig, ax = plot_bvm(np.asarray(self.bvm.array), counts, **kwargs) + if returnfig: + return fig, ax + + # ---- helpers ---- + + # ------------------------------------------------------------------ + # detector calibration + # ------------------------------------------------------------------ + + def measure_origins( + self, + search_radius: float = 6.0, + robust: bool = True, + plot: bool = False, + center=None, + ): + """Fit the diffraction origin at every probe position. + + The direct beam wanders with the probe (descan). The brightest peak + within `search_radius` of `center` gives the origin at each + position, and a plane fit over the scan smooths the result. + + Parameters + ---------- + search_radius : float, default=6.0 + Radius in detector pixels searched around `center`. + center : tuple of float, optional + ``(row, col)`` to search around; defaults to the detector + centre. Pass the beam position from + :func:`~quantem.diffraction.disk_detection.estimate_central_beam` + when the beam is not centred on the detector. + robust : bool, default=True + Reject outliers before the plane fit. + plot : bool, default=False + Show the measured and fitted origins. + + Returns + ------- + np.ndarray + ``(scan_row, scan_col, 2)`` origins in detector pixels. + """ + from quantem.diffraction import calibration + + return calibration.measure_origins( + self, search_radius=search_radius, robust=robust, plot=plot, center=center + ) + + def plot_origin_fit(self, origins, search_radius: float = 6.0, center=None): + """Measured origins, the plane fit, and their residual. + + Parameters + ---------- + origins : np.ndarray + ``(scan_row, scan_col, 2)`` origins from :meth:`measure_origins`. + search_radius : float, default=6.0 + The radius used to measure them. + center : tuple of float, optional + The centre used to measure them. + + Returns + ------- + tuple + ``(fig, axs)``. + """ + from quantem.diffraction import calibration + + return calibration.plot_origin_fit( + self, origins, search_radius=search_radius, center=center + ) + + def measure_scan_rotation( + self, + origins=None, + mask_radius: float | None = None, + plot: bool = False, + ) -> float: + """Rotation between the detector and the scan, from the direct beam. + + The center of mass of the direct beam traces the projected potential + gradient over the scan, which is a curl-free field in the scan frame. + Rotating the detector axes until the curl vanishes recovers the angle. + It sets the in-plane orientations and nothing else: zone axes, phases + and the out-of-plane maps do not depend on it. + + The curl is unchanged by a 180 degree rotation, so the answer is this + angle or this angle plus 180. + + Parameters + ---------- + origins : np.ndarray, optional + ``(scan_row, scan_col, 2)`` origins from :meth:`measure_origins`; + defaults to the detector centre. + mask_radius : float, optional + Radius in detector pixels around the origin used for the center + of mass, which keeps the Bragg disks out of it. Set it a little + beyond the direct beam. + plot : bool, default=False + Show the curl and divergence against the trial angle. + + Returns + ------- + float + Counter-clockwise rotation in degrees, in [0, 180). + """ + from quantem.diffraction import calibration + + out = calibration.measure_scan_rotation( + self.dataset, origins=origins, mask_radius=mask_radius, plot=plot + ) + return out[0] if isinstance(out, tuple) else out + + def calibrate( + self, + crystal, + pixel_size_guess: float, + rotation_ccw_deg: float = 0.0, + plot: bool = False, + **kwargs, + ): + """Measure the reciprocal pixel size and the elliptic distortion. + + Matches the radial distribution of the detected peaks against the + ring positions of a known crystal: a coarse scan over broadened + rings fixes the ring assignment even from a poor starting guess, + then the ellipse and the scale are refined in turn, and every ring is + checked separately. The result carries the evidence behind it and can + be saved and applied to another dataset, which is the route for a + strained sample: measure on a standard such as nanocrystalline gold, + save, and apply there. + + Call this on origin-corrected peaks -- the rings must be concentric + before their radii mean anything. + + Parameters + ---------- + crystal : Crystal or list of Crystal + Reference structure(s) with structure factors calculated. + pixel_size_guess : float + Starting reciprocal pixel size, 1/Angstroms per detector pixel. + Recovered from a factor of two out in either direction. + rotation_ccw_deg : float, default=0.0 + Counter-clockwise diffraction-to-scan rotation in degrees, recorded + on the calibration for later use. + plot : bool, default=False + Show the ring comparison before and after, for the first + reference crystal. :meth:`CrystalMap.plot_calibration` checks + every candidate phase afterwards. + **kwargs + Further arguments of + :func:`~quantem.diffraction.calibration.calibrate`, e.g. + `zone_axis` to fit only the rings of one zone, `fit_ellipse`, + `scale_search`, `k_min`, `k_max`, `marker_size` and `figsize`. + + Returns + ------- + DiffractionCalibration + The measured calibration, also stored on :attr:`calibration`. + Save it with ``cal.save(path)`` and reload it with + :func:`~quantem.core.io.serialize.load`. + + Raises + ------ + ValueError + If no peaks have been detected. + """ + from quantem.diffraction import calibration as _cal + + if self.peaks is None: + raise ValueError("Run detect_disks() before calibrate().") + cal = _cal.calibrate( + self.peaks, + crystal, + pixel_size_guess, + rotation_ccw_deg=rotation_ccw_deg, + plot=plot, + **kwargs, + ) + self.calibration = cal + return cal + + def apply_calibration(self, calibration=None, name: str = "bragg_peaks_calibrated"): + """Calibrated peaks in 1/Angstroms from the detected pixel peaks. + + Parameters + ---------- + calibration : DiffractionCalibration, optional + The calibration to apply. None (default) uses the one measured by + :meth:`calibrate`, so a calibration loaded from a standard is + passed here instead. + name : str, default="bragg_peaks_calibrated" + Name of the returned Vector. + + Returns + ------- + Vector + Peaks in 1/Angstroms, carrying the pixel size, the fitted origins + and the scan calibration, so plots downstream need none of them + passed. + + Raises + ------ + ValueError + If no peaks have been detected, or no calibration is available. + """ + if self.peaks is None: + raise ValueError("Run detect_disks() before apply_calibration().") + cal = calibration if calibration is not None else getattr(self, "calibration", None) + if cal is None: + raise ValueError( + "No calibration available: run calibrate() first, or pass one loaded from a standard." + ) + return cal.apply(self.peaks, name=name) + + def _scan_calibration(self) -> dict: + """Calibration of the dataset, to travel with the peaks. + + The scan step and the detector pixel size are properties of the + measurement, so they are set once on the dataset and carried by the + peaks rather than passed to every plot. Axes still in pixels record + nothing for that half. + """ + out: dict = {} + sampling = np.atleast_1d(np.asarray(self.dataset.sampling, dtype=float)) + units = [str(u).strip("b'\"") for u in self.dataset.units] + + scan_unit = units[0] if units else "" + if scan_unit.lower() not in ("pixels", "px", "pixel", ""): + out["scan_sampling"] = tuple(float(v) for v in sampling[:2]) + out["scan_units"] = scan_unit + + if sampling.size >= 4 and len(units) >= 4: + dp_unit = units[2] + if dp_unit.lower() not in ("pixels", "px", "pixel", ""): + out["pixel_size"] = float(sampling[2]) + out["pixel_size_units"] = dp_unit + return out + + def _resolve_background_sigma(self, background_sigma: float | str | None) -> float | None: + """Resolve the ``background_sigma`` argument to a value in pixels. + + ``None`` (the default everywhere) disables the background subtraction. + ``"auto"`` maps to twice the central-beam radius: wide enough that the disk-scale correlation peaks pass + untouched, narrow enough to remove the zero-sum template's negative + moat around a bright unscattered beam -- which otherwise pushes weak + disk peaks below zero, where the correlation clamp erases them before + peak finding. A float sets the width explicitly. + + Parameters + ---------- + background_sigma : float, "auto" or None + The value passed by the caller. + + Returns + ------- + float or None + Width in pixels, or ``None`` for no background subtraction. + """ + if background_sigma is None: + return None + if isinstance(background_sigma, str): + if background_sigma != "auto": + raise ValueError("background_sigma must be a float, None, or 'auto'.") + radius = self.metadata.get("template", {}).get("radius") + if radius is None: + dp_mean = np.asarray(self.dataset.dp_mean.array) + _, radius = estimate_central_beam(dp_mean) + return 2.0 * float(radius) + return float(background_sigma) + + def _set_template( + self, + probe: torch.Tensor, + center: tuple[float, float] | None, + subtract_mean: bool, + ) -> None: + """Validate the probe shape, then store the template and its conjugate FT. + + Parameters + ---------- + probe : torch.Tensor + Probe image; must match the diffraction-pattern shape. + center : tuple of float or None + ``(row, col)`` probe center rolled to the origin, or ``None`` for the + geometric center. + subtract_mean : bool + If ``True``, make the template zero-sum. + """ + dp_shape = tuple(self.dataset.shape[-2:]) + probe_t = torch.as_tensor(probe, dtype=torch.float, device=self.device) + if tuple(probe_t.shape) != dp_shape: + raise ValueError( + f"probe shape {tuple(probe_t.shape)} does not match diffraction pattern " + f"shape {dp_shape}." + ) + self._template = make_template(probe_t, center=center, subtract_mean=subtract_mean) + self._template_ft = template_fourier(self._template) + + def _sample_positions(self) -> list[tuple[int, int]]: + """Four scan positions spread across the field (quadrant centers). + + Returns + ------- + list of tuple of int + Up to four ``(row, col)`` scan positions at the quadrant centers. + """ + R, C = int(self.dataset.shape[0]), int(self.dataset.shape[1]) + rs = sorted({min(max(R // 4, 0), R - 1), min(max(3 * R // 4, 0), R - 1)}) + cs = sorted({min(max(C // 4, 0), C - 1), min(max(3 * C // 4, 0), C - 1)}) + return [(r, c) for r in rs for c in cs] + + def _detect_positions( + self, + coords: list[tuple[int, int]], + detect_kwargs: dict[str, Any], + batch_size: int | None, + *, + progressbar: bool, + ) -> list[NDArray]: + """Batched detection over a list of ``(row, col)`` scan positions. + + Patterns are stacked into chunks of ``batch_size`` and passed to + :func:`~quantem.diffraction.disk_detection.detect_disks_batch`. + + Parameters + ---------- + coords : list of tuple of int + ``(row, col)`` scan positions to detect on. + detect_kwargs : dict + Keyword arguments forwarded to + :func:`~quantem.diffraction.disk_detection.detect_disks_batch`. + batch_size : int or None + Number of patterns per batch. ``None`` picks a size from the detector + dimensions. + progressbar : bool + If ``True``, show a tqdm progress bar over the patterns. + + Returns + ------- + list of np.ndarray + One ``(M, 3)`` array of ``[q_row, q_col, intensity]`` per position, in + ``coords`` order. + """ + H, W = int(self.dataset.shape[-2]), int(self.dataset.shape[-1]) + if batch_size is None: + batch_size = int(min(1024, max(1, 16_000_000 // (H * W)))) + + use_cache = getattr(self, "_gpu_cache", None) is not None + + it = range(0, len(coords), batch_size) + if progressbar: + try: + from tqdm.auto import tqdm + + bar = tqdm(total=len(coords), desc="detect_disks", leave=True) + except Exception: + bar = None + else: + bar = None + + results: list[NDArray] = [] + for start in it: + chunk = coords[start : start + batch_size] + rows = [r for r, c in chunk] + cols = [c for r, c in chunk] + if use_cache: + dps = self._gpu_cache[rows, cols] + else: + # one fancy-indexed read per batch rather than one per pattern + dps_np = np.asarray(self.dataset.array[np.asarray(rows), np.asarray(cols)]) + dps = torch.as_tensor(dps_np, dtype=torch.float32, device=self.device) + + out = detect_disks_batch(dps, self._template_ft, **detect_kwargs) + results.extend( + arr if arr.shape[0] else np.empty((0, len(PEAK_FIELDS)), dtype=float) + for arr in out + ) + if bar is not None: + bar.update(len(chunk)) + + if bar is not None: + bar.close() + return results + + def _bvm_candidates( + self, num_candidates: int, min_spacing: float, min_abs_intensity: float + ) -> tuple[NDArray, NDArray]: + """Find the brightest, well-separated local maxima in the Bragg vector map. + + Parameters + ---------- + num_candidates : int + Maximum number of candidate peaks to return. + min_spacing : float + Minimum spacing in pixels between candidates. + min_abs_intensity : float + Drop candidates below this absolute intensity. + + Returns + ------- + cand_rc : np.ndarray + ``(N, 2)`` ``[row, col]`` candidate positions, brightest first. + cand_int : np.ndarray + ``(N,)`` candidate intensities. + """ + from quantem.diffraction.disk_detection import _filter_maxima, _local_maxima + + bvm = torch.as_tensor(self.bvm.array, dtype=torch.float) + peaks = _local_maxima(bvm, edge_boundary=1) + peaks = _filter_maxima(peaks, min_abs_intensity, min_spacing, num_candidates) + arr = peaks.detach().cpu().numpy() + return arr[:, :2].astype(float), arr[:, 2].astype(float) + + +def _fractional_indices(q: NDArray, origin: NDArray, g1: NDArray, g2: NDArray) -> NDArray: + """Continuous (unrounded) ``(a, b)`` lattice coordinates of peaks ``q``. + + Solves ``B [a, b]^T = q - origin`` (least squares), with + ``B = [[g1_row, g2_row], [g1_col, g2_col]]``. + + Parameters + ---------- + q : np.ndarray + ``(N, 2)`` ``[row, col]`` peak positions. + origin : np.ndarray + ``(2,)`` lattice origin ``[row, col]``. + g1 : np.ndarray + ``(2,)`` first lattice vector ``[row, col]``. + g2 : np.ndarray + ``(2,)`` second lattice vector ``[row, col]``. + + Returns + ------- + np.ndarray + ``(N, 2)`` float array of continuous ``[a, b]`` coordinates. + """ + if q.shape[0] == 0: + return np.empty((0, 2), dtype=float) + beta = np.array([[g1[0], g2[0]], [g1[1], g2[1]]], dtype=float) + alpha = (q - origin[None, :]).T # (2, N) + return np.linalg.lstsq(beta, alpha, rcond=None)[0].T # (N, 2) + + +def _index_directions(q: NDArray, origin: NDArray, g1: NDArray, g2: NDArray) -> NDArray: + """Integer Miller indices for peaks ``q`` against basis ``(g1, g2)`` about ``origin``. + + Rounds the continuous coordinates from :func:`_fractional_indices`. + + Parameters + ---------- + q : np.ndarray + ``(N, 2)`` ``[row, col]`` peak positions. + origin : np.ndarray + ``(2,)`` lattice origin ``[row, col]``. + g1 : np.ndarray + ``(2,)`` first lattice vector ``[row, col]``. + g2 : np.ndarray + ``(2,)`` second lattice vector ``[row, col]``. + + Returns + ------- + np.ndarray + ``(N, 2)`` int array of ``[a, b]`` Miller indices. + """ + if q.shape[0] == 0: + return np.empty((0, 2), dtype=int) + return np.round(_fractional_indices(q, origin, g1, g2)).astype(int) + + +def _fit_lattice_vectors( + q_row: NDArray, + q_col: NDArray, + a: NDArray, + b: NDArray, + intensity: NDArray, +) -> tuple[NDArray | None, float]: + """Intensity-weighted lattice fit ``q = x0 + a*g1 + b*g2`` for one pattern. + + Parameters + ---------- + q_row : np.ndarray + ``(N,)`` peak row positions. + q_col : np.ndarray + ``(N,)`` peak column positions. + a : np.ndarray + ``(N,)`` Miller index along ``g1``. + b : np.ndarray + ``(N,)`` Miller index along ``g2``. + intensity : np.ndarray + ``(N,)`` peak intensities, used as fit weights (``sqrt`` of the clamped + intensity). + + Returns + ------- + beta : np.ndarray or None + ``(3, 2)`` fit ``[x0; g1; g2]`` (row/col components per row), or ``None`` if + the fit is rank-deficient (e.g. all peaks share one lattice row). + rms : float + RMS fit residual in pixels (``nan`` when ``beta`` is ``None``). + """ + design = np.stack([np.ones_like(a), a, b], axis=1) # (N, 3) + target = np.stack([q_row, q_col], axis=1) # (N, 2) + w = np.sqrt(np.clip(intensity, 0.0, None))[:, None] + + if np.linalg.matrix_rank(design * w) < 3: + return None, float("nan") + + beta = np.linalg.lstsq(design * w, target * w, rcond=None)[0] # (3, 2): x0, g1, g2 + resid = target - design @ beta + rms = float(np.sqrt(np.mean(np.sum(resid**2, axis=1)))) + return beta, rms diff --git a/src/quantem/diffraction/bragg_vectors_visualization.py b/src/quantem/diffraction/bragg_vectors_visualization.py new file mode 100644 index 000000000..a105909ac --- /dev/null +++ b/src/quantem/diffraction/bragg_vectors_visualization.py @@ -0,0 +1,801 @@ +from __future__ import annotations + +import matplotlib.pyplot as plt +import numpy as np + + +def plot_template( + dp_mean: np.ndarray, + template: np.ndarray, + corr_map: np.ndarray, + position: tuple[int, int], + *, + crop: tuple[float, float, float] | None = None, + figsize: tuple[float, float] = (13, 4), +): + """Mean diffraction pattern, the (centered) template, and one correlation map. + + Parameters + ---------- + dp_mean : np.ndarray + Mean diffraction pattern. + template : np.ndarray + The correlation template, centered for display. + corr_map : np.ndarray + Correlation map computed at ``position``. + position : tuple of int + ``(row, col)`` scan position the correlation map was computed at. + crop : tuple of float, optional + ``(center_row, center_col, half_width)`` zoom window (in pixels). The mean + diffraction and correlation panels are centered on ``(center_row, + center_col)`` -- the central-beam position -- while the template panel is + centered on its own array center (it is displayed fftshifted to there). The + view spans ``half_width`` either side of the center, clamped to each image's + bounds, so an over-large ``half_width`` just shows the full image. ``None`` + (default) shows the full panels. + figsize : tuple of float, default=(13, 4) + Figure size in inches. + + Returns + ------- + tuple + ``(fig, ax)`` with ``ax`` a length-3 array of axes. + """ + fig, ax = plt.subplots(1, 3, figsize=figsize) + ax[0].imshow(dp_mean, cmap="gray") + ax[0].set_title("mean diffraction") + ax[1].imshow(template, cmap="gray") + ax[1].set_title("template (centered)") + ax[2].imshow(corr_map, cmap="viridis") + ax[2].set_title(f"correlation @ {tuple(position)} (real space pixels)") + for a in ax: + a.set_xticks([]) + a.set_yticks([]) + if crop is not None: + cr, cc, hw = float(crop[0]), float(crop[1]), float(crop[2]) + th, tw = template.shape[:2] + # The mean-diffraction beam and the correlation peak sit at the beam center + # (cr, cc); the displayed template is fftshifted to its own array center, so + # zoom that panel about its center instead. + panel_centers = ((cr, cc), (th / 2.0, tw / 2.0), (cr, cc)) + for a, img, (ecr, ecc) in zip(ax, (dp_mean, template, corr_map), panel_centers): + h, w = img.shape[:2] + a.set_xlim(max(ecc - hw, -0.5), min(ecc + hw, w - 0.5)) + a.set_ylim(min(ecr + hw, h - 0.5), max(ecr - hw, -0.5)) + fig.tight_layout() + return fig, ax + + +def _mark_positions(ax, positions, *, radius=None, linewidth=0.5): + """Overlay numbered red markers at each ``(row, col)`` scan position. + + Parameters + ---------- + ax : matplotlib.axes.Axes + Axis to draw the markers on. + positions : sequence of tuple of int + ``(row, col)`` scan positions to mark, numbered in order. + radius : float, optional + Marker radius in image pixels. If given, each marker is a ring of that + radius drawn as a circle patch in data coordinates; if ``None``, a fixed + screen-size scatter marker is used instead. + linewidth : float, default=0.5 + Ring stroke width. + """ + from matplotlib.patches import Circle + + for i, (r, c) in enumerate(positions): + if radius is None: + ax.scatter(c, r, s=70, facecolors="none", edgecolors="red", linewidths=linewidth) + else: + ax.add_patch( + Circle((c, r), radius=radius, fill=False, edgecolor="red", linewidth=linewidth) + ) + ax.annotate( + str(i), + (c, r), + color="red", + fontsize=11, + fontweight="bold", + xytext=(4, 4), + textcoords="offset points", + ) + + +def _blur(dp: np.ndarray, sigma: float | None) -> np.ndarray: + """Gaussian-blur a diffraction pattern for display only (passthrough if ``sigma`` falsy). + + Smoothing the *displayed* pattern (the detection still runs on the raw data) + makes it easier to judge by eye which correlation peaks sit on real disks and + which are noise. + + Parameters + ---------- + dp : np.ndarray + Diffraction pattern to blur. + sigma : float or None + Gaussian blur width in pixels. If falsy (``None``, ``0``, or negative), + ``dp`` is returned unchanged. + + Returns + ------- + np.ndarray + The blurred pattern, or ``dp`` unchanged when ``sigma`` is falsy. + """ + if not sigma or sigma <= 0: + return dp + from scipy.ndimage import gaussian_filter + + return gaussian_filter(np.asarray(dp, dtype=float), float(sigma)) + + +def _grid_axes(n, ncols, *, show_image, axsize=None, figsize=None): + """Figure with an optional left nav-image axis and a right ``nrows x ncols`` tile grid. + + With ``show_image`` a navigation-image axis is placed to the left of the tile + grid; otherwise the figure is just the grid. + + Parameters + ---------- + n : int + Number of tiles to lay out. + ncols : int + Number of tile columns, clamped to ``n``. + show_image : bool + If ``True``, add a navigation-image axis to the left of the tile grid. + axsize : tuple of float, optional + Per-tile size in inches, ``(w, h)``. Sizes the figure when ``figsize`` is + not given, so a larger ``axsize`` zooms every tile. + figsize : tuple of float, optional + Explicit figure size in inches; overrides the ``axsize``-derived size. + + Returns + ------- + tuple + ``(fig, ax_image, dp_axes, ncols)`` where ``ax_image`` is ``None`` when + ``show_image`` is ``False``, ``dp_axes`` is an ``(nrows, ncols)`` object + array of axes (trailing unused tiles already turned off), and ``ncols`` is + clamped to ``n``. + """ + ncols = max(1, min(ncols, n)) + nrows = int(np.ceil(n / ncols)) + + tile_w, tile_h = (float(axsize[0]), float(axsize[1])) if axsize is not None else (3.0, 3.2) + nav_w = (tile_w if axsize is not None else 4.0) if show_image else 0.0 + if figsize is None: + figsize = (nav_w + tile_w * ncols, max(tile_h, tile_h * nrows)) + + fig = plt.figure(figsize=figsize) + if show_image: + outer = fig.add_gridspec(1, 2, width_ratios=[nav_w, tile_w * ncols], wspace=0.12) + ax_image = fig.add_subplot(outer[0, 0]) + grid = outer[0, 1].subgridspec(nrows, ncols, wspace=0.08, hspace=0.2) + else: + ax_image = None + grid = fig.add_gridspec(nrows, ncols, wspace=0.08, hspace=0.2) + + dp_axes = np.empty((nrows, ncols), dtype=object) + for i in range(nrows * ncols): + a = fig.add_subplot(grid[i // ncols, i % ncols]) + dp_axes[i // ncols, i % ncols] = a + if i >= n: + a.axis("off") + return fig, ax_image, dp_axes, ncols + + +def plot_diffraction_grid( + image: np.ndarray | None, + dps: list[np.ndarray], + positions: list[tuple[int, int]], + *, + ncols: int = 4, + image_title: str = "navigation image", + image_kwargs: dict | None = None, + marker_radius: float | None = None, + linewidth: float = 0.5, + sigma_plot: float | None = None, + axsize: tuple[float, float] | None = None, + figsize: tuple[float, float] | None = None, + **show_kwargs, +): + """A tiled grid of the diffraction patterns at ``positions``, optionally beside a nav image. + + When ``image`` is given (e.g. a virtual dark-field image) it is drawn on the + left with each scan position marked as a numbered red marker at ``(x=col, + y=row)``; pass ``image=None`` to omit it and show only the grid. Right: the + diffraction patterns tiled ``ncols`` wide. Both are rendered with + :func:`~quantem.core.visualization.show_2d`, and the navigation image and the + diffraction tiles take separate, independent styling (``image_kwargs`` vs + ``show_kwargs``). + + Parameters + ---------- + image : np.ndarray or None + Navigation image (e.g. a virtual dark-field image) drawn at left, with the + scan positions marked. Pass ``None`` to show only the diffraction grid. + dps : list of np.ndarray + Diffraction patterns to tile, one per entry in ``positions``. + positions : list of tuple of int + ``(row, col)`` scan positions, used for the tile titles and the nav-image + markers. + ncols : int, default=4 + Number of columns in the diffraction-pattern tile grid. + image_title : str, default="navigation image" + Title for the navigation image. + image_kwargs : dict, optional + Extra keyword arguments (e.g. ``norm``, ``cmap``, ``scalebar``) for the + navigation image's :func:`show_2d` call, styling it independently of the + tiles. + marker_radius : float, optional + Scan-position marker radius in image pixels; ``None`` uses a fixed + screen-size marker. + linewidth : float, default=0.5 + Stroke width of the scan-position markers. + sigma_plot : float, optional + Gaussian blur width (pixels) applied to the *displayed* patterns only (the + data is untouched) to ease judging real features. + axsize : tuple of float, optional + Per-tile size in inches; a larger value zooms the tiles. + figsize : tuple of float, optional + Explicit figure size in inches. + **show_kwargs + Forwarded to :func:`show_2d` for the diffraction-pattern tiles (e.g. + ``norm``, ``cmap``, ``cbar``). + + Returns + ------- + tuple + ``(fig, (ax_image, dp_axes))`` with ``ax_image`` ``None`` when no image is + shown. + """ + from quantem.core.visualization import show_2d + + n = len(positions) + show_image = image is not None + fig, ax_image, dp_axes, ncols = _grid_axes( + n, ncols, show_image=show_image, axsize=axsize, figsize=figsize + ) + + if show_image: + image_show_kwargs = {"cmap": "gray", "title": image_title, **(image_kwargs or {})} + show_2d(np.asarray(image), figax=(fig, ax_image), **image_show_kwargs) + _mark_positions(ax_image, positions, radius=marker_radius, linewidth=linewidth) + + user_title = show_kwargs.pop("title", None) + for i in range(n): + r, c = positions[i] + title = user_title if user_title is not None else f"{i}: ({r},{c})" + show_2d( + _blur(dps[i], sigma_plot), + figax=(fig, dp_axes[i // ncols, i % ncols]), + title=title, + **show_kwargs, + ) + return fig, (ax_image, dp_axes) + + +def plot_detection( + image: np.ndarray | None, + dps: list[np.ndarray], + peaks: list[np.ndarray], + positions: list[tuple[int, int]], + *, + ncols: int = 4, + peak_radius: float = 6.0, + marker_radius: float | None = None, + linewidth: float = 0.5, + sigma_plot: float | None = None, + image_title: str = "virtual image", + image_kwargs: dict | None = None, + axsize: tuple[float, float] | None = None, + figsize: tuple[float, float] | None = None, + **show_kwargs, +): + """The diffraction patterns with detected peaks overlaid, optionally beside a nav image. + + When ``image`` is given (e.g. a virtual dark-field image) it is drawn on the + left with each chosen scan position drawn as a numbered red marker at ``(x=col, + y=row)``; pass ``image=None`` to omit it and show only the tiles. Right: the + diffraction patterns tiled ``ncols`` wide, each rendered with + :func:`~quantem.core.visualization.show_2d`, with the detected peaks overlaid as + cyan rings that trace the disks rather than obscuring them. + + Parameters + ---------- + image : np.ndarray or None + Navigation image (e.g. a virtual dark-field image) drawn at left, with the + scan positions marked. Pass ``None`` to show only the diffraction tiles. + dps : list of np.ndarray + Diffraction patterns to tile, one per entry in ``positions``. + peaks : list of np.ndarray + Detected peaks per pattern; ``peaks[i]`` is an ``(M, 3)`` array of + ``[q_row, q_col, intensity]``. Each peak is drawn as a cyan ring at + ``(x=q_col, y=q_row)``. + positions : list of tuple of int + ``(row, col)`` scan positions, used for the tile titles and the nav-image + markers. + ncols : int, default=4 + Number of columns in the diffraction-pattern tile grid. + peak_radius : float, default=6.0 + Radius of the cyan peak rings, in diffraction pixels. + marker_radius : float, optional + Scan-position marker radius in image pixels; ``None`` uses a fixed + screen-size marker. + linewidth : float, default=0.5 + Stroke width of both the scan-position markers and the peak rings. + sigma_plot : float, optional + Gaussian blur width (pixels) applied to the *displayed* patterns only + (detection still uses the raw data) to ease telling real disks from false + positives. + image_title : str, default="virtual image" + Title for the navigation image. + image_kwargs : dict, optional + Extra keyword arguments for the navigation image's :func:`show_2d` call, + styling it independently of the tiles. + axsize : tuple of float, optional + Per-tile size in inches; a larger value zooms the tiles. + figsize : tuple of float, optional + Explicit figure size in inches. + **show_kwargs + Forwarded to :func:`show_2d` for the diffraction-pattern tiles (e.g. + ``norm``, ``cmap``, ``cbar``). + + Returns + ------- + tuple + ``(fig, (ax_image, dp_axes))`` with ``ax_image`` ``None`` when no image is + shown. + """ + from matplotlib.patches import Circle + + from quantem.core.visualization import show_2d + + n = len(positions) + show_image = image is not None + fig, ax_image, dp_axes, ncols = _grid_axes( + n, ncols, show_image=show_image, axsize=axsize, figsize=figsize + ) + + if show_image: + image_show_kwargs = {"cmap": "gray", "title": image_title, **(image_kwargs or {})} + show_2d(np.asarray(image), figax=(fig, ax_image), **image_show_kwargs) + _mark_positions(ax_image, positions, radius=marker_radius, linewidth=linewidth) + + user_title = show_kwargs.pop("title", None) + for i in range(n): + a = dp_axes[i // ncols, i % ncols] + r, c = positions[i] + pk = peaks[i] + title = user_title if user_title is not None else f"{i}: ({r},{c}) n={pk.shape[0]}" + show_2d(_blur(dps[i], sigma_plot), figax=(fig, a), title=title, **show_kwargs) + for q in pk: + a.add_patch( + Circle( + (q[1], q[0]), + radius=peak_radius, + fill=False, + edgecolor="cyan", + linewidth=linewidth, + ) + ) + + return fig, (ax_image, dp_axes) + + +def plot_basis_vectors( + bvm: np.ndarray, + cand_rc: np.ndarray, + cand_int: np.ndarray, + origin: np.ndarray, + g1: np.ndarray, + g2: np.ndarray, + *, + cmap: str = "gray", + norm: str | dict = "log_auto", + zoom: bool = True, + figsize: tuple[float, float] = (6, 6), + **show_kwargs, +): + """The Bragg vector map with the candidate peaks, origin, and basis vectors overlaid. + + Every candidate peak is drawn as a numbered cyan ring; those numbers are the + indices accepted by + :meth:`~quantem.diffraction.bragg_vectors.BraggVectors.choose_basis_vectors` + for overriding ``origin``/``g1``/``g2`` by peak. The chosen origin is a green + marker and ``g1`` (red) / ``g2`` (blue) are arrows drawn from it, labelled at + their midpoints so the labels never sit on top of the candidate numbers. + + Parameters + ---------- + bvm : np.ndarray + Bragg vector map; rendered through + :func:`~quantem.core.visualization.show_2d` so its display scaling is + controlled by ``norm`` and ``show_kwargs``. + cand_rc : np.ndarray + ``(N, 2)`` ``[row, col]`` candidate peak positions, brightest first; the + ring labels are their row indices. + cand_int : np.ndarray + ``(N,)`` candidate intensities (unused for drawing; kept for parity with + the candidate API). + origin : np.ndarray + ``(row, col)`` chosen lattice origin. + g1 : np.ndarray + First lattice vector as a ``(row, col)`` offset from ``origin``. + g2 : np.ndarray + Second lattice vector as a ``(row, col)`` offset from ``origin``. + cmap : str, default="gray" + Colormap for the Bragg vector map; gray keeps the colored overlays legible. + norm : str or dict, default="log_auto" + Intensity scaling forwarded to :func:`show_2d` (e.g. ``"linear_auto"``, + ``"log_auto"``, ``"power_sqrt"``, or ``{"power": 0.5}``). + zoom : bool, default=True + If ``True``, frame the view to the candidate bounding box (plus a margin) + so the numbered peaks are large enough to read. + figsize : tuple of float, default=(6, 6) + Figure size in inches. + **show_kwargs + Extra keyword arguments forwarded to :func:`show_2d` for fine display + control (e.g. ``vmin``, ``vmax``, ``lower_quantile``, ``upper_quantile``). + + Returns + ------- + tuple + ``(fig, ax)``. + """ + import matplotlib.patheffects as path_effects + + from quantem.core.visualization import show_2d + + stroke = [path_effects.withStroke(linewidth=2.5, foreground="black")] + origin_color = (0.0, 0.7, 0.0) + g1_color = (1.0, 0.0, 0.0) + g2_color = (0.0, 0.7, 1.0) + + fig, ax = plt.subplots(figsize=figsize) + show_2d(np.asarray(bvm), figax=(fig, ax), cmap=cmap, norm=norm, **show_kwargs) + + cand_rc = np.asarray(cand_rc, dtype=float).reshape(-1, 2) + for i, (r, c) in enumerate(cand_rc): + ax.scatter(c, r, s=60, facecolors="none", edgecolors="cyan", linewidths=1.0, zorder=3) + ax.annotate( + str(i), + (c, r), + color="cyan", + fontsize=9, + fontweight="bold", + xytext=(4, 4), + textcoords="offset points", + path_effects=stroke, + zorder=4, + ) + + o = np.asarray(origin, dtype=float).reshape(2) + ax.scatter( + o[1], + o[0], + s=160, + marker="P", + facecolors=[origin_color], + edgecolors="white", + linewidths=1.5, + zorder=6, + ) + for g, label, color in ( + (np.asarray(g1, float), "g1", g1_color), + (np.asarray(g2, float), "g2", g2_color), + ): + tip = (o[1] + g[1], o[0] + g[0]) + ax.annotate( + "", + xy=tip, + xytext=(o[1], o[0]), + arrowprops=dict(arrowstyle="-|>", color=color, lw=2.4, shrinkA=0, shrinkB=0), + zorder=5, + ) + # Label at the arrow midpoint, nudged perpendicular to the shaft (in screen + # space) so it clears both the arrow and the candidate numbers on the peaks. + gnorm = float(np.hypot(g[0], g[1])) + 1e-12 + perp = (g[0] / gnorm * 15.0, g[1] / gnorm * 15.0) + mid = (o[1] + g[1] / 2.0, o[0] + g[0] / 2.0) + ax.annotate( + label, + mid, + color=color, + fontsize=14, + fontweight="bold", + xytext=perp, + textcoords="offset points", + ha="center", + va="center", + path_effects=stroke, + zorder=7, + ) + + if zoom and cand_rc.shape[0]: + rmin, cmin = cand_rc.min(axis=0) + rmax, cmax = cand_rc.max(axis=0) + margin = 0.12 * max(rmax - rmin, cmax - cmin, 1.0) + 6.0 + h, w = bvm.shape[:2] + ax.set_xlim(max(cmin - margin, -0.5), min(cmax + margin, w - 0.5)) + ax.set_ylim(min(rmax + margin, h - 0.5), max(rmin - margin, -0.5)) + + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_title("lattice basis (origin +, g1, g2; cyan = candidate index)") + fig.tight_layout() + return fig, ax + + +def plot_bvm( + bvm: np.ndarray, + counts: np.ndarray, + *, + figsize: tuple[float, float] = (10, 4), + norm: str | dict = "log_auto", + cmap: str = "inferno", + counts_kwargs: dict | None = None, + **plotting_kwargs, +): + """The Bragg vector map beside the per-position peak count. + + Parameters + ---------- + bvm : np.ndarray + ``(H, W)`` Bragg vector map in detector pixels. + counts : np.ndarray + Per-position peak count, shape ``(scan_row, scan_col)``. + figsize : tuple of float, default=(10, 4) + Figure size in inches. + norm : str or dict, default="log_auto" + Intensity normalization of the Bragg vector map, passed to + :func:`~quantem.core.visualization.show_2d`. + cmap : str, default="inferno" + Colormap of the Bragg vector map. + counts_kwargs : dict, optional + Keyword arguments for :func:`~quantem.core.visualization.show_2d` on the + peak-count panel (defaults: ``cmap="viridis"``, ``cbar=True``). + **plotting_kwargs + Further keyword arguments for :func:`~quantem.core.visualization.show_2d` + on the Bragg vector map panel. + + Returns + ------- + tuple + ``(fig, ax)`` with ``ax`` a length-2 array of axes. + """ + from quantem.core.visualization import show_2d + + fig, ax = plt.subplots(1, 2, figsize=figsize) + + show_2d( + np.asarray(bvm), + figax=(fig, ax[0]), + cmap=cmap, + norm=norm, + title="Bragg vector map", + **plotting_kwargs, + ) + counts_plot_kwargs = { + "cmap": "viridis", + "cbar": True, + "title": "peaks per position", + **(counts_kwargs or {}), + } + show_2d(np.asarray(counts), figax=(fig, ax[1]), **counts_plot_kwargs) + fig.tight_layout() + return fig, ax + + +def plot_reference_lattice( + bvm: np.ndarray, + ref_qpos: np.ndarray, + ref_ab: np.ndarray, + origin: np.ndarray, + g1: np.ndarray, + g2: np.ndarray, + *, + cmap: str = "gray", + norm: str | dict = "log_auto", + zoom: bool = True, + figsize: tuple[float, float] = (6, 6), + **show_kwargs, +): + """The reference lattice from :meth:`BraggVectors.index_peaks`, drawn over the BVM. + + Each indexed reference site is a ring labelled with its ``(a, b)`` Miller index; + the chosen origin is a green marker and ``g1`` (red) / ``g2`` (blue) are arrows + from it. The ring color encodes how far the picked candidate sits from its *ideal* + lattice site ``origin + a*g1 + b*g2`` (the colorbar reads pixels), scaled to half + the shorter lattice spacing — the default :meth:`fit_lattice` match radius. A + mis-picked or duplicate candidate stands out as a ring far from zero offset (bright + color), sitting off the regular grid, or carrying an index that breaks the pattern. + + Parameters + ---------- + bvm : np.ndarray + Bragg vector map; rendered through + :func:`~quantem.core.visualization.show_2d` so its display scaling is + controlled by ``norm`` and ``show_kwargs``. + ref_qpos : np.ndarray + ``(N, 2)`` ``[row, col]`` reference site positions. + ref_ab : np.ndarray + ``(N, 2)`` integer ``[a, b]`` Miller indices for each site. + origin : np.ndarray + ``(row, col)`` chosen lattice origin. + g1 : np.ndarray + First lattice vector as a ``(row, col)`` offset from ``origin``. + g2 : np.ndarray + Second lattice vector as a ``(row, col)`` offset from ``origin``. + cmap : str, default="gray" + Colormap for the Bragg vector map; gray keeps the colored overlays legible. + norm : str or dict, default="log_auto" + Intensity scaling forwarded to :func:`show_2d` (e.g. ``"linear_auto"``, + ``"log_auto"``, ``"power_sqrt"``, or ``{"power": 0.5}``). + zoom : bool, default=True + If ``True``, frame the view to the reference bounding box (plus a margin) + so the labelled sites are large enough to read. + figsize : tuple of float, default=(6, 6) + Figure size in inches. + **show_kwargs + Extra keyword arguments forwarded to :func:`show_2d` for fine display + control (e.g. ``vmin``, ``vmax``, ``lower_quantile``, ``upper_quantile``). + + Returns + ------- + tuple + ``(fig, ax)``. + """ + import matplotlib.patheffects as path_effects + from matplotlib.cm import ScalarMappable + from matplotlib.colors import Normalize + + from quantem.core.visualization import show_2d + + stroke = [path_effects.withStroke(linewidth=2.5, foreground="black")] + origin_color = (0.0, 0.7, 0.0) + g1_color = (1.0, 0.0, 0.0) + g2_color = (0.0, 0.7, 1.0) + + ref_qpos = np.asarray(ref_qpos, dtype=float).reshape(-1, 2) + ref_ab = np.asarray(ref_ab, dtype=int).reshape(-1, 2) + o = np.asarray(origin, dtype=float).reshape(2) + g1 = np.asarray(g1, dtype=float).reshape(2) + g2 = np.asarray(g2, dtype=float).reshape(2) + + fig, ax = plt.subplots(figsize=figsize) + show_2d(np.asarray(bvm), figax=(fig, ax), cmap=cmap, norm=norm, **show_kwargs) + + # offset of each picked candidate from its ideal lattice site origin + a*g1 + b*g2; + # this is the QC indicator -- a mis-picked or strongly strained candidate rings + # far from zero, while a clean pick rings near it. + ideal_qpos = o[None, :] + ref_ab[:, 0:1] * g1[None, :] + ref_ab[:, 1:2] * g2[None, :] + offset = np.linalg.norm(ref_qpos - ideal_qpos, axis=1) + + # color the rings by that offset, scaled to half the shorter lattice spacing (the + # default fit_lattice match radius) so the colorbar previews which candidates sit + # near the inclusion tolerance. + radius = 0.5 * float(min(np.hypot(*g1), np.hypot(*g2))) + cmap_offset = plt.get_cmap("plasma") + norm_offset = Normalize(vmin=0.0, vmax=radius if radius > 0 else 1.0) + + ax.scatter( + ref_qpos[:, 1], + ref_qpos[:, 0], + s=80, + facecolors="none", + edgecolors=cmap_offset(norm_offset(offset)), + linewidths=1.8, + zorder=3, + ) + for (r, c), (a, b) in zip(ref_qpos, ref_ab): + ax.annotate( + f"{int(a)},{int(b)}", + (c, r), + color="white", + fontsize=9, + fontweight="bold", + xytext=(4, 4), + textcoords="offset points", + path_effects=stroke, + zorder=4, + ) + + ax.scatter( + o[1], + o[0], + s=160, + marker="P", + facecolors=[origin_color], + edgecolors="white", + linewidths=1.5, + zorder=6, + ) + for g, label, color in ((g1, "g1", g1_color), (g2, "g2", g2_color)): + ax.annotate( + "", + xy=(o[1] + g[1], o[0] + g[0]), + xytext=(o[1], o[0]), + arrowprops=dict(arrowstyle="-|>", color=color, lw=2.4, shrinkA=0, shrinkB=0), + zorder=5, + ) + gnorm = float(np.hypot(g[0], g[1])) + 1e-12 + ax.annotate( + label, + (o[1] + g[1] / 2.0, o[0] + g[0] / 2.0), + color=color, + fontsize=13, + fontweight="bold", + xytext=(g[0] / gnorm * 15.0, g[1] / gnorm * 15.0), + textcoords="offset points", + ha="center", + va="center", + path_effects=stroke, + zorder=7, + ) + + if zoom and ref_qpos.shape[0]: + rmin, cmin = ref_qpos.min(axis=0) + rmax, cmax = ref_qpos.max(axis=0) + margin = 0.12 * max(rmax - rmin, cmax - cmin, 1.0) + 6.0 + h, w = bvm.shape[:2] + ax.set_xlim(max(cmin - margin, -0.5), min(cmax + margin, w - 0.5)) + ax.set_ylim(min(rmax + margin, h - 0.5), max(rmin - margin, -0.5)) + + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_title("reference lattice (ring color = offset from ideal index; +origin, g1, g2)") + + sm = ScalarMappable(norm=norm_offset, cmap=cmap_offset) + sm.set_array([]) + cbar = fig.colorbar(sm, ax=ax, fraction=0.046, pad=0.04, extend="max") + cbar.set_label("peak offset from ideal index (px)") + + fig.tight_layout() + return fig, ax + + +def plot_lattice_fit( + mask_weight: np.ndarray, + fit_error: np.ndarray, + *, + figsize: tuple[float, float] = (11, 4.5), +): + """Per-position diagnostics from :meth:`BraggVectors.fit_lattice`. + + Left: the mask weight — a lattice *order parameter* per position (how well all + detected intensity snaps to the fitted lattice, intensity-weighted) — so ``0`` is + a position dominated by off-lattice intensity and ``1`` a clean single crystal. + This is the weighting handed to the strain reference. Right: the RMS lattice-fit + residual over the matched peaks, in pixels (low = a clean fit; high = a poorly + fit, overlapping, or strongly strained position). + + Parameters + ---------- + mask_weight : np.ndarray + ``(scan_row, scan_col)`` lattice-order-parameter weight in ``[0, 1]``. + fit_error : np.ndarray + ``(scan_row, scan_col)`` RMS fit residual in pixels (``nan`` where no fit + was made). + figsize : tuple of float, default=(11, 4.5) + Figure size in inches. + + Returns + ------- + tuple + ``(fig, ax)`` with ``ax`` a length-2 array of axes. + """ + fig, ax = plt.subplots(1, 2, figsize=figsize) + + im0 = ax[0].imshow(np.asarray(mask_weight), cmap="viridis", vmin=0.0, vmax=1.0) + ax[0].set_title("mask weight (lattice order)") + fig.colorbar(im0, ax=ax[0], fraction=0.046, pad=0.04) + + im1 = ax[1].imshow(np.asarray(fit_error), cmap="magma") + ax[1].set_title("fit RMS error (px)") + fig.colorbar(im1, ax=ax[1], fraction=0.046, pad=0.04) + + for a in ax: + a.set_xticks([]) + a.set_yticks([]) + fig.tight_layout() + return fig, ax diff --git a/src/quantem/diffraction/calibration.py b/src/quantem/diffraction/calibration.py new file mode 100644 index 000000000..8ce1857c0 --- /dev/null +++ b/src/quantem/diffraction/calibration.py @@ -0,0 +1,1841 @@ +"""Diffraction-space calibration of 4D-STEM Bragg peaks. + +Each step works on detected Bragg peaks (a quantem Vector): + +- Origins: :func:`measure_origins` fits a plane to the direct-beam position + over the scan, for ``BraggVectors.correct_peak_origins``. +- Scan rotation: :func:`measure_scan_rotation` finds the detector-to-scan + rotation from the curl of the center-of-mass field. +- Pixel size and elliptic distortion: :func:`calibrate` matches the radial + peak histogram against the rings of one or more reference crystals and + returns a :class:`DiffractionCalibration`. +- :class:`DiffractionCalibration` holds the pixel size (1/Angstroms per + pixel), the ellipse and the rotation, converts pixel peaks to calibrated + (qx, qy) peaks with :meth:`DiffractionCalibration.apply`, and can be saved + from a standard and reused on another dataset. + +:func:`calibrate` is the main entry point. The other functions are +lower-level steps that act on peaks directly: :func:`peaks_to_calibrated` +and :func:`calibrate_ellipse` (both used by :func:`calibrate`), +:func:`calibrate_pixel_size` (radial histogram fit of the scale), +:func:`calibrate_pixel_size_matching` (scale fit by full orientation +matching), :func:`scale_peaks` and :func:`apply_ellipse`. +:func:`refine_calibration` measures the remaining calibration error from +strain maps after orientation matching. +""" + +from __future__ import annotations + +import warnings + +import numpy as np +import torch + +from quantem.core.datastructures.vector import Vector +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.defaults import MIN_NUMBER_PEAKS + + +def _measure_raw_origins(bragg_vectors, search_radius: float, center=None) -> np.ndarray: + """Brightest-peak origin per position, NaN where nothing is found.""" + peaks = bragg_vectors.peaks + scan_r, scan_c = peaks.shape[0], peaks.shape[1] + H, W = int(bragg_vectors.dataset.shape[-2]), int(bragg_vectors.dataset.shape[-1]) + c0 = np.array([H / 2, W / 2]) if center is None else np.asarray(center, dtype=float) + meas = np.full((scan_r, scan_c, 2), np.nan) + for r in range(scan_r): + for c in range(scan_c): + arr = peaks[r, c].numpy().astype(np.float64) + if arr.shape[0] == 0: + continue + d = np.hypot(arr[:, 0] - c0[0], arr[:, 1] - c0[1]) + near = d < search_radius + if not near.any(): + continue + sub = arr[near] + meas[r, c] = sub[np.argmax(sub[:, 2]), :2] + return meas + + +def _plot_origin_panels(meas: np.ndarray, origins: np.ndarray): + """Measured origins, plane fit and residual for both detector axes.""" + import matplotlib.pyplot as plt + + fig, axs = plt.subplots(2, 3, figsize=(13.5, 5.6)) + names = ["row", "col"] + for k in range(2): + m, f = meas[..., k], origins[..., k] + resid = m - f + mean_m = np.nanmean(m) + span = max(np.nanstd(m) * 3, 1e-3) + for j, (img, title) in enumerate( + [ + (m, f"measured origin {names[k]} (px)"), + (f, f"plane fit {names[k]} (px)"), + (resid, f"residual {names[k]} (px)"), + ] + ): + c0 = 0.0 if j == 2 else mean_m + sp = max(np.nanstd(resid) * 3, 1e-3) if j == 2 else span + im = axs[k, j].imshow( + img, + cmap="RdBu_r", + vmin=c0 - sp, + vmax=c0 + sp, + interpolation="nearest", + ) + axs[k, j].set_title(title, fontsize=10) + axs[k, j].set_xticks([]) + axs[k, j].set_yticks([]) + fig.colorbar(im, ax=axs[k, j], shrink=0.85) + fig.tight_layout() + return fig, axs + + +def plot_origin_fit(bragg_vectors, origins: np.ndarray, search_radius: float = 6.0, center=None): + """Measured origins against the plane fit of :func:`measure_origins`. + + Parameters + ---------- + bragg_vectors : BraggVectors + With detected peaks. + origins : np.ndarray + (scan_row, scan_col, 2) fitted origins from :func:`measure_origins`. + search_radius : float, default=6.0 + Radius in detector pixels searched around `center`, as passed to + :func:`measure_origins`. + center : tuple of float, optional + ``(row, col)`` detector position searched around, as passed to + :func:`measure_origins`. Defaults to the detector center. + + Returns + ------- + tuple + ``(fig, axs)``, axs of shape (2, 3): rows are the detector row and + column, columns are measured, fit and residual, all in pixels. + """ + meas = _measure_raw_origins(bragg_vectors, search_radius, center) + return _plot_origin_panels(meas, origins) + + +def measure_origins( + bragg_vectors, + search_radius: float = 6.0, + robust: bool = True, + plot: bool = False, + center=None, + min_coverage: float = 0.1, +): + """Per-position diffraction origin from the brightest central peak. + + At each scan position the most intense detected peak within + `search_radius` pixels of `center` is taken as the direct beam; a plane + is fit over the scan (least squares, optionally with one + outlier-rejection pass) to model the descan. + + Parameters + ---------- + bragg_vectors : BraggVectors + With detected peaks, fields (q_row, q_col, intensity) in detector + pixels, not yet origin-corrected. + search_radius : float, default=6.0 + Radius in detector pixels searched around `center`. + robust : bool, default=True + Reject outliers before the plane fit. + plot : bool, default=False + Show the fitted origin planes and the residuals of the measured + origins against the fit. + center : tuple of float, optional + ``(row, col)`` detector position to search around, e.g. from + :func:`~quantem.diffraction.disk_detection.estimate_central_beam`. + Defaults to the detector centre, which misses a beam that sits + further than `search_radius` from it. + min_coverage : float, default=0.1 + Fraction of scan positions that must yield an origin. Below it the + search has missed the beam, and the fit would be meaningless, so an + error is raised rather than returning a plane through nothing. + + Returns + ------- + np.ndarray + (scan_row, scan_col, 2) plane-fit origins, ready for + BraggVectors.correct_peak_origins(). With plot=True, also returns + (fig, axs). + + Raises + ------ + ValueError + If fewer than `min_coverage` of the positions have a peak within + `search_radius` of `center`. + """ + meas = _measure_raw_origins(bragg_vectors, search_radius, center) + scan_r, scan_c = meas.shape[0], meas.shape[1] + coverage = float(np.isfinite(meas[..., 0]).mean()) + if coverage < min_coverage: + H, W = int(bragg_vectors.dataset.shape[-2]), int(bragg_vectors.dataset.shape[-1]) + c0 = (H / 2, W / 2) if center is None else tuple(float(v) for v in center) + raise ValueError( + f"only {coverage:.1%} of scan positions have a peak within " + f"{search_radius:g} px of ({c0[0]:.1f}, {c0[1]:.1f}); the direct beam is " + "elsewhere. Pass center= from estimate_central_beam(dataset.dp_mean), " + "or widen search_radius." + ) + ry, rx = np.mgrid[0:scan_r, 0:scan_c] + + def plane(z, ok): + A = np.stack([np.ones(ok.sum()), ry[ok], rx[ok]], axis=1) + coef, *_ = np.linalg.lstsq(A, z[ok], rcond=None) + return coef[0] + coef[1] * ry + coef[2] * rx + + out = np.zeros((scan_r, scan_c, 2)) + for k in range(2): + z = meas[..., k] + ok = np.isfinite(z) + fit = plane(z, ok) + if robust: + resid = np.abs(z - fit) + thresh = 3 * np.nanmedian(resid[ok]) + 1e-9 + ok = ok & (resid < thresh) + fit = plane(z, ok) + out[..., k] = fit + + if plot: + fig, axs = _plot_origin_panels(meas, out) + return out, fig, axs + return out + + +def peaks_to_calibrated( + peaks_px, + pixel_size_inv_A: float, + rotation_ccw_deg: float = 0.0, + ellipse=None, + name: str = "bragg_peaks_calibrated", +): + """Convert origin-corrected pixel peaks to a calibrated (qx, qy) Vector. + + Parameters + ---------- + peaks_px : Vector + Peaks with fields (q_row, q_col, intensity) in detector pixels, + already origin-corrected so (0, 0) is the direct beam. + pixel_size_inv_A : float + Reciprocal pixel size in 1/Angstroms. + rotation_ccw_deg : float, default=0.0 + Diffraction-to-scan rotation: the detector coordinates are rotated + by this angle so the pattern axes align with the scan axes. With + this applied, in-plane orientations and strain axes are reported in + the image frame. + ellipse : array-like | None + [e11, e12] elliptic distortion correction from calibrate_ellipse(), + applied in the detector frame before the rotation. + name : str, default="bragg_peaks_calibrated" + Name of the returned Vector. + + Returns + ------- + Vector + Fields (qx, qy, intensity) in 1/Angstroms. + """ + scan_r, scan_c = peaks_px.shape[0], peaks_px.shape[1] + flat = peaks_px.select_fields("q_row", "q_col", "intensity").numpy().astype(np.float64) + row_counts = np.asarray(peaks_px.row_counts(), dtype=int) + qrc = flat[:, :2] * pixel_size_inv_A + if ellipse is not None: + qrc = qrc @ _ellipse_matrix(ellipse).T + if rotation_ccw_deg != 0.0: + th = np.deg2rad(rotation_ccw_deg) + rot = np.array([[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]]) + qrc = qrc @ rot.T + data = np.column_stack([qrc, flat[:, 2]]) + cells = np.split(data, np.cumsum(row_counts)[:-1]) + nested = [cells[r * scan_c : (r + 1) * scan_c] for r in range(scan_r)] + out = Vector.from_data( + nested, + fields=["qx", "qy", "intensity"], + units=["A^-1", "A^-1", "counts"], + name=name, + dtype=peaks_px.dtype, + ) + # carry the scan calibration through, so maps keep their scale bar, and + # record the detector-to-scan rotation so pattern-overlay plots can put + # peaks back into the raw detector frame + for key in ("scan_sampling", "scan_units", "origins", "origin_ref"): + if key in (peaks_px.metadata or {}): + out.metadata[key] = peaks_px.metadata[key] + out.metadata["rotation_ccw_deg"] = float(rotation_ccw_deg) + return out + + +def scale_peaks(peaks, scale: float): + """Return a copy of a (qx, qy, intensity) Vector with q scaled. + + Parameters + ---------- + peaks : Vector + Peaks with fields (qx, qy, intensity). + scale : float + Factor applied to qx and qy, e.g. from :func:`calibrate_pixel_size`. + + Returns + ------- + Vector + Scaled copy of `peaks`. + """ + out = peaks.copy() + flat = out.numpy().astype(np.float64) + flat[:, :2] *= scale + out.set_flattened(flat) + return out + + +def radial_histogram( + peaks: Vector, + k_min: float = 0.05, + k_max: float = 1.5, + k_step: float = 0.002, + bragg_k_power: float = 2.0, + bragg_intensity_power: float = 1.0, +) -> tuple[np.ndarray, np.ndarray]: + """Intensity-weighted histogram of Bragg peak radii over all positions. + + Parameters + ---------- + peaks : Vector + Calibrated peaks with fields (qx, qy, intensity) in 1/Angstroms. + k_min, k_max : float, default=0.05, 1.5 + Range of the bins, 1/Angstroms. + k_step : float, default=0.002 + Bin width, 1/Angstroms. + bragg_k_power : float, default=2.0 + Each peak is weighted by ``|q| ** bragg_k_power``, which offsets the + fall-off of the scattering factors with q. + bragg_intensity_power : float, default=1.0 + Each peak is weighted by ``intensity ** bragg_intensity_power``; 0 + counts every peak equally. + + Returns + ------- + k : np.ndarray + Bin centers (1/Angstroms). + hist : np.ndarray + Weighted counts, with linear interpolation between adjacent bins. + """ + flat = peaks.select_fields("qx", "qy", "intensity").numpy().astype(np.float64) + qr = np.hypot(flat[:, 0], flat[:, 1]) + weight = flat[:, 2] ** bragg_intensity_power * qr**bragg_k_power + + k = np.arange(k_min, k_max, k_step) + frac = (qr - k_min) / k_step + i0 = np.floor(frac).astype(int) + w1 = frac - i0 + ok = (i0 >= 0) & (i0 < k.size - 1) + hist = np.bincount(i0[ok], weights=weight[ok] * (1 - w1[ok]), minlength=k.size) + hist += np.bincount(i0[ok] + 1, weights=weight[ok] * w1[ok], minlength=k.size) + return k, hist + + +def simulated_ring_profile( + crystal: Crystal, + k: np.ndarray, + k_broadening: float = 0.01, + bragg_k_power: float = 2.0, +) -> np.ndarray: + """1D ring profile of a crystal: Gaussians at |g| weighted by intensity. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. + k : np.ndarray + Scattering vectors to evaluate at, 1/Angstroms. + k_broadening : float, default=0.01 + Gaussian standard deviation of each ring, 1/Angstroms. + bragg_k_power : float, default=2.0 + Each ring is weighted by ``|g| ** bragg_k_power`` times its intensity. + + Returns + ------- + np.ndarray + Profile, same shape as `k`. + """ + g = crystal.g_len.numpy() + w = crystal.struct_factors_int.numpy() * g**bragg_k_power + prof = (w[None, :] * np.exp(-((k[:, None] - g[None, :]) ** 2) / (2 * k_broadening**2))).sum( + axis=1 + ) + return prof + + +def calibrate_pixel_size_matching( + peaks, + crystal: Crystal | list[Crystal], + energy_ev: float, + scales: np.ndarray | None = None, + subsample: int = 8, + angle_step_deg: float = 3.0, + corr_kernel_size: float = 0.02, + min_number_peaks: int = MIN_NUMBER_PEAKS, + plot: bool = False, + return_scores: bool = False, + returnfig: bool = False, +): + """Refine the pixel size by maximizing the orientation-match correlation. + + The 1D radial fit can be fooled by ring-ratio degeneracies (e.g. the + hexagonal-net radii shared by hcp prismatic rings and bcc {110}-family + rings). Full-pattern matching is not: for each candidate scale, a + subsampled grid of patterns is orientation-matched against the crystal + and the median normalized correlation scored. Peaks in the score curve + identify the true calibration. + + Parameters + ---------- + peaks : Vector + Calibrated peaks (qx, qy, intensity) in 1/Angstroms. + crystal : Crystal | list[Crystal] + Reference crystal(s). Pass ALL candidate phases for multi-phase + samples: with a single reference, a scale that maps the majority + phase's net onto the reference's (e.g. the bcc {110} ring onto the + hcp prismatic ring, ratio 0.90 for Ti) can win the scan. Scoring the + mean over phases of the per-phase median correlation removes the + false optimum, since the other phases index nothing at the impostor + scale. + energy_ev : float + Beam energy in eV. + scales : np.ndarray | None + Candidate scale factors; defaults to 0.90 ... 1.10 in 2% steps. + subsample : int, default=8 + Stride of the probe-position grid used for scoring. + angle_step_deg : float, default=3.0 + Zone-axis and in-plane angular step of the orientation plan, degrees. + corr_kernel_size : float, default=0.02 + Matching kernel width, 1/Angstroms. Keep it near the peak position + noise: a wide kernel gives partial credit to near-coincident rings + at a wrong scale. + min_number_peaks : int, default=MIN_NUMBER_PEAKS + Positions with fewer peaks are not matched. + plot : bool, default=False + Plot the score against the scale factor. + return_scores : bool, default=False + Also return the scale factors and their scores. + returnfig : bool, default=False + With `plot`, also return the figure and axes. + + Returns + ------- + scale : float + Best scale factor (parabolic refinement over the score maximum). + Multiply the pixel size by it. + scales, scores : np.ndarray + Only with `return_scores`. The score is the mean over phases of the + median correlation of the matched positions, 0 for a phase where + nothing matched. + fig, ax + Only with `plot` and `returnfig`. + """ + from quantem.diffraction.orientation import OrientationMap + + crystals = crystal if isinstance(crystal, (list, tuple)) else [crystal] + if scales is None: + scales = np.arange(0.90, 1.101, 0.02) + sub = peaks[::subsample, ::subsample] + scores = np.zeros(len(scales)) + for i, s in enumerate(scales): + test = scale_peaks(sub, float(s)) + per_phase = [] + for xtl in crystals: + om = OrientationMap.from_vectors(test, xtl, energy_ev=energy_ev) + # detector_q_max must stay OFF here: the auto footprint shrinks + # with the candidate scale, silently removing the unexplained + # high-q template shells that penalize too-small scales -- the + # score then rises monotonically as the pattern is compressed. + # a tight kernel is essential for a wide scale scan: the default + # matching kernel (0.05) hands partial credit to near-miss ring + # coincidences of an impostor scale, while matches at the true + # scale are exact to the detection noise (~0.005) + om.build_plan( + angle_step_zone_axis_deg=angle_step_deg, + angle_step_in_plane_deg=angle_step_deg, + corr_kernel_size=corr_kernel_size, + detector_q_max=None, + verbose=False, + ) + om.match_orientations(progress_bar=False, min_number_peaks=min_number_peaks) + corr = om.corr[..., 0] + matched = corr[corr > 0] + # a phase that indexes nothing at this scale scores zero + per_phase.append(float(matched.median()) if matched.numel() > 0 else 0.0) + scores[i] = float(np.mean(per_phase)) + + i_best = int(np.nanargmax(scores)) if np.isfinite(scores).any() else 0 + scale = float(scales[i_best]) + if 0 < i_best < len(scales) - 1: + c0, c1, c2 = scores[i_best - 1 : i_best + 2] + denom = 4 * c1 - 2 * c0 - 2 * c2 + step = scales[1] - scales[0] + if abs(denom) > 1e-12: + scale += (c2 - c0) / denom * step + out: list = [scale] + if return_scores: + out += [scales, scores] + if plot: + import matplotlib.pyplot as plt + + fig, ax = plt.subplots(figsize=(6, 4)) + ax.plot(scales, scores, "k.-") + ax.axvline(scale, color="r", ls="--") + ax.set_xlabel("pixel size scale factor") + ax.set_ylabel("median correlation") + ax.set_title(f"best scale = {scale:.4f}") + if returnfig: + out += [fig, ax] + return tuple(out) if len(out) > 1 else out[0] + + +def calibrate_pixel_size( + peaks: Vector, + crystal: Crystal, + scale_range: tuple[float, float] = (0.8, 1.25), + scale_step: float = 5e-4, + k_min: float = 0.05, + k_max: float = 1.3, + k_broadening: float = 0.01, + bragg_k_power: float = 2.0, + plot: bool = False, + returnfig: bool = False, +): + """Refine the reciprocal pixel size against a reference crystal. + + Scans a multiplicative scale factor applied to the measured peak radii and + maximizes the normalized overlap between the measured radial histogram and + the crystal's simulated ring profile. Parabolic sub-step refinement of the + best scale. + + Parameters + ---------- + peaks : Vector + Calibrated peaks with fields (qx, qy, intensity) in 1/Angstroms. + crystal : Crystal + Reference crystal with structure factors calculated. Choose the + majority phase of the scan. + scale_range : tuple, default=(0.8, 1.25) + Search range of the scale factor. + scale_step : float, default=5e-4 + Step of the scale factor scan. + k_min, k_max : float, default=0.05, 1.3 + Range of scattering vectors compared, 1/Angstroms, after scaling. + k_broadening : float, default=0.01 + Gaussian standard deviation of the simulated rings, 1/Angstroms. + bragg_k_power : float, default=2.0 + Rings are weighted by ``|g| ** bragg_k_power``; see + :func:`simulated_ring_profile`. + plot : bool, default=False + Show the scaled histogram against the crystal ring profile. + returnfig : bool, default=False + With `plot`, also return the figure and axes. + + Returns + ------- + scale : float + Multiply existing q values (and the pixel size) by this factor, + e.g. with scale_peaks(). + fig, ax + Only with `plot` and `returnfig`. + """ + k, hist = radial_histogram(peaks, k_min=k_min * scale_range[0], k_max=k_max / scale_range[0]) + scales = np.arange(scale_range[0], scale_range[1], scale_step) + score = np.zeros_like(scales) + for i, s in enumerate(scales): + prof = simulated_ring_profile(crystal, k * s, k_broadening, bragg_k_power) + keep = (k * s > k_min) & (k * s < k_max) + h, p = hist[keep], prof[keep] + denom = np.linalg.norm(h) * np.linalg.norm(p) + score[i] = (h * p).sum() / denom if denom > 0 else 0.0 + + i_best = int(np.argmax(score)) + scale = float(scales[i_best]) + if 0 < i_best < scales.size - 1: + c0, c1, c2 = score[i_best - 1 : i_best + 2] + denom = 4 * c1 - 2 * c0 - 2 * c2 + if abs(denom) > 1e-12: + scale += (c2 - c0) / denom * scale_step + + out: list = [scale] + if plot: + import matplotlib.pyplot as plt + + fig, ax = plt.subplots(figsize=(10, 4)) + prof = simulated_ring_profile(crystal, k * scale, k_broadening / 2, bragg_k_power) + ax.fill_between( + k * scale, + hist / hist.max(), + color="r", + alpha=0.75, + lw=0, + label="measured (scaled)", + ) + ax.plot( + k * scale, + prof / prof.max(), + "k-", + lw=1.0, + label=f"{crystal.name} rings", + ) + ax.set_ylabel("intensity (norm.)") + ax.set_xlabel(r"scattering vector (1/$\mathrm{\AA}$)") + ax.set_title(f"1D radial fit, scale = {scale:.4f}", fontsize=10) + ax.legend(loc="upper right", fontsize=9) + if returnfig: + out += [fig, ax] + return tuple(out) if len(out) > 1 else out[0] + + +def measure_scan_rotation( + dataset, + origins: np.ndarray | None = None, + mask_radius: float | None = None, + plot: bool = False, + returnfig: bool = False, +): + """Detector-to-scan rotation from the curl of the center-of-mass field. + + The center of mass of each diffraction pattern (about the fitted origin) + forms a vector field over the scan. In the correct common frame that + field is (approximately) a gradient field, so its curl vanishes; rotating + the detector axes by the unknown scan rotation and minimizing the summed + squared curl recovers the angle. + + The curl is invariant under 180-degree rotation, so the sign of the + measured field cannot distinguish theta from theta + 180. Only the + candidate in [0, 180) is returned; the other is that angle + 180. Pick + the one consistent with a known feature (e.g. a Burgers orientation + relationship, or the divergence sign convention of DPC). The returned + angle is ready to pass to peaks_to_calibrated() as rotation_ccw_deg. + + Parameters + ---------- + dataset : Dataset4dstem + The 4D-STEM scan. + origins : np.ndarray | None + (scan_r, scan_c, 2) diffraction origins from measure_origins(); + defaults to the pattern center. + mask_radius : float | None + Restrict the center of mass to within this radius (pixels) of the + origin -- i.e. the DPC signal of the direct beam only, excluding the + Bragg disks. Recommended for crystalline data. + plot : bool, default=False + Plot the curl and divergence measures against the rotation angle. + returnfig : bool, default=False + With `plot`, also return the figure and axes. + + Returns + ------- + rotation_ccw_deg : float + Curl-minimizing rotation in [0, 180), degrees, on a 0.25 degree + grid; the physical answer is either this angle or this angle + 180. + fig, ax + Only with `plot` and `returnfig`. + """ + arr = dataset.array + scan_r, scan_c, H, W = arr.shape + rows = np.arange(H, dtype=float)[:, None] + cols = np.arange(W, dtype=float)[None, :] + if origins is None: + origins = np.zeros((scan_r, scan_c, 2)) + origins[..., 0] = H / 2 + origins[..., 1] = W / 2 + origins = np.asarray(origins, dtype=float) + # one scan row at a time: a float64 copy of a full scan is several times + # the size of the data (19 GB for 256 x 256 x 192 x 192), and the center + # of mass only ever needs one pattern + com_r = np.zeros((scan_r, scan_c)) + com_c = np.zeros((scan_r, scan_c)) + for i in range(scan_r): + block = np.asarray(arr[i], dtype=float) # (scan_c, H, W) + if mask_radius is not None: + rr = rows[None] - origins[i, :, 0][:, None, None] + cc = cols[None] - origins[i, :, 1][:, None, None] + block = block * (rr**2 + cc**2 <= mask_radius**2) + tot = block.sum(axis=(-2, -1)) + tot[tot <= 0] = 1.0 + com_r[i] = (block * rows).sum(axis=(-2, -1)) / tot - origins[i, :, 0] + com_c[i] = (block * cols).sum(axis=(-2, -1)) / tot - origins[i, :, 1] + + # spatial derivatives of both components over the scan + d_rr = np.gradient(com_r, axis=0) + d_rc = np.gradient(com_r, axis=1) + d_cr = np.gradient(com_c, axis=0) + d_cc = np.gradient(com_c, axis=1) + + theta = np.deg2rad(np.arange(0, 180, 0.25)) + ct, st = np.cos(theta)[:, None, None], np.sin(theta)[:, None, None] + # rotated field: (r', c') = (ct * r - st * c, st * r + ct * c) + curl = (st * d_rr + ct * d_cr) - (ct * d_rc - st * d_cc) + div = (ct * d_rr - st * d_cr) + (st * d_rc + ct * d_cc) + curl_sq = (curl**2).mean(axis=(1, 2)) + div_sq = (div**2).mean(axis=(1, 2)) + + i_best = int(np.argmin(curl_sq)) + rotation_ccw_deg = float(np.rad2deg(theta[i_best])) + + out: list = [rotation_ccw_deg] + if plot: + import matplotlib.pyplot as plt + + fig, ax = plt.subplots(figsize=(7, 4)) + deg = np.rad2deg(theta) + ax.plot(deg, curl_sq, "k-", label="mean squared curl") + ax.plot(deg, div_sq, "-", color="0.6", label="mean squared divergence") + ax.axvline(rotation_ccw_deg, color="r", ls="--") + ax.set_xlabel("rotation (degrees)") + ax.set_ylabel("field measure") + ax.set_title( + "scan rotation = %.1f deg (or %.1f)" % (rotation_ccw_deg, rotation_ccw_deg + 180) + ) + ax.legend() + if returnfig: + out += [fig, ax] + return tuple(out) if len(out) > 1 else out[0] + + +def calibrate_ellipse( + peaks, + k_min: float = 0.15, + k_max: float = 1.4, + n_bins: int = 800, + bragg_k_power: float = 2.0, + bragg_intensity_power: float = 1.0, + plot: bool = False, + returnfig: bool = False, +): + """Elliptic distortion (e11, e12) from radial histogram sharpness. + + Fits the traceless linear distortion A = [[1 + e11, e12], [e12, 1 - e11]] + that, applied to the measured peaks, maximizes the sharpness of the + radial peak histogram: an elliptic distortion smears every diffraction + ring, and undoing it re-focuses them. The histogram is accumulated in + log-radius bins, where a pure scale change is only a translation -- so + the ellipse fit is independent of the pixel size, and the calibration + workflow stays sequential: rough scale, then ellipse, then absolute + scale by pattern matching. + + Parameters + ---------- + peaks : Vector + Calibrated peaks (qx, qy, intensity), approximate scale is fine. + k_min, k_max : float, default=0.15, 1.4 + Radial range (1/Angstroms) included in the sharpness measure. + n_bins : int, default=800 + Number of log-radius bins between `k_min` and `k_max`. + bragg_k_power : float, default=2.0 + Each peak is weighted by ``|q| ** bragg_k_power``. + bragg_intensity_power : float, default=1.0 + Each peak is weighted by ``intensity ** bragg_intensity_power``. + plot : bool, default=False + Show the radial histogram before and after the correction. + returnfig : bool, default=False + With `plot`, also return the figure and axes. + + Returns + ------- + ellipse : np.ndarray + [e11, e12], the correction matrix [[1 + e11, e12], [e12, 1 - e11]] + that undoes the distortion; pass to peaks_to_calibrated(ellipse=...) + or apply_ellipse(). + fig, ax + Only with `plot` and `returnfig`. + """ + from scipy.optimize import minimize + + flat = peaks.select_fields("qx", "qy", "intensity").numpy().astype(np.float64) + q = flat[:, :2] + log_lo, log_hi = np.log(k_min), np.log(k_max) + bin_w = (log_hi - log_lo) / n_bins + + def histogram(e): + A = np.array([[1 + e[0], e[1]], [e[1], 1 - e[0]]]) + qe = q @ A.T + r = np.hypot(qe[:, 0], qe[:, 1]) + ok = (r > k_min) & (r < k_max) + w = flat[ok, 2] ** bragg_intensity_power * r[ok] ** bragg_k_power + f = (np.log(r[ok]) - log_lo) / bin_w + i0 = np.floor(f).astype(int) + w1 = f - i0 + h = np.bincount(i0, weights=w * (1 - w1), minlength=n_bins + 1) + h += np.bincount(i0 + 1, weights=w * w1, minlength=n_bins + 1) + return h + + def cost(e): + h = histogram(e) + total = h.sum() + if total <= 0: + return 0.0 + return -float((h**2).sum()) / total**2 + + res = minimize( + cost, + x0=np.zeros(2), + method="Nelder-Mead", + options={"xatol": 1e-5, "fatol": 1e-12, "maxiter": 400}, + ) + ellipse = res.x + + out: list = [ellipse] + if plot: + import matplotlib.pyplot as plt + + k_bins = np.exp(np.linspace(log_lo, log_hi, n_bins + 1)) + h0 = histogram(np.zeros(2)) + h1 = histogram(ellipse) + fig, ax = plt.subplots(figsize=(10, 4)) + ax.fill_between(k_bins, h0 / h0.max(), color="0.7", lw=0, label="measured") + ax.plot(k_bins, h1 / h1.max(), "r-", lw=1.0, label="ellipse corrected") + ax.set_xlabel(r"scattering vector (1/$\mathrm{\AA}$)") + ax.set_ylabel("intensity (norm.)") + mag = np.hypot(*ellipse) + ax.set_title( + "e11 = %.2e, e12 = %.2e (%.2f%% ellipticity)" % (ellipse[0], ellipse[1], 200 * mag) + ) + ax.legend() + if returnfig: + out += [fig, ax] + return tuple(out) if len(out) > 1 else out[0] + + +def _ellipse_matrix(ellipse) -> np.ndarray: + """The correction matrix [[1 + e11, e12], [e12, 1 - e11]] of an ellipse.""" + e11, e12 = float(ellipse[0]), float(ellipse[1]) + return np.array([[1 + e11, e12], [e12, 1 - e11]]) + + +def _compose_ellipse(ellipse_new, ellipse_prev) -> np.ndarray: + """One ellipse equivalent to applying `ellipse_prev`, then `ellipse_new`. + + The product of the two correction matrices is reduced to the traceless + symmetric form [[1 + e11, e12], [e12, 1 - e11]]: its isotropic part is + a pixel size change, fit separately, and its antisymmetric part is a + rotation of second order in the ellipse components. + + Parameters + ---------- + ellipse_new : array-like + [e11, e12] fit on peaks that already have `ellipse_prev` applied. + ellipse_prev : array-like | None + [e11, e12] already applied, or None. + + Returns + ------- + np.ndarray + Combined [e11, e12]. + """ + if ellipse_prev is None: + return np.asarray(ellipse_new, dtype=float) + A = _ellipse_matrix(ellipse_new) @ _ellipse_matrix(ellipse_prev) + sym = 0.5 * (A + A.T) / (0.5 * np.trace(A)) + return np.array([0.5 * (sym[0, 0] - sym[1, 1]), sym[0, 1]]) + + +def apply_ellipse(peaks, ellipse): + """Return a copy of (qx, qy, intensity) peaks with the ellipse applied. + + Parameters + ---------- + peaks : Vector + Peaks with fields (qx, qy, intensity). + ellipse : array-like + [e11, e12] from :func:`calibrate_ellipse`. Each q is mapped by the + correction matrix [[1 + e11, e12], [e12, 1 - e11]]. + + Returns + ------- + Vector + Corrected copy of `peaks`. + """ + out = peaks.copy() + flat = out.numpy().astype(np.float64) + flat[:, :2] = flat[:, :2] @ _ellipse_matrix(ellipse).T + out.set_flattened(flat) + return out + + +def zone_reflections(crystal: Crystal, zone_axis) -> Crystal: + """The crystal restricted to the reflections of one zone. + + A specimen that sits near one zone axis everywhere -- a flake lying flat, + a textured film -- only ever shows the reflections of that zone, so its + radial histogram holds those rings and no others. Calibrating or + plotting against every ring of the crystal then compares the peaks with + rings that cannot appear. This keeps the reflections hkl with + h u + k v + l w = 0 for the zone axis [uvw] and drops the rest. + + Parameters + ---------- + crystal : Crystal + With structure factors calculated. + zone_axis : sequence of int + Zone axis direction in the crystal's own cell, [uvw], or [UVTW] for a + hexagonal or trigonal cell, e.g. (0, 0, 0, 1) for the basal plane. + + Returns + ------- + Crystal + A shallow copy whose reflection list holds only that zone, named + after it; the original is unchanged. + """ + import copy + + if crystal.g_vec is None: + raise RuntimeError(f"{crystal.name}: run calculate_structure_factors() first") + z = np.asarray(zone_axis, dtype=float).ravel() + if z.size == 4: + U, V, T, W = z + uvw = np.array([U - T, V - T, W]) + label = "[" + "".join(str(int(round(v))) for v in z) + "]" + elif z.size == 3: + uvw = z + label = "[" + "".join(str(int(round(v))) for v in z) + "]" + else: + raise ValueError(f"zone_axis must have 3 or 4 indices, got {zone_axis}") + keep = torch.as_tensor(np.abs(crystal.hkl.numpy() @ uvw) < 1e-6) + out = copy.copy(crystal) + for name in ("hkl", "g_vec", "g_len", "struct_factors", "struct_factors_int"): + setattr(out, name, getattr(crystal, name)[keep]) + out.name = f"{crystal.name} {label}" + return out + + +def _restrict(crystals, zone_axis): + """Crystal list, each restricted to `zone_axis` when one is given.""" + xtls = [crystals] if isinstance(crystals, Crystal) else list(crystals) + return xtls if zone_axis is None else [zone_reflections(x, zone_axis) for x in xtls] + + +def _hkl_label(hkl: np.ndarray, hexagonal: bool) -> str: + """Compact (hkl) / (hkil) plane label with unicode overbars.""" + + def digit(v: int) -> str: + v = int(round(v)) + txt = str(abs(v)) + return txt + "̅" if v < 0 else txt + + h, k, ll = (int(round(v)) for v in hkl) + if hexagonal: + return "(" + digit(h) + digit(k) + digit(-(h + k)) + digit(ll) + ")" + return "(" + digit(h) + digit(k) + digit(ll) + ")" + + +def _crystal_rings(crystal: Crystal, k_min: float, k_max: float) -> np.ndarray: + """Distinct ring radii of a crystal with non-zero structure factor.""" + g = np.linalg.norm(np.asarray(crystal.g_vec), axis=1) + f = np.asarray(crystal.struct_factors_int) + keep = (g > k_min) & (g < k_max) & (f > 1e-6 * f.max()) + return np.array(sorted(set(np.round(g[keep], 4)))) + + +def _histogram_maxima(peaks, k_min: float, k_max: float, bragg_k_power: float) -> np.ndarray: + """Radii of the local maxima of the measured radial histogram.""" + from scipy.ndimage import gaussian_filter1d, maximum_filter1d + + k, hist = radial_histogram(peaks, k_min=k_min, k_max=k_max, bragg_k_power=bragg_k_power) + h = gaussian_filter1d(hist, 2.0) + loc = (h == maximum_filter1d(h, 15)) & (h > 0.05 * h.max()) + return k[loc] + + +class DiffractionCalibration(AutoSerialize): + """Reciprocal-space calibration of a detector, measured once and reused. + + Holds the reciprocal pixel size, the elliptic distortion and the + diffraction-to-scan rotation, with the evidence behind them. Strained + samples cannot calibrate themselves, so the normal route is to measure + this on a standard such as nanocrystalline gold, save it, and apply it + to the peaks of the sample of interest:: + + cal = calibrate(peaks_au, gold, 0.01) + cal.save("detector_300kV_80cm.zip", mode="o") + ... + cal = load("detector_300kV_80cm.zip") + peaks = cal.apply(bv_centered.peaks) + + A calibration is tied to the detector binning it was measured at, which + is recorded in `metadata`; `rebin` converts it to another binning. + + Parameters + ---------- + pixel_size : float + Reciprocal pixel size, 1/Angstroms per detector pixel. + ellipse : array-like | None + [e11, e12] elliptic distortion correction, see + :func:`calibrate_ellipse`. None applies no correction. + rotation_ccw_deg : float, default=0.0 + Diffraction-to-scan rotation in degrees, see + :func:`measure_scan_rotation`. + metadata : dict | None + Evidence and provenance, e.g. the reference phases, the matched + rings and their residual, and the detector binning. + """ + + def __init__( + self, + pixel_size: float, + ellipse=None, + rotation_ccw_deg: float = 0.0, + metadata: dict | None = None, + ): + self.pixel_size = float(pixel_size) + self.ellipse = None if ellipse is None else np.asarray(ellipse, dtype=float) + self.rotation_ccw_deg = float(rotation_ccw_deg) + self.metadata: dict = dict(metadata or {}) + + def apply(self, peaks_px, name: str = "bragg_peaks_calibrated"): + """Calibrated (qx, qy) peaks from origin-corrected pixel peaks. + + Parameters + ---------- + peaks_px : Vector + Origin-corrected peaks in detector pixels, from + :meth:`~quantem.diffraction.BraggVectors.correct_peak_origins`. + name : str, default="bragg_peaks_calibrated" + Name of the returned Vector. + + Returns + ------- + Vector + Peaks in 1/Angstroms, carrying the measured pixel size and the + scan calibration forward so plots downstream need neither passed + to them. + """ + out = peaks_to_calibrated( + peaks_px, + self.pixel_size, + rotation_ccw_deg=self.rotation_ccw_deg, + ellipse=self.ellipse, + name=name, + ) + # the measured pixel size replaces whatever the dataset was carrying + out.metadata["pixel_size"] = float(self.pixel_size) + out.metadata["pixel_size_units"] = "A^-1" + return out + + def rebin(self, factor: float) -> "DiffractionCalibration": + """The same calibration for data binned by `factor` more than this one. + + Parameters + ---------- + factor : float + Additional detector binning. The pixel size is multiplied by it; + the ellipse and rotation do not depend on binning. + + Returns + ------- + DiffractionCalibration + New calibration with ``metadata["binning"]`` updated. + """ + md = dict(self.metadata) + md["binning"] = md.get("binning", 1) * factor + return DiffractionCalibration( + self.pixel_size * factor, self.ellipse, self.rotation_ccw_deg, md + ) + + def __repr__(self) -> str: + e = "none" if self.ellipse is None else "e11 %+.5f e12 %+.5f" % tuple(self.ellipse) + rms = self.metadata.get("residual_rms") + q = ( + "" + if rms is None + else ", %d rings, rms %.2f%%" + % ( + self.metadata.get("n_rings", 0), + 100 * rms, + ) + ) + return ( + f"DiffractionCalibration(pixel_size={self.pixel_size:.6f} 1/A/px, " + f"ellipse: {e}, rotation {self.rotation_ccw_deg:g} deg{q})" + ) + + +def _ring_profile_scores( + peaks, + crystals, + scales: np.ndarray, + k_broadening: float, + k_min: float, + k_max: float, + bragg_k_power: float, +) -> np.ndarray: + """Normalized overlap of the measured rings with the reference rings, for + each trial scale. + + The score is the fraction of the total measured peak weight that lands on + a reference ring. Two normalizations matter and both are traps. Scaling + the reference rather than the data moves the comparison window with the + scale, and re-normalizing to the weight left inside the window rewards a + scale for pushing peaks out of it: either one lets a wrong scale beat the + truth. Dividing by the total weight, counted once and independent of the + scale, makes a lost peak a loss. + """ + flat = peaks.select_fields("qx", "qy", "intensity").numpy().astype(np.float64) + qr = np.hypot(flat[:, 0], flat[:, 1]) + weight = flat[:, 2] * qr**bragg_k_power + k_step = 0.002 + k = np.arange(k_min, k_max, k_step) + prof = np.zeros_like(k) + for xtl in crystals: + prof += simulated_ring_profile(xtl, k, k_broadening, bragg_k_power) + prof = prof / max(prof.max(), 1e-12) + total = float(weight.sum()) + scores = np.zeros_like(scales) + for i, sc in enumerate(scales): + frac = (qr * sc - k_min) / k_step + i0 = np.floor(frac).astype(int) + w1 = frac - i0 + ok = (i0 >= 0) & (i0 < k.size - 1) + h = np.bincount(i0[ok], weights=weight[ok] * (1 - w1[ok]), minlength=k.size) + h += np.bincount(i0[ok] + 1, weights=weight[ok] * w1[ok], minlength=k.size) + scores[i] = float((h[: k.size] * prof).sum() / total) if total > 0 else 0.0 + return scores + + +def _fit_scale(peaks, crystals, lo, hi, k_broadening, k_min, k_max, bragg_k_power, n=241): + """Best scale in [lo, hi] with parabolic refinement of the maximum.""" + scales = np.linspace(lo, hi, n) + scores = _ring_profile_scores( + peaks, crystals, scales, k_broadening, k_min, k_max, bragg_k_power + ) + i = int(np.argmax(scores)) + best = float(scales[i]) + if 0 < i < n - 1: + c0, c1, c2 = scores[i - 1 : i + 2] + denom = 4 * c1 - 2 * c0 - 2 * c2 + if abs(denom) > 1e-12: + best += (c2 - c0) / denom * (scales[1] - scales[0]) + return best, scales, scores + + +def calibrate( + peaks_px, + crystal, + pixel_size_guess: float, + rotation_ccw_deg: float = 0.0, + fit_ellipse: bool = True, + n_iter: int = 3, + scale_search: tuple[float, float] = (0.6, 1.7), + k_min: float = 0.05, + k_max: float = 1.3, + k_broadening: float = 0.01, + bragg_k_power: float = 2.0, + residual_tol: float = 0.01, + plot: bool = False, + figsize: tuple[float, float] = (13.0, 6.4), + marker_size: float = 8.0, + zone_axis=None, + returnfig: bool = False, +): + """Measure the reciprocal pixel size and the elliptic distortion. + + Three stages, each one removing the reason the next could fail. A coarse + scan over `scale_search` with a deliberately broadened ring profile finds + the right ring assignment even when the starting pixel size is far out; + a broad profile has one maximum where a sharp one has many. The ellipse + and the scale are then refined in turn, the ellipse in log-radius bins + where it does not depend on the scale, and the scale against + progressively sharper rings. Finally every measured ring is matched to + its reference ring separately, which is the only check that can tell a + correct calibration from a plausible one: a single pixel size that + explains the pattern gives per-ring scale factors agreeing to a few + tenths of a percent with no trend in k. + + Parameters + ---------- + peaks_px : Vector + Origin-corrected peaks in detector pixels, (q_row, q_col, intensity). + crystal : Crystal | list[Crystal] + Reference phase or phases, structure factors calculated. + pixel_size_guess : float + Starting reciprocal pixel size (1/Angstroms per pixel). Only the + order of magnitude matters; the coarse scan covers `scale_search`. + rotation_ccw_deg : float, default=0.0 + Diffraction-to-scan rotation recorded on the calibration. + fit_ellipse : bool, default=True + Fit the elliptic distortion as well as the scale. + n_iter : int, default=3 + Ellipse and scale refinement rounds. Each round fits the ellipse + left over after the current correction and composes the two. + scale_search : tuple, default=(0.6, 1.7) + Capture range of the coarse scan, as a multiple of the guess. + k_min, k_max : float, default=0.05, 1.3 + Range of scattering vectors fit, 1/Angstroms. The ellipse fit uses + at least 0.15 as its lower limit. + k_broadening : float, default=0.01 + Gaussian standard deviation of the reference rings in the fine scale + fit, 1/Angstroms; the coarse scan uses six times this. Measured rings + further than four times this from any reference ring are not counted + in the per-ring check. + bragg_k_power : float, default=2.0 + Peaks and rings are weighted by ``|q| ** bragg_k_power``. + residual_tol : float, default=0.01 + Per-ring residual rms above which the fit is reported as unreliable. + plot : bool, default=False + Show the ring comparison before and after the fit, for the first + reference crystal. + figsize : tuple, default=(13, 6.4) + Figure size. + marker_size : float, default=8.0 + Area of the brightest peak in the azimuth panels, where every peak is + drawn with area proportional to its intensity. Raise it to bring out + weak spots, lower it when strong ones hide the reference lines. + zone_axis : sequence of int, optional + Fit only the rings of this zone, [uvw] or [UVTW] (see + :func:`zone_reflections`). For a specimen near one zone axis + everywhere, whose peaks hold no other rings. + returnfig : bool, default=False + With `plot`, also return the figure and axes. + + Returns + ------- + DiffractionCalibration + The pixel size, ellipse and rotation. ``metadata`` holds the matched + rings as (measured k, reference k, ratio), their residual rms and + whether the fit passed `residual_tol`. With `plot` and `returnfig`, + ``(cal, fig, axs)`` instead, axs of shape (2, 2). + + Warns + ----- + UserWarning + If fewer than three rings match or their residual rms exceeds + `residual_tol`. + + Notes + ----- + The figure shows the first reference crystal before and after the fit. + To check every candidate phase against the result, including phases the + calibration was not fit to, plot :func:`plot_calibration` on the + calibrated peaks, or :meth:`CrystalMap.plot_calibration`. + """ + crystals = [crystal] if isinstance(crystal, Crystal) else list(crystal) + for xtl in crystals: + if xtl.g_vec is None: + raise RuntimeError(f"{xtl.name}: run calculate_structure_factors() first") + crystals = _restrict(crystals, zone_axis) + + peaks_0 = peaks_to_calibrated(peaks_px, pixel_size_guess) + # coarse: broad rings so the score has a single maximum over a wide range + scale, _, _ = _fit_scale( + peaks_0, + crystals, + scale_search[0], + scale_search[1], + 6 * k_broadening, + k_min, + k_max, + bragg_k_power, + ) + ellipse = None + for it in range(max(1, n_iter)): + if fit_ellipse: + pk = peaks_to_calibrated(peaks_px, pixel_size_guess * scale, ellipse=ellipse) + # the fit sees peaks with the current ellipse already applied, so + # it returns the residual distortion: compose it with the current + # correction instead of replacing it + residual = calibrate_ellipse( + pk, k_min=max(k_min, 0.15), k_max=k_max, bragg_k_power=bragg_k_power + ) + ellipse = _compose_ellipse(residual, ellipse) + pk = peaks_to_calibrated(peaks_px, pixel_size_guess, ellipse=ellipse) + half = 0.08 / (it + 1) + scale, _, _ = _fit_scale( + pk, + crystals, + scale * (1 - half), + scale * (1 + half), + k_broadening, + k_min, + k_max, + bragg_k_power, + ) + + pixel_size = pixel_size_guess * scale + peaks = peaks_to_calibrated( + peaks_px, pixel_size, rotation_ccw_deg=rotation_ccw_deg, ellipse=ellipse + ) + + rings = np.concatenate([_crystal_rings(x, k_min, k_max * 1.15) for x in crystals]) + rings = np.array(sorted(set(np.round(rings, 4)))) + k_meas = _histogram_maxima(peaks, k_min, k_max * 1.15, bragg_k_power) + table = [] + for km in k_meas: + j = int(np.argmin(np.abs(rings - km))) + if abs(rings[j] - km) < 4 * k_broadening: + table.append((float(km), float(rings[j]), float(rings[j] / km))) + per_ring = np.array([t[2] for t in table]) if table else np.array([np.nan]) + residual_rms = float(np.std(per_ring / np.median(per_ring))) if table else float("nan") + reliable = len(table) >= 3 and residual_rms <= residual_tol + + cal = DiffractionCalibration( + pixel_size, + ellipse, + rotation_ccw_deg, + metadata=dict( + reference=[x.name for x in crystals], + pixel_size_guess=float(pixel_size_guess), + scale=float(scale), + n_rings=len(table), + residual_rms=residual_rms, + rings=table, + reliable=bool(reliable), + binning=1, + ), + ) + if not reliable: + warnings.warn( + f"calibration looks unreliable: {len(table)} rings matched, residual rms " + f"{100 * residual_rms:.2f}% (tolerance {100 * residual_tol:.2f}%). Check the " + "reference phase and the peak detection before using this pixel size.", + stacklevel=2, + ) + if not plot: + return cal + + import matplotlib.pyplot as plt + + k_hi = k_max * 1.15 + fig, axs = plt.subplots( + 2, + 2, + figsize=figsize, + sharex="col", + gridspec_kw={"height_ratios": [1, 1.3]}, + ) + for col, (pk, ttl, kb) in enumerate( + ( + (peaks_0, f"before: {pixel_size_guess:.5f} " + r"$\mathrm{\AA}^{-1}$/px", None), + (peaks, f"after: {pixel_size:.5f} " + r"$\mathrm{\AA}^{-1}$/px", k_broadening), + ) + ): + _calibration_panels( + axs[0, col], + axs[1, col], + pk, + crystals[0], + k_min=k_min, + k_max=k_hi, + k_broadening=kb, + bragg_k_power=bragg_k_power, + marker_size=marker_size, + ) + axs[0, col].set_title(ttl, fontsize=10) + axs[1, 1].set_ylabel("") + if ellipse is not None: + axs[1, 0].text( + 0.02, + 0.97, + "ellipse e11 %+.4f e12 %+.4f" % tuple(ellipse), + transform=axs[1, 0].transAxes, + va="top", + fontsize=9, + ) + axs[1, 1].text( + 0.02, + 0.97, + f"{len(table)} rings, rms {100 * residual_rms:.2f}%" + + ("" if reliable else " UNRELIABLE"), + transform=axs[1, 1].transAxes, + va="top", + fontsize=9, + color="k" if reliable else "tab:red", + ) + fig.tight_layout() + fig.subplots_adjust(hspace=0.08) + return (cal, fig, axs) if returnfig else cal + + +# reflections closer than this in |g| (1/Angstroms) are drawn as one ring +_SHELL_STEP = 0.005 + + +def _shells(crystal: Crystal, bragg_k_power: float): + """Group a crystal's reflections into rings. + + Returns + ------- + radius : np.ndarray + Mean |g| of each ring, 1/Angstroms, ascending. + intensity : np.ndarray + Summed intensity of each ring, weighted by ``|g| ** bragg_k_power``. + index : np.ndarray + Ring index of every reflection. + """ + g = crystal.g_len.numpy() + ints = crystal.struct_factors_int.numpy() * g**bragg_k_power + _, index = np.unique(np.round(g / _SHELL_STEP).astype(np.int64), return_inverse=True) + index = index.ravel() + count = np.bincount(index) + radius = np.bincount(index, weights=g) / count + intensity = np.bincount(index, weights=ints) + return radius, intensity, index + + +def _ring_shells(crystal: Crystal, k_min: float, k_max: float, bragg_k_power: float): + """Ring radii of a crystal in (k_min, k_max) and their summed intensity, + relative to the strongest ring in that range.""" + radius, intensity, _ = _shells(crystal, bragg_k_power) + keep = (radius > k_min) & (radius < k_max) + uniq, tot = radius[keep], intensity[keep] + return uniq, tot / max(float(tot.max()), 1e-30) if tot.size else tot + + +def _calibration_panels( + ax_hist, + ax_az, + peaks, + crystal: Crystal, + k_min: float, + k_max: float, + k_broadening: float | None = None, + bragg_k_power: float = 2.0, + marker_size: float = 8.0, + n_rings: int = 12, +) -> None: + """The radial histogram against one crystal's rings (top) and every peak + as azimuth against scattering vector with its strongest rings (bottom).""" + plot_ring_comparison( + peaks, + [crystal], + k_min=k_min, + k_max=k_max, + k_broadening=k_broadening, + bragg_k_power=bragg_k_power, + n_labels=n_rings, + figax=(ax_hist.figure, [ax_hist]), + ) + ax_hist.set_xlabel("") + # azimuth against scattering vector: a pixel size error shifts every + # ring, the elliptic distortion is a cos(2 phi) wobble of each one + flat = peaks.select_fields("qx", "qy", "intensity").numpy().astype(np.float64) + r = np.hypot(flat[:, 0], flat[:, 1]) + phi = np.degrees(np.arctan2(flat[:, 1], flat[:, 0])) + sel = (r > k_min) & (r < k_max) + w = flat[sel, 2] + hi = float(np.percentile(w, 99.5)) if w.size else 1.0 + ax_az.scatter( + r[sel], + phi[sel], + s=marker_size * np.clip(w / max(hi, 1e-12), 0.03, 1.0), + c="r", + alpha=0.5, + lw=0, + rasterized=True, + ) + radii, rel = _ring_shells(crystal, k_min, k_max, bragg_k_power) + for g0 in radii[np.argsort(-rel)[:n_rings]]: + ax_az.axvline(g0, color="k", lw=0.7, alpha=0.8) + ax_az.set_xlim(k_min, k_max) + ax_az.set_ylim(-180, 180) + ax_az.set_yticks([-180, -90, 0, 90, 180]) + ax_az.set_xlabel(r"scattering vector (1/$\mathrm{\AA}$)") + ax_az.set_ylabel("azimuth (deg)") + + +def plot_calibration( + peaks, + crystals, + k_min: float = 0.05, + k_max: float = 1.5, + k_broadening: float | None = None, + bragg_k_power: float = 2.0, + marker_size: float = 8.0, + n_rings: int = 12, + zone_axis=None, + axsize: tuple[float, float] = (6.0, 6.4), + figax=None, +): + """Calibrated peaks against the rings of every crystal, one column each. + + The calibration check for all candidate phases, however many of them the + calibration itself was fit to. Each column holds the radial histogram of + every peak against that crystal's rings (top) and every peak as azimuth + against scattering vector with the same rings (bottom). A ring beside the + measured peaks in every direction is a lattice parameter or pixel size + error; a ring that wobbles with azimuth is elliptic distortion. + + Parameters + ---------- + peaks : Vector + Calibrated peaks (qx, qy, intensity) in 1/Angstroms. + crystals : Crystal | list[Crystal] + Candidate phases, structure factors calculated. + k_min, k_max : float + Range of scattering vectors shown, 1/Angstroms. + k_broadening : float | None + None draws sharp ring lines in the histograms; a width (1/Angstroms) + draws the broadened ring profile instead. + marker_size : float, default=8.0 + Area of the brightest peak in the azimuth panels. + n_rings : int, default=12 + Rings drawn in the azimuth panels and labeled in the histograms, + strongest first; a large cell would otherwise fill both with its + weak superstructure rings. + zone_axis : sequence of int, optional + Show only the rings of this zone, [uvw] or [UVTW] (see + :func:`zone_reflections`). + axsize : tuple, default=(6.0, 6.4) + Size of one column. + figax : (fig, axs) | None + Existing figure and a (2, n_crystals) array of axes. + + Returns + ------- + tuple + ``(fig, axs)``, axs of shape (2, n_crystals). + """ + import matplotlib.pyplot as plt + + xtls = _restrict(crystals, zone_axis) + if figax is None: + fig, axs = plt.subplots( + 2, + len(xtls), + figsize=(axsize[0] * len(xtls), axsize[1]), + sharex=True, + squeeze=False, + gridspec_kw={"height_ratios": [1, 1.3]}, + ) + else: + fig, axs = figax + axs = np.asarray(axs).reshape(2, len(xtls)) + for col, xtl in enumerate(xtls): + _calibration_panels( + axs[0, col], + axs[1, col], + peaks, + xtl, + k_min=k_min, + k_max=k_max, + k_broadening=k_broadening, + bragg_k_power=bragg_k_power, + marker_size=marker_size, + n_rings=n_rings, + ) + axs[0, col].set_title(xtl.name, fontsize=10) + if col: + axs[0, col].set_ylabel("") + axs[1, col].set_ylabel("") + if figax is None: + fig.tight_layout() + fig.subplots_adjust(hspace=0.08) + return fig, axs + + +def plot_ring_comparison( + peaks, + crystals, + k_min: float = 0.1, + k_max: float = 1.5, + k_broadening: float | None = None, + bragg_k_power: float = 2.0, + label_hkl: bool = True, + label_min_intensity: float = 0.05, + n_labels: int = 12, + zone_axis=None, + figax=None, +): + """Measured radial peak histogram against crystal ring positions. + + One panel per crystal: the measured histogram is the red fill, the + crystal's rings are black -- sharp vertical lines by default, or a + Gaussian profile of width `k_broadening` when set (use after + calibration, where the rings should sit inside the measured peaks). The + strongest rings are labeled by (hkl), 4-index (hkil) for hexagonal + crystals. + + Parameters + ---------- + peaks : Vector + Calibrated peaks (qx, qy, intensity) in 1/Angstroms. + crystals : Crystal | list[Crystal] + Reference crystal(s) with structure factors calculated. + k_min, k_max : float, default=0.1, 1.5 + Range of scattering vectors shown, 1/Angstroms. + k_broadening : float | None + None draws sharp lines at the ring positions; a value (1/Angstroms) + draws the broadened ring profile instead. + bragg_k_power : float, default=2.0 + Measured peaks and reference rings are weighted by + ``|q| ** bragg_k_power``. + label_hkl : bool, default=True + Label the strongest rings by their Miller indices. + label_min_intensity : float, default=0.05 + Label rings whose summed intensity exceeds this fraction of the + strongest ring. + n_labels : int, default=12 + Label at most this many rings, strongest first, which keeps a large + cell with many rings readable. + zone_axis : sequence of int, optional + Show only the rings of this zone, [uvw] or [UVTW] (see + :func:`zone_reflections`). + figax : (fig, axs) | None + Existing figure and one axis per crystal. + + Returns + ------- + tuple + ``(fig, axs)``, axs a 1D array with one axis per crystal. + """ + import matplotlib.pyplot as plt + + xtls = _restrict(crystals, zone_axis) + k, hist = radial_histogram(peaks, k_min=k_min, k_max=k_max, bragg_k_power=bragg_k_power) + + n = len(xtls) + if figax is None: + fig, axs = plt.subplots(n, 1, figsize=(11, 3.6 * n), sharex=True, squeeze=False) + axs = axs[:, 0] + else: + fig, axs = figax + axs = np.atleast_1d(axs) + + for ax, xtl in zip(axs, xtls): + ax.fill_between(k, hist / hist.max(), color="r", alpha=0.75, lw=0, label="measured") + hexagonal = xtl.hexagonal_matching + g_len = xtl.g_len.numpy() + ints = xtl.struct_factors_int.numpy() * g_len**bragg_k_power + hkl_np = xtl.hkl.numpy() + uniq, shell_int, shells = _shells(xtl, bragg_k_power) + shell_int = shell_int / shell_int.max() + + if k_broadening is not None: + prof = simulated_ring_profile(xtl, k, k_broadening, bragg_k_power) + ax.plot(k, prof / prof.max(), "k-", lw=1.0, label=f"{xtl.name} rings") + else: + keep = (uniq > k_min) & (uniq < k_max) + ax.vlines( + uniq[keep], + 0, + shell_int[keep], + colors="k", + lw=1.0, + label=f"{xtl.name} rings", + ) + + if label_hkl: + # the strongest rings first, each at least 0.03 1/A from the + # last, then drawn left to right + labeled: list[int] = [] + for j in np.argsort(-shell_int): + u, si = uniq[j], shell_int[j] + if len(labeled) >= n_labels: + break + if u < k_min or u > k_max or si < label_min_intensity: + continue + if any(abs(u - uniq[j0]) < 0.03 for j0 in labeled): + continue + labeled.append(int(j)) + for rows, j in enumerate(sorted(labeled)): + u = uniq[j] + idx = np.nonzero(shells == j)[0] + idx = idx[ints[idx] > 0.99 * ints[idx].max()] + key = [tuple(-hkl_np[i]) for i in idx] + best = idx[int(np.lexsort(np.array(key).T[::-1])[0])] + # three staggered rows keep neighbouring labels apart + y = 1.04 + 0.1 * (rows % 3) + ax.text( + u, + y, + _hkl_label(hkl_np[best], hexagonal), + fontsize=8, + ha="center", + va="bottom", + ) + ax.set_ylabel("intensity (norm.)") + ax.set_ylim(0, 1.42) + # below the band of ring labels at the top + ax.legend(loc="upper right", bbox_to_anchor=(1.0, 0.72), fontsize=8) + axs[-1].set_xlabel(r"scattering vector (1/$\mathrm{\AA}$)") + if figax is None: + fig.tight_layout() + return fig, axs + + +def transform_peaks(peaks, M: np.ndarray): + """Return a copy of a (qx, qy, intensity) Vector with q mapped by M (2x2).""" + out = peaks.copy() + flat = out.numpy().astype(np.float64) + flat[:, :2] = flat[:, :2] @ np.asarray(M, dtype=float).T + out.set_flattened(flat) + return out + + +def refine_calibration( + strain_maps, + masks=None, + max_strain: float = 0.05, +): + """Global calibration residual from matched orientations. + + The per-position deformation A fitted by + OrientationMap.calculate_strain() maps ideal simulated peaks onto the + measured ones, so it contains both the local strain and any global + calibration error. The element-wise median of A over many differently + oriented grains (across all phases) averages the strain away and leaves + the calibration residual: scale and ellipticity. A global detector + rotation is NOT observable this way -- the in-plane refinement absorbs + it into every orientation, so the reported rotation_deg is ~0 by + construction; measure the scan rotation independently + (measure_scan_rotation, or a known texture). Apply the returned + correction with transform_peaks() and re-match to close the loop. + + Parameters + ---------- + strain_maps : list[StrainMap] + One per phase, from calculate_strain() on the SAME calibrated peaks. + masks : list[np.ndarray] | None + Per-phase inclusion masks (e.g. phase == i and reliable); defaults + to all positions where the strain fit succeeded. + max_strain : float, default=0.05 + Discard positions whose deformation differs from the identity by + more than this (failed fits, overlaps). + + Returns + ------- + dict with: + 'M' : the median deformation (2, 2), + 'correction' : inv(M), ready for transform_peaks(), + 'scale' : multiply the pixel size by this, + 'rotation_deg' : residual detector rotation, + 'ellipse' : (e11, e12) traceless ellipticity components, + 'num_positions' : positions used. + + Raises + ------ + ValueError + If `strain_maps` is empty or no position passes the filters. + """ + As = [] + for i, sm in enumerate(strain_maps): + A = np.stack([sm.g1_array, sm.g2_array], axis=-1) # (R, C, 2, 2) + ok = np.isfinite(A).all(axis=(-2, -1)) + dev = np.abs(A - np.eye(2)).max(axis=(-2, -1)) + ok &= dev < max_strain + if masks is not None and masks[i] is not None: + ok &= np.asarray(masks[i]) > 0 + As.append(A[ok]) + if not As: + raise ValueError("strain_maps is empty: pass at least one StrainMap.") + A_all = np.concatenate(As, axis=0) + if A_all.shape[0] == 0: + raise ValueError( + "no positions left for the calibration residual: every strain fit " + f"failed, was masked out, or deviates from the identity by more than " + f"max_strain={max_strain:g}." + ) + M = np.median(A_all, axis=0) + + scale = float(np.sqrt(np.abs(np.linalg.det(M)))) + theta = 0.5 * (M[1, 0] - M[0, 1]) / scale + sym = 0.5 * (M + M.T) / scale + e11 = float(0.5 * (sym[0, 0] - sym[1, 1])) + e12 = float(sym[0, 1]) + return { + "M": M, + "correction": np.linalg.inv(M), + "scale": scale, + "rotation_deg": float(np.rad2deg(theta)), + "ellipse": (e11, e12), + "num_positions": int(A_all.shape[0]), + } + + +def plot_bragg_rings( + peaks, + crystals, + n_rings: int = 8, + q_max: float | None = None, + bins: int = 400, + power: float = 0.25, + zone_axis=None, + figax=None, +): + """2D histogram of all Bragg peaks with crystal rings overlaid. + + The Bragg vector map (histogram of every detected peak over the scan) + shows the calibration directly in 2D: the crystal's strongest rings are + drawn as thin circles, which should thread through the measured spot + density -- a radius mismatch is a pixel size error, and a direction- + dependent mismatch is elliptic distortion. + + Parameters + ---------- + peaks : Vector + Calibrated peaks (qx, qy, intensity) in 1/Angstroms. + crystals : Crystal | list[Crystal] + Reference crystal(s); the n_rings strongest rings of each are drawn + (solid, then dashed line styles). + n_rings : int, default=8 + Number of rings per crystal, strongest first. + q_max : float | None + Half-width of the histogram, 1/Angstroms. Defaults to just beyond + the largest peak radius. + bins : int, default=400 + Number of histogram bins along each axis. + power : float, default=0.25 + The histogram is shown raised to this power, which brings out weak + rings next to the direct beam. + zone_axis : sequence of int, optional + Draw only the rings of this zone, [uvw] or [UVTW] (see + :func:`zone_reflections`). + figax : (fig, ax) | None + Existing figure and axis. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + import matplotlib.pyplot as plt + + xtls = _restrict(crystals, zone_axis) + flat = peaks.select_fields("qx", "qy", "intensity").numpy().astype(np.float64) + if q_max is None: + q_max = float(np.hypot(flat[:, 0], flat[:, 1]).max()) * 1.02 + H, xe, ye = np.histogram2d( + flat[:, 0], + flat[:, 1], + bins=bins, + range=[[-q_max, q_max], [-q_max, q_max]], + ) + + if figax is None: + fig, ax = plt.subplots(figsize=(7.5, 7.5)) + else: + fig, ax = figax + ax.imshow( + H**power, + cmap="gray_r", + extent=(ye[0], ye[-1], xe[-1], xe[0]), + interpolation="nearest", + ) + styles = ["-", "--", ":"] + colors = ["r", "b", "g"] + th = np.linspace(0, 2 * np.pi, 361) + for ci, xtl in enumerate(xtls): + uniq, shell_int, _ = _shells(xtl, 2.0) + keep = uniq < q_max + uniq, shell_int = uniq[keep], shell_int[keep] + order = np.argsort(shell_int)[::-1][:n_rings] + for k, u in enumerate(np.sort(uniq[order])): + ax.plot( + u * np.sin(th), + u * np.cos(th), + ls=styles[ci % 3], + color=colors[ci % 3], + lw=0.5, + alpha=0.6, + label=f"{xtl.name} rings" if k == 0 else None, + ) + ax.set_xlabel(r"$q_c$ (1/$\mathrm{\AA}$)") + ax.set_ylabel(r"$q_r$ (1/$\mathrm{\AA}$)") + ax.legend(loc="upper right", fontsize=9) + return fig, ax diff --git a/src/quantem/diffraction/crystal.py b/src/quantem/diffraction/crystal.py new file mode 100644 index 000000000..51f938365 --- /dev/null +++ b/src/quantem/diffraction/crystal.py @@ -0,0 +1,1396 @@ +"""Crystal structures and kinematical diffraction for orientation mapping. + +A Crystal wraps an ase.Atoms object and computes the reciprocal lattice, +kinematical structure factors, symmetry operators (via spglib), and simulated +diffraction patterns for arbitrary orientations. All numerical state is stored +as torch tensors (float64) so downstream matching and refinement can run on +GPU and differentiate through the calculation. + +Conventions +----------- +- Real lattice vectors are rows of `lat_real` (Angstroms). +- Reciprocal lattice vectors are rows of `lat_recip` (1/Angstroms, no 2*pi). +- Structure factors follow F_hkl = (1/V) * sum_n f_n * exp(-2*pi*i * hkl.p_n), + so intensities have units of scattering amplitude per unit volume. +- Orientations are unit quaternions rotating crystal Cartesian vectors into + the lab frame (see quantem.diffraction.rotations). +""" + +from __future__ import annotations + +import json +import warnings +from contextlib import contextmanager +from importlib import resources +from pathlib import Path + +import numpy as np +import torch +from ase import Atoms +from ase.data import chemical_symbols + +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.defaults import SIGMA_EXCITATION +from quantem.diffraction.rotations import qrotate, symmetry_quaternions + +# unicode combining overline, applies to the preceding character +_B = "\u0305" + +_EXCITATION_MODELS = ("gaussian", "slab") + + +@contextmanager +def _spglib_raises(): + """Within the block, spglib raises SpglibError on failure. + + spglib 2.7 reports failures by returning None and emitting a + DeprecationWarning on every call, success or not, unless + ``spglib.error.OLD_ERROR_HANDLING`` is False, which spglib 2.8 makes the + default. Every spglib call here is wrapped in try/except, so the flag is + switched for the block only and restored afterwards, leaving other + spglib users in the process unaffected. Older spglib versions without + the flag are left as they are. + """ + import spglib + + err = getattr(spglib, "error", None) + if err is None or not hasattr(err, "OLD_ERROR_HANDLING"): + yield + return + old = err.OLD_ERROR_HANDLING + err.OLD_ERROR_HANDLING = False + try: + yield + finally: + err.OLD_ERROR_HANDLING = old + + +def direction_indices( + lat_real: torch.Tensor | np.ndarray, d, max_multiple: int = 12, atol: float = 2e-3 +) -> np.ndarray | None: + """Smallest integer [uvw] along a Cartesian direction. + + Parameters + ---------- + lat_real : torch.Tensor | np.ndarray + Real-space lattice vectors as rows (3, 3), Angstroms. + d : array-like + Cartesian direction (3,) in the crystal frame; need not be + normalized. + max_multiple : int, default=12 + Largest multiplier tried to make the indices integer. + atol : float, default=2e-3 + Allowed deviation of the scaled indices from integers. Loosen it to + index the axes of a pseudo-symmetry, which are lattice directions of + an ideal parent but only nearly so in the real cell. + + Returns + ------- + np.ndarray | None + Integer [uvw] (3,) with no common factor, or None if `d` is not a + lattice direction with indices up to `max_multiple`. + """ + A_T_inv = np.linalg.inv(np.asarray(lat_real, dtype=float).T) + v = A_T_inv @ np.asarray(d, dtype=float) + v = v / np.abs(v).max() + for m in range(1, max_multiple + 1): + w = v * m + if np.allclose(w, np.round(w), atol=atol): + ints = np.round(w).astype(int) + g = np.gcd.reduce(np.abs(ints)) + return ints // max(g, 1) + return None + + +def format_direction(uvw, hexagonal: bool = False, mathtext: bool = True) -> str: + """Direction label such as [011] or [10-10], with overlines on negatives. + + Parameters + ---------- + uvw : array-like | None + Integer 3-index direction [uvw]. + hexagonal : bool, default=False + Write the 4-index [UVTW] symbol instead, see + :func:`miller_to_miller_bravais`. + mathtext : bool, default=True + Overlines as matplotlib mathtext (``$\\bar{1}$``) for figures; False + uses unicode combining overlines for plain text. + + Returns + ------- + str + The label, or an empty string for None. + """ + if uvw is None: + return "" + ks = miller_to_miller_bravais(uvw) if hexagonal else np.asarray(uvw) + ks = np.atleast_1d(ks) + if mathtext: + body = "".join(str(k) if k >= 0 else "$\\bar{%d}$" % -k for k in ks) + else: + body = "".join(str(k) if k >= 0 else "%d%s" % (-k, _B) for k in ks) + return "[" + body + "]" + + +def miller_to_miller_bravais(uvw: np.ndarray) -> np.ndarray: + """Convert 3-index [u'v'w'] direction indices to 4-index [u v t w]. + + u = (2u' - v') / 3, v = (2v' - u') / 3, t = -(u + v), w = w', cleared to + the smallest integer form. + + Parameters + ---------- + uvw : array-like + Integer 3-index directions, (3,) or (N, 3). + + Returns + ------- + np.ndarray + Integer [u v t w], (4,) or (N, 4). + """ + uvw = np.atleast_2d(np.asarray(uvw, dtype=float)) + u = (2 * uvw[:, 0] - uvw[:, 1]) / 3 + v = (2 * uvw[:, 1] - uvw[:, 0]) / 3 + out = np.stack([u, v, -(u + v), uvw[:, 2]], axis=1) + # clear fractions and common factors + out = out * 3 + gcd = np.gcd.reduce(np.abs(np.round(out)).astype(int), axis=1) + gcd[gcd == 0] = 1 + out = out / gcd[:, None] + return np.rint(out).astype(int).squeeze() + + +def miller_bravais_to_miller(uvtw: np.ndarray) -> np.ndarray: + """Convert 4-index [u v t w] direction indices to 3-index [u'v'w']. + + u' = 2u + v, v' = 2v + u, w' = w (t is redundant: t = -(u + v)), + cleared to the smallest integer form. + + Parameters + ---------- + uvtw : array-like + Integer 4-index directions, (4,) or (N, 4). + + Returns + ------- + np.ndarray + Integer [u'v'w'], (3,) or (N, 3). + """ + uvtw = np.atleast_2d(np.asarray(uvtw, dtype=float)) + out = np.stack([2 * uvtw[:, 0] + uvtw[:, 1], 2 * uvtw[:, 1] + uvtw[:, 0], uvtw[:, 3]], axis=1) + gcd = np.gcd.reduce(np.abs(np.round(out)).astype(int), axis=1) + gcd[gcd == 0] = 1 + return np.rint(out / gcd[:, None]).astype(int).squeeze() + + +# point group -> Laue class +_LAUE_CLASS = { + "1": "-1", + "-1": "-1", + "2": "2/m", + "m": "2/m", + "2/m": "2/m", + "222": "mmm", + "mm2": "mmm", + "mmm": "mmm", + "4": "4/m", + "-4": "4/m", + "4/m": "4/m", + "422": "4/mmm", + "4mm": "4/mmm", + "-42m": "4/mmm", + "4/mmm": "4/mmm", + "3": "-3", + "-3": "-3", + "32": "-3m", + "3m": "-3m", + "-3m": "-3m", + "6": "6/m", + "-6": "6/m", + "6/m": "6/m", + "622": "6/mmm", + "6mm": "6/mmm", + "-6m2": "6/mmm", + "6/mmm": "6/mmm", + "23": "m-3", + "m-3": "m-3", + "432": "m-3m", + "-43m": "m-3m", + "m-3m": "m-3m", +} + + +def _load_lobato_params() -> dict[str, np.ndarray]: + with resources.files("quantem.diffraction").joinpath("data/lobato.json").open() as f: + raw = json.load(f) + return {sym: np.array(p) for sym, p in raw.items()} + + +_LOBATO: dict[str, np.ndarray] | None = None + + +def electron_scattering_factor(numbers: torch.Tensor, g: torch.Tensor) -> torch.Tensor: + """Lobato & Van Dyck (2014) electron scattering factors. + + Parameters + ---------- + numbers : torch.Tensor + Atomic numbers (N,). + g : torch.Tensor + Scattering vector magnitudes (M,) in 1/Angstroms. + + Returns + ------- + torch.Tensor + f_e(g) of shape (N, M) in Angstroms. + """ + global _LOBATO + if _LOBATO is None: + _LOBATO = _load_lobato_params() + g2 = (g**2)[None, :, None] # (1, M, 5) + a = torch.stack( + [ + torch.as_tensor(_LOBATO[chemical_symbols[int(z)]][0], dtype=g.dtype, device=g.device) + for z in numbers + ] + )[:, None, :] # (N, 1, 5) + b = torch.stack( + [ + torch.as_tensor(_LOBATO[chemical_symbols[int(z)]][1], dtype=g.dtype, device=g.device) + for z in numbers + ] + )[:, None, :] + return (a * (2.0 + b * g2) / (1.0 + b * g2) ** 2).sum(dim=-1) + + +def _expand_partial_occupancy(atoms: Atoms) -> Atoms: + """One atom per species on every shared site, with its fractional occupancy. + + ASE's CIF reader keeps the majority species of a mixed site and records + the full composition in ``atoms.info['occupancy']``, keyed by the site + index held in ``atoms.arrays['spacegroup_kinds']``. Structures with only + fully occupied sites are returned unchanged. + + Parameters + ---------- + atoms : Atoms + As returned by ``ase.io.read`` on a CIF. + + Returns + ------- + Atoms + The expanded structure, with ``arrays['occupancy']`` set. + """ + occ = atoms.info.get("occupancy") + kinds = atoms.arrays.get("spacegroup_kinds") + if not occ or kinds is None: + return atoms + sites = [occ.get(str(int(k))) for k in kinds] + if all( + s is None or (len(s) == 1 and abs(float(next(iter(s.values()))) - 1.0) < 1e-9) + for s in sites + ): + return atoms + + frac = atoms.get_scaled_positions(wrap=False) + symbols, positions, occupancy = [], [], [] + for i, site in enumerate(sites): + if not site: + symbols.append(atoms[i].symbol) + positions.append(frac[i]) + occupancy.append(1.0) + continue + for element, fraction in site.items(): + symbols.append(element) + positions.append(frac[i]) + occupancy.append(float(fraction)) + out = Atoms(symbols=symbols, scaled_positions=positions, cell=atoms.cell, pbc=atoms.pbc) + out.set_array("occupancy", np.asarray(occupancy, dtype=float)) + out.info.update({k: v for k, v in atoms.info.items() if k != "occupancy"}) + return out + + +class Crystal(AutoSerialize): + """A crystal structure with kinematical diffraction methods. + + Build with `from_ase` or `from_cif`, then call + `calculate_structure_factors` before generating patterns or orientation + plans. + + Parameters + ---------- + atoms : ase.Atoms + The structure. Fractional site occupancies are read from + ``atoms.arrays['occupancy']`` when present (see :meth:`from_cif`). + name : str | None + Display name; defaults to the chemical formula. + symprec : float, default=1e-4 + spglib tolerance (Angstroms) for the cell's own symmetry. + pseudo_symmetry_tol : float | None, default=0.01 + Dimensionless distance tolerance for the symmetry used in + orientation matching: a fraction of the shortest lattice vector + within which atoms and lattice vectors are allowed to deviate from a + higher-symmetry parent (a 4 A cell with an atom at + (0.5, 0.5, 0.50001) is body centered at any tolerance above 1e-5). + Cells within it are matched with the parent group, so variants no + experiment can separate are never sampled as distinct orientations; + the library builders warn when the matching group differs from the + cell's own. None matches with the exact symmetry. + pseudo_symmetry_intensity_tol : float, default=0.05 + Largest intensity difference allowed between reflections that a + candidate pseudo-symmetry would make equivalent, as a fraction of + the strongest reflection's kinematical intensity |F|^2. Each extra + rotation is applied to every reflection within 2.0 1/A and each + intensity compared with its image's; if any pair differs by more + than this fraction, the orientations the rotation relates are + distinguishable and it is rejected. 0.05 merges only orientations + whose patterns differ by reflections at 5% of the strongest; 0.4 + also merges orientations told apart only by a reflection at 40%, + appropriate when that reflection is known to be weak or absent in + the data (stacking disorder, cation mixing). Candidates are the + relaxed-position group, the lattice's own holohedry, and the + holohedries of the parent lattices generated by the strong + reflections, so superstructure twin variants are tested as well. + The printout names the reflection pair that decides each candidate. + verbose : bool, default=True + Print :meth:`symmetry_summary` after the symmetry analysis. + """ + + def __init__( + self, + atoms: Atoms, + name: str | None = None, + symprec: float = 1e-4, + pseudo_symmetry_tol: float | None = 0.01, + pseudo_symmetry_intensity_tol: float = 0.05, + verbose: bool = True, + ): + self.atoms = atoms + self.name = name if name is not None else atoms.get_chemical_formula() + self._pseudo_symmetry_tol = pseudo_symmetry_tol + self._pseudo_symmetry_intensity_tol = float(pseudo_symmetry_intensity_tol) + self.pseudo_symmetry_report: dict = {} + self._wedge_cache: torch.Tensor | None | str = "unset" + + self.lat_real = torch.as_tensor(atoms.cell[:], dtype=torch.float64) + self.positions_frac = torch.as_tensor(atoms.get_scaled_positions(), dtype=torch.float64) + self.numbers = torch.as_tensor(atoms.numbers, dtype=torch.long) + occupancy = atoms.arrays.get("occupancy", np.ones(len(atoms))) + self.occupancy = torch.as_tensor(np.asarray(occupancy, dtype=float)) + + with _spglib_raises(): + self._setup_symmetry(symprec, pseudo_symmetry_tol, pseudo_symmetry_intensity_tol) + # the summary states any pseudo-symmetry adopted; when it has been + # shown, the orientation plan does not warn about it again + self._summary_shown = bool(verbose) + if verbose: + print(self.symmetry_summary()) + + # populated by calculate_structure_factors + self.k_max: float | None = None + self.hkl: torch.Tensor | None = None + self.g_vec: torch.Tensor | None = None + self.g_len: torch.Tensor | None = None + self.struct_factors: torch.Tensor | None = None + self.struct_factors_int: torch.Tensor | None = None + + # populated by calculate_dynamical_structure_factors + self.hkl_dyn: torch.Tensor | None = None + self.g_len_dyn: torch.Tensor | None = None + self.U_dyn: torch.Tensor | None = None + self.dyn_energy_ev: float | None = None + self.dyn_k_max: float | None = None + + @classmethod + def from_ase(cls, atoms: Atoms, name: str | None = None, **kwargs) -> "Crystal": + """Build a Crystal from an ase.Atoms object. + + Parameters + ---------- + atoms : ase.Atoms + The structure, e.g. from ``ase.build.bulk``. + name : str, optional + Display name; defaults to the chemical formula. + **kwargs + Passed to the Crystal constructor, e.g. `pseudo_symmetry_tol` or + `verbose`. + + Returns + ------- + Crystal + """ + return cls(atoms, name=name, **kwargs) + + @classmethod + def from_cif(cls, file_path: str | Path, name: str | None = None, **kwargs) -> "Crystal": + """Build a Crystal from a CIF file, keeping fractional site occupancies. + + ASE reads a mixed-occupancy site as a single atom of the majority + species, which silently deletes every minority element: a layered + oxide with Sb sharing a site with Fe loads with no Sb at all, and Sb + is by far its strongest scatterer. ASE does record the occupancies it + discarded, so each shared site is expanded here back into one atom + per species, each carrying its fraction in + ``atoms.arrays['occupancy']``, which the structure-factor sum uses. + + Parameters + ---------- + file_path : str or Path + Path to the CIF file. + name : str, optional + Display name; defaults to the chemical formula. + **kwargs + Passed to the Crystal constructor, e.g. `pseudo_symmetry_tol`. + + Returns + ------- + Crystal + """ + from ase.io import read + + atoms = read(file_path) + assert isinstance(atoms, Atoms) + return cls(_expand_partial_occupancy(atoms), name=name, **kwargs) + + @property + def volume(self) -> float: + """Unit cell volume, cubic Angstroms.""" + return float(torch.abs(torch.linalg.det(self.lat_real))) + + @property + def lat_recip(self) -> torch.Tensor: + """Reciprocal lattice vectors as rows, no 2*pi factor.""" + return torch.linalg.inv(self.lat_real).T + + def _quick_intensities(self, k_max: float = 2.0) -> tuple[torch.Tensor, torch.Tensor]: + """Kinematical |F|^2 of every reflection with |g| <= k_max (hkl, I), + for the pseudo-symmetry intensity check; no thermal factors.""" + recip = self.lat_recip + k_len = torch.linalg.norm(recip, dim=1) + n_max = torch.ceil(k_max / k_len * 2).to(torch.long) + ranges = [torch.arange(-int(n), int(n) + 1) for n in n_max] + hkl = torch.cartesian_prod(*ranges).to(torch.float64) + g_vec = hkl @ recip + g_len = torch.linalg.norm(g_vec, dim=1) + keep = (g_len <= k_max) & (g_len > 0) + hkl, g_len = hkl[keep], g_len[keep] + f_e = electron_scattering_factor(self.numbers, g_len) + phase = torch.exp(-2j * np.pi * (self.positions_frac @ hkl.T)) + F = (f_e * self.occupancy[:, None] * phase).sum(dim=0) / self.volume + return hkl.to(torch.long), torch.abs(F) ** 2 + + def _setup_symmetry( + self, symprec: float, pseudo_symmetry_tol: float | None, intensity_tol: float + ) -> None: + """Detect the true symmetry group, and optionally a pseudo-symmetry group. + + The true group (at `symprec`) is stored for reporting and refinement. + The pseudo-symmetry group is detected at a distance tolerance of + `pseudo_symmetry_tol` times the shortest lattice vector and kept + only if its extra operations relate reflections of equal kinematical + intensity to within `intensity_tol` of the strongest reflection: + two orientations are merged only when no experiment could tell + their patterns apart, in position or in intensity. Matching uses + that group, so nearly-degenerate cells are idealized to their + higher-symmetry parent. + """ + import spglib + + cell = ( + self.lat_real.numpy(), + self.positions_frac.numpy(), + self.numbers.numpy(), + ) + try: + dataset = spglib.get_symmetry_dataset(cell, symprec=symprec) + except Exception: + dataset = None + if dataset is None: + # no symmetry found at all (e.g. overlapping atoms): carry on in P1 + warnings.warn( + f"{self.name}: spglib found no symmetry at symprec={symprec:g}; " + "using P1. Check the structure for overlapping atoms.", + stacklevel=3, + ) + rotations = np.eye(3, dtype=np.intc)[None] + self.spacegroup: str = "P1 (1)" + else: + rotations = dataset.rotations + self.spacegroup = f"{dataset.international} ({dataset.number})" + pg = spglib.get_pointgroup(rotations)[0].strip() + self.pointgroup: str = pg + self.laue_group: str = _LAUE_CLASS.get(pg, "-1") + self.sym_quats = symmetry_quaternions(rotations, self.lat_real.numpy()) + + self.pointgroup_matching = pg + self.laue_group_matching = self.laue_group + self.sym_quats_matching = self.sym_quats + if pseudo_symmetry_tol is None: + return + a_min = float(torch.linalg.norm(self.lat_real, dim=1).min()) + symprec_pseudo = float(pseudo_symmetry_tol) * a_min + self.pseudo_symmetry_report = {"distance_A": symprec_pseudo} + if symprec_pseudo <= symprec: + return + + # Two independent routes to a higher matching symmetry. + # + # "relaxed positions": the group spglib finds when every atom is + # allowed to move by symprec_pseudo, which catches a cell that is a + # slightly distorted child of a higher-symmetry parent. + # + # "lattice": the point group of the lattice alone, ignoring the + # basis. A structure whose symmetry is broken only by weakly + # scattering atoms (lithium and oxygen against a transition metal) + # or by a faint superstructure sits exactly here: no relaxation of + # the positions recovers the parent, because the atoms are already + # where they belong, but the diffraction still has the symmetry of + # the heavy sublattice. Both candidate sets are filtered by the same + # intensity test, so an operation is adopted only when it leaves the + # kinematical pattern unchanged. + # + # "parent lattice": the lattice generated by the strong reflections + # alone. A superstructure (cation ordering on a rocksalt or layered + # frame) has a larger cell than its parent, and the parent's + # symmetries map the superstructure onto a twin of itself rather than + # onto itself, so neither route above can see them. When the + # superstructure reflections are weak the twins give the same + # pattern, and these are exactly the variants matching must merge. + lat = self.lat_real.numpy() + candidates: list[tuple[str, np.ndarray, np.ndarray]] = [] + try: + ds_relaxed = spglib.get_symmetry_dataset(cell, symprec=symprec_pseudo) + except Exception: + ds_relaxed = None + if ds_relaxed is not None: + candidates.append(("relaxed positions", ds_relaxed.rotations, lat)) + lattice_cell = (lat, np.zeros((1, 3)), np.ones(1, dtype=int)) + try: + ds_lattice = spglib.get_symmetry_dataset(lattice_cell, symprec=symprec_pseudo) + except Exception: + ds_lattice = None + if ds_lattice is not None: + candidates.append(("lattice", ds_lattice.rotations, lat)) + candidates += self._parent_lattice_candidates(pseudo_symmetry_tol, intensity_tol) + + best = None + closest = None # the rejected candidate that came nearest to passing + for route, rotations, lattice in candidates: + quats = symmetry_quaternions(rotations, lattice) + if quats.shape[0] <= self.sym_quats.shape[0]: + continue + pg_cand = spglib.get_pointgroup(rotations)[0].strip() + accepted, worst = self._intensity_preserving_subgroup(quats, intensity_tol) + if accepted is None: + if closest is None or worst < closest[0]: + closest = (worst, pg_cand, route, self._breaking_reflection) + continue + if best is None or accepted.shape[0] > best[1].shape[0]: + best = (route, accepted, pg_cand, quats.shape[0]) + + if best is None: + if closest is not None: + worst, pg_cand, route, pair = closest + self.pseudo_symmetry_report.update( + candidate=pg_cand, + intensity_mismatch=worst, + broken_by=pair, + route=route, + rejected=True, + ) + return + route, accepted, pg_cand, n_cand = best + # the largest difference among the rotations actually adopted + _, worst = self._intensity_preserving_subgroup(accepted, 1.0) + self.pseudo_symmetry_report.update( + candidate=pg_cand, + intensity_mismatch=worst, + broken_by=self._breaking_reflection, + route=route, + rejected=False, + ) + self.sym_quats_matching = accepted + # name the accepted group by the candidate symbol when every one of + # its rotations survived the intensity test, otherwise by its size + if accepted.shape[0] == n_cand: + self.pointgroup_matching = pg_cand + self.laue_group_matching = _LAUE_CLASS.get(pg_cand, self.laue_group) + else: + self.pointgroup_matching = f"{accepted.shape[0]} rotations" + self.laue_group_matching = self.laue_group + + def _parent_lattice_candidates( + self, pseudo_symmetry_tol: float, intensity_tol: float + ) -> list[tuple[str, np.ndarray, np.ndarray]]: + """Holohedries of the lattices generated by the strong reflections. + + For each intensity cut, the parent translations are the fractions t + of the cell with h.t integer for every reflection h stronger than + the cut; they form the real-space lattice dual to the strong + reflections. Its holohedry is a candidate group. Cuts run from very + low (the cell's own lattice once centring is removed) up to the + intensity tolerance, since reflections weaker than that are allowed + to break the pseudo-symmetry anyway. + + Returns + ------- + list of tuple + ``(route, rotations, lattice)``: integer rotations in the basis of + ``lattice``, whose rows are the parent vectors in the crystal's + own Cartesian frame. + """ + import itertools + + import spglib + + hkl, inten = self._quick_intensities() + if inten.numel() == 0 or float(inten.max()) <= 0: + return [] + inten = inten / inten.max() + lat = self.lat_real.numpy() + vol = abs(float(np.linalg.det(lat))) + # denominators 1, 2, 3, 4, 6, 12 cover the supercells met in practice + n_grid = 12 + grid = np.array(list(itertools.product(range(n_grid), repeat=3)), dtype=float) / n_grid + + out: list[tuple[str, np.ndarray, np.ndarray]] = [] + seen: set[int] = set() + cuts = sorted({0.02, 0.05, 0.1, 0.2, 0.3, 0.5, float(intensity_tol)}) + for cut in cuts: + if cut > 1.0: + continue + strong = hkl[inten > cut].numpy().astype(float) + if strong.shape[0] < 3 or np.linalg.matrix_rank(strong) < 3: + continue + phase = strong @ grid.T + t = grid[np.all(np.abs(phase - np.round(phase)) < 1e-6, axis=0)] + if t.shape[0] < 2: + continue # the cell is its own parent: nothing new + try: + parent = spglib.standardize_cell( + (lat, t, np.ones(t.shape[0], dtype=int)), + to_primitive=True, + no_idealize=True, + symprec=1e-5, + ) + except Exception: + parent = None + if parent is None: + continue + L = np.asarray(parent[0], dtype=float) + ratio = int(round(vol / abs(float(np.linalg.det(L))))) + if ratio in seen: + continue + seen.add(ratio) + a_min = float(np.linalg.norm(L, axis=1).min()) + try: + ds = spglib.get_symmetry_dataset( + (L, np.zeros((1, 3)), np.ones(1, dtype=int)), + symprec=float(pseudo_symmetry_tol) * a_min, + ) + except Exception: + ds = None + if ds is not None: + out.append((f"parent lattice, {ratio}x smaller cell", ds.rotations, L)) + return out + + def _intensity_preserving_subgroup( + self, quats: torch.Tensor, intensity_tol: float + ) -> tuple[torch.Tensor | None, float]: + """Largest subgroup of `quats` that leaves the kinematical intensities + invariant, or None when nothing beyond the true symmetry survives. + + Every candidate operation is applied to the reflection list and the + intensity of each reflection compared with the intensity of its + image, relative to the strongest reflection. Operations that pass + are kept; the survivors are then closed under composition (dropping + the worst offender until they are), because a set of operations that + is not a group cannot be used to fold orientations. + + Returns the accepted quaternions and the worst mismatch among the + operations that were tested. + """ + from quantem.diffraction.rotations import quat_to_matrix + + hkl, inten = self._quick_intensities() + lut = {tuple(h): i for i, h in enumerate(hkl.tolist())} + g = hkl.to(torch.float64) @ self.lat_recip + i_max = float(inten.max()) + Rs = quat_to_matrix(quats) + Rs_true = quat_to_matrix(self.sym_quats) + + # a pseudo-symmetry group comes from a lattice that is only nearly + # ideal, so its rotations and their products agree to the distortion + # (~1e-2 here, ~1e-6 for a hexagonal cell given to 5 decimals). + # Distinct crystallographic rotations are at least 60 degrees apart, + # with matrix entries differing by ~0.5, so 0.05 is unambiguous. + match_tol = 0.05 + + def is_true(R): + return any(float((R - Rt).abs().max()) < match_tol for Rt in Rs_true) + + mismatch = torch.zeros(quats.shape[0], dtype=torch.float64) + worst = 0.0 + self._breaking_reflection = None + # a pseudo-symmetry holds only approximately in the metric too, so an + # image lands near, not on, the reflection it maps onto: snap it to + # the nearest one within the same fractional tolerance allowed for + # the atom positions, and treat anything farther as absent + snap = max(float(self._pseudo_symmetry_tol or 0.0), 1e-3) + g_len = torch.linalg.norm(g, dim=1) + for i, R in enumerate(Rs): + if is_true(R): + continue + g_img = g @ R.T + hkl_img = torch.round(g_img @ self.lat_real.T).to(torch.long) + idx = torch.tensor([lut.get(tuple(h), -1) for h in hkl_img.tolist()]) + ok = idx >= 0 + if not bool(ok.any()): + continue + near = torch.linalg.norm(g_img - g[idx.clamp(min=0)], dim=1) <= snap * g_len + 1e-9 + i_img = torch.where(near, inten[idx.clamp(min=0)], torch.zeros_like(inten)) + diff = (inten[ok] - i_img[ok]).abs() + m = float(diff.max()) / i_max + mismatch[i] = m + if m > worst: + # the reflection pair responsible, reported so the user can + # judge whether the data actually resolve it + j = int(torch.argmax(diff)) + src = torch.nonzero(ok).squeeze(1)[j] + self._breaking_reflection = ( + tuple(int(v) for v in hkl[src].tolist()), + float(inten[src]) / i_max, + tuple(int(v) for v in hkl[idx[src]].tolist()), + float(i_img[src]) / i_max, + ) + worst = max(worst, m) + + keep = mismatch <= intensity_tol + # close under composition: a product of kept operations must also be + # kept, or the set is not a group + for _ in range(quats.shape[0]): + idx = torch.nonzero(keep).squeeze(1) + if idx.numel() <= self.sym_quats.shape[0]: + return None, worst + R_keep = Rs[idx] + prod = torch.einsum("aij,bjk->abik", R_keep, R_keep).reshape(-1, 3, 3) + d = (prod[:, None] - R_keep[None]).abs().amax(dim=(-1, -2)) + closed = bool((d.min(dim=1).values < match_tol).all()) + if closed: + return quats[idx], worst + drop = idx[int(torch.argmax(mismatch[idx]))] + keep[drop] = False + return None, worst + + def projected_rotation_order( + self, + zone_axis, + k_max: float | None = None, + tol_zone: float = 0.02, + intensity_tol: float = 0.05, + snap_deg: float = 4.0, + max_index: int = 3, + ): + """Apparent rotational symmetry of the zero-layer pattern, per zone axis. + + A zone-layer pattern can be more symmetric about the beam than the + crystal is, and where it is, the in-plane orientation cannot be + indexed. Body-centered cubic along <111> is the standard case: the + zero-layer net of {110} reflections is hexagonal, so the pattern + repeats every 60 degrees while the crystal repeats every 120, and the + two orientations 60 degrees apart give the same peak positions and the + same kinematical intensities. Only the higher-order Laue zones or the + dynamical intensities separate them. + + Returned is the largest n in (6, 4, 3, 2, 1) for which rotating the + zero-layer reflections by 360/n about the zone axis reproduces the + set, in position and in kinematical intensity. Fold an in-plane angle + or color by 360/n to get a map that is continuous across the + ambiguity, and use `n` against the crystal's own rotational order + about the same axis to see where indexing is degenerate. + + Parameters + ---------- + zone_axis : array-like + Cartesian zone axis (3,), or a stack of them (..., 3); need not + be normalized. + k_max : float | None + Only reflections within this scattering vector are tested; + defaults to the crystal's own k_max. + tol_zone : float, default=0.02 + Half-thickness of the zero layer (1/Angstroms): reflections with + |g . zone_axis| below this count as zero layer. + intensity_tol : float, default=0.05 + A reflection and its image must agree in |F|^2 to within this + fraction of the strongest zero-layer reflection. + snap_deg : float, default=4.0 + Zone axes within this angle of a low-index lattice direction are + evaluated at that direction. The extra symmetry is exact only on + the pole and decays away from it, but a beam a degree or two off + still produces a pattern whose positions carry it, which is + where a measured orientation normally sits; testing the exact + tilted axis would report no symmetry at all and miss the + ambiguity the indexing actually suffers. Set to 0 to test the + axis as given. + max_index : int, default=3 + Largest |u|, |v|, |w| considered when snapping. + + Returns + ------- + int | np.ndarray + The order n, scalar for a single zone axis. + """ + if self.g_vec is None: + raise RuntimeError("Run calculate_structure_factors() first.") + axes = torch.as_tensor(np.asarray(zone_axis, dtype=float), dtype=torch.float64) + single = axes.ndim == 1 + axes = axes.reshape(-1, 3) + axes = axes / torch.linalg.norm(axes, dim=1, keepdim=True).clamp_min(1e-12) + + g = self.g_vec + inten = self.struct_factors_int.to(torch.float64) + if k_max is not None: + sel = self.g_len <= float(k_max) + g, inten = g[sel], inten[sel] + + if snap_deg > 0: + rng = torch.arange(-max_index, max_index + 1, dtype=torch.float64) + uvw = torch.cartesian_prod(rng, rng, rng) + uvw = uvw[uvw.abs().sum(dim=1) > 0] + cart = uvw @ self.lat_real + cart = cart / torch.linalg.norm(cart, dim=1, keepdim=True).clamp_min(1e-12) + dots = torch.abs(axes @ cart.T) + best = dots.max(dim=1) + near = best.values > np.cos(np.deg2rad(snap_deg)) + snapped = cart[best.indices] + # keep the original sense so the returned axis still points along + # the beam, and only replace the ones close enough to snap + sign = torch.sign(torch.einsum("ni,ni->n", snapped, axes)).unsqueeze(1) + axes = torch.where(near.unsqueeze(1), snapped * sign, axes) + + out = np.ones(axes.shape[0], dtype=int) + eye = torch.eye(3, dtype=torch.float64) + for i, u in enumerate(axes): + zol = torch.abs(g @ u) <= tol_zone + gz, iz = g[zol], inten[zol] + if gz.shape[0] < 3: + continue + i_max = abs(float(iz.max())) or 1.0 + ux = torch.tensor( + [[0.0, -u[2], u[1]], [u[2], 0.0, -u[0]], [-u[1], u[0], 0.0]], + dtype=torch.float64, + ) + for n in (6, 4, 3, 2): + th = 2 * np.pi / n + R = eye + np.sin(th) * ux + (1 - np.cos(th)) * (ux @ ux) # Rodrigues + d = torch.cdist(gz @ R.T, gz) + dmin, j = d.min(dim=1) + if float(dmin.max()) > tol_zone: + continue + if float((iz - iz[j]).abs().max()) / i_max <= intensity_tol: + out[i] = n + break + return int(out[0]) if single else out + + def zone_axis_wedge(self) -> torch.Tensor | None: + """Fundamental zone-axis wedge corners (3, 3) Cartesian, or None. + + Built from the symmetry operations actually used for matching (the + pseudo-symmetry group when one was found), so the wedge is right + for every crystal setting. None means the Laue class (-1 or 2/m) + has no 3-corner wedge and libraries sample the full hemisphere. + """ + if isinstance(self._wedge_cache, str): + from quantem.diffraction.rotations import fundamental_zone_axis_wedge + + self._wedge_cache = fundamental_zone_axis_wedge(self.sym_quats_matching) + return self._wedge_cache + + @property + def hexagonal_matching(self) -> bool: + """Whether directions are written with 4-index symbols. + + Directions are always indexed in the crystal's own cell, so this + follows that cell's Laue class, not the matching group's: a + monoclinic superstructure matched with a trigonal parent group still + has a monoclinic cell, and Miller-Bravais indices would be wrong. + """ + return self.laue_group in ("6/m", "6/mmm", "-3", "-3m") + + def zone_axis_wedge_labels(self, mathtext: bool = True) -> list[str] | None: + """Direction labels of the wedge corners (4-index for hexagonal and + trigonal crystals), indexed from the corner directions themselves.""" + corners = self.zone_axis_wedge() + if corners is None: + return None + loose = max(2.0 * float(self._pseudo_symmetry_tol or 0.0), 0.02) + labels = [] + for c in corners: + uvw = direction_indices(self.lat_real, c.numpy()) + prefix = "" + if uvw is None: + # an axis of the pseudo-symmetry parent, a lattice direction + # only to within the distortion of the real cell + uvw = direction_indices(self.lat_real, c.numpy(), atol=loose) + prefix = "~" + if uvw is not None: + # v and -v are the same zone axis: name it with the first + # nonzero index of the printed symbol positive, [100] rather + # than [-100] (for 4-index symbols, [U V T W] with U = 2u - v, + # V = 2v - u, T = -(u + v), up to a common factor) + uvw = np.asarray(uvw) + u, v, w = uvw + shown = ( + np.array([2 * u - v, 2 * v - u, -(u + v), w]) + if self.hexagonal_matching + else uvw + ) + nz = np.flatnonzero(shown) + if nz.size and shown[nz[0]] < 0: + uvw = -uvw + labels.append( + prefix + + format_direction(uvw, hexagonal=self.hexagonal_matching, mathtext=mathtext) + if uvw is not None + else "(irrational)" + ) + return labels + + def matching_symmetry_warning(self) -> str | None: + """Message when the matching (pseudo) symmetry differs from the + cell's own symmetry, or None when they agree.""" + if self.pointgroup_matching == self.pointgroup: + return None + n_extra = self.sym_quats_matching.shape[0] // max(self.sym_quats.shape[0], 1) + # a partially accepted group has no Laue class of its own, and keeps + # the cell's: name it only when it differs + laue = ( + f"Laue class {self.laue_group_matching}, " + if self.laue_group_matching != self.laue_group + else "" + ) + return ( + f"{self.name}: orientation libraries are built with the " + f"pseudo-symmetry point group {self.pointgroup_matching} ({laue}" + f"found at pseudo_symmetry_tol = " + f"{self._pseudo_symmetry_tol:g} of the shortest lattice vector, " + f"intensities matching within {self.pseudo_symmetry_report.get('intensity_mismatch', 0.0):.2f} " + "of the strongest reflection), " + f"while the cell's own symmetry " + f"is {self.pointgroup} (Laue class {self.laue_group}). Orientations " + f"related by the extra operations give the same library entry, so " + f"the {n_extra} variants they generate are reported as one and the " + f"distortion between them is not resolved. To match with the exact " + f"symmetry, build the Crystal with pseudo_symmetry_tol=None (or a " + f"tolerance below the distortion)." + ) + + def symmetry_summary(self) -> str: + """Human-readable symmetry report, including any pseudo-symmetry.""" + import re + + # subscript the space group screw/glide digits: P6_3/mmc -> P6[sub3]/mmc + subs = str.maketrans("0123456789", "₀₁₂₃₄₅₆₇₈₉") + sg = re.sub(r"_(\d)", lambda m: m.group(1).translate(subs), self.spacegroup) + lines = [ + f"{self.name}", + f" space group {sg}", + f" point group {self.pointgroup} (Laue class {self.laue_group})", + ] + if self.pointgroup_matching != self.pointgroup: + laue = ( + f"(Laue class {self.laue_group_matching}) " + if self.laue_group_matching != self.laue_group + else "" + ) + lines += [ + f" pseudo-symmetry {self.pointgroup_matching} {laue}" + "-- used for orientation matching", + ] + rep = self.pseudo_symmetry_report + if rep.get("route"): + lines += [f" from the {rep['route']}"] + if rep.get("intensity_mismatch", 0.0) > 0 and rep.get("broken_by") is not None: + h0, i0, h1, i1 = rep["broken_by"] + lines += [ + " accepted at intensity tol %.2f: largest difference " + "(%s) at %.2f against (%s) at %.2f" + % ( + self._pseudo_symmetry_intensity_tol, + " ".join(map(str, h0)), + i0, + " ".join(map(str, h1)), + i1, + ) + ] + elif self._pseudo_symmetry_tol is not None: + rep = self.pseudo_symmetry_report + if rep.get("rejected"): + lines += [ + f" pseudo-symmetry {rep['candidate']} from the {rep.get('route', 'lattice')}, " + f"rejected: intensities differ by " + f"{rep['intensity_mismatch']:.2f} (tol {self._pseudo_symmetry_intensity_tol:.2f})", + ] + if rep.get("broken_by") is not None: + h0, i0, h1, i1 = rep["broken_by"] + lines += [ + " broken by (%s) at %.2f against (%s) at %.2f of the " + "strongest reflection" + % (" ".join(map(str, h0)), i0, " ".join(map(str, h1)), i1) + ] + else: + lines += [ + f" pseudo-symmetry none found at tol = {self._pseudo_symmetry_tol:g} " + f"({rep.get('distance_A', 0.0):.3f} A)", + ] + else: + lines += [" pseudo-symmetry not checked (set pseudo_symmetry_tol)"] + # matching line reflects the symmetry actually used, after any + # pseudo-symmetry reduction + labels = self.zone_axis_wedge_labels(mathtext=False) + wedge_txt = ( + f"zone axis wedge {labels[0]}, {labels[1]}, {labels[2]}" + if labels is not None + else "full hemisphere" + ) + lines += [ + f" matching {self.sym_quats_matching.shape[0]} proper rotations, {wedge_txt}" + ] + return "\n".join(lines) + + def calculate_structure_factors( + self, + k_max: float = 1.5, + tol_structure_factor: float = 1e-4, + thermal_sigma: float | dict[str, float] | None = None, + ) -> "Crystal": + """Kinematical structure factors for all reflections with |g| <= k_max. + + Parameters + ---------- + k_max : float, default=1.5 + Maximum scattering vector magnitude, 1/Angstroms. + tol_structure_factor : float, default=1e-4 + Discard reflections with |F| below this threshold. + thermal_sigma : float | dict[str, float] | None + RMS thermal displacement (Angstroms), scalar or per-element, + applied as a Debye-Waller factor. + + Returns + ------- + Crystal + self, for chaining. + """ + self.k_max = float(k_max) + recip = self.lat_recip + + # index range: project k_max onto each reciprocal cell direction + k_len = torch.linalg.norm(recip, dim=1) + n_max = torch.ceil(k_max / k_len * 2).to(torch.long) + ranges = [torch.arange(-int(n), int(n) + 1) for n in n_max] + hkl = torch.cartesian_prod(*ranges).to(torch.float64) + g_vec = hkl @ recip + g_len = torch.linalg.norm(g_vec, dim=1) + keep = (g_len <= k_max) & (g_len > 0) + hkl, g_vec, g_len = hkl[keep], g_vec[keep], g_len[keep] + + f_e = electron_scattering_factor(self.numbers, g_len) # (N_atoms, N_g) + + if thermal_sigma is not None: + if isinstance(thermal_sigma, dict): + sigma = torch.tensor( + [thermal_sigma[chemical_symbols[int(z)]] for z in self.numbers], + dtype=torch.float64, + ) + else: + sigma = torch.full((len(self.numbers),), float(thermal_sigma)) + dwf = torch.exp(-0.5 * (2 * np.pi * sigma[:, None] * g_len[None, :]) ** 2) + f_e = f_e * dwf + + phase = torch.exp(-2j * np.pi * (self.positions_frac @ hkl.T)) # (N_atoms, N_g) + F = (f_e * self.occupancy[:, None] * phase).sum(dim=0) / self.volume + + keep = torch.abs(F) > tol_structure_factor + self.hkl = hkl[keep].to(torch.long) + self.g_vec = g_vec[keep] + self.g_len = g_len[keep] + self.struct_factors = F[keep] + self.struct_factors_int = torch.abs(F[keep]) ** 2 + return self + + def calculate_dynamical_structure_factors( + self, + energy_ev: float, + thermal_sigma: float | dict[str, float] = 0.05, + k_max: float | None = None, + include_core: bool = True, + include_phonon: bool = True, + ) -> "Crystal": + """Absorptive structure factors for Bloch wave calculations. + + Uses the Weickenmeier-Kohl parameterization (Acta Cryst. A47, 590 + (1991)): the elastic part is Debye-Waller damped, and the imaginary + (absorptive) part includes core-loss and phonon/TDS contributions. + The returned factors are relativistically corrected and already carry + the 1/pi convention of the Bloch structure matrix, i.e. they are the + U_g of De Graef ch. 5 after division by the unit cell volume. + + All reflections up to k_max are kept, including kinematically + forbidden ones (their U_g can be nonzero through absorption and they + are required as coupling vectors g - h). + + Parameters + ---------- + energy_ev : float + Beam energy in eV. + thermal_sigma : float | dict[str, float], default=0.05 + RMS thermal displacement (Angstroms), scalar or per-element. + k_max : float | None + Maximum |g| of stored factors; defaults to the kinematical k_max. + For Bloch calculations with beams out to k, the couplings reach + 2k, but the factors fall off fast and 1.5k is enough. + include_core : bool, default=True + Include the core-loss (inner-shell ionization) absorptive part. + include_phonon : bool, default=True + Include the phonon (thermal diffuse scattering) absorptive part. + + Returns + ------- + Crystal + self, for chaining. Sets ``hkl_dyn`` (N, 3) and ``g_len_dyn`` (N,) + for every reflection with |g| <= k_max including (000), + ``U_dyn`` (N,) complex128 in 1/Angstroms^2, and the + ``dyn_energy_ev`` and ``dyn_k_max`` they were computed for. + + Raises + ------ + RuntimeError + If `k_max` is None and :meth:`calculate_structure_factors` has + not been run. + """ + from quantem.diffraction.wk_scattering_factors import compute_WK_factor + + if k_max is None: + if self.k_max is None: + raise RuntimeError("Provide k_max or run calculate_structure_factors.") + k_max = self.k_max + recip = self.lat_recip + k_len = torch.linalg.norm(recip, dim=1) + n_max = torch.ceil(k_max / k_len * 2).to(torch.long) + ranges = [torch.arange(-int(n), int(n) + 1) for n in n_max] + hkl = torch.cartesian_prod(*ranges).to(torch.float64) + g_vec = hkl @ recip + g_len = torch.linalg.norm(g_vec, dim=1) + keep = g_len <= k_max + hkl, g_len = hkl[keep], g_len[keep] + + g_np = g_len.numpy() + if isinstance(thermal_sigma, dict): + sigma_per_atom = np.array( + [thermal_sigma[chemical_symbols[int(z)]] for z in self.numbers] + ) + else: + sigma_per_atom = np.full(len(self.numbers), float(thermal_sigma)) + + # one WK evaluation per unique (Z, sigma) pair + f_atoms = np.zeros((len(self.numbers), g_np.size), dtype=np.complex128) + cache: dict[tuple[int, float], np.ndarray] = {} + for i, (z, sig) in enumerate(zip(self.numbers.tolist(), sigma_per_atom)): + key = (int(z), float(sig)) + if key not in cache: + cache[key] = compute_WK_factor( + g_np, + int(z), + energy_ev, + thermal_sigma=float(sig), + include_core=include_core, + include_phonon=include_phonon, + ) + f_atoms[i] = cache[key] + + phase = np.exp(-2j * np.pi * (self.positions_frac.numpy() @ hkl.numpy().T)) + occ = self.occupancy.numpy()[:, None] + U = (f_atoms * occ * phase).sum(axis=0) / self.volume + + self.hkl_dyn = hkl.to(torch.long) + self.g_len_dyn = g_len + self.U_dyn = torch.as_tensor(U, dtype=torch.complex128) + self.dyn_energy_ev = float(energy_ev) + self.dyn_k_max = float(k_max) + return self + + def direction_vector(self, direction) -> torch.Tensor: + """Unit Cartesian vector (3,) of a lattice direction in the crystal frame. + + Parameters + ---------- + direction : sequence of float + [uvw] in this cell, or [UVTW] for a hexagonal or trigonal cell. + """ + d = np.asarray(direction, dtype=float).ravel() + if d.size == 4: + U, V, T, W = d + d = np.array([U - T, V - T, W]) + elif d.size != 3: + raise ValueError(f"a direction has 3 or 4 indices, got {direction}") + v = d @ self.lat_real.numpy() + return torch.as_tensor(v / np.linalg.norm(v), dtype=torch.float64) + + def generate_pattern( + self, + orientation: torch.Tensor, + energy_ev: float = 300e3, + sigma_excitation: float = SIGMA_EXCITATION, + tol_excitation_mult: float = 3.0, + k_max: float | None = None, + precession_deg: float = 0.0, + semiconv_mrad: float = 0.0, + excitation_model: str = "gaussian", + thickness_A: float | None = None, + foil_normal=None, + ) -> dict[str, torch.Tensor]: + """Kinematical diffraction pattern for one orientation. + + The intensity of each reflection is |F_g|^2 times a Gaussian + excitation envelope of width sigma_excitation, averaged exactly over + the illumination when a precession angle or a convergence + semiangle is given (quantem.diffraction.illumination): the + precession ring sweeps the excitation error of reflection g by + +- a_g = r |g_xy| / |K - g_z| about its central value c_g, and the + averaged envelope is the Bessel transform G(c_g, a_g, b_g; sigma). + + Parameters + ---------- + orientation : torch.Tensor + Unit quaternion (4,) rotating crystal vectors into the lab frame. + energy_ev : float, default=300e3 + Beam energy in eV. + sigma_excitation : float, default=SIGMA_EXCITATION + Excitation error tolerance (1/Angstroms) in the shape-factor + envelope exp(-s_g^2 / 2 sigma^2); the default is + :data:`quantem.diffraction.defaults.SIGMA_EXCITATION`. + tol_excitation_mult : float, default=3.0 + Include reflections with |s_g| below this multiple of sigma. + k_max : float | None + Optionally trim the pattern below the structure-factor k_max. + precession_deg, semiconv_mrad : float + Precession semi-angle (degrees) and convergence semiangle + (mrad) of the illumination the intensities are averaged over. + excitation_model : {"gaussian", "slab"} + "gaussian" is the empirical envelope of width sigma_excitation + used by the orientation library. "slab" is the finite-thickness + first Born rocking curve, (pi |U_g| z / k0)^2 sinc(s_g z)^2 + with U_g = gamma_rel F_g / pi, averaged over the illumination + the same way; it needs thickness_A and is the kinematical limit + of the Bloch wave calculation for thin crystals. + thickness_A : float | None + Thickness for the slab model (Angstroms). + foil_normal : sequence of float, optional + Plate normal of the specimen as a direction in this crystal, + [uvw] or [UVTW], e.g. (0, 0, 0, 1) for a 2D material lying in its + basal plane. Every reflection is then a rod along that normal: it + is excited by its distance along the rod to the Ewald sphere, and + its spot sits where the rod meets the sphere rather than at the + projection of g. For a tilted flake, whose rods are long, the + spots shift by up to s_g tan(tilt). None (default) takes the + normal along the beam, the usual geometry. + + Returns + ------- + dict + 'qx', 'qy' (1/Angstroms), 'intensity', 'hkl', 's_g' (the central + excitation error, along the rod when `foil_normal` is given), 'a' + and 'b' (ring and disk sweep amplitudes), one entry per excited + reflection. + + Raises + ------ + RuntimeError + If :meth:`calculate_structure_factors` has not been run. + ValueError + If `excitation_model` is not "gaussian" or "slab", or the slab + model is asked for without `thickness_A`. + """ + if self.g_vec is None: + raise RuntimeError("Run calculate_structure_factors first.") + if excitation_model not in _EXCITATION_MODELS: + raise ValueError( + f"excitation_model must be one of {_EXCITATION_MODELS}, got {excitation_model!r}" + ) + from quantem.diffraction.illumination import ( + averaged_gaussian_intensity_envelope, + excitation_coefficients, + relrod_factor, + slab_envelope, + ) + + g = qrotate(orientation, self.g_vec) + n_lab = None + if foil_normal is not None: + n_lab = qrotate(orientation, self.direction_vector(foil_normal)[None])[0] + if excitation_model == "slab": + if thickness_A is None: + raise ValueError("the slab excitation model needs thickness_A") + c, a, b = excitation_coefficients(g, energy_ev, precession_deg, semiconv_mrad) + if n_lab is not None: + f = relrod_factor(g.numpy(), n_lab.numpy(), energy_ev, precession_deg) + c, a, b = c * f, a * np.abs(f), b * np.abs(f) + # the sinc^2 tails are algebraic: keep everything whose main + # lobe (width 1/z) plus illumination sweep is within the tolerance + width = tol_excitation_mult / float(thickness_A) + c_t = torch.as_tensor(c, dtype=torch.float64) + a_t = torch.as_tensor(a, dtype=torch.float64) + b_t = torch.as_tensor(b, dtype=torch.float64) + keep = torch.abs(c_t) < a_t + b_t + width + if k_max is not None: + keep &= self.g_len <= k_max + env = slab_envelope( + c[keep.numpy()], a[keep.numpy()], b[keep.numpy()], float(thickness_A) + ) + from quantem.core.utils.utils import electron_wavelength_angstrom + + lam = electron_wavelength_angstrom(energy_ev) + gamma_rel = 1.0 + float(energy_ev) / 510998.95 + u_abs = torch.abs(self.struct_factors[keep]) * (gamma_rel / np.pi) + intensity = (np.pi * u_abs * float(thickness_A) * lam) ** 2 * torch.as_tensor( + env, dtype=torch.float64 + ) + else: + env, c, a, b = averaged_gaussian_intensity_envelope( + g, + energy_ev, + sigma_excitation, + precession_deg, + semiconv_mrad, + foil_normal_lab=None if n_lab is None else n_lab.numpy(), + ) + c_t = torch.as_tensor(c, dtype=torch.float64) + a_t = torch.as_tensor(a, dtype=torch.float64) + b_t = torch.as_tensor(b, dtype=torch.float64) + # the full illumination support enters the selection, not only + # the central excitation error + keep = torch.abs(c_t) < a_t + b_t + sigma_excitation * tol_excitation_mult + if k_max is not None: + keep &= self.g_len <= k_max + intensity = ( + self.struct_factors_int[keep] * torch.as_tensor(env, dtype=torch.float64)[keep] + ) + qxy = g[keep, :2] + if n_lab is not None: + # the spot is where the rod meets the sphere: g + t n, t = -c + qxy = qxy - c_t[keep, None] * n_lab[None, :2] + return { + "qx": qxy[:, 0], + "qy": qxy[:, 1], + "intensity": intensity, + "hkl": self.hkl[keep], + "s_g": c_t[keep], + "a": a_t[keep], + "b": b_t[keep], + } + + def __repr__(self) -> str: + return ( + f"Crystal({self.name}, {len(self.numbers)} atoms, " + f"spacegroup {self.spacegroup}, pointgroup {self.pointgroup})" + ) diff --git a/src/quantem/diffraction/crystal_map.py b/src/quantem/diffraction/crystal_map.py new file mode 100644 index 000000000..2452a5204 --- /dev/null +++ b/src/quantem/diffraction/crystal_map.py @@ -0,0 +1,1164 @@ +"""Orientation and phase mapping over one or more candidate crystals. + +CrystalMap is the standard entry point for ACOM. It owns one +:class:`~quantem.diffraction.orientation.OrientationMap` per candidate +crystal, fans the matching and refinement stages out over all of them, and +holds the :class:`~quantem.diffraction.phase.PhaseMap` that decides which +crystal sits at each probe position:: + + cm = CrystalMap.from_vectors(peaks, [Cu_metal, Cu2O], energy_ev=300e3) + cm.build_plan(**plan_params) + cm.match_orientations(**match_params) + cm.refine_orientations(**refine_params) + cm.fit() + cm.plot_phase() + cm.plot_orientation() + +The individual maps stay available for anything asymmetric or per-crystal -- +``cm["Cu metal"]``, ``cm[0]`` or ``cm.orientation_maps`` -- and every method +on OrientationMap works there exactly as before. +""" + +from __future__ import annotations + +import numpy as np + +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.phase import SHADE_GAMMA, PhaseMap + + +def _common_k_max(crystals, k_max: float | None) -> float: + """The one k_max every crystal is simulated to, or an error saying why not.""" + if k_max is not None: + return float(k_max) + have = {xtl.name: xtl.k_max for xtl in crystals} + missing = [n for n, k in have.items() if k is None] + if missing: + raise ValueError( + f"no structure factors for {missing}: pass k_max= to CrystalMap.from_vectors, " + "which computes them for every crystal" + ) + if len({round(float(k), 9) for k in have.values()}) > 1: + raise ValueError( + f"crystals are simulated to different k_max {have}; pass k_max= to " + "CrystalMap.from_vectors to set one range for all of them" + ) + return float(next(iter(have.values()))) + + +class CrystalMap(AutoSerialize): + """Per-position crystal orientation and phase over a scan. + + Parameters + ---------- + orientation_maps : list of OrientationMap + One matched (or unmatched) map per candidate crystal. + + Attributes + ---------- + orientation_maps : list of OrientationMap + The per-crystal maps, in the order they were given. + names : list of str + Crystal names, used for indexing and plot labels. + phases : PhaseMap or None + The phase decision, populated by :meth:`fit`. + dynamical : dict or None + The last :meth:`refine_dynamical` result. + """ + + _token = object() + + def __init__( + self, orientation_maps: list[OrientationMap], _token=None, k_max: float | None = None + ): + """Private constructor; use :meth:`from_vectors` or :meth:`from_orientation_maps`. + + Parameters + ---------- + orientation_maps : list of OrientationMap + One per crystal, all sharing a scan shape. Names must be unique, + because crystals are indexed by name. + _token : object + Guard against direct construction. + k_max : float, optional + Scattering-vector limit shared by the crystals; taken from them + when None. + + Raises + ------ + RuntimeError + If called without the class token. + ValueError + If no crystals are given, names repeat, or scan shapes differ. + """ + if _token is not self._token: + raise RuntimeError( + "Use CrystalMap.from_vectors() or CrystalMap.from_orientation_maps()." + ) + if len(orientation_maps) == 0: + raise ValueError("CrystalMap needs at least one crystal.") + names = [om.crystal.name for om in orientation_maps] + if len(set(names)) != len(names): + raise ValueError( + f"crystal names must be unique for indexing by name, got {names}. " + "Set Crystal(name=...) to tell them apart." + ) + shapes = {tuple(om.peaks.shape[:2]) for om in orientation_maps} + if len(shapes) != 1: + raise ValueError(f"all crystals must share one scan shape, got {shapes}") + self.orientation_maps = orientation_maps + self.phases: PhaseMap | None = None + self.dynamical: dict | None = None + self.metadata: dict = {} + self.k_max = _common_k_max([om.crystal for om in orientation_maps], k_max) + + def __attrs_post_init__(self): + """After loading: the phase map shares this map's orientation maps. + + A saved file holds the phase map's copies of the orientation maps + separately, and without this they would load as independent objects: + anything computed on one (orientations written back by a refinement, + structure factors attached to a crystal) would be missed by the other. + """ + if self.phases is not None: + self.phases.orientation_maps = self.orientation_maps + + # ------------------------------------------------------------------ + # construction + # ------------------------------------------------------------------ + + @classmethod + def from_vectors( + cls, + peaks, + crystals: Crystal | list[Crystal], + energy_ev: float = 300e3, + precession_deg: float = 0.0, + semiconv_mrad: float = 0.0, + k_max: float | None = None, + foil_normal=None, + ) -> "CrystalMap": + """Build one OrientationMap per crystal from a shared peak table. + + Parameters + ---------- + peaks : Vector + Calibrated Bragg peaks, shared by every crystal. + crystals : Crystal or list of Crystal + Candidate phases. A single Crystal is accepted, so single-phase + work uses the same entry point. + energy_ev, precession_deg, semiconv_mrad + Passed to :meth:`OrientationMap.from_vectors`. + k_max : float, optional + Largest scattering vector (1/Angstroms) in the simulated patterns, + applied to every crystal: their structure factors are computed + here, so it is set once. Match it to the detector; reflections + beyond it cannot be paired. None keeps the structure factors the + crystals already have, which must then share one k_max, since two + phases simulated to different ranges are not compared fairly. + foil_normal : sequence or dict, optional + Plate normal of a 2D material or thin flake, as a direction in each + crystal, [uvw] or [UVTW]; a dict keyed by crystal name sets it per + crystal. Reflections are then rods along it, which places the + spots of a tilted flake where the rods meet the Ewald sphere (see + :meth:`OrientationMap.from_vectors`). None (default) is the usual + geometry. + + Raises + ------ + ValueError + If `k_max` is None and the crystals have no structure factors, or + have them to different ranges. + """ + xtls = list(crystals) if isinstance(crystals, (list, tuple)) else [crystals] + if k_max is not None: + for xtl in xtls: + xtl.calculate_structure_factors(k_max=float(k_max)) + else: + _common_k_max(xtls, None) # say what is wrong before anything is built + oms = [ + OrientationMap.from_vectors( + peaks, + xtl, + energy_ev=energy_ev, + precession_deg=precession_deg, + semiconv_mrad=semiconv_mrad, + foil_normal=( + foil_normal.get(xtl.name) if isinstance(foil_normal, dict) else foil_normal + ), + ) + for xtl in xtls + ] + return cls(oms, _token=cls._token, k_max=k_max) + + @classmethod + def from_orientation_maps(cls, orientation_maps: list[OrientationMap]) -> "CrystalMap": + """Wrap maps that were built and matched by hand. + + Parameters + ---------- + orientation_maps : list of OrientationMap + One per crystal, all sharing a scan shape and with unique + crystal names. The scattering-vector limit is taken from the + crystals, which must share one. + + Returns + ------- + CrystalMap + + Raises + ------ + ValueError + If the list is empty, names repeat, scan shapes differ, or the + crystals have no common k_max. + """ + return cls(list(orientation_maps), _token=cls._token) + + # ------------------------------------------------------------------ + # access + # ------------------------------------------------------------------ + + @property + def names(self) -> list[str]: + """Crystal names, in the order the maps were given.""" + return [om.crystal.name for om in self.orientation_maps] + + @property + def peaks(self): + """The shared peak list; every crystal was matched against these.""" + return self.orientation_maps[0].peaks + + @property + def shape(self) -> tuple[int, int]: + """Scan shape ``(rows, cols)`` in probe positions.""" + return tuple(self.orientation_maps[0].peaks.shape[:2]) + + def __len__(self) -> int: + """Number of crystals.""" + return len(self.orientation_maps) + + def __iter__(self): + """Iterate over the per-crystal OrientationMaps.""" + return iter(self.orientation_maps) + + def __getitem__(self, key) -> OrientationMap: + """`cm[0]` or `cm["Cu metal"]` -> the OrientationMap of that crystal.""" + if isinstance(key, str): + try: + return self.orientation_maps[self.names.index(key)] + except ValueError: + raise KeyError(f"no crystal named {key!r}; have {self.names}") from None + return self.orientation_maps[key] + + def __repr__(self) -> str: + """Scan shape, how far the analysis has run, and the crystal names.""" + R, C = self.shape + stage = "unmatched" + if self.orientation_maps[0].quats is not None: + stage = "matched" + if "refine" in self.orientation_maps[0].metadata: + stage = "refined" + if self.phases is not None and self.phases.phase_index is not None: + stage += ", phase fit" + return "CrystalMap(%d x %d, %s, [%s])" % (R, C, stage, ", ".join(self.names)) + + # ------------------------------------------------------------------ + # staged workflow, fanned out over the crystals + # ------------------------------------------------------------------ + + def build_plan(self, overrides: dict | None = None, **kwargs) -> "CrystalMap": + """Build the correlation plan for every crystal. + + Parameters + ---------- + overrides : dict, optional + Per-crystal keyword overrides, keyed by crystal name, e.g. + ``overrides={"Cu2O": dict(angle_step_zone_axis_deg=2.0)}``. + **kwargs + Passed to :meth:`OrientationMap.build_plan` for every crystal. + """ + return self._fanout("build_plan", overrides, **kwargs) + + def match_orientations(self, overrides: dict | None = None, **kwargs) -> "CrystalMap": + """Match every crystal against the measured peaks. + + Each crystal is correlated against its own plan, so the scores are + comparable: the library slices are unit vectors and the measured + polar image is divided by its norm, making the correlation a cosine + similarity in [0, 1]. With `num_matches` above 1 each further match + is fitted to what the earlier ones leave unexplained, so a probe + straddling two grains indexes both. Every match is still scored + against the pattern as measured, so comparing the matches tells the + two cases apart: two grains score alike, while a spurious second + match on a single grain scores well below the first. + + Parameters + ---------- + overrides : dict, optional + Per-crystal keyword overrides, keyed by crystal name. + **kwargs + Passed to :meth:`OrientationMap.match_orientations` for every + crystal. The ones usually set are `num_matches`, + `min_number_peaks` and `positions`, the last restricting the + match to a few probe positions for a staged test run. + + Returns + ------- + CrystalMap + Self, so stages chain. + + Notes + ----- + The correlation saturates on sparse patterns: a position carrying + the direct beam and two noise peaks scores about as well as a real + grain, because some library orientation almost always has a + reflection at that radius and angle. Judge a match by + :meth:`signal_confidence` or the peak count, not by the correlation + alone. + """ + return self._fanout("match_orientations", overrides, **kwargs) + + def refine_orientations( + self, + overrides: dict | None = None, + competitive_margin: float | None = 0.1, + **kwargs, + ) -> "CrystalMap": + """Refine every crystal off the library grid. + + Least squares on the paired peak positions removes the quantization + of the plan: the in-plane rotation comes from the pairing, the zone + axis tilt from the intensity envelope. Positions that disagree with + a neighbour are then retried from the candidates around them, and + ties go to the orientation the neighbours share. + + With several crystals, each is refined only where it is in the + running: where its library correlation is within + `competitive_margin` of the best crystal's. Elsewhere its + orientations are noise -- the crystal is not there -- and refining + them, and retrying them against their equally random neighbours, is + most of the work in a two-phase map for no change in the phase + decision. + + Parameters + ---------- + overrides : dict, optional + Per-crystal keyword overrides, keyed by crystal name. + competitive_margin : float or None, default=0.1 + Correlation margin behind the best crystal within which a crystal + is still refined. None refines every crystal everywhere. + **kwargs + Passed to :meth:`OrientationMap.refine_orientations` for every + crystal, commonly `num_iterations` and `zone_search_deg`. An + explicit `positions` is used as given. + + Returns + ------- + CrystalMap + Self, so stages chain. + """ + oms = self.orientation_maps + if competitive_margin is None or len(oms) < 2 or "positions" in kwargs: + return self._fanout("refine_orientations", overrides, **kwargs) + overrides = self._check_overrides(overrides) + corr = np.stack([om.corr[..., 0].numpy() for om in oms]) + best = corr.max(axis=0) + for i, om in enumerate(oms): + kw = {**kwargs, **overrides.get(om.crystal.name, {})} + kw.setdefault("positions", corr[i] >= best - competitive_margin) + om.refine_orientations(**kw) + return self + + def _check_overrides(self, overrides: dict | None) -> dict: + """Per-crystal overrides, checked against the crystal names. + + Parameters + ---------- + overrides : dict or None + Keyword arguments per crystal name. + + Returns + ------- + dict + `overrides`, or an empty dict for None. + + Raises + ------ + KeyError + If an override names a crystal not in this map. + """ + overrides = overrides or {} + unknown = set(overrides) - set(self.names) + if unknown: + raise KeyError(f"overrides name unknown crystals {sorted(unknown)}; have {self.names}") + return overrides + + def _fanout(self, method: str, overrides: dict | None, **kwargs) -> "CrystalMap": + """Call ``method`` on every OrientationMap, with per-crystal overrides. + + Parameters + ---------- + method : str + Name of the OrientationMap method to call. + overrides : dict or None + Keyword arguments per crystal name, merged over ``kwargs``. + **kwargs + Arguments common to every crystal. + + Returns + ------- + CrystalMap + Self, so stages chain. + + Raises + ------ + KeyError + If an override names a crystal not in this map. + """ + overrides = self._check_overrides(overrides) + for om in self.orientation_maps: + kw = {**kwargs, **overrides.get(om.crystal.name, {})} + getattr(om, method)(**kw) + return self + + def fit(self, **kwargs) -> "CrystalMap": + """Decide the phase at every position; see :meth:`PhaseMap.fit`. + + With a single crystal there is nothing to choose between, but the fit + still runs: it applies the null hypothesis, so positions with no + diffracted signal come out unindexed and the maps below fade them. + """ + self.phases = PhaseMap.from_orientation_maps(self.orientation_maps) + self.phases.fit(**kwargs) + return self + + def refine_dynamical( + self, + mask=None, + k_max: float | None = None, + k_max_coupling: float | None = None, + **kwargs, + ) -> "CrystalMap": + """Dynamical refinement of orientation, thickness, strain and phase. + + Bloch-wave intensities, averaged over the precession ring, are fit to + the measured peaks of every candidate at every position of `mask`, + and the candidate with the lowest cost decides the phase there; the + rest of the scan keeps its decision. See + :func:`~quantem.diffraction.bloch.refine_dynamical` for the model and + every argument. The refined orientations are written back, so calls + can be staged: thickness and tilt first with + ``refine_deformation=False``, then the in-plane strain from those + orientations with a narrower tilt search. + + Every crystal is given absorptive structure factors out to + `k_max_coupling`, which the couplings between beams need; they are + computed here when missing or too short. + + Parameters + ---------- + mask : np.ndarray or list of tuple, optional + Positions to refine, an (R, C) boolean mask or (row, col) list. + None refines every matched position, which takes hours. + k_max : float, optional + Largest |g| (1/Angstroms) of the beams in the Bloch calculation. + None takes the k_max of the kinematical simulation. Cutting it + low saves time but drops beams that carry real dynamical + coupling. + k_max_coupling : float, optional + Largest |g| (1/Angstroms) of the structure factors coupling the + beams. The couplings g - h reach twice `k_max`, but the factors + fall off fast; None uses 1.5 `k_max`. It sets accuracy, not run + time, which the number of beams sets. + **kwargs + Passed to :func:`~quantem.diffraction.bloch.refine_dynamical`. + `require_phase_weight` defaults to False here, so a crystal the + kinematical fit rejected still competes. + + Returns + ------- + CrystalMap + Self, with the result in :attr:`dynamical`. + """ + from quantem.diffraction import bloch + + pm = self._require_fit("refine_dynamical()") + k_max = float(k_max if k_max is not None else self._k_max_or_crystals()) + k_c = float(k_max_coupling if k_max_coupling is not None else 1.5 * k_max) + energy_ev = self.orientation_maps[0].energy_ev + for om in self.orientation_maps: + xtl = om.crystal + if ( + getattr(xtl, "U_dyn", None) is None + or getattr(xtl, "dyn_k_max", 0.0) < k_c - 1e-9 + or abs(getattr(xtl, "dyn_energy_ev", energy_ev) - energy_ev) > 1.0 + ): + xtl.calculate_dynamical_structure_factors(energy_ev, k_max=k_c) + kwargs["k_max"] = k_max + kwargs.setdefault("require_phase_weight", False) + self.dynamical = bloch.refine_dynamical(pm, mask=mask, **kwargs) + pm.apply_dynamical(self.dynamical) + return self + + def plot_dynamical(self, phase=None, strain: bool = False, crop: bool = True, **kwargs): + """Maps of the last :meth:`refine_dynamical`. + + Thickness, tilt correction, the cost gain of the tilt search and the + final cost, or with `strain` the six crystal-frame strain components. + + Parameters + ---------- + phase : int or str, optional + Show only positions this crystal won. None shows all of them. + strain : bool, default=False + Plot the strain components instead. + crop : bool, default=True + Crop to the refined positions. + **kwargs + Passed to :func:`~quantem.diffraction.bloch.plot_dynamical_maps` + or :func:`~quantem.diffraction.bloch.plot_strain_crystal_frame`. + + Returns + ------- + tuple + ``(fig, axs)``. + """ + from quantem.diffraction import bloch + + result = getattr(self, "dynamical", None) + if result is None: + raise ValueError("run refine_dynamical() before plot_dynamical().") + i = None if phase is None else self._phase_indices(phase)[0] + maps = bloch.dynamical_maps(result, self.phases, crystal_index=i) + m = maps["mask"].numpy() + sl = (slice(None), slice(None)) + if crop and m.any(): + rows, cols = np.nonzero(m) + sl = (slice(rows.min(), rows.max() + 1), slice(cols.min(), cols.max() + 1)) + kwargs.setdefault("scalebar", self.orientation_maps[0].scan_scalebar) + if strain: + comps = {k: v.numpy()[sl] for k, v in maps["strain"].items()} + return bloch.plot_strain_crystal_frame(comps, mask=m[sl], **kwargs) + cropped = dict(maps) + for k in ("thickness", "tilt_deg", "gain", "cost"): + cropped[k] = maps[k][sl] + return bloch.plot_dynamical_maps(cropped, **kwargs) + + def example_positions( + self, + phase=None, + num: int = 4, + ambiguous: bool = False, + min_distance: float = 16.0, + min_signal: float = 0.3, + ) -> list[tuple[int, int]]: + """Well-separated probe positions for inspecting the phase decision. + + The clearest examples of a crystal are where it won by the largest + margin (the phase reliability); the ambiguous ones are where the two + best crystals scored closest. With a single crystal there is no + margin, and positions are ranked by its correlation instead. Only positions that diffract are + considered, and each pick is at least `min_distance` from the others, + so the examples come from different parts of the scan. + + Parameters + ---------- + phase : int or str, optional + Crystal the positions must have been assigned to. None allows + any crystal. + num : int, default=4 + Number of positions. + ambiguous : bool, default=False + Pick the closest decisions instead of the clearest. + min_distance : float, default=16.0 + Smallest separation between picks, in probe positions. + min_signal : float, default=0.3 + Smallest :meth:`signal_confidence` a position needs. + + Returns + ------- + list of tuple of int + ``(row, col)`` positions, clearest (or closest) first. + """ + pm = self._require_fit("example_positions()") + ph = self.phase_index + rel = np.asarray(pm.reliability, dtype=float) + if not np.isfinite(rel[ph >= 0]).any(): + # one crystal: no runner-up to compare against, so rank by how + # well the crystal itself matches + corr = np.stack([om.corr[..., 0].numpy() for om in self.orientation_maps]) + rel = np.where( + ph >= 0, np.take_along_axis(corr, np.clip(ph, 0, None)[None], 0)[0], np.nan + ) + ok = (ph >= 0) & np.isfinite(rel) & (self.signal_confidence() >= min_signal) + if phase is not None: + ok &= ph == self._phase_indices(phase)[0] + rc = np.argwhere(ok) + order = np.argsort(rel[ok] if ambiguous else -rel[ok], kind="stable") + picks: list[tuple[int, int]] = [] + for r, c in rc[order]: + if all((r - a) ** 2 + (c - b) ** 2 >= min_distance**2 for a, b in picks): + picks.append((int(r), int(c))) + if len(picks) == num: + break + return picks + + # ------------------------------------------------------------------ + # derived quantities + # ------------------------------------------------------------------ + + def _require_fit(self, what: str) -> PhaseMap: + """Return the fitted PhaseMap, or explain which call needs it first. + + Parameters + ---------- + what : str + Name of the caller, used in the error message. + + Returns + ------- + PhaseMap + + Raises + ------ + ValueError + If :meth:`fit` has not been run. + """ + if self.phases is None or self.phases.phase_index is None: + raise ValueError(f"run fit() before {what}.") + return self.phases + + @property + def phase_index(self) -> np.ndarray: + """Winning crystal at each probe position. + + Returns + ------- + np.ndarray + ``(scan_row, scan_col)`` index into :attr:`names`, or -1 where + the null hypothesis in :meth:`fit` found too little diffracted + signal to name a crystal. + """ + return self._require_fit("phase_index").phase_index.numpy() + + def signal_confidence(self, signal_range="auto", gamma: float = 1.0) -> np.ndarray: + """Confidence in [0, 1] that a crystal is present, from the data alone. + + The measured intensity beyond the direct beam, scaled to [0, 1]. + Vacuum and amorphous support diffract nothing, so they score zero + however well some orientation happens to correlate -- which the + correlation itself cannot tell you, since it saturates on sparse + patterns. + + Parameters + ---------- + signal_range : tuple or "auto", default="auto" + Diffracted intensity mapped to 0 ... 1. "auto" spans zero to half + the median over the indexed positions. + gamma : float, default=1.0 + Exponent applied to the result. The default is the raw confidence, + which is what a threshold should be taken on. + + Returns + ------- + np.ndarray + ``(scan_row, scan_col)`` confidence in [0, 1]. + """ + return self._require_fit("signal_confidence()").signal_confidence(signal_range, gamma) + + def mask(self, phase=None, signal_range="auto", gamma: float = SHADE_GAMMA) -> np.ndarray: + """Display mask in [0, 1] for one crystal, or for all indexed positions. + + The phase decision times the diffracted-signal confidence: positions + of another crystal, and positions with nothing there, are zero. + + Parameters + ---------- + phase : int or str, optional + Crystal index or name. None (default) keeps every indexed + position, whichever crystal won. + signal_range : tuple or "auto" + Passed to :meth:`signal_confidence`. Set it to override the + automatic brightness range, e.g. (0, 500). + gamma : float, default=:data:`~quantem.diffraction.phase.SHADE_GAMMA` + Brightness exponent, the same one :meth:`plot_phase` shades with, + so a masked orientation map and the phase map agree. Diffracted + intensity is strongly skewed, so the default lifts the faint + positions; 1.0 gives the linear scale. Zero stays zero, so vacuum + is black at any value. + + Returns + ------- + np.ndarray + ``(scan_row, scan_col)`` mask in [0, 1]. + + Raises + ------ + KeyError + If `phase` names no crystal in this map. + """ + conf = self.signal_confidence(signal_range, gamma) + if phase is None: + return conf + i = self._phase_indices(phase)[0] + return (self.phase_index == i) * conf + + def phase_fractions(self) -> dict[str, float]: + """Fraction of the scan won by each crystal, plus the unindexed share. + + These are area fractions of the phase decision, counted over every + probe position. The per-position model weights are + `phases.crystal_weights`. + + Returns + ------- + dict[str, float] + "unindexed" and one entry per crystal name; the values sum to 1. + """ + ph = self.phase_index + out = {"unindexed": float((ph == -1).mean())} + for i, n in enumerate(self.names): + out[n] = float((ph == i).mean()) + return out + + # ------------------------------------------------------------------ + # plotting + # ------------------------------------------------------------------ + + def plot_phase(self, **kwargs): + """Map of which crystal won at each probe position. + + Color gives the crystal, brightness gives the evidence. Positions + the null hypothesis left unindexed in :meth:`fit` -- vacuum, + amorphous support, anything that diffracts nothing -- are black. + + Parameters + ---------- + shade_by : {"signal", "reliability", "none"}, default="signal" + What the brightness means. "signal" fades by the measured + diffracted intensity, so the map shows where crystals are. + "reliability" uses the cost gap to the best model without the + winning crystal, which answers which phase rather than whether + there is one. + shade_range : tuple or "auto", default="auto" + Values mapped to black ... full color. + shade_gamma : float, default=0.5 + Exponent applied to the brightness. Diffracted intensity is + strongly skewed, so below 1 lifts the faint positions and 1.0 is + the linear scale. Vacuum stays black at any value. + majority_filter : int, default=0 + Radius in probe positions of a majority filter on the phase + decision, for display only. 1 replaces each position by the most + common phase in its 3x3 neighbourhood, dropping isolated single + positions without moving a real boundary. + phase_colors : np.ndarray, optional + One RGB color per crystal. + scalebar : dict, "auto" or None, default="auto" + Real-space scale bar; "auto" takes the scan step carried by the + peaks. + **kwargs + Further arguments of :meth:`PhaseMap.plot_phase`. + + Returns + ------- + tuple + ``(fig, ax)``. + + Raises + ------ + ValueError + If :meth:`fit` has not been run. + """ + return self._require_fit("plot_phase()").plot_phase(**kwargs) + + def plot_orientation( + self, + direction=("z", "r"), + phase=None, + mask=None, + signal_range="auto", + shade_gamma: float = SHADE_GAMMA, + **kwargs, + ): + """Inverse pole figure maps of every crystal, masked by the phase decision. + + Parameters + ---------- + direction : str or sequence of str, default=("z", "r") + Out-of-plane, in-plane, or both. + phase : int or str, optional + Restrict to one crystal. None (default) plots all of them. + mask : np.ndarray, optional + Overrides the automatic phase-and-signal mask entirely. + signal_range : tuple or "auto", default="auto" + Diffracted intensity mapped to black ... full color. Set it to + override the automatic range, e.g. (0, 500). + shade_gamma : float, default=:data:`~quantem.diffraction.phase.SHADE_GAMMA` + Brightness exponent of that mask, the same one :meth:`plot_phase` + shades with. Below 1 lifts the faint positions, 1.0 is the linear + scale. Both are ignored when `mask` is given. + **kwargs + Passed to :meth:`OrientationMap.plot_orientation`, e.g. `smooth`, + and `saturation_power` and `chroma` for the color wedge. + + Returns + ------- + tuple or list of tuple + A single ``(fig, ax)`` when `phase` names one crystal and one + direction is given; otherwise one ``(fig, ax)`` per crystal and + direction, the same rule as :meth:`plot_pole_figure`. + """ + dirs = [direction] if isinstance(direction, str) else list(direction) + out = [] + for i in self._phase_indices(phase): + m = mask if mask is not None else self.mask(i, signal_range, shade_gamma) + for d in dirs: + out.append( + self.orientation_maps[i].plot_orientation(direction=d, mask=m, **kwargs) + ) + return out[0] if phase is not None and len(out) == 1 else out + + def plot_pole_figure( + self, + pole=(0, 0, 1), + phase=None, + mask=None, + signal_range="auto", + shade_gamma: float = SHADE_GAMMA, + **kwargs, + ): + """Stereographic pole figure of each crystal, masked by the phase decision. + + Parameters + ---------- + pole : tuple of int, default=(0, 0, 1) + Crystal direction plotted, in Miller indices [uvw] or [uvtw] of + each crystal. + phase : int or str, optional + Restrict to one crystal. None (default) plots all of them. + mask : np.ndarray, optional + Overrides the automatic phase-and-signal mask entirely. + signal_range : tuple or "auto", default="auto" + Diffracted intensity mapped to black ... full color. + shade_gamma : float, default=:data:`~quantem.diffraction.phase.SHADE_GAMMA` + Brightness exponent of that mask. Ignored when `mask` is given. + **kwargs + Passed to :meth:`OrientationMap.plot_pole_figure`, e.g. + `color_by`, `int_range` and `overlay`. + + Returns + ------- + tuple or list of tuple + A single ``(fig, ax)`` when `phase` names one crystal; otherwise + one ``(fig, ax)`` per crystal, the same rule as + :meth:`plot_orientation`. + """ + out = [] + for i in self._phase_indices(phase): + m = mask if mask is not None else self.mask(i, signal_range, shade_gamma) + out.append(self.orientation_maps[i].plot_pole_figure(pole=pole, mask=m, **kwargs)) + return out[0] if phase is not None and len(out) == 1 else out + + def plot_matches(self, positions, phase=None, **kwargs): + """Matched patterns at a few probe positions, over the measured peaks. + + One panel per candidate, with the measured peaks as gray disks and + the simulated pattern as colored markers, both sized by intensity. + Passing a `dataset` puts the recorded pattern behind them instead; + the pixel size and the fitted origins then come from the peaks, which + carry them from the dataset through the calibration, so neither needs + passing. A position matched by nothing is drawn with its peaks alone + and labelled "no match". + + Parameters + ---------- + positions : list of tuple of int + ``(row, col)`` probe positions, one panel row each. + phase : int or str, optional + Restrict to one crystal. None (default) shows all of them. + matches : tuple of int, default=(0, 1) + Which matches of each crystal to draw; indices a crystal does not + hold are skipped, so the default draws one panel per crystal after + `num_matches=1`. With `num_matches` of 2, (0, 1) shows the best + and the residual match side by side, which is how a probe + straddling two grains shows itself. + dataset : Dataset4dstem, optional + Show the recorded diffraction pattern behind the overlay. + norm : dict or str, optional + Passed to `show_2d`, which draws that pattern, e.g. + {"power": 0.5, "upper_quantile": 0.98}. + measured_scale, measured_power : float, optional + Size and intensity compression of the gray measured peaks. + q_max_plot : float, optional + Half-width of every panel, 1/Angstroms. The default fits it to + the peaks plotted. + q_max_quantile : float, default=0.98 + Quantile of the measured peak radii setting that automatic limit. + A few stray high-angle detections would otherwise set the scale + for every panel and leave the pattern surrounded by empty space. + Lower it to crop in further; 1.0 encloses every peak. + transpose_plots : bool, default=False + Panel layout only: rows are positions unless this is True. + **kwargs + Further arguments of + :func:`~quantem.diffraction.orientation_visualization.plot_pattern_matches`. + + Returns + ------- + tuple + ``(fig, axs)``. + """ + from quantem.diffraction.orientation_visualization import plot_pattern_matches + + oms = [self.orientation_maps[i] for i in self._phase_indices(phase)] + md = self.peaks.metadata or {} + if kwargs.get("dataset") is not None: + if md.get("pixel_size") is not None: + kwargs.setdefault("pixel_size", float(md["pixel_size"])) + if md.get("origins") is not None: + kwargs.setdefault("origins", np.asarray(md["origins"])) + return plot_pattern_matches(oms, positions=positions, **kwargs) + + def plot_ring_comparison(self, k_min: float = 0.1, k_max: float | None = None, **kwargs): + """Measured radial peak distribution against the rings of every crystal. + + One panel per candidate: the red fill is the histogram of every + calibrated peak, the black lines are that crystal's ring positions. + Run it before matching. With the scale fixed by a standard, a ring + that sits beside the measured peaks means either the reference + lattice parameter is wrong for this specimen or the calibration did + not transfer -- and matching cannot recover from either. + + Parameters + ---------- + k_min : float, default=0.1 + Smallest scattering vector shown, 1/Angstroms. + k_max : float, optional + Largest scattering vector shown; defaults to the map's own k_max. + k_broadening : float, optional + Broaden the rings into a simulated profile; None (default) draws + sharp lines. + **kwargs + Further arguments of + :func:`~quantem.diffraction.calibration.plot_ring_comparison`. + + Returns + ------- + tuple + ``(fig, axs)``. + """ + from quantem.diffraction import calibration + + return calibration.plot_ring_comparison( + self.peaks, + [om.crystal for om in self.orientation_maps], + k_min=k_min, + k_max=k_max if k_max is not None else self._k_max_or_crystals(), + **kwargs, + ) + + def plot_calibration(self, k_min: float = 0.05, k_max: float | None = None, **kwargs): + """Calibrated peaks against the rings of every crystal in the map. + + One column per crystal: the radial histogram of every peak against + that crystal's rings, and every peak as azimuth against scattering + vector with the same rings. Run it after calibrating, whichever + phase the calibration was fit to: every candidate should line up, + and one that does not has the wrong lattice parameter for this + specimen, which matching cannot recover from. + + Parameters + ---------- + k_min : float, default=0.05 + Smallest scattering vector shown, 1/Angstroms. + k_max : float, optional + Largest scattering vector shown; defaults to the map's own k_max. + **kwargs + Further arguments of + :func:`~quantem.diffraction.calibration.plot_calibration`, e.g. + `zone_axis` to show only the rings of one zone, `k_broadening` + and `marker_size`. + + Returns + ------- + tuple + ``(fig, axs)``, axs of shape (2, number of crystals). + """ + from quantem.diffraction import calibration + + return calibration.plot_calibration( + self.peaks, + [om.crystal for om in self.orientation_maps], + k_min=k_min, + k_max=k_max if k_max is not None else self._k_max_or_crystals(), + **kwargs, + ) + + def plot_bragg_rings(self, **kwargs): + """Bragg vector map of every peak with the rings of every crystal. + + Parameters + ---------- + **kwargs + Arguments of + :func:`~quantem.diffraction.calibration.plot_bragg_rings`, e.g. + `n_rings`, `q_max` and `zone_axis`. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + from quantem.diffraction import calibration + + return calibration.plot_bragg_rings( + self.peaks, [om.crystal for om in self.orientation_maps], **kwargs + ) + + def _k_max_or_crystals(self) -> float: + """Scattering-vector limit, falling back to the crystals\' own values. + + Maps saved before ``k_max`` lived on the CrystalMap carry it only on + their crystals, so this keeps those files loadable. + + Returns + ------- + float + """ + k = getattr(self, "k_max", None) + if k is None: + k = max(float(om.crystal.k_max or 1.5) for om in self.orientation_maps) + return float(k) + + def plot_correlation(self, **kwargs): + """Correlation and reliability of every crystal, one panel each. + + The top row is the best correlation of each crystal, the bottom row + its reliability. Read them with care: the correlation is a cosine + similarity, so it saturates on patterns carrying only a few peaks + and stays high on the substrate, and the reliability compares + crystals rather than testing whether one is there at all. Use + :meth:`plot_phase` or :meth:`signal_confidence` for that. + + Parameters + ---------- + mask : bool, default=False + If True, multiply every panel by :meth:`signal_confidence`, so + positions with no diffracted signal go to zero. + shared_scale : bool, default=True + Put every crystal on one scale, so the panels can be compared + directly. Each row keeps its own range, since correlation and + reliability are different quantities. False lets each panel + autoscale, which shows the structure within a weak crystal at + the cost of comparability. Passing `norm` overrides both. + **kwargs + Passed to :func:`~quantem.core.visualization.show_2d`. + + Returns + ------- + tuple + ``(fig, axs)``. + """ + from quantem.core.visualization import show_2d + + mask = kwargs.pop("mask", False) + shared_scale = kwargs.pop("shared_scale", True) + corr = [om.corr[..., 0].numpy() for om in self.orientation_maps] + rel = [om.reliability.numpy() for om in self.orientation_maps] + if mask: + conf = self.signal_confidence() + corr = [c * conf for c in corr] + rel = [r * conf for r in rel] + if shared_scale and "norm" not in kwargs: + # one scale per row, so the crystals are directly comparable; + # correlation and reliability keep their own ranges + kwargs["norm"] = [ + [ + { + "interval_type": "manual", + "vmin": float(min(np.nanmin(a) for a in row)), + "vmax": float(max(np.nanmax(a) for a in row)), + } + ] + * len(row) + for row in (corr, rel) + ] + kwargs.setdefault("cbar", True) + # panels shaped like the scan, so a wide map leaves no gap between rows + R, C = self.shape + kwargs.setdefault("axsize", (4.5, 4.5 * R / C)) + kwargs.setdefault( + "title", + [ + [f"{n} correlation" for n in self.names], + [f"{n} reliability" for n in self.names], + ], + ) + sb = self.orientation_maps[0].scan_scalebar + if sb is not None: + kwargs.setdefault("scalebar", [[sb] + [False] * (len(corr) - 1), [False] * len(corr)]) + return show_2d([corr, rel], **kwargs) + + def _phase_indices(self, phase) -> list[int]: + """Resolve a crystal selector to a list of indices. + + Parameters + ---------- + phase : int, str or None + One crystal by index or name, or None for all of them. + + Returns + ------- + list of int + + Raises + ------ + KeyError + If `phase` names no crystal in this map, by name or index. + """ + if phase is None: + return list(range(len(self.orientation_maps))) + if isinstance(phase, str): + if phase not in self.names: + raise KeyError(f"no crystal named {phase!r}; have {self.names}") + return [self.names.index(phase)] + i = int(phase) + if not -len(self.names) <= i < len(self.names): + raise KeyError(f"no crystal {i}; have {len(self.names)}") + return [i % len(self.names)] + + # ------------------------------------------------------------------ + # checkpointing + # ------------------------------------------------------------------ + + def save(self, path, mode: str = "w", include_plan: bool = False, **kwargs): + """Save the whole analysis to one file. + + Load it back with :func:`quantem.core.io.serialize.load`. + + Parameters + ---------- + path : str or Path + Target path; a ".zip" extension writes one zip file, anything + else a directory. + mode : {"w", "o"}, default="w" + "w" refuses to replace an existing file, "o" overwrites it. + include_plan : bool, default=False + The correlation plan dominates the file size and is rebuilt in + seconds by :meth:`build_plan`, so it is dropped by default. Pass + True to keep it and reload a map ready to match again. + **kwargs + Passed to :meth:`AutoSerialize.save`, e.g. `compression_level`. + """ + if include_plan: + return AutoSerialize.save(self, path, mode=mode, **kwargs) + stash = [(om, om.plan_fft) for om in self.orientation_maps] + try: + for om, _ in stash: + om.plan_fft = None + return AutoSerialize.save(self, path, mode=mode, **kwargs) + finally: + for om, plan in stash: + om.plan_fft = plan diff --git a/src/quantem/diffraction/data/lobato.json b/src/quantem/diffraction/data/lobato.json new file mode 100644 index 000000000..41487366d --- /dev/null +++ b/src/quantem/diffraction/data/lobato.json @@ -0,0 +1,1650 @@ +{ + "H": [ + [ + 0.00647384848835291, + -0.490192576780229, + 0.573284160390876, + -0.37940330148399, + 0.554426474774079 + ], + [ + 2.78519885379148, + 2.77620428330644, + 2.77538591050625, + 2.76759302867258, + 2.76511897642927 + ] + ], + "He": [ + [ + 3.05745116099835, + -62.0044779127325, + 64.0055537084614, + -5.0013257854278, + 0.151798828700526 + ], + [ + 1.08967248726078, + 0.939838798143121, + 0.925289034386265, + 0.82294749870865, + 0.577393110675402 + ] + ], + "Li": [ + [ + 3.92622272886147, + -4.54861962639998, + 2.19335312878658, + 0.0699451265033965, + 0.00209864224851937 + ], + [ + 8.1427601351728, + 4.98941077007855, + 4.1442899923941, + 0.40192231506568, + 0.156479034719823 + ] + ], + "Be": [ + [ + 3.39824970557054, + -1.90866886095696, + 0.0390702117539227, + -0.0111631010210714, + 0.00946204465357523 + ], + [ + 4.44270178622409, + 3.32451542526423, + 0.189772880348214, + 0.0871918614644603, + 0.082780906004134 + ] + ], + "B": [ + [ + 1.47279248639329, + -0.401933042199387, + 0.305998956982689, + 0.0196144217173168, + 0.0009771771060882 + ], + [ + 3.74974048281819, + 0.588066536139673, + 0.515639613103011, + 0.121377570080603, + 0.0680982412160313 + ] + ], + "C": [ + [ + 124.466088621343, + -220.352857078963, + 195.235352280479, + -98.1079361269799, + 0.0142023041213623 + ], + [ + 2.42120849256005, + 2.30537943752425, + 2.04851932106564, + 1.93352552917547, + 0.0768976818478339 + ] + ], + "N": [ + [ + 58.1327150702556, + -147.542409087812, + 130.143065649639, + -39.6195674084154, + 0.010595776333148 + ], + [ + 1.70044856413471, + 1.5590385260174, + 1.41576827473146, + 1.27841818205455, + 0.0565587798474805 + ] + ], + "O": [ + [ + 29.9474045242362, + -77.6101266255278, + 99.8817764623144, + -51.2127005505673, + 0.00819618954446032 + ], + [ + 1.3028398788001, + 1.15794105258309, + 1.00988549338025, + 0.943327971433266, + 0.0433197611321825 + ] + ], + "F": [ + [ + 0.948984894503524, + -30.1333923043554, + 52.7965078127338, + -22.7062703795272, + 0.00656997664531441 + ], + [ + 1.45882933198645, + 0.68877999318768, + 0.654239869346695, + 0.614836130811994, + 0.0342837419495011 + ] + ], + "Ne": [ + [ + 0.582741192220907, + 0.370676561841054, + -0.546744967350809, + 0.414052682480208, + 0.00519903080863993 + ], + [ + 1.28118573143877, + 0.444520897170477, + 0.198650875510481, + 0.185477246656276, + 0.0275738382033885 + ] + ], + "Na": [ + [ + 23.6700603946792, + -21.8531786159742, + 0.592499448108946, + -0.0244652290310244, + 0.00483950221706535 + ], + [ + 8.45148773514603, + 8.04096600474298, + 0.624996000526315, + 0.132450394947296, + 0.0233994362049878 + ] + ], + "Mg": [ + [ + 4.85501047687149, + -2.66220906476843, + 0.478001236085108, + -0.0702307064692064, + 0.00398905828104019 + ], + [ + 5.94639273842456, + 4.17130312520697, + 0.398269808150374, + 0.161886185837474, + 0.0195345056363103 + ] + ], + "Al": [ + [ + 2.83409561607507, + -4.28004133378261, + 4.42191680548311, + -0.03457744718964, + 0.00352385941406079 + ], + [ + 6.66235023980533, + 0.551294722224021, + 0.509328963445973, + 0.111784837425331, + 0.0167602351805257 + ] + ], + "Si": [ + [ + 2.87189142611612, + -2.06173501195173, + 2.17114024204478, + -0.0663073633058801, + 0.00301070709670513 + ], + [ + 5.08487103642989, + 0.429178185305126, + 0.366485434192162, + 0.119710611296903, + 0.0143994536128397 + ] + ], + "P": [ + [ + 2.79151840023151, + -4.36506837823822, + 4.43558455516699, + -0.0809635773399473, + 0.0026790001796644 + ], + [ + 3.90065961865466, + 0.32982596837715, + 0.306089956505888, + 0.108083232545972, + 0.0125894495331186 + ] + ], + "S": [ + [ + 2.67971415610199, + -0.474252822230755, + 0.514835948989687, + -0.0958360024990722, + 0.00248871963818935 + ], + [ + 3.06889121199971, + 0.378216702185809, + 0.188721811902548, + 0.0923370590031494, + 0.0111920877241144 + ] + ], + "Cl": [ + [ + 2.5662483998002, + -0.338876350828591, + 1.14584558755515, + -0.923109316547079, + 0.00229168002041042 + ], + [ + 2.41594920365612, + 0.421414239310216, + 0.10959240497583, + 0.0990955458226753, + 0.00999665948927521 + ] + ], + "Ar": [ + [ + 2.45981746414068, + -0.364198177076995, + 0.250584477222474, + -0.0577437029544345, + 0.00230143866828583 + ], + [ + 1.94004631988856, + 0.399241067884398, + 0.117472406274412, + 0.0567803726023621, + 0.00915579832920708 + ] + ], + "K": [ + [ + 5.81107878601454, + -50.2537096539422, + 48.8609412059842, + 0.0740628592048382, + 0.00072780273862622 + ], + [ + 12.6691483399036, + 3.95641039698166, + 3.68385059577154, + 0.107458517569562, + 0.00665576789391501 + ] + ], + "Ca": [ + [ + 21.1781161524159, + -339.043824317468, + 322.756958523296, + 0.0650077673896996, + 0.00065587436657859 + ], + [ + 6.39608619431736, + 3.74024713891749, + 3.64888449922605, + 0.0945090634514673, + 0.00598520619883758 + ] + ], + "Sc": [ + [ + 12.6035186572148, + -276.875382053702, + 268.871603907342, + 0.0556824178897458, + 0.0005770712551454 + ], + [ + 6.15625615363852, + 3.08873554266679, + 3.02727663298548, + 0.0818874748375215, + 0.00538289832050797 + ] + ], + "Ti": [ + [ + 8.57595775238129, + -210.331563465304, + 206.097172601541, + 0.0477773948977261, + 0.00050571648448032 + ], + [ + 6.00780668875581, + 2.60285856745213, + 2.55352345051105, + 0.0711429484024153, + 0.00485628439383726 + ] + ], + "V": [ + [ + 6.52768433234789, + -200.430576829172, + 198.015053889994, + 0.0413911518064005, + 0.00044745502432986 + ], + [ + 5.83552479352448, + 2.23255952381127, + 2.19786018594229, + 0.0623973876828547, + 0.00440383649079958 + ] + ], + "Cr": [ + [ + 3.02831784843691, + -95.5393933081432, + 96.1761562352198, + 0.0359777315957987, + 0.00039149289071842 + ], + [ + 8.35911504314631, + 1.80263790264103, + 1.77509488957047, + 0.0548144412143113, + 0.00399828968916034 + ] + ], + "Mn": [ + [ + 4.37417550633122, + -160.925510918779, + 160.273308060322, + 0.0312303861049162, + 0.00034696602102094 + ], + [ + 5.51031705504928, + 1.68798216402433, + 1.66614047777702, + 0.0483370390321277, + 0.00364746960053068 + ] + ], + "Fe": [ + [ + 3.79810090836859, + -91.6893549381687, + 91.4454252155429, + 0.0272754344027604, + 0.00030337985438377 + ], + [ + 5.31712645899433, + 1.49713094884748, + 1.4680924180441, + 0.0427247850128954, + 0.00332791855231874 + ] + ], + "Co": [ + [ + 3.33037874467544, + -77.0017596472967, + 77.0725221790516, + 0.0239904669040569, + 0.00026825666552394 + ], + [ + 5.18135964579543, + 1.32915122265139, + 1.30284928963228, + 0.0380645486826371, + 0.0030501000799157 + ] + ], + "Ni": [ + [ + 2.96908078725286, + -75.747706912904, + 76.0398287625307, + 0.0210162113692967, + 0.00023115175113515 + ], + [ + 5.0418094910938, + 1.18275507921629, + 1.16216545846629, + 0.0337479087883035, + 0.0027868086207074 + ] + ], + "Cu": [ + [ + 1.75207145212145, + -43.0410523492124, + 44.0705915543573, + 0.0186876154088144, + 0.00020172732479162 + ], + [ + 6.18750497986187, + 1.00266263628976, + 0.98538431135303, + 0.0302984703916117, + 0.00255855598748879 + ] + ], + "Zn": [ + [ + 2.46637110499459, + -61.4678541332537, + 62.0176945237481, + 0.0164160173931416, + 0.00017248711759538 + ], + [ + 4.91028078493815, + 0.967898520322992, + 0.951283834775355, + 0.0269600967667566, + 0.00234109611046253 + ] + ], + "Ga": [ + [ + 2.76010203108428, + -34.4452614207467, + 35.2262267244016, + 0.0132067196999419, + 0.00012594556092825 + ], + [ + 6.10128224537662, + 0.765143313553464, + 0.751328623382859, + 0.0224879634317251, + 0.00206737374278712 + ] + ], + "Ge": [ + [ + 3.1824163526, + -52.4514037811166, + 52.9690827162218, + 0.0114096168591864, + 9.509543581451e-05 + ], + [ + 5.01719040860914, + 0.712395764437798, + 0.702280192528194, + 0.0196747295694015, + 0.00184146614494011 + ] + ], + "As": [ + [ + 3.45642969119604, + -33.3176044431721, + 33.5712193855332, + 0.00979002295640342, + 6.534348622652e-05 + ], + [ + 4.01358016032945, + 0.662355778050629, + 0.645771941056073, + 0.0170919353231011, + 0.0016030160283942 + ] + ], + "Se": [ + [ + 3.64905047801926, + -43.6851662221238, + 43.6920288601154, + 0.00844902284199144, + 3.7861147411e-05 + ], + [ + 3.25043267112593, + 0.609666201650289, + 0.596971300802123, + 0.0148554512740362, + 0.00133625587356563 + ] + ], + "Br": [ + [ + 3.83846312224289, + -52.2723471011233, + 51.9861279495659, + 0.00733955989338044, + 1.646942079513e-05 + ], + [ + 2.61189470573232, + 0.566195062874718, + 0.555279326699873, + 0.012974646943244, + 0.00102986536855875 + ] + ], + "Kr": [ + [ + 4.02541030310827, + -46.3042332078929, + 45.7213681904101, + 0.00635337959625315, + 1.33477845256e-06 + ], + [ + 2.1364838144374, + 0.526591166454131, + 0.514136784464577, + 0.0112807242242006, + 0.00048808985794097 + ] + ], + "Rb": [ + [ + 3.38975351595394, + 2.14348348679172, + 0.354322603510978, + 0.00374009340085664, + 3.0034250234e-07 + ], + [ + 20.5744814368171, + 1.91079945218516, + 0.197410589396603, + 0.00813459465343988, + 0.00029268579105695 + ] + ], + "Sr": [ + [ + 4.77092509299826, + 1.47597850155283, + 0.304451355544188, + 0.00359474981916728, + 3.0008554923e-07 + ], + [ + 13.3668881330412, + 1.33738379577196, + 0.177532369437782, + 0.00779105033001826, + 0.00028225513964854 + ] + ], + "Y": [ + [ + 4.60721019875234, + 1.4280185103984, + 0.295581045577753, + 0.00338997840852895, + 2.6686297497e-07 + ], + [ + 10.8686905557151, + 1.31137455873124, + 0.168022870784711, + 0.00735964545350478, + 0.00026234328102139 + ] + ], + "Zr": [ + [ + 4.31175453406871, + 1.49331578039339, + 0.281236050128884, + 0.00309337738455747, + 2.5802444851e-07 + ], + [ + 9.4589658051749, + 1.33063622845245, + 0.156506871444453, + 0.00682410523811735, + 0.0002498188791107 + ] + ], + "Nb": [ + [ + 3.11179039113495, + 2.20259060945803, + 0.27033077491002, + 0.00268799194751098, + 2.3254948335e-07 + ], + [ + 10.6903141343023, + 1.65316356158933, + 0.145115185718969, + 0.0061395635158285, + 0.00023224023593767 + ] + ], + "Mo": [ + [ + 2.83105968432053, + 2.34858137489674, + 0.245105888429696, + 0.00235282108218515, + 2.3127083834e-07 + ], + [ + 10.435719575895, + 1.60482868674597, + 0.131696934774606, + 0.00554977901383345, + 0.0002227474637649 + ] + ], + "Tc": [ + [ + 2.57179859323352, + 2.45633741957203, + 0.220658440881371, + 0.00205531418012438, + 2.3213294779e-07 + ], + [ + 10.1643117131377, + 1.5344191924455, + 0.11913861975411, + 0.00501853240882988, + 0.00021459037912086 + ] + ], + "Ru": [ + [ + 2.33230030353946, + 2.53578025489049, + 0.198208042136026, + 0.00176120130460053, + 1.9812941155e-07 + ], + [ + 9.92167460159959, + 1.45585668808788, + 0.107682187873349, + 0.00447243545654967, + 0.0001965747162325 + ] + ], + "Rh": [ + [ + 2.1135253485947, + 2.58636316155013, + 0.177063913863686, + 0.00149737051003056, + 2.0548144556e-07 + ], + [ + 9.65913725859762, + 1.37106656934329, + 0.0970353027584648, + 0.00397128509088792, + 0.0001913551907378 + ] + ], + "Pd": [ + [ + 0.642159796182687, + 2.97914814426328, + 0.16815442600427, + 0.00133744213878407, + 1.9141096825e-07 + ], + [ + 5.97479750263406, + 1.43359432541277, + 0.0909868401172677, + 0.00362410137116162, + 0.00018068891441605 + ] + ], + "Ag": [ + [ + 1.55317216680389, + 2.63930364699988, + 0.142015486956788, + 0.0010085046009976, + 1.9463842997e-07 + ], + [ + 8.15620235758956, + 1.21600887481801, + 0.0790098864915873, + 0.00296547901363949, + 0.00017459309591915 + ] + ], + "Cd": [ + [ + 61.530785199286, + -78.6016741201582, + 21.550129260277, + 0.137685015664191, + 0.00032464493094652 + ], + [ + 3.11468102533247, + 2.76016983388574, + 1.93551312324722, + 0.0722468347260293, + 0.00117001629624684 + ] + ], + "In": [ + [ + 4.22232177901524, + -26.4121318353202, + 27.2852852708746, + 0.121617901884138, + 0.00030688354637153 + ], + [ + 6.07265510403227, + 1.64550178959387, + 1.52257074939506, + 0.0656507914695225, + 0.00112093851473035 + ] + ], + "Sn": [ + [ + 5.14222074642053, + -25.4945413762576, + 25.7414487548601, + 0.111778227592199, + 0.00029364738473774 + ], + [ + 5.27272636474484, + 1.53194959209148, + 1.40257525040087, + 0.0611685387263086, + 0.00107560815394522 + ] + ], + "Sb": [ + [ + 6.24164031831883, + -93.3868724419552, + 92.6332875829884, + 0.103412692221298, + 0.00028184842689439 + ], + [ + 4.26984108082687, + 1.3940774614025, + 1.35566854910434, + 0.0572667949652199, + 0.00103307954282067 + ] + ], + "Te": [ + [ + 7.37743301813403, + -126.025106931628, + 124.128405004095, + 0.0959978371818569, + 0.00027107221639888 + ], + [ + 3.46917757783897, + 1.29759810886736, + 1.26771067646066, + 0.0537771750179423, + 0.00099308457702918 + ] + ], + "I": [ + [ + 9.64400666272123, + -122.924435350112, + 118.682564816567, + 0.0895025947025539, + 0.00026127612149048 + ], + [ + 2.72645545366544, + 1.23723425850866, + 1.20062036872249, + 0.0506678682285199, + 0.0009554378833835 + ] + ], + "Xe": [ + [ + 15.545174967486, + -118.241027856744, + 108.009524963197, + 0.0836259341992221, + 0.00025199186242907 + ], + [ + 2.10637340865492, + 1.20860376129512, + 1.15395270567214, + 0.0478189391274824, + 0.00091988625759261 + ] + ], + "Cs": [ + [ + 4.28708739181692, + 3.23250665422965, + 0.674029533561706, + 0.0618908083387697, + 0.00023561205295219 + ], + [ + 22.6587870798541, + 2.23797386470064, + 0.368995568664316, + 0.0402266575298484, + 0.00088376189087804 + ] + ], + "Ba": [ + [ + 6.24475187390461, + 2.35172271416388, + 0.474279373222252, + 0.0638113874286849, + 0.00023465128056279 + ], + [ + 15.1431354190985, + 1.45379000604814, + 0.320835646387801, + 0.0404354532094699, + 0.00085408113128003 + ] + ], + "La": [ + [ + 6.09788179599509, + 2.19495164736675, + 0.548172791960022, + 0.0616669573222662, + 0.00022680735586471 + ], + [ + 12.4288544315854, + 1.50535992352023, + 0.333738039717124, + 0.0387744553499998, + 0.00082404988406487 + ] + ], + "Ce": [ + [ + 5.7952687964724, + 2.37022664107843, + 0.471398756901114, + 0.0574368260587866, + 0.00021897948925966 + ], + [ + 14.2801055058256, + 1.35969015719134, + 0.302017349663514, + 0.0366436798147007, + 0.00079543452651837 + ] + ], + "Pr": [ + [ + 5.60406255377525, + 2.35796259561812, + 0.476001098572882, + 0.0544123374315247, + 0.00021141460220515 + ], + [ + 13.9517490230574, + 1.31239754917439, + 0.294933701192956, + 0.0348601542462324, + 0.00076826168266181 + ] + ], + "Nd": [ + [ + 5.4290839196977, + 2.3368732536088, + 0.483373554104881, + 0.0514654964557479, + 0.00020377613286376 + ], + [ + 13.650364942766, + 1.26759841390319, + 0.288610681576929, + 0.0331291807446609, + 0.00074234848380127 + ] + ], + "Pm": [ + [ + 5.26774445089444, + 2.30855813326305, + 0.493265479026012, + 0.0486356279741696, + 0.00019630884232068 + ], + [ + 13.3602096827394, + 1.22585856586941, + 0.282919632127866, + 0.0314735397127137, + 0.00071765802506962 + ] + ], + "Sm": [ + [ + 5.12680428517077, + 2.2692553400838, + 0.504209300453302, + 0.0456923992416222, + 0.00018867505049389 + ], + [ + 13.1581501026129, + 1.18129508243337, + 0.277107038219062, + 0.0297957604549544, + 0.0006940110991991 + ] + ], + "Eu": [ + [ + 4.97962359749809, + 2.24183087455669, + 0.512933961402893, + 0.0429801848498334, + 0.00018138169248579 + ], + [ + 12.8392663078033, + 1.14705446420486, + 0.270387160935243, + 0.0282418714952602, + 0.00067148660317399 + ] + ], + "Gd": [ + [ + 5.07835830045611, + 1.95744027181659, + 0.592825983221375, + 0.0419502034143925, + 0.00017524109152832 + ], + [ + 10.5132725469047, + 1.11764941281524, + 0.284386741833675, + 0.0272663327641347, + 0.0006503108607073 + ] + ], + "Tb": [ + [ + 4.711616366573, + 2.17261950786578, + 0.539717601342865, + 0.0379796004547811, + 0.00016692376356484 + ], + [ + 12.3109415948539, + 1.08743796767492, + 0.259804965509275, + 0.0253289923895189, + 0.00062927856696191 + ] + ], + "Dy": [ + [ + 4.59075504485146, + 2.13572373036567, + 0.551355560094152, + 0.0356058406718513, + 0.00015982401685679 + ], + [ + 12.0656740623369, + 1.05811737115864, + 0.253994438918661, + 0.0239508739649649, + 0.00060948274843329 + ] + ], + "Ho": [ + [ + 4.4840011062671, + 2.08904347078539, + 0.565932618316104, + 0.0332703590211063, + 0.00015244561028289 + ], + [ + 11.874925069157, + 1.02828493708789, + 0.248907807908076, + 0.0225844065532147, + 0.00059037444515942 + ] + ], + "Er": [ + [ + 4.37665140786963, + 2.04645150383634, + 0.580662860591587, + 0.0312388528327766, + 0.00014537486966254 + ], + [ + 11.6653080717139, + 1.00314769675495, + 0.244153292686312, + 0.0213618580000731, + 0.00057211033845477 + ] + ], + "Tm": [ + [ + 4.28308318247419, + 1.9953800145373, + 0.597053571332504, + 0.0290951668869559, + 0.00013806476903752 + ], + [ + 11.4961905956683, + 0.97703910201882, + 0.239429132975829, + 0.0200912233327255, + 0.00055437713849227 + ] + ], + "Yb": [ + [ + 4.19563840737123, + 1.94333285978645, + 0.612446617285465, + 0.0271515207694932, + 0.00013059478735193 + ], + [ + 11.4107750456652, + 0.949011924454877, + 0.234950326535304, + 0.0189051576280975, + 0.00053723364920971 + ] + ], + "Lu": [ + [ + 4.35692593296384, + 1.69589204777315, + 0.66390452006847, + 0.0263020018760281, + 0.00012549731849982 + ], + [ + 9.29434514718568, + 0.91050004595742, + 0.238746595911376, + 0.0182098542546462, + 0.00052155937196356 + ] + ], + "Hf": [ + [ + 4.33138405664923, + 1.5276486472868, + 0.735795922891238, + 0.0249526323201402, + 0.00011874085258106 + ], + [ + 7.87684433810288, + 0.942515642627748, + 0.241699480039806, + 0.0172899894448509, + 0.00050583463131353 + ] + ], + "Ta": [ + [ + 4.19726001319642, + 1.46854806774708, + 0.783931248211088, + 0.0232494094864395, + 0.00011126135860729 + ], + [ + 6.93674024894457, + 1.01725276618348, + 0.23894893971507, + 0.0162332433075381, + 0.00049023900402115 + ] + ], + "W": [ + [ + 3.97629715819077, + 1.52292684257341, + 0.797706246390561, + 0.0213666741558685, + 0.00010307868938572 + ], + [ + 6.29685779228774, + 1.11289951157691, + 0.231057040678473, + 0.0151013545567204, + 0.00047468247710503 + ] + ], + "Re": [ + [ + 3.75144381439811, + 1.62768802977632, + 0.781756755454271, + 0.0195170803912291, + 9.431998042934e-05 + ], + [ + 5.79754636082953, + 1.1822363107771, + 0.21991358602477, + 0.0139810455481456, + 0.00045912712582984 + ] + ], + "Os": [ + [ + 3.48401517338711, + 1.79377920416965, + 0.74487835669192, + 0.0175427785162434, + 8.448723515689e-05 + ], + [ + 5.43998846080953, + 1.22792134140128, + 0.206208856992909, + 0.0127899401721146, + 0.00044303222678748 + ] + ], + "Ir": [ + [ + 1.599565781988442, + 2.975344521941205, + 6.950926783668822e-01, + 1.487796274586880e-02, + 6.905495791015833e-05 + ], + [ + 5.792444473855675, + 1.553009829732260, + 1.886263359366482e-01, + 1.117634706624533e-02, + 4.227723616555514e-04 + ] + ], + "Pt": [ + [ + 2.04021563975728, + 2.89922634624825, + 0.636344083815797, + 0.0132071919595408, + 5.673821925001e-05 + ], + [ + 6.65819429609627, + 1.41337923778932, + 0.174000104502179, + 0.010068780453021, + 0.00040307710569222 + ] + ], + "Au": [ + [ + 1.6759346706487, + 3.00486602969729, + 0.595340013161635, + 0.0117163186623094, + 4.296782976398e-05 + ], + [ + 5.52231093211402, + 1.38007223007196, + 0.162229237655945, + 0.00901814890416575, + 0.00037927766747767 + ] + ], + "Hg": [ + [ + 2.23522850443105, + 2.68276638651994, + 0.555194926212433, + 0.0107273354366258, + 3.284739920463e-05 + ], + [ + 5.0203098896024, + 1.23077590583777, + 0.152248122992863, + 0.0082839911690621, + 0.00035624193893176 + ] + ], + "Tl": [ + [ + 2.8034273741251, + 2.7188278796602, + 0.522475915391999, + 0.00984537836971042, + 2.345245265083e-05 + ], + [ + 6.55876872804534, + 1.16972422516667, + 0.143556691800178, + 0.00761976526230039, + 0.00032962767390299 + ] + ], + "Pb": [ + [ + 3.60861020997771, + 2.45056774737198, + 0.478639500107257, + 0.00887214235169215, + 1.040019329095e-05 + ], + [ + 6.58162594621962, + 1.02772852660588, + 0.13353368063472, + 0.00684861238948429, + 0.00027638887545667 + ] + ], + "Bi": [ + [ + 4.24209901057315, + 2.09994342055827, + 0.432836631379922, + 0.00802021570381059, + 7.2178503085e-07 + ], + [ + 5.7528011553608, + 0.873901489195461, + 0.123599956180524, + 0.0061760033305857, + 0.00014149517761723 + ] + ], + "Po": [ + [ + 4.63620095668962, + 1.78063311426987, + 0.403730443952738, + 0.00768539554668674, + 8.954159028e-08 + ], + [ + 4.88825299780982, + 0.756310526880074, + 0.117302364452249, + 0.00593034394242894, + 7.66668663873e-05 + ] + ], + "At": [ + [ + 4.9659225059507, + 1.43815561507033, + 0.371299200579298, + 0.00747256428450201, + 1.1411528415e-07 + ], + [ + 4.09129378787496, + 0.629228996602551, + 0.111096678695535, + 0.00577235010347761, + 8.085118938794e-05 + ] + ], + "Rn": [ + [ + 5.30615614475033, + 1.11733160236072, + 0.315858723110854, + 0.0072034224231021, + 1.0735533055e-07 + ], + [ + 3.48800735481655, + 0.481190741149771, + 0.101874467945753, + 0.00559245374417078, + 7.795066735407e-05 + ] + ], + "Fr": [ + [ + 4.52053399042336, + 4.10695397909124, + 0.713946878503728, + 0.0169294027687453, + 8.574921291917e-05 + ], + [ + 19.4482234232948, + 1.89824673155996, + 0.169553563595341, + 0.0114819567548934, + 0.00034612203825798 + ] + ], + "Ra": [ + [ + 6.52401072001873, + 3.20787080745661, + 0.540478774349637, + 0.00878278806898853, + 6.91010602949e-06 + ], + [ + 14.0092554298974, + 1.32635035961665, + 0.131408568386036, + 0.00628647434520409, + 0.00022783998145895 + ] + ], + "Ac": [ + [ + 6.89602853559619, + 2.83514154536523, + 0.503506881888815, + 0.00802294510966706, + 9.204009548e-08 + ], + [ + 11.0763825670398, + 1.17132616290503, + 0.123494651391799, + 0.00573515595790497, + 7.197075418025e-05 + ] + ], + "Th": [ + [ + 7.0937490016263, + 2.52912373903194, + 0.482107819888889, + 0.00771936674145111, + 7.271141994e-08 + ], + [ + 9.09473795165944, + 1.06391667033161, + 0.118619446621877, + 0.00553945708895034, + 6.584947305285e-05 + ] + ], + "Pa": [ + [ + 6.43401324797284, + 2.97099970535788, + 0.460796651787848, + 0.00724033125775455, + 6.362366643e-08 + ], + [ + 10.2596851307297, + 1.13177452306163, + 0.112775935362382, + 0.00527459583353132, + 6.192638240888e-05 + ] + ], + "U": [ + [ + 6.21070826784019, + 3.03934453825623, + 0.437339984439162, + 0.00680714764147316, + 6.182293277e-08 + ], + [ + 10.0214105859162, + 1.10399880155883, + 0.107218973162453, + 0.00502432130850682, + 6.005069941972e-05 + ] + ], + "Np": [ + [ + 6.00431598308986, + 3.09431654531418, + 0.412808487716417, + 0.00640891657880066, + 6.730073762e-08 + ], + [ + 9.81793046719106, + 1.06978705068029, + 0.101635291448645, + 0.00479524846888725, + 6.028065987265e-05 + ] + ], + "Pu": [ + [ + 5.20061716810304, + 3.49849440367144, + 0.40541493113132, + 0.00622343700826529, + 6.008593295e-08 + ], + [ + 11.2038310194406, + 1.12846369546928, + 0.0989333472686131, + 0.0046591743197795, + 5.725906389623e-05 + ] + ], + "Am": [ + [ + 5.02533860036029, + 3.51843988285154, + 0.381950349446272, + 0.00582110208096211, + 6.526093388e-08 + ], + [ + 10.9772190697628, + 1.0847721483263, + 0.0936746874853835, + 0.00442682346873485, + 5.738568370505e-05 + ] + ], + "Cm": [ + [ + 5.34656100260619, + 3.2246846656091, + 0.349461752625445, + 0.00529251756748815, + 6.159176302e-08 + ], + [ + 9.23118379752449, + 0.972836229193835, + 0.0870596220966562, + 0.00413011964465032, + 5.494638944497e-05 + ] + ], + "Bk": [ + [ + 5.22582350334004, + 3.22818873995295, + 0.327098878823851, + 0.00488882575592239, + 5.212722694e-08 + ], + [ + 9.07135706618718, + 0.931124277456634, + 0.0820720602059845, + 0.00389029655993416, + 5.103679766366e-05 + ] + ], + "Cf": [ + [ + 4.58641247753547, + 3.51869595665101, + 0.319261714201205, + 0.00472979647985625, + 5.513244808e-08 + ], + [ + 10.3186102092335, + 0.957278837829232, + 0.079680743144388, + 0.00377193199996848, + 5.09750860376e-05 + ] + ], + "Es": [ + [ + 4.45799480675412, + 3.50867212637662, + 0.301375380528315, + 0.00440763746214481, + 4.887879032e-08 + ], + [ + 10.0890693383401, + 0.919446447361622, + 0.0756419177786329, + 0.00357026465124809, + 4.804951743958e-05 + ] + ], + "Fm": [ + [ + 4.3389757640117, + 3.49185098896403, + 0.281471099016286, + 0.00400209271499046, + 5.529298082e-08 + ], + [ + 9.96330937223836, + 0.878083956982974, + 0.0711637115768777, + 0.00332152404215077, + 4.851031794604e-05 + ] + ], + "Md": [ + [ + 4.22729470403101, + 3.47249227510661, + 0.264822229496898, + 0.00369072877833192, + 6.258714262e-08 + ], + [ + 9.72400602040031, + 0.84287377596021, + 0.0673534743974851, + 0.0031236460626332, + 4.912170970176e-05 + ] + ], + "No": [ + [ + 4.1095170244302, + 3.4579913252275, + 0.247087351222386, + 0.00330423920940995, + 5.991049262e-08 + ], + [ + 9.6773599451012, + 0.806940042470817, + 0.0632815436954167, + 0.00287544639641209, + 4.706791536976e-05 + ] + ], + "Lr": [ + [ + 4.52147421198378, + 3.20212985587804, + 0.223028726956451, + 0.00281716453892029, + 4.06427954e-08 + ], + [ + 8.28309906861142, + 0.731918958125393, + 0.0580942773018655, + 0.00256168016047444, + 4.03816515529e-05 + ] + ] +} diff --git a/src/quantem/diffraction/defaults.py b/src/quantem/diffraction/defaults.py new file mode 100644 index 000000000..df94766fa --- /dev/null +++ b/src/quantem/diffraction/defaults.py @@ -0,0 +1,63 @@ +"""Shared defaults of the matching, phase and refinement chain. + +Every function takes these as keyword arguments, so they can be changed per +call; keeping one value for each across the chain makes the matching, the +phase fit, the strain and the dynamical refinement compare the same peaks +in the same way. +""" + +# peak pairing distance between simulated and measured peaks, and the width +# of the polar correlation kernel (1/Angstroms) +PAIR_DISTANCE = 0.05 + +# intensities are compared as I ** POWER_INTENSITY; 0.25 flattens the +# dynamic range so weak reflections count, 0 compares positions only +POWER_INTENSITY = 0.25 + +# excitation-error envelope of the kinematical patterns (1/Angstroms); +# about twice the physical width, so orientations halfway between library +# zones keep their intensities +SIGMA_EXCITATION = 0.04 + +# excitation-error cutoff of the Bloch wave beam list (1/Angstroms) +SG_MAX = 0.1 + +# simulated reflections weaker than this fraction of the strongest one are +# treated as unobservable and do not count against a candidate +MIN_SIM_INTENSITY_REL = 0.02 + +# detected peaks (including the direct beam) a pattern needs to be matched +# or refined, and paired reflections a least-squares refinement needs +MIN_NUMBER_PEAKS = 5 +MIN_PAIRS = 4 + + +def resolve(value, key: str, *sources: dict | None, default=None): + """First non-None of an explicit value, metadata entries and a default. + + The refinement stages call this so a parameter left as None inherits + the value the previous stage used. + + Parameters + ---------- + value : object + Explicit value; returned whenever it is not None. + key : str + Key looked up in each source. + *sources : dict | None + Metadata dicts, searched in order; None sources are skipped, as are + entries that are None. + default : object, optional + Returned when nothing else is found. + + Returns + ------- + object + The resolved value. + """ + if value is not None: + return value + for src in sources: + if src is not None and src.get(key) is not None: + return src[key] + return default diff --git a/src/quantem/diffraction/digital_dark_field.py b/src/quantem/diffraction/digital_dark_field.py new file mode 100644 index 000000000..3297056e6 --- /dev/null +++ b/src/quantem/diffraction/digital_dark_field.py @@ -0,0 +1,1034 @@ +"""Digital dark field imaging from detected Bragg peaks. + +A digital dark field (DDF) image is the summed intensity of a selected subset +of the detected Bragg peaks at each probe position. These functions select +the peaks in three ways, following MacLaren and co-workers +(https://doi.org/10.1093/mam/ozae104): + +- Virtual apertures: peaks within a radius of a set of aperture positions, + usually a lattice built from two reciprocal lattice vectors + (refine_lattice_vectors, lattice_distance, aperture_array, + aperture_array_subtract, aperture_ddf_image). +- Polar selection: peaks within a ring of radius q, optionally restricted to + a range of azimuthal angles (add_polar_fields, polar_mask, + radial_ddf_image). +- Clustering: DBSCAN in the joint (diffraction, scan) space groups the peaks + into single-spot, single-crystallite clusters (L1), clustering either the + real-space centers of mass or the dark field images of those clusters + groups the spots of each grain (L2), and clustering the remaining peaks in + diffraction space alone isolates ring-like nanocrystalline or amorphous + components (L3) (ddf_images, cluster_coms, cluster_centers, + group_ddf_images, assign_grain_labels). + +The aperture, polar and DDF image functions are ports of Ian MacLaren's +digital dark field functions in py4DSTEM (aperture_array_generator, +aperture_array_subtract, DDFimage, pointlist_to_array with rphi=True, +DDF_radial_image and DDFradialazimuthimage), rewritten to read the peaks +from a Vector. + +All functions read the peaks from a Vector with one cell per probe position. +The diffraction coordinates are given by `q_fields`, which defaults to +("qx", "qy") for calibrated peaks or ("q_row", "q_col") for peaks in detector +pixels. Boolean masks are aligned with the flattened rows of the Vector, so +they can be combined with & and | and passed to ddf_image or to +quantem.core.utils.clustering.filter_rows. + +Where a function needs the diffraction origin (`center=None`), calibrated +("qx", "qy") peaks are taken to be relative to the direct beam, so the origin +is (0, 0). Peaks in detector pixels ("q_row", "q_col") use the "origin_ref" +stored in the Vector metadata by BraggVectors.correct_peak_origins; without +it, pass `center` explicitly. + +The azimuth qphi of the polar functions follows py4DSTEM's DDF functions: +qphi = atan2(-q0, q1) in degrees, measured anticlockwise from the +col +(right) direction as the pattern is displayed with rows increasing +downward, so phi ranges written for py4DSTEM carry over unchanged. This is +not the azimuth used by quantem.diffraction.calibration, atan2(q1, q0), +which is measured from the +row axis toward +col. +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.collections import EllipseCollection +from matplotlib.colors import hsv_to_rgb + +from quantem.core.utils.clustering import cluster_vector, dbscan # noqa: F401 + + +def _scan_cells(vector) -> np.ndarray: + """(N, 2) scan (row, col) of every flattened row of a ragged Vector.""" + counts = np.asarray(vector.row_counts(), dtype=int) + shape = vector.shape[:2] + cell_r, cell_c = np.divmod(np.arange(counts.size), shape[1]) + return np.stack([np.repeat(cell_r, counts), np.repeat(cell_c, counts)], axis=1) + + +def _resolve_q_fields(vector, q_fields) -> tuple[str, str]: + """The two diffraction coordinate fields of a peak Vector.""" + if q_fields is not None: + return tuple(q_fields) + for candidate in (("qx", "qy"), ("q_row", "q_col")): + if all(f in vector.fields for f in candidate): + return candidate + raise KeyError( + f"No diffraction coordinate fields found in {vector.fields}; pass q_fields explicitly." + ) + + +def _resolve_center(vector, q_fields, center) -> np.ndarray: + """Diffraction origin: `center` if given, else (0, 0) or the stored origin_ref. + + Calibrated fields are relative to the direct beam already. Detector-pixel + fields ("q_row", "q_col") need the common origin that + BraggVectors.correct_peak_origins stores as "origin_ref". + """ + if center is not None: + return np.asarray(center, dtype=np.float64).reshape(2) + if tuple(q_fields) == ("q_row", "q_col"): + origin_ref = vector.metadata.get("origin_ref") + if origin_ref is None: + raise ValueError( + "Peaks are in detector pixels and carry no 'origin_ref'; pass " + "center=(row, col) of the direct beam, or correct the origins with " + "BraggVectors.correct_peak_origins first." + ) + return np.asarray(origin_ref, dtype=np.float64).reshape(2) + return np.zeros(2) + + +def _q_coordinates(vector, q_fields=None, center=(0.0, 0.0)) -> np.ndarray: + """(N, 2) diffraction coordinates of every flattened row, relative to center.""" + q_fields = _resolve_q_fields(vector, q_fields) + q = vector.select_fields(*q_fields).numpy().astype(np.float64) + return q - np.asarray(center, dtype=np.float64)[None, :] + + +# --------------------------------------------------------------------------- # +# DDF images +# --------------------------------------------------------------------------- # + + +def ddf_image( + peaks, + mask=None, + intensity_field: str = "intensity", +) -> np.ndarray: + """Digital dark field image from a subset of the peaks. + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + mask : array-like of bool, optional + (N,) selection aligned with the flattened rows of `peaks`, for + example from aperture_mask or polar_mask. None uses every peak. + intensity_field : str, default="intensity" + Field summed at each probe position. Negative values are clipped to 0. + + Returns + ------- + np.ndarray + (scan_row, scan_col) image. + """ + inten = peaks.select_fields(intensity_field).numpy()[:, 0].astype(np.float64).clip(min=0) + rc = _scan_cells(peaks) + if mask is not None: + mask = np.asarray(mask, dtype=bool) + inten, rc = inten[mask], rc[mask] + R, C = peaks.shape[:2] + image = np.zeros((R, C)) + np.add.at(image, (rc[:, 0], rc[:, 1]), inten) + return image + + +# --------------------------------------------------------------------------- # +# Virtual apertures +# --------------------------------------------------------------------------- # + + +def aperture_array( + g1, + g2=None, + mode: str = "array", + center=(0.0, 0.0), + shift=(0.0, 0.0), + n1_range: tuple[int, int] = (-5, 5), + n2_range: tuple[int, int] = (-5, 5), + radius_range: tuple[float, float] = (0.0, np.inf), + shape=None, + edge: float = 0.0, +) -> np.ndarray: + """Virtual aperture positions on a lattice of diffraction vectors. + + Each aperture sits at center + (n1 + s1) g1 + (n2 + s2) g2, where (s1, s2) + is `shift`. We keep the positions whose distance from `center` falls + inside `radius_range`, which is how the direct beam is usually excluded. + When the peaks were shifted to a common origin with + BraggVectors.correct_peak_origins, setting `center` to that origin puts + the apertures in detector pixels, so they can be drawn over the mean + pattern or the Bragg vector map. + + Parameters + ---------- + g1, g2 : array-like of float + (2,) lattice vectors in the same coordinates as the peaks. g2 is + required for mode="array", ignored for mode="line", and optional for + mode="single" (where it is only used with a nonzero s2). + mode : {"array", "line", "single"}, default="array" + "array" places a 2D lattice of apertures over n1_range and n2_range, + "line" places a row of apertures along g1 over n1_range (a systematic + row, such as a two-beam condition), and "single" places one aperture + at (s1 g1 + s2 g2). + center : array-like of float, default=(0, 0) + (2,) origin of the lattice, normally the direct beam. + shift : (float, float), default=(0, 0) + Lattice offset (s1, s2) in multiples of g1 and g2; fractional values + place apertures between lattice points, for example (0.5, 0.5) for a + centered superlattice. + n1_range, n2_range : (int, int), default=(-5, 5) + Inclusive range of lattice multiples of g1 and g2. + radius_range : (float, float), default=(0, inf) + Inner and outer distance from `center` of the apertures kept. + shape : (int, int), optional + Detector shape. When given, apertures closer than `edge` to the + detector boundary are removed, which requires `center` to be in + detector pixels. + edge : float, default=0 + Boundary width in pixels, used only with `shape`. + + Returns + ------- + np.ndarray + (N, 2) aperture positions. + """ + if mode == "array" and g2 is None: + raise ValueError('mode="array" needs both g1 and g2.') + g1 = np.asarray(g1, dtype=np.float64) + g2 = np.zeros(2) if g2 is None else np.asarray(g2, dtype=np.float64) + s1, s2 = (float(v) for v in shift) + if mode == "single": + n1 = np.array([0]) + n2 = np.array([0]) + elif mode == "line": + n1 = np.arange(n1_range[0], n1_range[1] + 1) + n2 = np.zeros_like(n1) + elif mode == "array": + n1, n2 = np.meshgrid( + np.arange(n1_range[0], n1_range[1] + 1), + np.arange(n2_range[0], n2_range[1] + 1), + indexing="ij", + ) + n1, n2 = n1.ravel(), n2.ravel() + else: + raise ValueError(f"mode must be 'array', 'line' or 'single', got {mode!r}.") + + offsets = (n1[:, None] + s1) * g1[None, :] + (n2[:, None] + s2) * g2[None, :] + r = np.hypot(offsets[:, 0], offsets[:, 1]) + keep = (r >= radius_range[0]) & (r <= radius_range[1]) + positions = offsets + np.asarray(center, dtype=np.float64)[None, :] + if shape is not None: + keep &= ( + (positions[:, 0] > edge) + & (positions[:, 0] < shape[0] - edge) + & (positions[:, 1] > edge) + & (positions[:, 1] < shape[1] - edge) + ) + return positions[keep] + + +def refine_lattice_vectors( + peaks, + g1, + g2, + center=None, + radius: float = 6.0, + n1_range: tuple[int, int] = (-5, 5), + n2_range: tuple[int, int] = (-5, 5), + num_iterations: int = 3, + q_fields=None, + intensity_field: str = "intensity", +): + """Refine one pair of lattice vectors against the peaks of the whole scan. + + This is a single global refinement, used to place virtual apertures; it is + not the per-position lattice fit of BraggVectors.fit_lattice used for + strain mapping. For each lattice point n1 g1 + n2 g2 (excluding the + origin) we take the + intensity-weighted mean position of all peaks within `radius` of it, + summed over the scan, then solve for g1 and g2 by weighted least squares + with each point weighted by its summed intensity. Repeating this a few + times lets the fit follow lattice points that start near the edge of + `radius`. The lattice with the most total intensity dominates, so the + starting vectors should be close to the orientation of interest. + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + g1, g2 : array-like of float + (2,) starting lattice vectors, in the units of the peak coordinates. + center : array-like of float, optional + (2,) lattice origin, held fixed. None uses (0, 0) for calibrated peaks + and the stored "origin_ref" for peaks in detector pixels. + radius : float, default=6.0 + Search radius around each lattice point, in the units of the peak + coordinates. + n1_range, n2_range : (int, int), default=(-5, 5) + Inclusive range of lattice multiples used in the fit. + num_iterations : int, default=3 + Number of search and fit passes. + q_fields : (str, str), optional + Diffraction coordinate fields. + intensity_field : str, default="intensity" + Field used as the weight of each peak. + + Returns + ------- + g1, g2 : np.ndarray + (2,) refined lattice vectors. + """ + q_fields = _resolve_q_fields(peaks, q_fields) + center = _resolve_center(peaks, q_fields, center) + q = _q_coordinates(peaks, q_fields, center) + w = peaks.select_fields(intensity_field).numpy()[:, 0].astype(np.float64).clip(min=0) + n1, n2 = np.meshgrid( + np.arange(n1_range[0], n1_range[1] + 1), + np.arange(n2_range[0], n2_range[1] + 1), + indexing="ij", + ) + n = np.stack([n1.ravel(), n2.ravel()], axis=1) + n = n[np.any(n != 0, axis=1)].astype(np.float64) + + g = np.stack([np.asarray(g1, dtype=np.float64), np.asarray(g2, dtype=np.float64)]) + order = np.argsort(q[:, 0]) + q_sorted, w_sorted = q[order], w[order] + for _ in range(num_iterations): + targets = n @ g + means = np.full_like(targets, np.nan) + weights = np.zeros(targets.shape[0]) + for k, t in enumerate(targets): + # peaks sorted along the first coordinate, so each search is a slice + i0, i1 = np.searchsorted(q_sorted[:, 0], [t[0] - radius, t[0] + radius]) + qs, ws = q_sorted[i0:i1], w_sorted[i0:i1] + m = ((qs - t[None, :]) ** 2).sum(axis=1) <= radius**2 + if ws[m].sum() > 0: + weights[k] = ws[m].sum() + means[k] = (qs[m] * ws[m, None]).sum(axis=0) / weights[k] + ok = weights > 0 + if ok.sum() < 2: + raise ValueError("Fewer than two lattice points have peaks within radius.") + sw = np.sqrt(weights[ok])[:, None] + g, *_ = np.linalg.lstsq(n[ok] * sw, means[ok] * sw, rcond=None) + return g[0], g[1] + + +def lattice_distance( + positions, + g1, + g2, + center=(0.0, 0.0), +) -> np.ndarray: + """Distance from each position to the nearest point of a 2D lattice. + + Parameters + ---------- + positions : array-like of float + (N, 2) diffraction positions, for example from cluster_centers. + g1, g2 : array-like of float + (2,) lattice vectors. + center : array-like of float, default=(0, 0) + (2,) lattice origin. + + Returns + ------- + np.ndarray + (N,) distances, in the units of the positions. + """ + basis = np.stack([np.asarray(g1, dtype=np.float64), np.asarray(g2, dtype=np.float64)]) + q = np.atleast_2d(np.asarray(positions, dtype=np.float64)) - np.asarray(center)[None, :] + frac = q @ np.linalg.inv(basis) + return np.linalg.norm((frac - np.round(frac)) @ basis, axis=1) + + +def aperture_array_subtract( + positions, + positions_remove, + tol: float = 1.0, +) -> np.ndarray: + """Remove the apertures that coincide with a second set. + + Subtracting the fundamental lattice from a finer lattice leaves only the + superlattice positions, for example. + + Parameters + ---------- + positions : array-like of float + (N, 2) aperture positions. + positions_remove : array-like of float + (M, 2) aperture positions to remove from `positions`. + tol : float, default=1.0 + Apertures within this distance of any position in + `positions_remove` are removed. + + Returns + ------- + np.ndarray + (N', 2) remaining aperture positions. + """ + positions = np.atleast_2d(np.asarray(positions, dtype=np.float64)) + positions_remove = np.atleast_2d(np.asarray(positions_remove, dtype=np.float64)) + if positions_remove.size == 0: + return positions + d2 = ((positions[:, None, :] - positions_remove[None, :, :]) ** 2).sum(axis=-1) + return positions[d2.min(axis=1) > tol**2] + + +def aperture_mask( + peaks, + positions, + radius: float = 1.0, + q_fields=None, +) -> np.ndarray: + """Peaks that fall inside any of a set of virtual apertures. + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + positions : array-like of float + (M, 2) aperture positions, in the coordinates of `q_fields`. + radius : float, default=1.0 + Aperture radius. + q_fields : (str, str), optional + Diffraction coordinate fields. Defaults to ("qx", "qy") or + ("q_row", "q_col"), whichever are present. + + Returns + ------- + np.ndarray + (N,) bool mask aligned with the flattened rows of `peaks`. A peak + inside two overlapping apertures is selected once. + """ + q = _q_coordinates(peaks, q_fields) + positions = np.atleast_2d(np.asarray(positions, dtype=np.float64)) + mask = np.zeros(q.shape[0], dtype=bool) + r2 = radius**2 + for p in positions: + mask |= ((q - p[None, :]) ** 2).sum(axis=1) <= r2 + return mask + + +def aperture_ddf_image( + peaks, + positions, + radius: float = 1.0, + q_fields=None, + intensity_field: str = "intensity", +) -> np.ndarray: + """Digital dark field image through a set of virtual apertures. + + Equivalent to ddf_image(peaks, aperture_mask(peaks, positions, radius)). + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + positions : array-like of float + (M, 2) aperture positions, from aperture_array for example. + radius : float, default=1.0 + Aperture radius, in the units of the peak coordinates. + q_fields : (str, str), optional + Diffraction coordinate fields. + intensity_field : str, default="intensity" + Field summed at each probe position. + + Returns + ------- + np.ndarray + (scan_row, scan_col) image. + """ + mask = aperture_mask(peaks, positions, radius=radius, q_fields=q_fields) + return ddf_image(peaks, mask, intensity_field=intensity_field) + + +def plot_apertures( + positions, + image=None, + radius: float | None = None, + positions_removed=None, + color="tab:green", + color_removed="tab:red", + marker_size: float = 60.0, + figax=None, + **show_kwargs, +): + """Aperture positions drawn over a diffraction image. + + Parameters + ---------- + positions : array-like of float + (N, 2) aperture positions in (row, col) pixels of `image`. + image : array-like, optional + Background image, such as the mean pattern or the Bragg vector map. + radius : float, optional + Draw each aperture as a circle of this radius in pixels. None draws + markers of size `marker_size`. + positions_removed : array-like of float, optional + (M, 2) positions drawn in `color_removed`, for example the apertures + removed by aperture_array_subtract. + color, color_removed : matplotlib color + Colors of the kept and removed apertures. + marker_size : float, default=60 + Marker size in points squared, used when radius is None. + figax : (Figure, Axes), optional + Axes to draw into. + **show_kwargs : + Passed to quantem.core.visualization.show_2d, for example + norm={"power": 0.5}. + + Returns + ------- + fig, ax + """ + from quantem.core.visualization import show_2d + + if image is not None: + show_kwargs.setdefault("axsize", (6, 6)) + fig, ax = show_2d(np.asarray(image), figax=figax, **show_kwargs) + elif figax is not None: + fig, ax = figax + else: + fig, ax = plt.subplots(figsize=(6, 6)) + ax.set_aspect("equal") + ax.invert_yaxis() + + def draw(p, c): + p = np.atleast_2d(np.asarray(p, dtype=np.float64)) + if p.size == 0: + return + if radius is None: + ax.scatter(p[:, 1], p[:, 0], s=marker_size, color=c, alpha=0.5, lw=0) + else: + ax.add_collection( + EllipseCollection( + widths=2.0 * radius, + heights=2.0 * radius, + angles=0, + units="xy", + facecolors=c, + alpha=0.4, + offsets=p[:, ::-1], + offset_transform=ax.transData, + ) + ) + + if positions_removed is not None: + draw(positions_removed, color_removed) + draw(positions, color) + return fig, ax + + +# --------------------------------------------------------------------------- # +# Polar selection +# --------------------------------------------------------------------------- # + + +def _polar_coordinates(peaks, q_fields=None, center=None): + """(qr, qphi) of every flattened row, with qphi = atan2(-q0, q1) in degrees.""" + q_fields = _resolve_q_fields(peaks, q_fields) + center = _resolve_center(peaks, q_fields, center) + q = _q_coordinates(peaks, q_fields, center) + qr = np.hypot(q[:, 0], q[:, 1]) + qphi = np.degrees(np.arctan2(-q[:, 0], q[:, 1])) + return qr, qphi + + +def add_polar_fields( + peaks, + q_fields=None, + center=None, + names: tuple[str, str] = ("qr", "qphi"), +): + """Copy of the peaks with polar coordinate fields added. + + The radius qr has the units of the diffraction coordinates. The angle + qphi = atan2(-q0, q1) is in degrees, measured anticlockwise from the +col + (right) direction as the pattern is displayed with rows increasing + downward, over the range (-180, 180]. This is py4DSTEM's DDF convention + (see the module docstring); it differs from the azimuth used in + quantem.diffraction.calibration. + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + q_fields : (str, str), optional + Diffraction coordinate fields. + center : array-like of float, optional + (2,) origin of the polar coordinates, normally the direct beam. None + uses (0, 0) for calibrated peaks and the stored "origin_ref" for peaks + in detector pixels. + names : (str, str), default=("qr", "qphi") + Names of the new fields. + + Returns + ------- + Vector + """ + q_fields = _resolve_q_fields(peaks, q_fields) + qr, qphi = _polar_coordinates(peaks, q_fields, center) + q_unit = peaks.units[peaks.fields.index(q_fields[0])] + out = peaks.copy() + out.add_fields(list(names), values=np.stack([qr, qphi], axis=1), units=[q_unit, "deg"]) + return out + + +def polar_mask( + peaks, + q_radius: float, + tol: float = 1.0, + phi_range: tuple[float, float] | None = None, + q_fields=None, + center=None, +) -> np.ndarray: + """Peaks inside a ring, optionally restricted to a range of angles. + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + q_radius : float + Ring radius, in the units of the diffraction coordinates. + tol : float, default=1.0 + Half width of the ring: peaks with |qr - q_radius| <= tol are kept. + phi_range : (float, float), optional + Angular range (phi_0, phi_1) in degrees, using the qphi convention of + add_polar_fields. Peaks with phi_0 <= qphi < phi_1 are kept. When + phi_0 > phi_1 the range wraps through 180 degrees. + q_fields : (str, str), optional + Diffraction coordinate fields. + center : array-like of float, optional + (2,) origin of the polar coordinates. None uses (0, 0) for calibrated + peaks and the stored "origin_ref" for peaks in detector pixels. + + Returns + ------- + np.ndarray + (N,) bool mask aligned with the flattened rows of `peaks`. + """ + qr, qphi = _polar_coordinates(peaks, q_fields, center) + mask = np.abs(qr - q_radius) <= tol + if phi_range is not None: + phi_0, phi_1 = phi_range + if phi_0 <= phi_1: + mask &= (qphi >= phi_0) & (qphi < phi_1) + else: + mask &= (qphi >= phi_0) | (qphi < phi_1) + return mask + + +def radial_ddf_image( + peaks, + q_radius: float, + tol: float = 1.0, + phi_range: tuple[float, float] | None = None, + q_fields=None, + center=None, + intensity_field: str = "intensity", +) -> np.ndarray: + """Digital dark field image from a ring of diffraction space. + + Equivalent to ddf_image(peaks, polar_mask(peaks, q_radius, tol, + phi_range, q_fields, center)). + + Parameters + ---------- + peaks : Vector + Peaks with one cell per probe position. + q_radius : float + Ring radius, in the units of the diffraction coordinates. + tol : float, default=1.0 + Half width of the ring, in the same units. + phi_range : (float, float), optional + Angular range (phi_0, phi_1) in degrees, as in polar_mask. + q_fields : (str, str), optional + Diffraction coordinate fields. + center : array-like of float, optional + (2,) origin of the polar coordinates. None uses (0, 0) for calibrated + peaks and the stored "origin_ref" for peaks in detector pixels. + intensity_field : str, default="intensity" + Field summed at each probe position. + + Returns + ------- + np.ndarray + (scan_row, scan_col) image. + """ + mask = polar_mask( + peaks, q_radius, tol=tol, phi_range=phi_range, q_fields=q_fields, center=center + ) + return ddf_image(peaks, mask, intensity_field=intensity_field) + + +# --------------------------------------------------------------------------- # +# Clustering +# --------------------------------------------------------------------------- # + + +def cluster_coms( + labeled, + label_field: str = "cluster", + intensity_field: str = "intensity", + weighted: bool = True, +): + """Real-space center of mass of every cluster. + + Parameters + ---------- + labeled : Vector + Vector carrying a cluster label field (from cluster_vector). + label_field : str, default="cluster" + Field holding the cluster labels; negative labels are ignored. + intensity_field : str, default="intensity" + Field used as the weight of each peak when `weighted` is True. + weighted : bool, default=True + Weight the center of mass by peak intensity. + + Returns + ------- + coms : np.ndarray + (K, 2) centers of mass in scan (row, col) pixels, ordered by cluster + id. K is the largest label plus one; (0, 2) when no peak is labeled. + sizes : np.ndarray + (K,) number of peaks per cluster. + """ + fields = labeled.fields + flat = labeled.numpy().astype(np.float64) + labels = flat[:, fields.index(label_field)].astype(int) + w = flat[:, fields.index(intensity_field)].clip(min=0) if weighted else None + rc = _scan_cells(labeled).astype(float) + + n = max(int(labels.max()) + 1, 0) if labels.size else 0 + coms = np.zeros((n, 2)) + sizes = np.zeros(n, dtype=int) + for k in range(n): + m = labels == k + sizes[k] = int(m.sum()) + if sizes[k] == 0: + coms[k] = np.nan + continue + wk = w[m] if w is not None else np.ones(sizes[k]) + wk = wk / max(wk.sum(), 1e-12) + coms[k] = (rc[m] * wk[:, None]).sum(axis=0) + return coms, sizes + + +def cluster_centers( + labeled, + q_fields=None, + label_field: str = "cluster", + intensity_field: str = "intensity", +) -> np.ndarray: + """Intensity-weighted diffraction position of every cluster. + + Parameters + ---------- + labeled : Vector + Vector carrying a cluster label field (from cluster_vector). + q_fields : (str, str), optional + Diffraction coordinate fields. + label_field : str, default="cluster" + Field holding the cluster labels. + intensity_field : str, default="intensity" + Field used as the weight of each peak. + + Returns + ------- + np.ndarray + (K, 2) mean diffraction positions, ordered by cluster id; (0, 2) when + no peak is labeled. + """ + q = _q_coordinates(labeled, q_fields) + labels = labeled.select_fields(label_field).numpy()[:, 0].astype(int) + w = labeled.select_fields(intensity_field).numpy()[:, 0].astype(np.float64).clip(min=0) + m = labels >= 0 + n = int(labels[m].max()) + 1 if m.any() else 0 + if n == 0: + return np.zeros((0, 2)) + wsum = np.maximum(np.bincount(labels[m], weights=w[m], minlength=n), 1e-12) + return np.stack( + [np.bincount(labels[m], weights=w[m] * q[m, k], minlength=n) / wsum for k in range(2)], + axis=1, + ) + + +def ddf_images( + labeled, + cluster_ids, + label_field: str = "cluster", + intensity_field: str = "intensity", +) -> np.ndarray: + """Digital dark field images: per-cluster summed intensity per position. + + Parameters + ---------- + labeled : Vector + Peaks carrying a cluster label field (from cluster_vector). + cluster_ids : int or array-like of int + Cluster labels to image, one image each, in this order. + label_field : str, default="cluster" + Field holding the cluster labels. + intensity_field : str, default="intensity" + Field summed at each probe position. Negative values are clipped to 0. + + Returns + ------- + np.ndarray + (len(cluster_ids), scan_row, scan_col) images. + """ + fields = labeled.fields + flat = labeled.numpy().astype(np.float64) + labels = flat[:, fields.index(label_field)].astype(int) + inten = flat[:, fields.index(intensity_field)].clip(min=0) + rc = _scan_cells(labeled) + R, C = labeled.shape[:2] + + cluster_ids = np.atleast_1d(cluster_ids) + out = np.zeros((len(cluster_ids), R, C)) + for i, k in enumerate(cluster_ids): + m = labels == k + np.add.at(out[i], (rc[m, 0], rc[m, 1]), inten[m]) + return out + + +def assign_grain_labels( + labeled, + grain_labels, + label_field: str = "cluster", + grain_field: str = "grain_label", +): + """Copy of the L1-labeled peaks with the L2 grain of every peak added. + + Parameters + ---------- + labeled : Vector + Peaks carrying L1 cluster labels (from cluster_vector). + grain_labels : array-like of int + (K,) L2 grain label of each L1 cluster, for example + dbscan(cluster_coms(labeled)[0], ...). + label_field : str, default="cluster" + Field holding the L1 cluster labels. + grain_field : str, default="grain_label" + Name of the new field. + + Returns + ------- + Vector + Copy of `labeled` with `grain_field` added. Peaks outside every L1 + cluster get -2, and peaks whose L1 cluster joined no grain get -1. + """ + l1 = labeled.select_fields(label_field).numpy()[:, 0].astype(int) + grain_labels = np.asarray(grain_labels, dtype=int) + grains = np.where(l1 >= 0, grain_labels[l1.clip(min=0)], -2) + + out = labeled.copy() + out.add_fields(grain_field, values=grains[:, None], units="index") + return out + + +def group_ddf_images( + images: np.ndarray, + min_correlation: float = 0.7, + min_samples: int = 2, + device: str = "cpu", +) -> np.ndarray: + """Group DDF images that show the same region of the sample. + + The spots of one grain or lath share the same dark field image, so we + cluster the images by their cosine similarity. Each image is normalized + to unit length, and DBSCAN runs with eps = sqrt(2 (1 - min_correlation)), + so that neighbors have a cosine similarity of at least min_correlation. + Unlike clustering the centers of mass (cluster_coms), this separates + grains that extend across the whole field of view, such as a matrix + phase. + + Parameters + ---------- + images : np.ndarray + (K, R, C) DDF images, for example ddf_images(labeled, range(K)). + min_correlation : float, default=0.7 + Cosine similarity between neighboring images, from 0 to 1. + min_samples : int, default=2 + DBSCAN min_samples, counting the image itself. + device : str, default="cpu" + Torch device for the distance computations. + + Returns + ------- + np.ndarray + (K,) group label of each image, -1 for images that joined no group. + Groups are numbered largest first. + """ + K = images.shape[0] + v = images.reshape(K, -1).astype(np.float64) + v = v / np.maximum(np.linalg.norm(v, axis=1, keepdims=True), 1e-12) + eps = float(np.sqrt(2.0 * (1.0 - min_correlation))) + return dbscan(v, eps=eps, min_samples=min_samples, device=device) + + +# --------------------------------------------------------------------------- # +# Display +# --------------------------------------------------------------------------- # + + +def composite_ddf( + images: np.ndarray, + colors=None, + gamma: float = 0.33, + normalize: str = "each", +) -> np.ndarray: + """Blend a stack of DDF images into one RGB composite. + + Parameters + ---------- + images : np.ndarray + (K, R, C) cluster images. + colors : array-like | None + (K, 3) RGB color per image; defaults to evenly spaced hues. + gamma : float, default=0.33 + Power scaling applied to each normalized image before coloring. + normalize : {"each", "global"} + Normalize each image to its own maximum, or all to the stack max. + + Returns + ------- + np.ndarray + (R, C, 3) RGB image in [0, 1]. + """ + K = images.shape[0] + if colors is None: + hues = np.linspace(0, 1, K, endpoint=False) + colors = hsv_to_rgb(np.stack([hues, np.ones(K), np.ones(K)], axis=1)) + colors = np.asarray(colors, dtype=float) + + if normalize == "global": + norm = np.full(K, max(float(images.max()), 1e-12)) + else: + norm = np.maximum(images.reshape(K, -1).max(axis=1), 1e-12) + scaled = (images / norm[:, None, None]) ** gamma + rgb = np.einsum("krc,kj->rcj", scaled, colors) + return np.clip(rgb, 0, 1) + + +def color_wheel(n: int = 256, saturation: float = 1.0) -> np.ndarray: + """Hue wheel for labeling composite images. + + Parameters + ---------- + n : int, default=256 + Size of the output image in pixels. + saturation : float, default=1.0 + Saturation of the hues, from 0 (gray) to 1 (full color). + + Returns + ------- + np.ndarray + (n, n, 4) RGBA image in [0, 1], transparent outside the wheel. + """ + y, x = np.mgrid[-1 : 1 : n * 1j, -1 : 1 : n * 1j] + r = np.hypot(x, y) + hue = (np.arctan2(y, x) / (2 * np.pi)) % 1.0 + hsv = np.stack([hue, np.full_like(hue, saturation), np.clip(r, 0, 1)], axis=-1) + rgba = np.concatenate([hsv_to_rgb(hsv), (r <= 1.0)[..., None].astype(float)], axis=-1) + return rgba + + +def plot_cluster_scatter( + labeled, + q_fields=None, + label_field: str = "cluster", + specific_cluster: int | Sequence[int] | None = None, + max_clusters: int | None = None, + show_unclustered: bool = True, + point_size: float = 2.0, + alpha: float = 0.2, + center=None, + figax=None, +): + """All peaks in diffraction space, colored by cluster. + + Parameters + ---------- + labeled : Vector + Peaks carrying a label field. + q_fields : (str, str), optional + Diffraction coordinate fields, drawn as (vertical, horizontal). + label_field : str, default="cluster" + Field holding the labels, for example "grain_label" from + assign_grain_labels. + specific_cluster : int or sequence of int, optional + Draw only these clusters, in one color, without the unclustered peaks. + max_clusters : int, optional + Draw only the first max_clusters clusters (the largest, since dbscan + sorts clusters by size). + show_unclustered : bool, default=True + Draw the peaks with a negative label in gray. + point_size : float, default=2.0 + Marker size. + alpha : float, default=0.2 + Marker opacity. Values well below 1 show the dense regions of large + datasets. + center : array-like of float, optional + (2,) diffraction origin at the middle of the plot. None uses (0, 0) + for calibrated peaks and the "origin_ref" stored by + BraggVectors.correct_peak_origins for peaks in detector pixels (or the + middle of the peak positions when there is none). + figax : (Figure, Axes), optional + Axes to draw into. + + Returns + ------- + fig, ax + """ + q_fields = _resolve_q_fields(labeled, q_fields) + q = labeled.select_fields(*q_fields).numpy().astype(np.float64) + try: + center = _resolve_center(labeled, q_fields, center) + except ValueError: + # pixel peaks without a stored origin: frame the peaks themselves + center = 0.5 * (q.min(axis=0) + q.max(axis=0)) if q.size else np.zeros(2) + labels = labeled.select_fields(label_field).numpy()[:, 0].astype(int) + q0, q1 = q[:, 0], q[:, 1] + + if figax is None: + fig, ax = plt.subplots(figsize=(6.5, 6.5)) + else: + fig, ax = figax + q_max = 1.05 * float(np.abs(q - center[None, :]).max()) if q.size else 1.0 + ax.set_xlim(center[1] - q_max, center[1] + q_max) + ax.set_ylim(center[0] + q_max, center[0] - q_max) + ax.set_aspect("equal") + ax.set_xlabel(q_fields[1]) + ax.set_ylabel(q_fields[0]) + + if specific_cluster is not None: + m = np.isin(labels, np.atleast_1d(specific_cluster)) + ax.scatter(q1[m], q0[m], s=point_size, color="C0", lw=0, alpha=alpha) + return fig, ax + + if show_unclustered: + m = labels < 0 + ax.scatter(q1[m], q0[m], s=point_size, color="0.85", lw=0) + n = max(int(labels.max()) + 1, 0) if labels.size else 0 + n_show = n if max_clusters is None else min(n, max_clusters) + cmap = plt.get_cmap("hsv") + rng = np.random.default_rng(0) + hues = rng.permutation(np.linspace(0, 1, n_show, endpoint=False)) + for k in range(n_show): + m = labels == k + ax.scatter(q1[m], q0[m], s=point_size, color=cmap(hues[k]), lw=0, alpha=alpha) + return fig, ax diff --git a/src/quantem/diffraction/disk_detection.py b/src/quantem/diffraction/disk_detection.py new file mode 100644 index 000000000..14bcf3599 --- /dev/null +++ b/src/quantem/diffraction/disk_detection.py @@ -0,0 +1,1278 @@ +"""Template-matching Bragg disk detection for 4D-STEM, in torch. + +Each diffraction pattern is cross-correlated with a probe template in Fourier +space (optionally as a hybrid or phase correlation, with a Fourier high-pass +background removal and low-pass smoothing), local maxima of the correlation map +are kept as candidate disks, and each one is refined to subpixel precision by a +parabolic fit and then by DFT upsampling of the Fourier product. The approach +follows the multicorr / ``find_Bragg_disks`` routines of py4DSTEM, which use the +single-step DFT upsampling of Guizar-Sicairos, Thurman and Fienup, "Efficient +subpixel image registration algorithms", Optics Letters 33, 156 (2008). +Functions here are used by :class:`~quantem.diffraction.bragg_vectors.BraggVectors`. +Peak coordinates are ``(row, col)`` in detector pixels. +""" + +from __future__ import annotations + +import numpy as np +import torch + +SUBPIXEL_MODES = ("none", "parabolic", "upsample") + + +def make_template( + probe: torch.Tensor, + center: tuple[float, float] | None = None, + subtract_mean: bool = False, +) -> torch.Tensor: + """Build a cross-correlation template from a (vacuum) probe image. + + The probe is normalized to unit sum and rolled so its center sits at the array + origin ``[0, 0]`` (FFT corner), so correlation peaks land at absolute disk + positions. + + Parameters + ---------- + probe : torch.Tensor + ``(H, W)`` probe / vacuum disk image. + center : tuple of float, optional + ``(row, col)`` probe center rolled to the origin; defaults to the geometric + center ``(H // 2, W // 2)``. + subtract_mean : bool, default=False + If ``True``, make the template zero-sum — a band-pass kernel that suppresses + the uniform background in the correlation. + + Returns + ------- + torch.Tensor + ``(H, W)`` template, corner-centered (and zero-sum when ``subtract_mean``). + """ + probe = torch.as_tensor(probe) + total = probe.sum() + if total != 0: + probe = probe / total + + H, W = probe.shape + if center is None: + cr, cc = H // 2, W // 2 + else: + cr, cc = int(round(float(center[0]))), int(round(float(center[1]))) + + template = torch.roll(probe, shifts=(-cr, -cc), dims=(0, 1)) + if subtract_mean: + template = template - template.mean() + return template + + +def synthetic_probe( + shape: tuple[int, int], + radius: float, + edge: float = 1.0, + center: tuple[float, float] | None = None, +) -> torch.Tensor: + """Soft-edged disk for a synthetic correlation template. + + Returns ``0.5 - 0.5*tanh((r - radius)/edge)``: a disk of the given ``radius`` + (pixels) with a ``tanh`` falloff over ``edge`` pixels. + + Parameters + ---------- + shape : tuple of int + ``(H, W)`` output shape in pixels. + radius : float + Disk radius in pixels. + edge : float, default=1.0 + Width in pixels of the ``tanh`` edge falloff. + center : tuple of float, optional + ``(row, col)`` disk center; defaults to the geometric center + ``((H - 1) / 2, (W - 1) / 2)``. + + Returns + ------- + torch.Tensor + ``(H, W)`` soft-edged disk image. + """ + H, W = int(shape[0]), int(shape[1]) + if center is None: + cr, cc = (H - 1) / 2.0, (W - 1) / 2.0 + else: + cr, cc = float(center[0]), float(center[1]) + rows = torch.arange(H, dtype=torch.float).view(H, 1) + cols = torch.arange(W, dtype=torch.float).view(1, W) + rr = torch.sqrt((rows - cr) ** 2 + (cols - cc) ** 2) + edge = max(float(edge), 1e-6) + return 0.5 - 0.5 * torch.tanh((rr - float(radius)) / edge) + + +def _central_blob(image: torch.Tensor, threshold: float) -> tuple[np.ndarray | None, np.ndarray]: + """Mask of the connected bright region containing the brightest pixel. + + Thresholds ``image`` at ``threshold * max`` and keeps only the connected + component holding the brightest pixel — normally the central (unscattered) + disk of a mean diffraction pattern or vacuum probe — so other diffracted disks + are excluded. + + Parameters + ---------- + image : torch.Tensor + ``(H, W)`` image, e.g. a mean diffraction pattern or vacuum probe. + threshold : float + Fraction of the peak intensity (after min-subtraction) used to threshold + the image before connected-component labeling. + + Returns + ------- + blob_mask : np.ndarray or None + Boolean ``(H, W)`` mask of the central component, or ``None`` for an empty + / flat image. + img : np.ndarray + The min-subtracted image. + """ + from scipy import ndimage + + img = np.asarray(torch.as_tensor(image, dtype=torch.float).detach().cpu()) + img = img - img.min() + peak = float(img.max()) + if peak <= 0: + return None, img + labels, n = ndimage.label(img >= threshold * peak) + if n == 0: + return None, img + peak_label = int(labels[np.unravel_index(int(np.argmax(img)), img.shape)]) + return labels == peak_label, img + + +def estimate_central_beam( + image: torch.Tensor, + threshold: float = 0.5, + plot_result: bool = False, + **kwargs, +) -> tuple[tuple[float, float], float]: + """Center ``(row, col)`` and radius (pixels) of the central (direct) beam. + + Locates the connected bright region containing the brightest pixel (see + :func:`_central_blob`) — the unscattered / direct beam of a mean diffraction + pattern — and returns its intensity-weighted center together with an + area-equivalent radius (``A = pi r^2``). Other diffracted disks are excluded, so + the estimate holds whether the pattern shows one disk or many. Falls back to the + geometric center and unit radius for an empty / flat image. + + Parameters + ---------- + image : torch.Tensor or Dataset2d + ``(H, W)`` image — a raw array / tensor or a :class:`Dataset2d` (e.g. + ``dataset.dp_mean``). + threshold : float, default=0.5 + Fraction of the peak intensity used to threshold the image when isolating + the central beam (see :func:`_central_blob`). + plot_result : bool, default=False + If ``True``, show the image in greyscale with the fitted beam drawn as a red + circle. + **kwargs + Extra keyword arguments (e.g. ``norm``, ``cbar``, ``scalebar``) forwarded to + :func:`~quantem.core.visualization.show_2d` when ``plot_result=True``. + + Returns + ------- + center : tuple of float + ``(row, col)`` intensity-weighted center of the central beam. + radius : float + Area-equivalent radius in pixels (``A = pi r^2``). + """ + arr = image.array if hasattr(image, "array") else image + blob, img = _central_blob(arr, threshold) + if blob is None: + center = (img.shape[0] / 2.0, img.shape[1] / 2.0) + radius = 1.0 + else: + radius = float(np.sqrt(max(float(blob.sum()), 1.0) / np.pi)) + w = img * blob + total = float(w.sum()) + if total <= 0: + rr, cc = np.nonzero(blob) + center = (float(rr.mean()), float(cc.mean())) + else: + rows = np.arange(img.shape[0])[:, None] + cols = np.arange(img.shape[1])[None, :] + center = (float((w * rows).sum() / total), float((w * cols).sum() / total)) + + if plot_result: + from matplotlib.patches import Circle + + from quantem.core.visualization import show_2d + + show_kwargs = {"cmap": "gray", "title": "central beam", **kwargs} + _fig, ax = show_2d(arr, **show_kwargs) + ax.add_patch( + Circle((center[1], center[0]), radius, fill=False, edgecolor="red", linewidth=1.5) + ) + + return center, radius + + +def probe_centroid(probe: torch.Tensor) -> tuple[float, float]: + """Intensity-weighted ``(row, col)`` centroid of a probe image. + + Parameters + ---------- + probe : torch.Tensor + ``(H, W)`` probe image. Negative values are clamped to zero before + weighting. + + Returns + ------- + tuple of float + ``(row, col)`` intensity-weighted centroid; the geometric center for a + non-positive image. + """ + p = torch.clamp(torch.as_tensor(probe, dtype=torch.float), min=0.0) + total = p.sum() + if total <= 0: + return (p.shape[0] / 2.0, p.shape[1] / 2.0) + rows = torch.arange(p.shape[0], dtype=torch.float, device=p.device).view(-1, 1) + cols = torch.arange(p.shape[1], dtype=torch.float, device=p.device).view(1, -1) + return (float((p * rows).sum() / total), float((p * cols).sum() / total)) + + +def template_fourier(template: torch.Tensor) -> torch.Tensor: + """Pre-compute the conjugate FT of a template for repeated correlation. + + Parameters + ---------- + template : torch.Tensor + ``(H, W)`` corner-centered correlation template. + + Returns + ------- + torch.Tensor + ``(H, W)`` complex ``conj(fft2(template))``, ready to multiply against + ``fft2(dp)``. + """ + return torch.conj(torch.fft.fft2(template)) + + +def _background_highpass( + shape: tuple[int, int], + background_sigma: float, + device, + rfft: bool = False, +) -> torch.Tensor: + """Fourier-domain high-pass ``1 - G`` removing correlation background. + + ``G`` is the transform of a real-space Gaussian of standard deviation + ``background_sigma`` pixels, so multiplying the Fourier product by ``1 - G`` + subtracts a Gaussian-smoothed copy of the correlation map -- the slowly + varying background (the zero-sum template's negative moat around the bright + central beam) that otherwise pushes weak disk peaks below zero, where the + ``relu`` clamp erases them before peak finding. + + Parameters + ---------- + shape : tuple of int + ``(H, W)`` detector shape. + background_sigma : float + Real-space standard deviation of the subtracted background, in pixels. + device + Torch device for the filter tensor. + rfft : bool, default=False + If ``True``, return the ``(H, W // 2 + 1)`` half-plane filter for + ``rfft2`` products instead of the full ``(H, W)`` filter. + + Returns + ------- + torch.Tensor + The ``1 - G`` filter in the requested Fourier layout. + """ + H, W = int(shape[0]), int(shape[1]) + qr = torch.fft.fftfreq(H, device=device, dtype=torch.float)[:, None] + if rfft: + qc = torch.fft.rfftfreq(W, device=device, dtype=torch.float)[None, :] + else: + qc = torch.fft.fftfreq(W, device=device, dtype=torch.float)[None, :] + g = torch.exp(-2.0 * (torch.pi**2) * (float(background_sigma) ** 2) * (qr**2 + qc**2)) + return 1.0 - g + + +def _smoothing_lowpass( + shape: tuple[int, int], + sigma: float, + device, + rfft: bool = False, +) -> torch.Tensor: + """Fourier-domain Gaussian low-pass of width ``sigma`` pixels. + + Smoothing the correlation map before peak finding merges the speckle of a + noisy disk into one maximum, which is what stops a weak disk from being + split into several sub-threshold peaks. + """ + H, W = int(shape[0]), int(shape[1]) + qr = torch.fft.fftfreq(H, device=device, dtype=torch.float)[:, None] + qc = ( + torch.fft.rfftfreq(W, device=device, dtype=torch.float)[None, :] + if rfft + else torch.fft.fftfreq(W, device=device, dtype=torch.float)[None, :] + ) + return torch.exp(-2.0 * (torch.pi**2) * (float(sigma) ** 2) * (qr**2 + qc**2)) + + +def _apply_corr_power(m: torch.Tensor, corr_power: float) -> torch.Tensor: + """Hybrid correlation: keep the phase, raise the magnitude to ``corr_power``. + + ``corr_power=1`` is the plain cross-correlation, ``0`` the phase + correlation, and values in between the hybrid correlation. Dividing out + part of the magnitude equalizes the weak and strong reflections, so a + faint disk on a bright background produces a peak of the same height as a + strong one; this is what makes the phase and hybrid correlations find far + more weak disks than the plain product. The cost is that the peak height + is no longer proportional to the disk intensity -- with ``corr_power`` + below 1 the reported intensities are compressed, which matters when they + are used downstream as weights (orientation matching, strain). + """ + if corr_power == 1.0: + return m + mag = torch.abs(m) + return m * mag.clamp_min(1e-12) ** (float(corr_power) - 1.0) + + +def _fourier_filter( + m: torch.Tensor, + shape: tuple[int, int], + background_sigma: float | None, + sigma_cc: float | None, + rfft: bool = False, +) -> torch.Tensor: + """Apply the background high-pass and the smoothing low-pass to a product.""" + if background_sigma is not None and background_sigma > 0: + m = m * _background_highpass(shape, background_sigma, m.device, rfft=rfft) + if sigma_cc is not None and sigma_cc > 0: + m = m * _smoothing_lowpass(shape, sigma_cc, m.device, rfft=rfft) + return m + + +def cross_correlation( + dp: torch.Tensor, + template_ft: torch.Tensor, + background_sigma: float | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Cross-correlate a diffraction pattern with a template. + + Parameters + ---------- + dp : torch.Tensor + ``(H, W)`` diffraction pattern. + template_ft : torch.Tensor + ``(H, W)`` pre-computed template FT from :func:`template_fourier`. + background_sigma : float, optional + Width in pixels of the Gaussian whose smoothed copy of the correlation is + subtracted (a Fourier high-pass). ``None`` or ``0`` disables it. + corr_power : float, default=1.0 + Exponent applied to the magnitude of the Fourier product: 1 is the plain + cross-correlation, 0 the phase correlation, and values in between the + hybrid correlation. + sigma_cc : float, optional + Width in pixels of a Gaussian smoothing of the correlation map (a Fourier + low-pass). ``None`` or ``0`` disables it. + + Returns + ------- + corr_map : torch.Tensor + ``(H, W)`` real-space correlation map ``relu(real(ifft2(m)))`` (used for + peak finding). + m : torch.Tensor + ``(H, W)`` Fourier-domain product ``fft2(dp) * template_ft`` after the + ``corr_power`` and the filters (used for DFT subpixel refinement). + """ + dp = torch.as_tensor(dp) + m = _apply_corr_power(torch.fft.fft2(dp) * template_ft, corr_power) + m = _fourier_filter(m, m.shape[-2:], background_sigma, sigma_cc) + corr_map = torch.clamp(torch.fft.ifft2(m).real, min=0.0) + return corr_map, m + + +def cross_correlation_batch( + dps: torch.Tensor, + template_ft: torch.Tensor, + background_sigma: float | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Cross-correlate a stack of diffraction patterns with one template. + + Batched form of :func:`cross_correlation`. ``fft2`` acts on the trailing two + axes and the ``(H, W)`` ``template_ft`` broadcasts over the batch, so the result + for each pattern is bit-identical to :func:`cross_correlation`. + + Parameters + ---------- + dps : torch.Tensor + ``(B, H, W)`` stack of diffraction patterns. + template_ft : torch.Tensor + ``(H, W)`` pre-computed template FT from :func:`template_fourier`. + background_sigma : float, optional + Width in pixels of the Gaussian whose smoothed copy of the correlation is + subtracted (a Fourier high-pass). ``None`` or ``0`` disables it. + corr_power : float, default=1.0 + Exponent applied to the magnitude of the Fourier product: 1 is the plain + cross-correlation, 0 the phase correlation, and values in between the + hybrid correlation. + sigma_cc : float, optional + Width in pixels of a Gaussian smoothing of the correlation map (a Fourier + low-pass). ``None`` or ``0`` disables it. + + Returns + ------- + corr_map : torch.Tensor + ``(B, H, W)`` real-space correlation maps. + m : torch.Tensor + ``(B, H, W)`` Fourier-domain products. + """ + dps = torch.as_tensor(dps) + m = _apply_corr_power(torch.fft.fft2(dps) * template_ft, corr_power) + m = _fourier_filter(m, m.shape[-2:], background_sigma, sigma_cc) + corr_map = torch.clamp(torch.fft.ifft2(m).real, min=0.0) + return corr_map, m + + +def _corr_map_rfft( + dps: torch.Tensor, + template_ft: torch.Tensor, + background_sigma: float | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, +) -> torch.Tensor: + """Real-FFT correlation map(s), used when no Fourier product is needed downstream. + + For real ``dps`` and a real template the Fourier product is conjugate-symmetric, + so ``relu(real(ifft2(fft2(dps) * template_ft)))`` is reproduced exactly by an + ``rfft2`` / ``irfft2`` pair at roughly half the FFT cost. Only the correlation map + is returned — the complex product (needed solely for DFT upsampling) is skipped. + + Parameters + ---------- + dps : torch.Tensor + ``(H, W)`` or ``(B, H, W)`` diffraction pattern(s). + template_ft : torch.Tensor + ``(H, W)`` pre-computed template FT from :func:`template_fourier`. + background_sigma : float, optional + Width in pixels of the Gaussian whose smoothed copy of the correlation is + subtracted (a Fourier high-pass). ``None`` or ``0`` disables it. + corr_power : float, default=1.0 + Exponent applied to the magnitude of the Fourier product: 1 is the plain + cross-correlation, 0 the phase correlation, and values in between the + hybrid correlation. + sigma_cc : float, optional + Width in pixels of a Gaussian smoothing of the correlation map (a Fourier + low-pass). ``None`` or ``0`` disables it. + + Returns + ------- + torch.Tensor + ``(H, W)`` or ``(B, H, W)`` real-space correlation map(s), ``relu``-clamped. + """ + dps = torch.as_tensor(dps) + H, W = dps.shape[-2], dps.shape[-1] + prod = _apply_corr_power(torch.fft.rfft2(dps) * template_ft[..., : W // 2 + 1], corr_power) + prod = _fourier_filter(prod, (H, W), background_sigma, sigma_cc, rfft=True) + corr_map = torch.fft.irfft2(prod, s=(H, W)) + return torch.clamp(corr_map, min=0.0) + + +def detect_disks( + dp: torch.Tensor, + template_ft: torch.Tensor, + *, + min_abs_intensity: float = 0.0, + min_spacing: float = 0.0, + edge_boundary: int = 1, + subpixel: str = "upsample", + upsample_factor: int = 16, + max_num_peaks: int = 1000, + background_sigma: float | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, +) -> np.ndarray: + """Detect Bragg disks in one diffraction pattern by template matching. + + Parameters + ---------- + dp : torch.Tensor + ``(H, W)`` diffraction pattern. + template_ft : torch.Tensor + ``(H, W)`` pre-computed template FT from :func:`template_fourier`. + min_abs_intensity : float, default=0.0 + Drop correlation peaks below this absolute intensity. + min_spacing : float, default=0.0 + Minimum spacing in pixels between kept peaks; closer / dimmer peaks are + suppressed. + edge_boundary : int, default=1 + Width in pixels of the border in which peaks are ignored. + subpixel : {"none", "parabolic", "upsample"}, default="upsample" + ``"none"`` returns pixel-resolution peaks; ``"parabolic"`` adds a 3-point + quadratic refinement; ``"upsample"`` further refines each peak by + Guizar-Sicairos DFT upsampling. + upsample_factor : int, default=16 + Upsampling factor for the ``"upsample"`` subpixel refinement. + max_num_peaks : int, default=1000 + Maximum number of peaks to keep (after intensity sorting). + background_sigma : float, optional + Width in pixels of the smoothed correlation background subtracted + before peak finding. ``None`` (default) disables it. + corr_power : float, default=1.0 + Correlation type: 1 the plain cross-correlation, 0 the phase + correlation, in between the hybrid correlation. Below 1 the weak disks + are amplified relative to the strong ones, which finds many more of + them on a bright background, at the cost of compressing the reported + intensities. + sigma_cc : float | None + Width in pixels of a Gaussian smoothing of the correlation map before + peak finding; merges the speckle of a noisy disk into one maximum. + + Returns + ------- + np.ndarray + ``(M, 3)`` array of ``[q_row, q_col, intensity]`` rows, sorted by descending + intensity. + """ + if subpixel not in SUBPIXEL_MODES: + raise ValueError(f"subpixel must be in {SUBPIXEL_MODES}, got {subpixel!r}") + + corr_map, m = cross_correlation(dp, template_ft, background_sigma, corr_power, sigma_cc) + + peaks = _local_maxima(corr_map, edge_boundary) + peaks = _filter_maxima(peaks, min_abs_intensity, min_spacing, max_num_peaks) + + if peaks.shape[0] == 0 or subpixel == "none": + return _to_numpy(peaks) + + peaks = _refine_parabolic(corr_map, peaks) + + if subpixel == "parabolic": + return _to_numpy(peaks) + + peaks = _refine_dft(m, peaks, upsample_factor) + return _to_numpy(peaks) + + +def detect_disks_batch( + dps: torch.Tensor, + template_ft: torch.Tensor, + *, + min_abs_intensity: float = 0.0, + min_spacing: float = 0.0, + edge_boundary: int = 1, + subpixel: str = "upsample", + upsample_factor: int = 16, + max_num_peaks: int = 1000, + background_sigma: float | None = None, + corr_power: float = 1.0, + sigma_cc: float | None = None, +) -> list[np.ndarray]: + """Detect Bragg disks across a stack of diffraction patterns (batched). + + Batched equivalent of :func:`detect_disks`: cross-correlation, peak extraction, + ``min_spacing`` suppression, and subpixel refinement are all batched across + patterns. The local-maxima search and greedy ``min_spacing`` suppression are + vectorized over the whole stack (no per-pattern Python loop) but reproduce the + per-pattern greedy result bit-for-bit, so each output matches :func:`detect_disks` + for that pattern. When ``subpixel`` is not ``"upsample"`` the correlation maps are + formed with a real FFT (``rfft2``), skipping the complex Fourier product that only + DFT upsampling needs. + + Parameters + ---------- + dps : torch.Tensor + ``(B, H, W)`` stack of diffraction patterns. + template_ft : torch.Tensor + ``(H, W)`` pre-computed template FT from :func:`template_fourier`. + min_abs_intensity : float, default=0.0 + Drop correlation peaks below this absolute intensity. + min_spacing : float, default=0.0 + Minimum spacing in pixels between kept peaks; closer / dimmer peaks are + suppressed. + edge_boundary : int, default=1 + Width in pixels of the border in which peaks are ignored. + subpixel : {"none", "parabolic", "upsample"}, default="upsample" + ``"none"`` returns pixel-resolution peaks; ``"parabolic"`` adds a 3-point + quadratic refinement; ``"upsample"`` further refines each peak by + Guizar-Sicairos DFT upsampling. + upsample_factor : int, default=16 + Upsampling factor for the ``"upsample"`` subpixel refinement. + max_num_peaks : int, default=1000 + Maximum number of peaks to keep per pattern (after intensity sorting). + background_sigma : float, optional + Width in pixels of the smoothed correlation background subtracted + before peak finding. ``None`` (default) disables it. + corr_power : float, default=1.0 + Correlation type: 1 cross-correlation, 0 phase correlation, in between + hybrid (see :func:`detect_disks`). + sigma_cc : float | None + Width in pixels of a Gaussian smoothing of the correlation map. + + Returns + ------- + list of np.ndarray + Length-``B`` list of ``(M, 3)`` arrays of ``[q_row, q_col, intensity]`` + rows, each sorted by descending intensity. + """ + if subpixel not in SUBPIXEL_MODES: + raise ValueError(f"subpixel must be in {SUBPIXEL_MODES}, got {subpixel!r}") + + if subpixel == "upsample": + corr_map, m = cross_correlation_batch( + dps, template_ft, background_sigma, corr_power, sigma_cc + ) + else: + corr_map = _corr_map_rfft(dps, template_ft, background_sigma, corr_power, sigma_cc) + m = None + + peaks_all, bidx, counts = _detect_peaks_batched( + corr_map, edge_boundary, min_abs_intensity, min_spacing, max_num_peaks + ) + + if subpixel == "none" or peaks_all.shape[0] == 0: + return _split_by_counts(peaks_all, counts) + + peaks_all = _refine_parabolic_batched(corr_map, peaks_all, bidx) + if subpixel == "upsample": + peaks_all = _refine_dft_batched(m, peaks_all, bidx, upsample_factor) + + return _split_by_counts(peaks_all, counts) + + +# ---- helpers ---- + + +def _local_maxima(corr_map: torch.Tensor, edge_boundary: int) -> torch.Tensor: + """Find 8-neighbor local maxima, sorted by descending intensity. + + Parameters + ---------- + corr_map : torch.Tensor + ``(H, W)`` correlation map. + edge_boundary : int + Width in pixels of the border in which maxima are ignored. + + Returns + ------- + torch.Tensor + ``(K, 3)`` tensor of ``[row, col, intensity]`` maxima, sorted by descending + intensity. + """ + is_max = _local_maxima_mask(corr_map, edge_boundary) + return _extract_maxima(corr_map, is_max) + + +def _local_maxima_mask(a: torch.Tensor, edge_boundary: int) -> torch.Tensor: + """Boolean 8-neighbor local-maxima mask of ``a``. + + Works on a single ``(H, W)`` map or a batch ``(B, H, W)`` — the neighbor + comparisons and edge masking use the trailing two axes — so the same code drives + the single-pattern and batched detection paths bit-identically. + + Parameters + ---------- + a : torch.Tensor + ``(H, W)`` or ``(B, H, W)`` correlation map(s). + edge_boundary : int + Width in pixels of the border (clamped to at least 1) set to ``False``. + + Returns + ------- + torch.Tensor + Boolean mask the same shape as ``a``, ``True`` at 8-neighbor local maxima. + """ + is_max = ( + (a >= torch.roll(a, (-1, 0), dims=(-2, -1))) + & (a > torch.roll(a, (1, 0), dims=(-2, -1))) + & (a >= torch.roll(a, (0, -1), dims=(-2, -1))) + & (a > torch.roll(a, (0, 1), dims=(-2, -1))) + & (a >= torch.roll(a, (-1, -1), dims=(-2, -1))) + & (a > torch.roll(a, (-1, 1), dims=(-2, -1))) + & (a >= torch.roll(a, (1, -1), dims=(-2, -1))) + & (a > torch.roll(a, (1, 1), dims=(-2, -1))) + ) + + eb = max(1, int(edge_boundary)) + is_max[..., :eb, :] = False + is_max[..., -eb:, :] = False + is_max[..., :, :eb] = False + is_max[..., :, -eb:] = False + return is_max + + +def _extract_maxima(a: torch.Tensor, is_max: torch.Tensor) -> torch.Tensor: + """Gather masked maxima of one ``(H, W)`` map into a descending-sorted ``(K, 3)``. + + Rows are taken in row-major ``nonzero`` order, then stably sorted by descending + intensity — matching the original single-pattern behaviour. + + Parameters + ---------- + a : torch.Tensor + ``(H, W)`` correlation map. + is_max : torch.Tensor + Boolean ``(H, W)`` local-maxima mask from :func:`_local_maxima_mask`. + + Returns + ------- + torch.Tensor + ``(K, 3)`` tensor of ``[row, col, intensity]`` maxima, sorted by descending + intensity. + """ + rows, cols = torch.nonzero(is_max, as_tuple=True) + intensity = a[rows, cols] + order = torch.argsort(intensity, descending=True) + return torch.stack((rows[order].to(a.dtype), cols[order].to(a.dtype), intensity[order]), dim=1) + + +def _filter_maxima( + peaks: torch.Tensor, + min_abs_intensity: float, + min_spacing: float, + max_num_peaks: int, +) -> torch.Tensor: + """Drop dim peaks, suppress peaks closer than ``min_spacing``, cap the count. + + Parameters + ---------- + peaks : torch.Tensor + ``(K, 3)`` ``[row, col, intensity]`` peaks, sorted by descending intensity. + min_abs_intensity : float + Drop peaks below this absolute intensity (ignored when ``<= 0``). + min_spacing : float + Minimum spacing in pixels; for each kept peak, dimmer peaks within this + distance are suppressed (ignored when ``<= 0``). + max_num_peaks : int + Maximum number of peaks to keep; the brightest are retained. + + Returns + ------- + torch.Tensor + ``(M, 3)`` filtered peaks. + """ + if peaks.shape[0] == 0: + return peaks + + if min_abs_intensity > 0: + peaks = peaks[peaks[:, 2] >= min_abs_intensity] + + if min_spacing > 0 and peaks.shape[0] > 1: + keep = torch.ones(peaks.shape[0], dtype=torch.bool, device=peaks.device) + rc = peaks[:, :2] + for i in range(peaks.shape[0]): + if not keep[i]: + continue + d2 = ((rc - rc[i]) ** 2).sum(dim=1) + too_close = d2 < min_spacing**2 + too_close[: i + 1] = False + keep[too_close] = False + peaks = peaks[keep] + + if max_num_peaks is not None and peaks.shape[0] > max_num_peaks: + peaks = peaks[:max_num_peaks] + + return peaks + + +def _detect_peaks_batched( + corr: torch.Tensor, + edge_boundary: int, + min_abs_intensity: float, + min_spacing: float, + max_num_peaks: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Extract, suppress, and cap peaks across a whole ``(B, H, W)`` stack at once. + + Vectorized replacement for the per-pattern ``_extract_maxima`` + ``_filter_maxima`` + loop. Local maxima are found and intensity-thresholded for the whole stack, then + the greedy ``min_spacing`` suppression runs as a single loop over the *rank* axis + (brightest-first) that is batched over every pattern. Suppressing peak ``k``'s + fainter neighbours only when ``k`` is still kept reproduces the per-pattern greedy + result of :func:`_filter_maxima` exactly, with no Python loop over patterns. + + Parameters + ---------- + corr : torch.Tensor + ``(B, H, W)`` correlation maps. + edge_boundary : int + Width in pixels of the border in which maxima are ignored. + min_abs_intensity : float + Drop peaks below this absolute intensity (ignored when ``<= 0``). + min_spacing : float + Minimum spacing in pixels; for each kept peak, fainter peaks within this + distance are suppressed (ignored when ``<= 0``). + max_num_peaks : int + Maximum number of peaks kept per pattern (the brightest are retained). + + Returns + ------- + peaks : torch.Tensor + ``(T, 3)`` ``[row, col, intensity]`` peaks pooled over all patterns, grouped + by ascending pattern index and sorted by descending intensity within a pattern. + bidx : torch.Tensor + ``(T,)`` pattern index of each peak. + counts : torch.Tensor + ``(B,)`` number of peaks kept per pattern. + """ + B = corr.shape[0] + device = corr.device + dtype = corr.dtype + + mask = _local_maxima_mask(corr, edge_boundary) + if min_abs_intensity > 0: + mask = mask & (corr >= min_abs_intensity) + + idx = torch.nonzero(mask) # (T0, 3): [b, row, col] + counts = torch.zeros(B, dtype=torch.long, device=device) + if idx.shape[0] == 0: + return corr.new_zeros((0, 3)), torch.zeros(0, dtype=torch.long, device=device), counts + + bcand, rcand, ccand = idx[:, 0], idx[:, 1], idx[:, 2] + inten = corr[bcand, rcand, ccand] + + # Group candidates by pattern, descending intensity within each pattern. A global + # descending-intensity sort followed by a stable sort on the pattern index keeps the + # within-pattern order identical to the per-pattern argsort in _extract_maxima. + o1 = torch.argsort(inten, descending=True) + order = o1[torch.argsort(bcand[o1], stable=True)] + bcand, rcand, ccand, inten = bcand[order], rcand[order], ccand[order], inten[order] + + counts = torch.bincount(bcand, minlength=B) + kmax = int(counts.max()) + + # Pack candidates into a dense (B, kmax) grid indexed by [pattern, brightness rank]. + starts = torch.zeros(B, dtype=torch.long, device=device) + starts[1:] = torch.cumsum(counts, 0)[:-1] + rank = torch.arange(bcand.shape[0], device=device) - starts[bcand] + flat = bcand * kmax + rank + + rows = torch.zeros(B * kmax, dtype=dtype, device=device) + cols = torch.zeros(B * kmax, dtype=dtype, device=device) + vals = torch.zeros(B * kmax, dtype=dtype, device=device) + valid = torch.zeros(B * kmax, dtype=torch.bool, device=device) + rows[flat] = rcand.to(dtype) + cols[flat] = ccand.to(dtype) + vals[flat] = inten + valid[flat] = True + rows, cols, valid = rows.view(B, kmax), cols.view(B, kmax), valid.view(B, kmax) + + keep = valid.clone() + if min_spacing > 0 and kmax > 1: + s2 = float(min_spacing) ** 2 + for k in range(kmax - 1): + active = keep[:, k] & valid[:, k] # (B,) brightest unsuppressed peak at rank k + dr = rows[:, k + 1 :] - rows[:, k : k + 1] + dc = cols[:, k + 1 :] - cols[:, k : k + 1] + suppress = ((dr * dr + dc * dc) < s2) & active[:, None] & valid[:, k + 1 :] + keep[:, k + 1 :] &= ~suppress + + if max_num_peaks is not None: + kept_rank = torch.cumsum(keep.to(torch.long), dim=1) - 1 + keep &= kept_rank < max_num_peaks + + final = (keep & valid).view(-1) + sel = torch.nonzero(final, as_tuple=True)[0] # row-major: pattern-major, rank-minor + peaks = torch.stack((rows.view(-1)[sel], cols.view(-1)[sel], vals[sel]), dim=1) + bidx = torch.div(sel, kmax, rounding_mode="floor") + counts = torch.bincount(bidx, minlength=B) + return peaks, bidx, counts + + +def _split_by_counts(peaks: torch.Tensor, counts: torch.Tensor) -> list[np.ndarray]: + """Split a pattern-grouped ``(T, 3)`` peak stack into one ``(M, 3)`` array per pattern. + + Parameters + ---------- + peaks : torch.Tensor + ``(T, 3)`` peaks ordered by ascending pattern index (contiguous per pattern). + counts : torch.Tensor + ``(B,)`` number of peaks belonging to each pattern, in order. + + Returns + ------- + list of np.ndarray + Length-``B`` list of ``(M, 3)`` arrays. + """ + out = [] + start = 0 + for cnt in counts.tolist(): + out.append(_to_numpy(peaks[start : start + cnt])) + start += cnt + return out + + +def _refine_parabolic(corr_map: torch.Tensor, peaks: torch.Tensor) -> torch.Tensor: + """3-point quadratic subpixel refinement of every peak (vectorized over peaks). + + Parameters + ---------- + corr_map : torch.Tensor + ``(H, W)`` correlation map the peaks were found in. + peaks : torch.Tensor + ``(M, 3)`` ``[row, col, intensity]`` peaks to refine. + + Returns + ------- + torch.Tensor + ``(M, 3)`` peaks with subpixel ``[row, col]`` and bilinearly interpolated + intensity. + """ + if peaks.shape[0] == 0: + return peaks + a = corr_map + H, W = a.shape + out = peaks.clone() + zero = torch.zeros(out.shape[0], device=a.device, dtype=a.dtype) + + r = out[:, 0].round().long().clamp(0, H - 1) + c = out[:, 1].round().long().clamp(0, W - 1) + + r_in = (r > 0) & (r < H - 1) + ix0, ix1, ix2 = a[(r - 1).clamp(0, H - 1), c], a[r, c], a[(r + 1).clamp(0, H - 1), c] + denom_r = 4.0 * ix1 - 2.0 * ix2 - 2.0 * ix0 + dr = torch.where(r_in & (denom_r != 0), (ix2 - ix0) / denom_r, zero) + + c_in = (c > 0) & (c < W - 1) + iy0, iy1, iy2 = a[r, (c - 1).clamp(0, W - 1)], a[r, c], a[r, (c + 1).clamp(0, W - 1)] + denom_c = 4.0 * iy1 - 2.0 * iy2 - 2.0 * iy0 + dc = torch.where(c_in & (denom_c != 0), (iy2 - iy0) / denom_c, zero) + + r_sub = r.to(a.dtype) + dr + c_sub = c.to(a.dtype) + dc + out[:, 0] = r_sub + out[:, 1] = c_sub + out[:, 2] = _bilinear(a, r_sub, c_sub) + return out + + +def _refine_parabolic_batched( + corr: torch.Tensor, peaks: torch.Tensor, bidx: torch.Tensor +) -> torch.Tensor: + """3-point quadratic refinement of peaks pooled across a batch of correlation maps. + + Batched form of :func:`_refine_parabolic`. Reading ``corr[bidx, r, c]`` gathers + exactly the values the single-pattern path would read, so each peak refines + identically. + + Parameters + ---------- + corr : torch.Tensor + ``(B, H, W)`` correlation maps. + peaks : torch.Tensor + ``(T, 3)`` ``[row, col, intensity]`` peaks pooled over all patterns. + bidx : torch.Tensor + ``(T,)`` batch index selecting each peak's correlation map. + + Returns + ------- + torch.Tensor + ``(T, 3)`` peaks with subpixel ``[row, col]`` and bilinearly interpolated + intensity. + """ + if peaks.shape[0] == 0: + return peaks + H, W = corr.shape[-2], corr.shape[-1] + dtype = corr.dtype + out = peaks.clone() + zero = torch.zeros(out.shape[0], device=corr.device, dtype=dtype) + + r = out[:, 0].round().long().clamp(0, H - 1) + c = out[:, 1].round().long().clamp(0, W - 1) + + r_in = (r > 0) & (r < H - 1) + ix0 = corr[bidx, (r - 1).clamp(0, H - 1), c] + ix1 = corr[bidx, r, c] + ix2 = corr[bidx, (r + 1).clamp(0, H - 1), c] + denom_r = 4.0 * ix1 - 2.0 * ix2 - 2.0 * ix0 + dr = torch.where(r_in & (denom_r != 0), (ix2 - ix0) / denom_r, zero) + + c_in = (c > 0) & (c < W - 1) + iy0 = corr[bidx, r, (c - 1).clamp(0, W - 1)] + iy1 = corr[bidx, r, c] + iy2 = corr[bidx, r, (c + 1).clamp(0, W - 1)] + denom_c = 4.0 * iy1 - 2.0 * iy2 - 2.0 * iy0 + dc = torch.where(c_in & (denom_c != 0), (iy2 - iy0) / denom_c, zero) + + r_sub = r.to(dtype) + dr + c_sub = c.to(dtype) + dc + out[:, 0] = r_sub + out[:, 1] = c_sub + out[:, 2] = _bilinear_batched(corr, bidx, r_sub, c_sub) + return out + + +def _refine_dft(m: torch.Tensor, peaks: torch.Tensor, upsample_factor: int) -> torch.Tensor: + """Guizar-Sicairos DFT upsampling refinement of every peak (vectorized over peaks). + + Each peak is rounded to half-pixel precision (matching py4DSTEM multicorr) before + upsampling, then all peaks are refined together with one batched DFT upsampling. + + Parameters + ---------- + m : torch.Tensor + ``(H, W)`` Fourier-domain correlation product from :func:`cross_correlation`. + peaks : torch.Tensor + ``(M, 3)`` ``[row, col, intensity]`` peaks to refine. + upsample_factor : int + DFT upsampling factor. + + Returns + ------- + torch.Tensor + ``(M, 3)`` peaks with DFT-refined ``[row, col]`` (intensity unchanged). + """ + if peaks.shape[0] == 0: + return peaks + out = peaks.clone() + xy = torch.round(out[:, :2] * 2.0) / 2.0 + refined = _upsampled_correlation_batch(m, int(upsample_factor), xy) + out[:, :2] = refined + return out + + +def _refine_dft_batched( + m: torch.Tensor, peaks: torch.Tensor, bidx: torch.Tensor, upsample_factor: int +) -> torch.Tensor: + """DFT-upsampling refinement of peaks pooled across a batch of correlation products. + + Batched form of :func:`_refine_dft`. Peaks are processed in sub-chunks so the + gathered per-peak product ``m[bidx]`` (``(chunk, H, W)`` complex) stays within a + fixed memory budget regardless of how many peaks were found. + + Parameters + ---------- + m : torch.Tensor + ``(B, H, W)`` Fourier-domain correlation products. + peaks : torch.Tensor + ``(T, 3)`` ``[row, col, intensity]`` peaks pooled over all patterns. + bidx : torch.Tensor + ``(T,)`` batch index selecting each peak's correlation product. + upsample_factor : int + DFT upsampling factor. + + Returns + ------- + torch.Tensor + ``(T, 3)`` peaks with DFT-refined ``[row, col]`` (intensity unchanged). + """ + if peaks.shape[0] == 0: + return peaks + out = peaks.clone() + M, N = m.shape[-2], m.shape[-1] + xy = torch.round(out[:, :2] * 2.0) / 2.0 + cap = max(1, 8_000_000 // (M * N)) + total = peaks.shape[0] + for start in range(0, total, cap): + stop = min(start + cap, total) + b = bidx[start:stop] + out[start:stop, :2] = _upsampled_correlation_batch( + m[b], int(upsample_factor), xy[start:stop] + ) + return out + + +def _upsampled_correlation_batch( + m: torch.Tensor, upsample_factor: int, xy: torch.Tensor +) -> torch.Tensor: + """Batched DFT upsampling of the correlation peak for many shifts at once. + + Vectorizes :func:`~quantem.core.utils.imaging_utils.upsampled_correlation_torch` + over the shifts ``xy`` with batched matmuls. Two exact speedups are applied: (1) ``m`` + is conjugate-symmetric (the diffraction pattern and template are real), so only the + non-negative half of the row frequencies is contracted — halving the dominant matmul — + with a rank-1 correction on the Nyquist column; (2) the per-peak DFT kernels are + factored into shared base kernels times per-peak phases, so ``torch.exp`` runs on far + fewer elements. Both are algebraically identical to the full-spectrum DFT upsample. + + Parameters + ---------- + m : torch.Tensor + ``(M, N)`` or ``(K, M, N)`` Fourier-domain correlation product(s). A single + ``(M, N)`` product broadcasts over all ``K`` shifts. + upsample_factor : int + DFT upsampling factor. + xy : torch.Tensor + ``(K, 2)`` ``[row, col]`` peak shifts at half-pixel precision. + + Returns + ------- + torch.Tensor + ``(K, 2)`` DFT-refined ``[row, col]`` positions. + """ + import math + + device = m.device + dtype = torch.get_default_dtype() + uf = float(upsample_factor) + M, N = m.shape[-2], m.shape[-1] + + xy = torch.round(xy * uf) / uf + global_shift = math.floor(math.ceil(uf * 1.5) / 2.0) + upsample_center = global_shift - uf * xy # (K, 2): [row, col] + + num = int(math.ceil(1.5 * uf)) + half = M // 2 + 1 + col_freq = (torch.fft.ifftshift(torch.arange(N, device=device)) - math.floor(N / 2)).to(dtype) + row_freq = (torch.fft.ifftshift(torch.arange(M, device=device)) - math.floor(M / 2)).to(dtype) + + # ``m = F(dp) * conj(F(template))`` is 2D conjugate-symmetric for real ``dp`` and + # template, so the upper-half row frequencies are redundant. Contract only the + # non-negative half (rows ``u = 0 .. M // 2``) and fold the conjugate upper rows + # back in with a weight of 2 — 1 for the self-paired DC row and, when ``M`` is even, + # the Nyquist row. This halves the dominant matmul over the row axis. + row_freq = row_freq[:half] + row_weight = torch.full((half,), 2.0, device=device, dtype=dtype) + row_weight[0] = 1.0 + if M % 2 == 0: + row_weight[-1] = 1.0 + + base = torch.arange(num, device=device, dtype=dtype) + factor_col = -2j * math.pi / (N * uf) + factor_row = -2j * math.pi / (M * uf) + + # Factor each kernel ``exp(f · freq · (base - center))`` into a peak-independent base + # kernel ``exp(f · freq · base)`` times a per-peak phase ``exp(-f · freq · center)``. + # The base kernels are tiny and built once; the phases fold into ``m`` and the + # intermediate product, so ``torch.exp`` runs on ~20x fewer elements. + base_col = torch.exp(factor_col * (col_freq[:, None] * base[None, :])) # (N, num) + base_row = torch.exp(factor_row * (base[:, None] * row_freq[None, :]))[None] # (1, num, half) + phase_col = torch.exp(-factor_col * (col_freq[None, :] * upsample_center[:, 1:2])) # (K, N) + phase_row = torch.exp(-factor_row * (upsample_center[:, 0:1] * row_freq[None, :])) # (K, half) + + mc_half = m.conj()[..., :half, :] # (K, half, N) — non-negative row frequencies only + # Fold the row phase and conjugate-symmetry weight into ``m`` so the shared base row + # kernel is reused across peaks; then contract rows, apply the column phase, and + # contract columns against the shared base column kernel. + prod = torch.matmul(base_row, mc_half * (row_weight * phase_row)[:, :, None]) # (K, num, N) + up = torch.matmul(prod * phase_col[:, None, :], base_col) # (K, num, num) + + if N % 2 == 0: + # Row-only folding is exact for every column except the Nyquist column ``v = N/2``, + # whose conjugate partner stays in the same column. Correct it with an exact rank-1 + # update ``pr_diff ⊗ col_kern[:, N/2]``, where ``pr_diff = -2i Im(S)`` and ``S`` sums + # only the kept interior rows. + nyq = N // 2 + interior = row_weight - 1.0 # 1 on interior rows, 0 on the DC / Nyquist rows + s = torch.matmul(base_row, (mc_half[..., nyq] * (interior * phase_row)).unsqueeze(-1)) + pr_diff = -2j * s[..., 0].imag # (K, num) + col_nyq = base_col[nyq, :][None, :] * phase_col[:, nyq : nyq + 1] # (K, num) + up = up + pr_diff[:, :, None] * col_nyq[:, None, :] + + image_up = up.real + + K = xy.shape[0] + kidx = torch.arange(K, device=device) + idx = torch.argmax(image_up.reshape(K, -1), dim=1) + sub_r = torch.div(idx, num, rounding_mode="floor") + sub_c = idx % num + + # 3-point parabolic refinement around the upsampled maximum (interior only) + interior = (sub_r > 0) & (sub_r < num - 1) & (sub_c > 0) & (sub_c < num - 1) + rr = sub_r.clamp(1, num - 2) + cc = sub_c.clamp(1, num - 2) + c11 = image_up[kidx, rr, cc] + c21, c01 = image_up[kidx, rr + 1, cc], image_up[kidx, rr - 1, cc] + c12, c10 = image_up[kidx, rr, cc + 1], image_up[kidx, rr, cc - 1] + zero = torch.zeros(K, device=device, dtype=dtype) + denom_x = 4.0 * c11 - 2.0 * c21 - 2.0 * c01 + denom_y = 4.0 * c11 - 2.0 * c12 - 2.0 * c10 + dx = torch.where(interior & (denom_x != 0), (c21 - c01) / denom_x, zero) + dy = torch.where(interior & (denom_y != 0), (c12 - c10) / denom_y, zero) + + sub = torch.stack([sub_r.to(dtype), sub_c.to(dtype)], dim=1) - global_shift + return xy + (sub + torch.stack([dx, dy], dim=1)) / uf + + +def _bilinear(a: torch.Tensor, r: torch.Tensor, c: torch.Tensor) -> torch.Tensor: + """Bilinear interpolation of ``a`` at fractional ``(r, c)`` (vectorized over peaks). + + Parameters + ---------- + a : torch.Tensor + ``(H, W)`` map to sample. + r : torch.Tensor + ``(K,)`` fractional row coordinates. + c : torch.Tensor + ``(K,)`` fractional column coordinates. + + Returns + ------- + torch.Tensor + ``(K,)`` interpolated values. + """ + H, W = a.shape + r0 = torch.floor(r).long() + c0 = torch.floor(c).long() + r1 = (r0 + 1).clamp(max=H - 1) + c1 = (c0 + 1).clamp(max=W - 1) + r0 = r0.clamp(0, H - 1) + c0 = c0.clamp(0, W - 1) + dr = r - r0.to(a.dtype) + dc = c - c0.to(a.dtype) + return ( + (1 - dr) * (1 - dc) * a[r0, c0] + + (1 - dr) * dc * a[r0, c1] + + dr * (1 - dc) * a[r1, c0] + + dr * dc * a[r1, c1] + ) + + +def _bilinear_batched( + corr: torch.Tensor, bidx: torch.Tensor, r: torch.Tensor, c: torch.Tensor +) -> torch.Tensor: + """Bilinear interpolation of a batch ``corr`` ``(B, H, W)`` at per-peak ``(r, c)``. + + Batched form of :func:`_bilinear`. The four corner samples are gathered as + ``corr[bidx, r0, c0]`` etc., matching the single-pattern interpolation + value-for-value. + + Parameters + ---------- + corr : torch.Tensor + ``(B, H, W)`` maps to sample. + bidx : torch.Tensor + ``(T,)`` batch index selecting each peak's map. + r : torch.Tensor + ``(T,)`` fractional row coordinates. + c : torch.Tensor + ``(T,)`` fractional column coordinates. + + Returns + ------- + torch.Tensor + ``(T,)`` interpolated values. + """ + H, W = corr.shape[-2], corr.shape[-1] + dtype = corr.dtype + r0 = torch.floor(r).long() + c0 = torch.floor(c).long() + r1 = (r0 + 1).clamp(max=H - 1) + c1 = (c0 + 1).clamp(max=W - 1) + r0 = r0.clamp(0, H - 1) + c0 = c0.clamp(0, W - 1) + dr = r - r0.to(dtype) + dc = c - c0.to(dtype) + return ( + (1 - dr) * (1 - dc) * corr[bidx, r0, c0] + + (1 - dr) * dc * corr[bidx, r0, c1] + + dr * (1 - dc) * corr[bidx, r1, c0] + + dr * dc * corr[bidx, r1, c1] + ) + + +def _to_numpy(peaks: torch.Tensor) -> np.ndarray: + """Convert a peaks tensor to a contiguous ``(M, 3)`` float64 numpy array. + + Parameters + ---------- + peaks : torch.Tensor + ``(M, 3)`` (or flat) peaks tensor on any device. + + Returns + ------- + np.ndarray + ``(M, 3)`` ``float64`` array of ``[q_row, q_col, intensity]`` rows. + """ + return peaks.detach().cpu().numpy().astype(np.float64).reshape(-1, 3) diff --git a/src/quantem/diffraction/illumination.py b/src/quantem/diffraction/illumination.py new file mode 100644 index 000000000..54f60aeff --- /dev/null +++ b/src/quantem/diffraction/illumination.py @@ -0,0 +1,457 @@ +"""Illumination-averaged excitation envelopes for kinematical patterns. + +Precession rotates the incident beam on a cone and a convergent probe fills +a disk of directions; the recorded intensity of a reflection is the +average of its rocking curve over that support. For a ring centered on the +optic axis the excitation error of reflection g is exactly +s_g(phi) = c_g + a_g cos(phi - delta_g), and for a small disk it is affine +in the incident direction to leading order, so the average of a Gaussian +excitation envelope over ring and disk reduces to one scalar transform, + + G(c, a, b; sigma) = sqrt(2/pi) int_0^inf exp(-x^2/2) cos(c x / sigma) + J0(a x / sigma) jinc(b x / sigma) dx, + +with jinc(x) = 2 J1(x) / x; G(c, 0, 0) = exp(-c^2 / 2 sigma^2) recovers the +static envelope. The functions here evaluate G for whole reflection lists +at once (vectorized Gauss-Legendre quadrature of the transform, exact to +~1e-9 over the parameter range of electron diffraction), give the ring and +disk coefficients (c, a, b) from the geometry, and are shared by the +pattern simulation, the orientation library and the refinements. +""" + +from __future__ import annotations + +import numpy as np +import torch +from scipy.special import j0, j1 + +from quantem.core.utils.utils import electron_wavelength_angstrom + +_X_NODES, _X_WEIGHTS = np.polynomial.legendre.leggauss(400) +_X_MAX = 9.0 +_X = 0.5 * _X_MAX * (_X_NODES + 1) +_W = 0.5 * _X_MAX * _X_WEIGHTS * np.exp(-0.5 * _X**2) * np.sqrt(2 / np.pi) + + +def _jinc(x: np.ndarray) -> np.ndarray: + out = np.ones_like(x) + nz = np.abs(x) > 1e-12 + out[nz] = 2 * j1(x[nz]) / x[nz] + return out + + +def gaussian_envelope(c, a, b, sigma: float) -> np.ndarray: + """Illumination-averaged Gaussian excitation envelope G(c, a, b; sigma). + + Parameters + ---------- + c, a, b : array-like + Central excitation error, ring amplitude and disk amplitude of each + reflection (1/Angstroms), from `excitation_coefficients`. + sigma : float + Width of the excitation envelope (1/Angstroms). + + Returns + ------- + np.ndarray + The averaged envelope in [0, 1], same shape as c. + """ + c = np.asarray(c, dtype=float) + a = np.broadcast_to(np.asarray(a, dtype=float), c.shape) + b = np.broadcast_to(np.asarray(b, dtype=float), c.shape) + x = _X / sigma + integrand = np.cos(c[..., None] * x) * j0(a[..., None] * x) * _jinc(b[..., None] * x) + return np.clip(integrand @ _W, 0.0, 1.0) + + +def excitation_coefficients( + g_lab: torch.Tensor | np.ndarray, + energy_ev: float, + precession_deg: float = 0.0, + semiconv_mrad: float = 0.0, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Central excitation error and illumination amplitudes of each reflection. + + For a precession ring of radius r = k0 sin(theta_p) centered on the optic + axis, with K = sqrt(k0^2 - r^2), + + c_g = (2 K g_z - |g|^2) / (2 (K - g_z)), a_g = r |g_xy| / |K - g_z|, + + which is exact: s_g(phi) = c_g + a_g cos(phi - delta_g). The convergence + disk of radius R = k0 sin(alpha) adds b_g = R |g_xy| / |K - g_z| under + the same affine model. Without illumination, c_g is the static + excitation error and a_g = b_g = 0. + + Parameters + ---------- + g_lab : torch.Tensor | np.ndarray + Lab-frame reciprocal vectors (N, 3), 1/Angstroms, with the beam + along -z. + energy_ev : float + Beam energy, eV. + precession_deg : float, default=0.0 + Precession semi-angle, degrees. + semiconv_mrad : float, default=0.0 + Convergence semi-angle, mrad. + + Returns + ------- + c, a, b : np.ndarray + Central excitation error, ring amplitude and disk amplitude (N,), + 1/Angstroms. + """ + g = np.asarray( + g_lab.detach().cpu().numpy() if isinstance(g_lab, torch.Tensor) else g_lab, dtype=float + ) + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + r = k0 * np.sin(np.deg2rad(precession_deg)) + R = k0 * np.sin(semiconv_mrad * 1e-3) + K = np.sqrt(k0**2 - r**2) + den = K - g[:, 2] + g2 = (g**2).sum(axis=1) + gxy = np.hypot(g[:, 0], g[:, 1]) + c = (2 * K * g[:, 2] - g2) / (2 * den) + a = r * gxy / np.abs(den) + b = R * gxy / np.abs(den) + return c, a, b + + +def relrod_factor(g_lab, n_lab, energy_ev: float, precession_deg: float = 0.0): + """How far along a relrod the Ewald sphere is, per unit excitation error. + + A plate-shaped crystal spreads every reciprocal lattice point into a rod + along the plate normal n, a long one for a 2D material. The sphere meets + the rod through g at g + t n, with, to first order, + + t = -f s_g, f = (K - g_z) / (K n_z - n . g), + + where s_g is the excitation error measured along the beam and + K = sqrt(k0^2 - r^2) as in :func:`excitation_coefficients`. For n along + the beam f = 1. A rod nearly tangent to the sphere (an edge-on plate) is + never excited; f is set huge there so the reflection drops out. + + Parameters + ---------- + g_lab : torch.Tensor | np.ndarray + Lab-frame reciprocal vectors (..., 3), 1/Angstroms. + n_lab : torch.Tensor | np.ndarray + Lab-frame unit plate normal, broadcastable to `g_lab`. + energy_ev : float + Beam energy, eV. + precession_deg : float, default=0.0 + Precession semi-angle, degrees; sets K as above. + + Returns + ------- + torch.Tensor | np.ndarray + f (...,), dimensionless, the same type as `g_lab`. 1e6 where the rod + is within 0.05 K of tangent to the sphere. + """ + k0 = 1.0 / electron_wavelength_angstrom(energy_ev) + r = k0 * np.sin(np.deg2rad(precession_deg)) + K = float(np.sqrt(k0**2 - r**2)) + if isinstance(g_lab, torch.Tensor): + n = torch.as_tensor(n_lab, dtype=g_lab.dtype, device=g_lab.device) + den_n = K * n[..., 2] - (g_lab * n).sum(-1) + tangent = den_n.abs() < 0.05 * K + f = (K - g_lab[..., 2]) / torch.where(tangent, torch.ones_like(den_n), den_n) + return torch.where(tangent, torch.full_like(f, 1e6), f) + g = np.asarray(g_lab, dtype=float) + n = np.asarray(n_lab, dtype=float) + den_n = K * n[..., 2] - (g * n).sum(-1) + tangent = np.abs(den_n) < 0.05 * K + f = (K - g[..., 2]) / np.where(tangent, 1.0, den_n) + return np.where(tangent, 1e6, f) + + +def averaged_gaussian_intensity_envelope( + g_lab, + energy_ev: float, + sigma: float, + precession_deg: float = 0.0, + semiconv_mrad: float = 0.0, + foil_normal_lab=None, +) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + """Illumination-averaged Gaussian envelope of every reflection. + + With `foil_normal_lab` the excitation is measured along the relrod + instead of along the beam: c, a and b are the distances along the rod + (see :func:`relrod_factor`). + + Parameters + ---------- + g_lab : torch.Tensor | np.ndarray + Lab-frame reciprocal vectors (N, 3), 1/Angstroms. + energy_ev : float + Beam energy, eV. + sigma : float + Width of the excitation envelope, 1/Angstroms. + precession_deg : float, default=0.0 + Precession semi-angle, degrees. + semiconv_mrad : float, default=0.0 + Convergence semi-angle, mrad. + foil_normal_lab : array-like, optional + Lab-frame unit plate normal (3,). + + Returns + ------- + envelope : np.ndarray + Averaged envelope (N,) in [0, 1]. + c, a, b : np.ndarray + Central excitation error and ring and disk amplitudes (N,), + 1/Angstroms; a + b is the half-width of the swept range about c. + """ + c, a, b = excitation_coefficients(g_lab, energy_ev, precession_deg, semiconv_mrad) + if foil_normal_lab is not None: + g_np = g_lab.detach().cpu().numpy() if isinstance(g_lab, torch.Tensor) else g_lab + f = relrod_factor( + np.asarray(g_np, dtype=float), foil_normal_lab, energy_ev, precession_deg + ) + c, a, b = c * f, a * np.abs(f), b * np.abs(f) + if precession_deg <= 0 and semiconv_mrad <= 0: + return np.exp(-0.5 * (c / sigma) ** 2), c, a, b + # A reflection farther from the Ewald sphere than the illumination sweeps + # it, plus six envelope widths, is never excited (the envelope there is + # below 1e-8). At k_max = 2 that is nearly all of them, and the quadrature + # below costs 400 Bessel evaluations per reflection. + env = np.zeros_like(c) + live = np.abs(c) < a + b + 6.0 * sigma + if live.any(): + if semiconv_mrad <= 0: + # a pure precession ring has the fast Bessel series the orientation + # plan uses; the disk quadrature is needed only with convergence + env[live] = ( + gaussian_envelope_ring_torch( + torch.as_tensor(c[live]), torch.as_tensor(a[live]), sigma + ) + .cpu() + .numpy() + ) + else: + env[live] = gaussian_envelope(c[live], a[live], b[live], sigma) + return env, c, a, b + + +def ring_disk_quadrature(r: float, R: float, n_phi: int = 128, n_r: int = 8, n_psi: int = 32): + """Positive quadrature of the ring x disk illumination. + + The reference against which the analytic envelopes are checked. + + Parameters + ---------- + r : float + Ring radius, in the units the tilts are wanted in. + R : float + Disk radius, same units. + n_phi : int, default=128 + Equally spaced points on the ring. + n_r : int, default=8 + Gauss-Legendre radial nodes of the disk (in r^2). + n_psi : int, default=32 + Equally spaced azimuths of the disk. + + Returns + ------- + tilts : np.ndarray + In-plane tilts (M, 2), ring point plus disk point. + weights : np.ndarray + Positive weights (M,) summing to 1. + """ + phi = 2 * np.pi * np.arange(n_phi) / n_phi if r > 0 else np.zeros(1) + ring = r * np.column_stack((np.cos(phi), np.sin(phi))) + if R > 0: + x, w = np.polynomial.legendre.leggauss(n_r) + radius = R * np.sqrt((x + 1) / 2) + psi = 2 * np.pi * np.arange(n_psi) / n_psi + disk = (radius[:, None, None] * np.column_stack((np.cos(psi), np.sin(psi)))[None]).reshape( + -1, 2 + ) + wd = np.repeat(w / 2 / n_psi, n_psi) + else: + disk, wd = np.zeros((1, 2)), np.ones(1) + t = (ring[:, None] + disk[None]).reshape(-1, 2) + weights = np.tile(wd, len(ring)) / len(ring) + return t, weights + + +def gaussian_envelope_ring_series(c, a, sigma: float, n_terms: int = 6) -> np.ndarray: + """Ring-averaged Gaussian envelope, b = 0, by the Bessel series + + G = exp(-c^2/2 sigma^2 - v) [I0(u) I0(v) + 2 sum_n (-1)^n I_2n(u) I_n(v)], + u = c a / sigma^2, v = a^2 / 4 sigma^2, + + which converges in a few terms for v < 1 (the precession sweep of the + excitation error smaller than the envelope width, the electron + diffraction regime); larger v falls back to the transform. + + Parameters + ---------- + c : np.ndarray | torch.Tensor + Central excitation errors, any shape, 1/Angstroms. + a : np.ndarray | torch.Tensor | float + Ring amplitudes, broadcast to `c`, 1/Angstroms. + sigma : float + Width of the excitation envelope, 1/Angstroms. + n_terms : int, default=6 + Terms of the sum over n. + + Returns + ------- + np.ndarray | torch.Tensor + Envelope in [0, 1], shape of `c`; a float64 tensor when `c` is a + tensor. + """ + from scipy.special import iv + + is_torch = isinstance(c, torch.Tensor) + c_np = np.asarray(c.detach().cpu().numpy() if is_torch else c, dtype=float) + a_np = np.broadcast_to( + np.asarray(a.detach().cpu().numpy() if isinstance(a, torch.Tensor) else a, dtype=float), + c_np.shape, + ) + u = c_np * a_np / sigma**2 + v = a_np**2 / (4 * sigma**2) + total = iv(0, u) * iv(0, v) + for n in range(1, n_terms + 1): + total = total + 2 * (-1) ** n * iv(2 * n, u) * iv(n, v) + out = np.exp(-0.5 * (c_np / sigma) ** 2 - v) * total + big = v > 1.0 + if np.any(big): + out[big] = gaussian_envelope(c_np[big], a_np[big], 0.0, sigma) + out = np.clip(out, 0.0, 1.0) + return torch.as_tensor(out, dtype=torch.float64) if is_torch else out + + +def excitation_amplitudes( + g_lab: torch.Tensor, energy_ev: float, precession_deg: float, semiconv_mrad: float +): + """Ring and disk amplitudes of lab-frame reflections, in torch. + + The a and b of :func:`excitation_coefficients`, without c, for any + leading shape and differentiable in `g_lab`. + + Parameters + ---------- + g_lab : torch.Tensor + Lab-frame reciprocal vectors (..., 3), 1/Angstroms. + energy_ev : float + Beam energy, eV. + precession_deg : float + Precession semi-angle, degrees. + semiconv_mrad : float + Convergence semi-angle, mrad. + + Returns + ------- + a, b : torch.Tensor + Ring and disk amplitudes (...,), 1/Angstroms; zero without + precession or convergence. + """ + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + r = k0 * np.sin(np.deg2rad(precession_deg)) + R = k0 * np.sin(semiconv_mrad * 1e-3) + K = np.sqrt(k0**2 - r**2) + den = (K - g_lab[..., 2]).abs().clamp_min(1e-12) + gxy = torch.hypot(g_lab[..., 0], g_lab[..., 1]) + return r * gxy / den, R * gxy / den + + +_V_NODES, _V_WEIGHTS = np.polynomial.legendre.leggauss(200) +_V = 0.5 * (_V_NODES + 1) +_WV = 0.5 * _V_WEIGHTS * 2 * (1 - _V) + + +def slab_envelope(c, a, b, thickness_A: float) -> np.ndarray: + """Illumination-averaged finite-thickness (first Born) rocking curve, + + S(c, a, b; z) = 2 int_0^1 (1 - v) cos(2 pi c z v) J0(2 pi a z v) + jinc(2 pi b z v) dv, + + which reduces to sinc(c z)^2 without illumination; the Born intensity + of reflection g is (pi |U_g| z / k0)^2 S. Vectorized Gauss-Legendre + quadrature over v, exact to ~1e-10 for the phase ranges of electron + diffraction (c z below ~20). + + Parameters + ---------- + c, a, b : array-like + Central excitation error, ring amplitude and disk amplitude + (1/Angstroms), from :func:`excitation_coefficients`; a and b are + broadcast to the shape of c. + thickness_A : float + Specimen thickness z, Angstroms. + + Returns + ------- + np.ndarray + S in [0, 1], same shape as c. + """ + c = np.asarray(c, dtype=float) + a = np.broadcast_to(np.asarray(a, dtype=float), c.shape) + b = np.broadcast_to(np.asarray(b, dtype=float), c.shape) + x = 2 * np.pi * thickness_A * _V + integrand = np.cos(c[..., None] * x) * j0(a[..., None] * x) * _jinc(b[..., None] * x) + return np.clip(integrand @ _WV, 0.0, 1.0) + + +def gaussian_envelope_ring_torch(c: torch.Tensor, a: torch.Tensor, sigma: float) -> torch.Tensor: + """Ring-averaged Gaussian envelope (b = 0) in torch, for the refinements. + + The Bessel series of :func:`gaussian_envelope_ring_series` truncated + after the I_4 term, with I_n(u) from the I_0 / I_1 recurrences (and + their small-argument series where the recurrence would cancel). Below + v = a^2 / 4 sigma^2 = 0.3 the truncation error is under 1e-5; larger + sweeps are averaged over the ring directly by quadrature. Stays on the + device of `c` and is differentiable. + + Parameters + ---------- + c : torch.Tensor + Central excitation errors, any shape, 1/Angstroms. + a : torch.Tensor + Ring amplitudes, broadcastable to `c`, 1/Angstroms. + sigma : float + Width of the excitation envelope, 1/Angstroms. + + Returns + ------- + torch.Tensor + float64 envelope in [0, 1], the broadcast shape of `c` and `a`. + """ + c = c.to(torch.float64) + a = torch.as_tensor(a, dtype=torch.float64) + u = c * a / sigma**2 + v = a * a / (4 * sigma**2) + i0u = torch.special.i0(u) + i1u = torch.special.i1(u) + small2 = u.abs() < 1e-3 + u_safe = torch.where(small2, torch.ones_like(u), u) + i2u = torch.where(small2, u * u / 8, i0u - 2 * i1u / u_safe) + small4 = u.abs() < 5e-2 + i3u = torch.where(small4, u**3 / 48, i1u - 4 * i2u / u_safe) + i4u = torch.where(small4, u**4 / 384, i2u - 6 * i3u / u_safe) + i0v = torch.special.i0(v) + i1v = torch.special.i1(v) + smallv = v < 1e-3 + v_safe = torch.where(smallv, torch.ones_like(v), v) + i2v = torch.where(smallv, v * v / 8, i0v - 2 * i1v / v_safe) + out = torch.exp(-0.5 * (c / sigma) ** 2 - v) * (i0u * i0v - 2 * i2u * i1v + 2 * i4u * i2v) + # v follows the shape of a, which may be narrower than the output: the + # fallback mask has to be taken in the broadcast shape or it indexes the + # wrong elements (a narrow sigma or a large sweep reaches this branch) + big = torch.broadcast_to(v > 0.3, out.shape) + if bool(big.any()): + out = out.clone() + c_b = torch.broadcast_to(c, out.shape)[big] + a_b = torch.broadcast_to(a, out.shape)[big] + # the ring average itself, (1/pi) int_0^pi exp(-(c - a cos phi)^2 / + # 2 sigma^2) dphi, by the midpoint rule: exact to rounding for this + # periodic integrand once the nodes resolve a / sigma, and it stays + # in torch on the input's device + n = int(min(256, max(16, np.ceil(8 + 4 * float(a_b.max()) / sigma)))) + phi = (torch.arange(n, dtype=torch.float64, device=c_b.device) + 0.5) * (np.pi / n) + d = c_b[:, None] - a_b[:, None] * torch.cos(phi) + out[big] = torch.exp(-0.5 * (d / sigma) ** 2).mean(dim=-1).to(out.dtype) + return out.clamp(0.0, 1.0) diff --git a/src/quantem/diffraction/orientation.py b/src/quantem/diffraction/orientation.py new file mode 100644 index 000000000..e4d4010a1 --- /dev/null +++ b/src/quantem/diffraction/orientation.py @@ -0,0 +1,2851 @@ +"""Orientation mapping of crystalline 4D-STEM data. + +OrientationMap matches measured Bragg peaks (a quantem Vector) against a +library of simulated kinematical patterns from a Crystal, using sparse polar +correlation (Ophus et al., Microsc. Microanal. 28, 390 (2022)) implemented as +batched torch operations. + +The method: + +1. Sample zone axes over the symmetry-reduced fundamental wedge (or the + hemisphere), build a polar-coordinate reference library P(zone, shell, + gamma) where shells are the reciprocal-lattice radii of the crystal. +2. Convert measured peaks at each probe position into the same sparse polar + representation X(shell, gamma). +3. Correlate over in-plane angle gamma by FFT, over all zones at once, using + one batched matrix multiplication per gamma frequency. The mirror channel + (conjugate FFT) tests inversion-related orientations at no library cost. +4. Optionally refine the best zone axes on a finer local grid. + +Orientations are unit quaternions; see quantem.diffraction.rotations. +""" + +from __future__ import annotations + +import warnings + +import numpy as np +import torch +from tqdm import tqdm + +from quantem.core.datastructures.vector import Vector +from quantem.core.io.serialize import AutoSerialize +from quantem.core.utils.utils import electron_wavelength_angstrom +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.defaults import ( + MIN_NUMBER_PEAKS, + MIN_PAIRS, + PAIR_DISTANCE, + POWER_INTENSITY, + SIGMA_EXCITATION, + resolve, +) +from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qconj, + qmult, + qnormalize, + qrotate, + quat_from_axis_angle, + quat_from_zone_axis, + sample_zone_axes, +) + + +def position_mask(positions, shape: tuple[int, int]) -> torch.Tensor: + """Normalize a `positions` argument into an (R, C) boolean mask. + + Used by the staged workflow: run matching or refinement on a handful of + positions, look at the fits, then run the whole scan with the same + arguments. + + Parameters + ---------- + positions : None | list[tuple[int, int]] | np.ndarray + None for every position, a list of (row, col) scan positions, or an + (R, C) boolean array. + shape : tuple[int, int] + Scan shape (R, C) in probe positions. + + Returns + ------- + torch.Tensor + (R, C) boolean mask. + + Raises + ------ + ValueError + If a boolean mask has the wrong shape, the input is neither a mask + nor a list of (row, col), or a position lies outside the scan. + """ + R, C = shape + if positions is None: + return torch.ones((R, C), dtype=torch.bool) + arr = np.asarray(positions) + if arr.dtype == bool: + if arr.shape != (R, C): + raise ValueError(f"boolean positions mask must have shape {(R, C)}, got {arr.shape}") + return torch.as_tensor(arr, dtype=torch.bool) + arr = np.atleast_2d(arr) + if arr.ndim != 2 or arr.shape[1] != 2: + raise ValueError("positions must be None, an (R, C) boolean mask, or a list of (row, col)") + mask = torch.zeros((R, C), dtype=torch.bool) + for r, c in arr.astype(int): + if not (0 <= r < R and 0 <= c < C): + raise ValueError(f"position ({r}, {c}) is outside the scan {(R, C)}") + mask[r, c] = True + return mask + + +def scan_scalebar(metadata: dict) -> dict | None: + """Scale bar arguments from the scan calibration recorded on the peaks. + + Parameters + ---------- + metadata : dict + Peak metadata carrying "scan_sampling" and "scan_units". + + Returns + ------- + dict | None + {"sampling": step, "units": units} when the scan was calibrated, or + None when it is still in pixels, which tells a plot to draw no + scale bar. + """ + step = (metadata or {}).get("scan_sampling") + units = (metadata or {}).get("scan_units") + if step is None or units is None: + return None + step = float(np.mean(np.atleast_1d(np.asarray(step, dtype=float)))) + units = str(units) + if not np.isfinite(step) or step <= 0 or units.lower() in ("pixels", "px", "pixel"): + return None + return {"sampling": step, "units": units} + + +def smooth_quaternions( + quats: torch.Tensor, + active: torch.Tensor, + sym_quats: torch.Tensor, + sigma_px: float = 1.0, + sigma_deg: float = 1.0, + max_angle_deg: float = 5.0, +) -> torch.Tensor: + """Bilateral average of an orientation field, (R, C, 4). + + Each position is replaced by the weighted mean of the orientations around + it, with weight exp(-r^2 / 2 sigma_px^2) * exp(-theta^2 / 2 sigma_deg^2) + for a neighbour r probe positions away and theta degrees misoriented, and + with neighbours beyond `max_angle_deg` dropped. The angular term is what + keeps a grain boundary or a second variant out of the average. + + This is an average, not a fit: it moves each orientation away from the one + that best explains its own pattern. Use it to display a map, not to + produce the orientations a later step will measure from. + + Parameters + ---------- + quats : torch.Tensor + (R, C, 4) orientation quaternions. + active : torch.Tensor + (R, C) boolean mask of positions to smooth and to average over. + sym_quats : torch.Tensor + (S, 4) symmetry rotations used to reduce the misorientations. + sigma_px : float, default=1.0 + Spatial width of the kernel in probe positions. + sigma_deg : float, default=1.0 + Angular width of the kernel in degrees. + max_angle_deg : float, default=5.0 + Neighbours misoriented by more than this many degrees are dropped. + + Returns + ------- + torch.Tensor + (R, C, 4) smoothed quaternions; inactive positions, and positions + with no neighbour inside `max_angle_deg`, are returned unchanged. + """ + q = torch.as_tensor(quats, dtype=torch.float64) + R, C = q.shape[:2] + active = torch.as_tensor(active, dtype=torch.bool) + sym = torch.as_tensor(sym_quats, dtype=torch.float64).reshape(-1, 4) + rad = max(1, int(np.ceil(3 * sigma_px))) + cos_max = float(np.cos(np.deg2rad(max_angle_deg) / 2)) + # accumulate the weighted outer products one neighbour offset at a time, + # over the whole map at once + M = torch.zeros((R, C, 4, 4), dtype=torch.float64) + count = torch.zeros((R, C), dtype=torch.long) + for dr in range(-rad, rad + 1): + for dc in range(-rad, rad + 1): + r0, r1 = max(0, -dr), min(R, R - dr) + c0, c1 = max(0, -dc), min(C, C - dc) + if r1 <= r0 or c1 <= c0: + continue + q0 = q[r0:r1, c0:c1] + qn = q[r0 + dr : r1 + dr, c0 + dc : c1 + dc] + ok = active[r0:r1, c0:c1] & active[r0 + dr : r1 + dr, c0 + dc : c1 + dc] + # the symmetry image of the neighbour nearest each centre + cand = qmult(qn[..., None, :], sym) # (r, c, S, 4) + dots = torch.einsum("rcsi,rci->rcs", cand, q0) + best = dots.abs().argmax(dim=-1) + qk = torch.gather(cand, 2, best[..., None, None].expand(*best.shape, 1, 4))[..., 0, :] + dot = (qk * q0).sum(-1) + qk = qk * torch.sign(dot)[..., None] + cos_half = dot.abs().clamp(max=1.0) + ok &= cos_half >= cos_max + ang = torch.rad2deg(2 * torch.acos(cos_half)) + w = np.exp(-(dr * dr + dc * dc) / (2 * sigma_px**2)) * torch.exp( + -(ang**2) / (2.0 * sigma_deg**2) + ) + w = torch.where(ok, w, torch.zeros_like(w)) + M[r0:r1, c0:c1] += w[..., None, None] * qk[..., :, None] * qk[..., None, :] + count[r0:r1, c0:c1] += ok.to(torch.long) + out = q.clone() + # a position needs itself and at least one neighbour inside the angle + upd = active & (count >= 2) + if bool(upd.any()): + _, evecs = torch.linalg.eigh(M[upd]) + out[upd] = qnormalize(evecs[..., -1]) + return out + + +def fibonacci_hemisphere(n_points: int, dtype=torch.float64) -> torch.Tensor: + """Spherical Fibonacci sampling of the upper hemisphere, (N, 3).""" + i = torch.arange(n_points, dtype=dtype) + 0.5 + z = i / n_points # (0, 1): upper hemisphere + phi = i * (np.pi * (3 - np.sqrt(5))) + r = torch.sqrt(1 - z**2) + return torch.stack((r * torch.cos(phi), r * torch.sin(phi), z), dim=-1) + + +def _zone_peak_parabolic( + za: torch.Tensor, + n_pos: torch.Tensor, + c_n: torch.Tensor, + n_ok: torch.Tensor, + step_rad: float, +) -> torch.Tensor: + """Sub-grid zone axis from the correlations of a zone and its neighbors. + + A quadratic surface c(x, y) is fit by least squares to the correlation + over the neighborhood in the tangent plane of the best zone (x, y in + radians); its vertex is the refined zone axis when it lies within one + grid step of the node and the surface is concave. Otherwise the + correlation-weighted centroid of the neighbors above 70 % of the best + value is used, and the node itself when neither applies. + + Parameters + ---------- + za : (B, 3) best zone axes; n_pos : (B, K, 3) neighbor directions + (the best zone included); c_n : (B, K) their correlations; n_ok : (B, K) + validity; step_rad : zone grid step. + """ + B, K = c_n.shape + # tangent frame at the node + ref = torch.where( + za[:, 2:3].abs() < 0.9, + torch.tensor([0.0, 0.0, 1.0], dtype=za.dtype).expand(B, 3), + torch.tensor([1.0, 0.0, 0.0], dtype=za.dtype).expand(B, 3), + ) + e1 = torch.cross(za, ref, dim=-1) + e1 = e1 / torch.linalg.norm(e1, dim=-1, keepdim=True).clamp_min(1e-12) + e2 = torch.cross(za, e1, dim=-1) + d = n_pos - za[:, None, :] + x = (d * e1[:, None, :]).sum(-1) + y = (d * e2[:, None, :]).sum(-1) + c_best = c_n.amax(dim=1, keepdim=True) + w = n_ok.to(za.dtype) + out = za.clone() + # centroid fallback (the previous estimator) + wgt = (c_n - 0.7 * c_best).clamp_min(0) * w + cen = (wgt[:, :, None] * n_pos).sum(1) + cen_ok = torch.linalg.norm(cen, dim=-1) > 1e-12 + cen = cen / torch.linalg.norm(cen, dim=-1, keepdim=True).clamp_min(1e-12) + out[cen_ok] = cen[cen_ok] + # quadratic fit where at least 6 valid neighbors exist + A = torch.stack([torch.ones_like(x), x, y, x * x, x * y, y * y], dim=-1) * w[:, :, None] + b = (c_n - c_best) * w + enough = w.sum(1) >= 6 + if bool(enough.any()): + At = A.transpose(1, 2) + AtA = At @ A + 1e-12 * torch.eye(6, dtype=za.dtype) + coef = torch.linalg.solve(AtA, (At @ b[:, :, None]))[..., 0] # (B, 6) + cb, cc, cd, ce, cf = coef[:, 1], coef[:, 2], coef[:, 3], coef[:, 4], coef[:, 5] + H = torch.stack([torch.stack([2 * cd, ce], -1), torch.stack([ce, 2 * cf], -1)], -2) + det = 4 * cd * cf - ce * ce + concave = (cd < 0) & (cf < 0) & (det > 0) + grad = torch.stack([cb, cc], -1) + vert = torch.zeros_like(grad) + ok = enough & concave + if bool(ok.any()): + vert[ok] = -torch.linalg.solve(H[ok], grad[ok][..., None])[..., 0] + inside = ok & (torch.linalg.norm(vert, dim=-1) <= step_rad) + if bool(inside.any()): + v = za + vert[:, 0:1] * e1 + vert[:, 1:2] * e2 + v = v / torch.linalg.norm(v, dim=-1, keepdim=True).clamp_min(1e-12) + out[inside] = v[inside] + return out + + +class OrientationMap(AutoSerialize): + """Match crystal orientations to Bragg peaks at every probe position. + + Workflow:: + + om = OrientationMap.from_vectors(peaks, crystal, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0) + om.match_orientations(num_matches=1) + om.plot_orientation() + + The object is both the engine and the result: after + `match_orientations()`, `quats` holds (R, C, M, 4) orientation + quaternions, `corr` the correlation scores, and `mirror` the inversion + flags. + """ + + _token = object() + + def __init__( + self, + peaks: Vector, + crystal: Crystal, + energy_ev: float, + _token: object | None = None, + ): + """Private constructor; use :meth:`from_vectors`. + + Parameters + ---------- + peaks : Vector + Calibrated Bragg peaks over the scan. + crystal : Crystal + Candidate crystal with structure factors already calculated. + energy_ev : float + Beam energy in eV. + _token : object + Guard against direct construction. + + Raises + ------ + RuntimeError + If called without the class token. + """ + if _token is not self._token: + raise RuntimeError("Use OrientationMap.from_vectors() to construct.") + self.peaks = peaks + self.crystal = crystal + self.energy_ev = float(energy_ev) + self.wavelength = electron_wavelength_angstrom(energy_ev) + # processing hyperparameters of every stage, recorded as they run; + # later stages inherit from these when an argument is left as None + self.metadata: dict = { + "energy_ev": self.energy_ev, + "peaks": dict(getattr(peaks, "metadata", {}) or {}), + } + + # plan state + self.zone_axes: torch.Tensor | None = None + self.zone_quats: torch.Tensor | None = None + self.plan_fft: torch.Tensor | None = None + self.shell_radii: torch.Tensor | None = None + + # results + self.quats: torch.Tensor | None = None + self.corr: torch.Tensor | None = None + self.corr_residual: torch.Tensor | None = None + self.score: torch.Tensor | None = None + self.corr_second: torch.Tensor | None = None + self.reliability: torch.Tensor | None = None + self.mirror: torch.Tensor | None = None + # positions carrying a result; a subset after a staged test run + self.computed: torch.Tensor | None = None + + @classmethod + def from_vectors( + cls, + peaks: Vector, + crystal: Crystal, + energy_ev: float = 300e3, + precession_deg: float = 0.0, + semiconv_mrad: float = 0.0, + foil_normal=None, + ) -> "OrientationMap": + """Create from detected Bragg peaks. + + Parameters + ---------- + peaks : Vector + Ragged peak table over scan positions with fields including + ('qx', 'qy', 'intensity') in calibrated 1/Angstrom units. + crystal : Crystal + Candidate crystal with structure factors already calculated. + precession_deg, semiconv_mrad : float + Precession semi-angle and convergence semiangle of the + experiment, recorded for the dynamical refinements (which + average the intensities over them) and inherited by them. + energy_ev : float, default=300e3 + Beam energy in eV. + foil_normal : sequence of float, optional + Plate normal of the specimen as a direction in the crystal, [uvw] + or [UVTW], e.g. (0, 0, 0, 1) for a 2D material lying in its basal + plane. Reflections are then rods along it, which moves the + simulated spots of a tilted flake to where the rods meet the + Ewald sphere (see :meth:`Crystal.generate_pattern`); the library, + the refinement and every simulated pattern use it. None + (default) is the usual geometry. + """ + if crystal.g_vec is None: + raise RuntimeError("Run crystal.calculate_structure_factors() first.") + om = cls(peaks, crystal, energy_ev, _token=cls._token) + om.metadata["precession_deg"] = float(precession_deg) + om.metadata["semiconv_mrad"] = float(semiconv_mrad) + om.metadata["foil_normal"] = ( + None if foil_normal is None else [float(v) for v in np.ravel(foil_normal)] + ) + return om + + def _foil_normal_crystal(self) -> torch.Tensor | None: + """Unit plate normal in the crystal frame, or None for the usual + geometry (normal along the beam).""" + fn = self.metadata.get("foil_normal") + return None if fn is None else self.crystal.direction_vector(fn) + + # ------------------------------------------------------------------ + # orientation plan + # ------------------------------------------------------------------ + + def build_plan( + self, + angle_step_zone_axis_deg: float = 1.0, + angle_step_in_plane_deg: float = 5.0, + zone_axis_range="auto", + fiber_axis=None, + fiber_angle_deg: float = 0.0, + corr_kernel_size: float = PAIR_DISTANCE, + sigma_excitation: float = SIGMA_EXCITATION, + power_radial: float = 1.0, + power_intensity: float = POWER_INTENSITY, + power_intensity_experiment: float | None = None, + tol_shell_distance: float = 0.01, + detector_q_max: float | tuple[float, float] | str | None = "auto", + device: str | torch.device = "cpu", + verbose: bool = True, + progress_bar: bool = True, + ) -> "OrientationMap": + """Build the polar correlation library over the fundamental wedge. + + Parameters + ---------- + angle_step_zone_axis_deg : float, default=1.0 + Angular step between sampled zone axes. The zone axis is the + coordinate the correlation search cannot refine continuously + (only by the neighbor-weighted centroid), so it is sampled + finely; the wedge sampling is isotropic at this step. + angle_step_in_plane_deg : float, default=5.0 + Angular step of the in-plane (gamma) axis; the number of gamma + samples is round(360 / step). The in-plane angle is refined + continuously (parabolic sub-bin interpolation, then least + squares on the paired peaks in refine_orientations), so a + coarse step costs little accuracy and keeps the library small. + zone_axis_range : {"auto", "full", "fiber"} | array-like, default="auto" + Which zone axes the library covers. + + - "auto": the fundamental wedge of the matching point group + (the pseudo-symmetry group when one was detected), falling + back to the hemisphere for triclinic and monoclinic cells. + This is the right choice for an unknown texture. + - "full": the whole hemisphere, whatever the symmetry. Use when + the symmetry the cell reports is not the symmetry of its + diffraction, so the wedge would fold distinct orientations + onto each other. + - "fiber": a cap of half angle `fiber_angle_deg` about + `fiber_axis`, for a known texture (a 2D material or a + textured film). `fiber_angle_deg=0` samples the fiber axis + alone, so the match is over the in-plane angle only, which + makes the library tiny and the match far more robust. + - an array of 2 or 3 lattice directions [uvw] (or [uvtw] for a + hexagonal cell): the spherical triangle they span. With two + rows the wedge runs from [001] through both. + + The in-plane angle is always searched over the full 360 degrees; + the correlation is circular in it, so restricting it saves + nothing. + fiber_axis : array-like | None + Lattice direction [uvw] (or [uvtw]) of the fiber axis, required + by zone_axis_range="fiber". + fiber_angle_deg : float, default=0.0 + Half angle of the fiber cap, degrees. + corr_kernel_size : float, default=0.05 + Correlation kernel size delta (1/Angstroms): azimuthal extent of + each reference peak and radial tolerance for shell assignment. + sigma_excitation : float, default=0.04 + Excitation error envelope of the library (1/Angstroms). Keep this + about 2x the physical excitation tolerance: orientations halfway + between sampled zones shift s_g by ~ (step/2) * k, and a wider + envelope keeps their library intensities from collapsing. With a + precession angle or convergence semiangle recorded by + from_vectors, the envelope of every library reflection is its + exact average over that illumination + (quantem.diffraction.illumination). + power_radial, power_intensity : float + Weighting prefactor q^power_radial * |V_g|^power_intensity for + library peaks. power_intensity=0 matches on positions only + (best for strongly dynamical data). + power_intensity_experiment : float | None + Exponent applied to the *measured* peak intensities; defaults to + `power_intensity`. Lower it than the library exponent when the + measured intensities are less trustworthy than the simulated + ones (saturation, a beam stop, strong dynamical transfer). + tol_shell_distance : float, default=0.01 + Reciprocal lattice radii closer than this merge into one shell. + detector_q_max : float | tuple | "auto" | None, default="auto" + Half-width of the square detector (1/Angstroms), scalar or + (row_max, col_max). Library reflections beyond the detector edge + cannot be measured, and which ones fall off depends on the + in-plane rotation; the correlation is normalized by the masked + template norm at every in-plane angle, so orientations with + strong reflections outside the detector are not penalized. + "auto" measures the detector footprint from the peaks themselves + (largest |q| along each detector axis, undoing any + detector-to-scan rotation recorded on the peaks). None disables + the correction. + device : str | torch.device, default="cpu" + Device for the library and the correlation compute. On Apple + silicon 'mps' runs the correlation in float32 (about 1.5x faster + than the CPU); the refinements that follow stay on the CPU. + verbose : bool, default=True + Print the plan size and the group used for matching. The crystal + prints its full symmetry, pseudo-symmetry included, when built. + progress_bar : bool, default=True + Show a progress bar while the library is deposited, which is + most of the time: ~20 s per crystal at k_max = 2 over a trigonal + wedge, several times that over a full hemisphere. + """ + crystal = self.crystal + self.device = torch.device(device) + # MPS has no float64: the correlation runs in float32 there (the + # cosine similarities are insensitive to it); results are returned + # in float64 either way + self.dtype = torch.float32 if self.device.type == "mps" else torch.float64 + self.cdtype = torch.complex64 if self.dtype == torch.float32 else torch.complex128 + self.corr_kernel_size = float(corr_kernel_size) + self.sigma_excitation = float(sigma_excitation) + self.power_radial = float(power_radial) + self.power_intensity = float(power_intensity) + self.power_intensity_experiment = float( + power_intensity if power_intensity_experiment is None else power_intensity_experiment + ) + + # zone axis sampling: the symmetry wedge, the hemisphere, a fiber + # cap, or an explicit spherical triangle of lattice directions + za = self._sample_zone_axes( + zone_axis_range, fiber_axis, fiber_angle_deg, angle_step_zone_axis_deg + ) + self.zone_axis_range = zone_axis_range + self.zone_axes = za + self.zone_quats = quat_from_zone_axis(za) + self.zone_step_deg = float(angle_step_zone_axis_deg) + + # symmetry-complete neighbor sets for sub-grid zone refinement: a + # zone on the wedge boundary only has in-wedge grid neighbors on one + # side, and a one-sided correlation centroid would drag it inward; + # the symmetry images of the grid across the boundary restore the + # missing side. + from quantem.diffraction.rotations import quat_to_matrix + + Rs = quat_to_matrix(crystal.sym_quats_matching) + images = torch.einsum("sij,zj->szi", Rs, za) + images = torch.cat([images, -images], dim=0).reshape(-1, 3) # (S2*Z, 3) + img_zone = torch.arange(za.shape[0]).repeat(2 * crystal.sym_quats_matching.shape[0]) + # deduplicate coincident image positions (keep one per position/zone) + key = torch.cat( + [torch.round(images / 1e-6) * 1e-6, img_zone[:, None].to(images.dtype)], + dim=1, + ) + _, first = np.unique(key.numpy(), axis=0, return_index=True) + images = images[torch.as_tensor(np.sort(first))] + img_zone = img_zone[torch.as_tensor(np.sort(first))] + + cos_lim = np.cos(np.deg2rad(1.6 * self.zone_step_deg)) + nbr_idx_list, nbr_pos_list = [], [] + dots = images @ za.T # (M, Z) + for i in range(za.shape[0]): + sel = torch.nonzero(dots[:, i] > cos_lim).squeeze(1) + nbr_idx_list.append(img_zone[sel]) + nbr_pos_list.append(images[sel]) + K = max(len(v) for v in nbr_idx_list) + Z = za.shape[0] + self.zone_nbr_idx = torch.zeros((Z, K), dtype=torch.long) + self.zone_nbr_pos = torch.zeros((Z, K, 3), dtype=torch.float64) + self.zone_nbr_valid = torch.zeros((Z, K), dtype=torch.bool) + for i, (idx, pos) in enumerate(zip(nbr_idx_list, nbr_pos_list)): + k = len(idx) + self.zone_nbr_idx[i, :k] = idx + self.zone_nbr_pos[i, :k] = pos + self.zone_nbr_valid[i, :k] = True + + # radial shells from unique reciprocal lattice vector lengths + g_len = crystal.g_len + radii = torch.unique(torch.round(g_len / tol_shell_distance) * tol_shell_distance) + self.shell_radii = radii + self.num_gamma = int(round(360 / angle_step_in_plane_deg)) + self.gamma = torch.linspace(0, 2 * np.pi, self.num_gamma + 1, dtype=torch.float64)[:-1] + + plan = self._build_reference(self.zone_quats, progress_bar=progress_bar) + # store conj(fft) along gamma so matching is a single complex matmul + self.plan_fft = torch.conj(torch.fft.fft(plan, dim=-1)).to(self.cdtype).to(self.device) + + # square-detector aperture correction: the masked template norm at + # every in-plane shift is the circular correlation of the squared + # plan with the polar detector mask. The mask lives in the DETECTOR + # frame: any detector-to-scan rotation recorded on the peaks rotates + # the square aperture in the calibrated (qx, qy) frame. + rot_deg = float(self.peaks.metadata.get("rotation_ccw_deg", 0.0) or 0.0) + if isinstance(detector_q_max, str) and detector_q_max == "auto": + flat = self.peaks.select_fields("qx", "qy", "intensity").numpy().astype(np.float64) + if flat.shape[0] == 0: + detector_q_max = None + else: + th_b = np.deg2rad(-rot_deg) + rb = np.array([[np.cos(th_b), -np.sin(th_b)], [np.sin(th_b), np.cos(th_b)]]) + det_rc = flat[:, :2] @ rb.T + detector_q_max = ( + float(np.abs(det_rc[:, 0]).max()) + self.corr_kernel_size, + float(np.abs(det_rc[:, 1]).max()) + self.corr_kernel_size, + ) + if detector_q_max is not None: + if np.isscalar(detector_q_max): + qx_max = qy_max = float(detector_q_max) + else: + qx_max, qy_max = (float(v) for v in detector_q_max) + r = self.shell_radii[:, None] + g = self.gamma[None, :] - np.deg2rad(rot_deg) + mask = ( + (torch.abs(r * torch.cos(g)) <= qx_max) & (torch.abs(r * torch.sin(g)) <= qy_max) + ).to(torch.float64) + self.detector_mask = mask # (S, G) + plan_sq_fft = torch.conj(torch.fft.fft(plan**2, dim=-1)) + mask_fft = torch.fft.fft(mask, dim=-1) + # norm^2 per (zone, shift), direct and mirrored channels + n2 = torch.fft.ifft( + torch.einsum("zsg,sg->zg", plan_sq_fft, mask_fft), dim=-1 + ).real.clamp_min(0) + n2_m = torch.fft.ifft( + torch.einsum("zsg,sg->zg", plan_sq_fft, torch.conj(mask_fft)), dim=-1 + ).real.clamp_min(0) + # fraction of template weight on the detector; used to suppress + # zones that are mostly unmeasurable at a given rotation + full = (plan**2).sum(dim=(1, 2))[:, None].clamp_min(1e-12) + self.plan_norm_shift = ( + torch.stack([torch.sqrt(n2), torch.sqrt(n2_m)]).to(self.dtype).to(self.device) + ) # (2, Z, G) + self.plan_frac_shift = ( + torch.stack([n2 / full, n2_m / full]).to(self.dtype).to(self.device) + ) + else: + self.detector_mask = None + self.plan_norm_shift = None + self.plan_frac_shift = None + self.metadata["plan"] = dict( + excitation_model="gaussian", + precession_deg=float(self.metadata.get("precession_deg", 0.0) or 0.0), + semiconv_mrad=float(self.metadata.get("semiconv_mrad", 0.0) or 0.0), + angle_step_zone_axis_deg=float(angle_step_zone_axis_deg), + angle_step_in_plane_deg=float(angle_step_in_plane_deg), + zone_axis_range=zone_axis_range + if isinstance(zone_axis_range, str) + else np.asarray(zone_axis_range).tolist(), + fiber_axis=None if fiber_axis is None else np.asarray(fiber_axis).tolist(), + fiber_angle_deg=float(fiber_angle_deg), + corr_kernel_size=self.corr_kernel_size, + pair_distance=self.corr_kernel_size, + sigma_excitation=self.sigma_excitation, + power_radial=self.power_radial, + power_intensity=self.power_intensity, + power_intensity_experiment=self.power_intensity_experiment, + tol_shell_distance=float(tol_shell_distance), + detector_q_max=None + if detector_q_max is None + else tuple(np.atleast_1d(detector_q_max).tolist()), + ) + if verbose: + # the crystal printed its own symmetry when it was built; the plan + # adds only what it sampled + print( + "%s: orientation plan %d zone axes x %d in-plane angles, " + "%d radial shells, matching %s" + % ( + crystal.name, + self.zone_axes.shape[0], + self.gamma.shape[0], + self.shell_radii.shape[0], + crystal.pointgroup_matching, + ) + ) + return self + + def _sample_zone_axes( + self, + zone_axis_range, + fiber_axis, + fiber_angle_deg: float, + step_deg: float, + ) -> torch.Tensor: + """Zone-axis sampling requested by build_plan, (Z, 3) Cartesian.""" + from quantem.diffraction.rotations import sample_zone_axis_cap + + crystal = self.crystal + + def cartesian(uvw) -> torch.Tensor: + v = np.asarray(uvw, dtype=float).reshape(-1) + if v.shape[0] == 4: # Miller-Bravais [uvtw] + from quantem.diffraction.crystal import miller_bravais_to_miller + + v = np.asarray(miller_bravais_to_miller(v)).reshape(-1) + if v.shape[0] != 3: + raise ValueError("a direction must have 3 indices [uvw] or 4 [uvtw]") + d = torch.as_tensor(v, dtype=torch.float64) @ crystal.lat_real + return d / torch.linalg.norm(d).clamp_min(1e-12) + + if isinstance(zone_axis_range, str): + mode = zone_axis_range.lower() + if mode == "fiber": + if fiber_axis is None: + raise ValueError('zone_axis_range="fiber" needs a fiber_axis') + return sample_zone_axis_cap(cartesian(fiber_axis), fiber_angle_deg, step_deg) + if mode in ("full", "hemisphere"): + n_zones = int(np.ceil(2 * np.pi / np.deg2rad(step_deg) ** 2)) + return fibonacci_hemisphere(n_zones) + if mode != "auto": + raise ValueError( + 'zone_axis_range must be "auto", "full", "fiber", or an array of directions' + ) + msg = crystal.matching_symmetry_warning() + # a crystal built with verbose=True already said this in its summary + if msg is not None and not getattr(crystal, "_summary_shown", False): + warnings.warn(msg, stacklevel=3) + wedge = crystal.zone_axis_wedge() + if wedge is None: # triclinic / monoclinic: not a spherical triangle + n_zones = int(np.ceil(2 * np.pi / np.deg2rad(step_deg) ** 2)) + return fibonacci_hemisphere(n_zones) + za, _ = sample_zone_axes(wedge, step_deg) + return za + + rows = np.atleast_2d(np.asarray(zone_axis_range, dtype=float)) + dirs = [cartesian(r) for r in rows] + if len(dirs) == 2: + dirs = [cartesian([0, 0, 1]), *dirs] + if len(dirs) != 3: + raise ValueError("zone_axis_range as an array needs 2 or 3 directions") + za, _ = sample_zone_axes(torch.stack(dirs), step_deg) + return za + + def _deposit_polar( + self, + qr: torch.Tensor, + qphi: torch.Tensor, + amp: torch.Tensor, + out: torch.Tensor, + image: torch.Tensor | None = None, + progress: str | None = None, + ) -> torch.Tensor: + """Deposit peaks into polar images with the shared correlation kernel. + + Every peak spreads as a Gaussian of width delta in both the radial + direction (across shells) and arc length (along gamma). The library + and the experimental patterns use this same kernel, so the normalized + correlation of a pattern with itself is exactly 1. + + Parameters + ---------- + qr, qphi, amp : torch.Tensor + Peak radii, azimuths, amplitudes, flat (K,). + out : torch.Tensor + (S, G) accumulator, or (N, S, G) when `image` is given; modified + in place. + image : torch.Tensor | None + (K,) image index of every peak, so a whole batch of patterns + (or a whole library) is deposited in one call. + progress : str | None + Description for a progress bar over the chunks; None shows none. + """ + radii = self.shell_radii.to(qr.dtype) + delta = self.corr_kernel_size + S = radii.shape[0] + + dr = qr[:, None] - radii[None, :] # (K, S) + k_all, s_all = torch.nonzero(dr.abs() < 3 * delta, as_tuple=True) + if k_all.numel() == 0: + return out + gamma = self.gamma.to(qr.dtype) + flat = out.view(-1, out.shape[-1]) + # chunked so the (entries, G) weight array stays a few tens of MB + chunk = max(1, 4_000_000 // gamma.shape[0]) + starts = range(0, k_all.numel(), chunk) + if progress is not None: + starts = tqdm(starts, desc=progress) + for c0 in starts: + k_idx = k_all[c0 : c0 + chunk] + s_idx = s_all[c0 : c0 + chunk] + w_r = torch.exp(-(dr[k_idx, s_idx] ** 2) / (2 * delta**2)) * amp[k_idx] + dg = qphi[k_idx, None] - gamma[None, :] + dg = (dg + np.pi) % (2 * np.pi) - np.pi + arc = dg * qr[k_idx, None] + w = w_r[:, None] * torch.exp(-(arc**2) / (2 * delta**2)) + rows = s_idx if image is None else image[k_idx] * S + s_idx + flat.index_add_(0, rows, w) + return out + + def _build_reference( + self, zone_quats: torch.Tensor, progress_bar: bool = False + ) -> torch.Tensor: + """Polar reference library (Z, S, G) for the given zone-axis quats.""" + crystal = self.crystal + lam = self.wavelength + g = crystal.g_vec # (N, 3) + delta = self.corr_kernel_size + + gr = qrotate(zone_quats[:, None, :], g[None, :, :]) # (Z, N, 3) + gz = gr[..., 2] + g2 = (gr**2).sum(-1) + s_g = (2 * gz - lam * g2) / (2 - 2 * lam * gz) + prec = float(self.metadata.get("precession_deg", 0.0) or 0.0) + conv = float(self.metadata.get("semiconv_mrad", 0.0) or 0.0) + spot_xy = gr[..., :2] + n_c = self._foil_normal_crystal() + f_rod = None + if n_c is not None: + # plate geometry: excitation along the rod, spot where the rod + # meets the sphere; the normal turns with the crystal, so the + # template still rolls in gamma + from quantem.diffraction.illumination import relrod_factor + + n_lab = qrotate(zone_quats, n_c[None].expand(zone_quats.shape[0], 3)) # (Z, 3) + f_rod = relrod_factor(gr, n_lab[:, None, :], self.energy_ev, 0.0) + s_g = s_g * f_rod + spot_xy = spot_xy - s_g[..., None] * n_lab[:, None, :2] + if prec > 0 or conv > 0: + # excitation envelope averaged over the illumination + # (quantem.diffraction.illumination); the peak weighting below + # keeps the established correlation-library semantics + from quantem.diffraction.illumination import ( + excitation_amplitudes, + gaussian_envelope, + gaussian_envelope_ring_torch, + ) + + a_r, b_r = excitation_amplitudes(gr, self.energy_ev, prec, conv) + if f_rod is not None: + a_r, b_r = a_r * f_rod.abs(), b_r * f_rod.abs() + if conv <= 0: + amp = gaussian_envelope_ring_torch(s_g, a_r, self.sigma_excitation) + else: + amp = torch.as_tensor( + gaussian_envelope( + s_g.numpy(), a_r.numpy(), b_r.numpy(), self.sigma_excitation + ), + dtype=torch.float64, + ) + amp = amp * (s_g.abs() < a_r + b_r + delta * 4) + else: + amp = torch.exp(-(s_g**2) / (2 * self.sigma_excitation**2)) + amp = amp * (s_g.abs() < delta * 4) + + weight = ( + crystal.g_len**self.power_radial * crystal.struct_factors_int**self.power_intensity + ) + vals = amp * weight[None, :] # (Z, N) + qr = torch.hypot(spot_xy[..., 0], spot_xy[..., 1]) + qphi = torch.atan2(spot_xy[..., 1], spot_xy[..., 0]) + + Z = zone_quats.shape[0] + plan = torch.zeros((Z, self.shell_radii.shape[0], self.num_gamma), dtype=torch.float64) + z_idx, n_idx = torch.nonzero(vals > 1e-8, as_tuple=True) + self._deposit_polar( + qr[z_idx, n_idx], + qphi[z_idx, n_idx], + vals[z_idx, n_idx], + plan, + image=z_idx, + progress=f"orientation plan {crystal.name}" if progress_bar else None, + ) + + norm = torch.linalg.norm(plan.reshape(Z, -1), dim=1).clamp_min(1e-12) + return plan / norm[:, None, None] + + # ------------------------------------------------------------------ + # experimental polar images + # ------------------------------------------------------------------ + + def _grid_quats(self, flat_idx: torch.Tensor, Z: int, G: int, n_ch: int) -> torch.Tensor: + """Library orientations at flat (channel, zone, gamma) indices. + + No subpixel refinement: these are used only to test whether two + candidates are the same orientation, where the grid step is far + finer than the separation being tested. + + Parameters + ---------- + flat_idx : torch.Tensor + ``(..., )`` indices into the flattened ``(ch, Z, G)`` correlation. + + Returns + ------- + torch.Tensor + ``(..., 4)`` quaternions. + """ + idx = flat_idx.cpu() + ch = idx // (Z * G) + z = (idx // G) % Z + g = idx % G + q_zone = self.zone_quats[z] + gamma = self.gamma[g].clone() + is_mirror = ch == 1 + if n_ch > 1: + q_flip = torch.tensor([0.0, 1.0, 0.0, 0.0], dtype=torch.float64) + q_zone = torch.where(is_mirror[..., None], qmult(q_flip, q_zone), q_zone) + gamma = torch.where(is_mirror, -gamma - np.pi, gamma) + half = gamma / 2 + zeros = torch.zeros_like(half) + q_spin = torch.stack((torch.cos(half), zeros, zeros, torch.sin(half)), dim=-1) + return qmult(q_spin, q_zone) + + def _polar_images(self, arrays: list[np.ndarray], ix: list[int]) -> torch.Tensor: + """Sparse polar images (B, S, G) of a batch of measured patterns, + deposited in one call.""" + data = np.concatenate([a[:, ix] for a in arrays], axis=0) + image = torch.repeat_interleave( + torch.arange(len(arrays)), torch.tensor([a.shape[0] for a in arrays]) + ) + qx = torch.as_tensor(data[:, 0], dtype=torch.float64) + qy = torch.as_tensor(data[:, 1], dtype=torch.float64) + intensity = torch.as_tensor(data[:, 2], dtype=torch.float64) + qr = torch.hypot(qx, qy) + qphi = torch.atan2(qy, qx) + amp = intensity.clamp_min(0) ** self.power_intensity_experiment * qr**self.power_radial + out = torch.zeros( + (len(arrays), self.shell_radii.shape[0], self.num_gamma), dtype=torch.float64 + ) + return self._deposit_polar(qr, qphi, amp, out, image=image) + + # ------------------------------------------------------------------ + # matching + # ------------------------------------------------------------------ + + def match_orientations( + self, + num_matches: int = 1, + positions=None, + include_mirror: bool = True, + min_number_peaks: int | None = None, + min_angle_between_matches_deg: float = 15.0, + suppress_matched: float = 1.0, + top_k_matches: int = 256, + subpixel_gamma: bool = True, + subpixel_zone: bool = True, + min_detector_fraction: float = 0.3, + batch_size: int = 128, + progress_bar: bool = True, + ) -> "OrientationMap": + """Match probe positions against the orientation plan. + + Patterns are processed in batches: the polar images are stacked, and + the correlation over all zones and in-plane angles reduces to one + complex matrix product per gamma frequency plus a batched inverse FFT. + + Correlation scores are normalized to [0, 1]: the library slices are + unit vectors and the experimental polar image is divided by its own + norm, so `corr` is a cosine similarity comparable across patterns, + crystals, and datasets. After the best match, the highest correlation + among zone axes at least `min_angle_between_matches_deg` away is + stored in `corr_second`; `reliability = corr - corr_second` is the + primary confidence metric. + + Parameters + ---------- + num_matches : int, default=1 + Number of orientations to return per probe position; matches + after the first suppress zones within + `min_angle_between_matches_deg` of earlier matches. + positions : list[tuple[int, int]] | np.ndarray | None + Scan positions to match: a list of (row, col), or an (R, C) + boolean mask. None (default) matches the whole scan. Pass a + handful of positions to check the plan and these parameters + with `plot_pattern_matches` before committing to the full + scan; the positions carrying a result are recorded in + `computed`, which the refinements and the phase fit follow. + include_mirror : bool, default=True + Also correlate against the in-plane mirrored pattern, testing + inversion-related (opposite hemisphere) zone axes at no library + cost. Exact in the flat-Ewald / Friedel limit. + min_number_peaks : int | None + Skip positions with fewer detected peaks (including the direct + beam); defaults to MIN_NUMBER_PEAKS (5). + suppress_matched : float, default=1.0 + Fraction of each accepted match subtracted from the measured + polar image before the next match is sought, so later matches + fit the peaks earlier ones leave unexplained. 1.0 removes the + matched component exactly; 0 disables the deflation and leaves + `min_angle_between_matches_deg` as the only thing separating the + matches, which lets two of them index the same peaks. Only + meaningful with `num_matches` above 1. The deflation steers the + search only: `corr` scores every match against the pattern as + measured, so the matches stay comparable, and `corr_residual` + holds the score against the residual each one actually saw. + min_angle_between_matches_deg : float, default=15.0 + Minimum separation between matches, applied to the zone axis and + the in-plane angle together: a candidate is rejected only when it + is within this angle of an earlier match in both. Two grains + sharing a zone axis but rotated in plane past this angle are + therefore kept as separate matches. The same test picks the + second-best score used in `reliability`. + top_k_matches : int, default=256 + Number of highest-scoring library entries searched, in order, + for a candidate that passes the separation test, both for the + matches after the first and for `corr_second`. Most of the top + entries are symmetry copies or near neighbours of the best one, + so too small a value leaves later matches empty. + subpixel_gamma : bool, default=True + Parabolic sub-bin refinement of the in-plane angle. + subpixel_zone : bool, default=True + Sub-grid zone axis from a quadratic fit of the correlation over + the best zone and its grid neighbors (centroid fallback). + min_detector_fraction : float, default=0.3 + With a detector footprint in the plan (`detector_q_max`), library + orientations that put less than this fraction of their template + weight on the detector at a given in-plane angle score zero + there. Ignored when the plan has no detector correction. + batch_size : int, default=128 + Number of patterns correlated at once. + progress_bar : bool, default=True + Show a progress bar over the batches. + + Returns + ------- + OrientationMap + Self, with `quats` (R, C, M, 4), `corr` and `corr_residual` + (R, C, M), `corr_second` and `reliability` (R, C), `mirror` + (R, C, M) and `computed` (R, C) filled in. Positions not matched + keep the identity orientation and a correlation of zero. + + Raises + ------ + RuntimeError + If the plan has not been built, or no requested position has + `min_number_peaks` peaks. + """ + if self.plan_fft is None: + raise RuntimeError("Run build_plan() first.") + min_number_peaks = resolve(min_number_peaks, "min_number_peaks", default=MIN_NUMBER_PEAKS) + self.metadata["match"] = dict( + num_matches=int(num_matches), + positions=None if positions is None else "subset", + include_mirror=bool(include_mirror), + min_number_peaks=int(min_number_peaks), + min_angle_between_matches_deg=float(min_angle_between_matches_deg), + suppress_matched=float(suppress_matched), + top_k_matches=int(top_k_matches), + subpixel_gamma=bool(subpixel_gamma), + subpixel_zone=bool(subpixel_zone), + min_detector_fraction=float(min_detector_fraction), + ) + peaks = self.peaks + shape = peaks.shape + R, C = shape[0], shape[1] + M = num_matches + device = self.device + G = self.num_gamma + Z = self.zone_axes.shape[0] + + quats = torch.zeros((R, C, M, 4), dtype=torch.float64) + quats[..., 0] = 1.0 + corr_out = torch.zeros((R, C, M), dtype=torch.float64) + corr_res = torch.zeros((R, C, M), dtype=torch.float64) + corr_second = torch.zeros((R, C), dtype=torch.float64) + mirror_out = torch.zeros((R, C, M), dtype=torch.bool) + + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + + dtype = getattr(self, "dtype", torch.float64) + + plan_fft = self.plan_fft # (Z, S, G) complex + wanted = position_mask(positions, (R, C)) + valid_rc = [ + (rx, ry) + for rx, ry in np.ndindex(R, C) + if wanted[rx, ry] + and peaks[rx, ry].numpy().astype(np.float64).shape[0] >= min_number_peaks + ] + if not valid_rc: + raise RuntimeError( + "no requested scan position has at least min_number_peaks = %d detected peaks" + % min_number_peaks + ) + computed = torch.zeros((R, C), dtype=torch.bool) + for rx, ry in valid_rc: + computed[rx, ry] = True + self.computed = computed + batches = [valid_rc[i : i + batch_size] for i in range(0, len(valid_rc), batch_size)] + if progress_bar: + batches = tqdm(batches, desc=f"matching {self.crystal.name}") + + gamma_grid = self.gamma + for batch in batches: + im_stack = ( + self._polar_images( + [peaks[rx, ry].numpy().astype(np.float64) for rx, ry in batch], ix + ) + .to(dtype) + .to(device) + ) + B = im_stack.shape[0] + with warnings.catch_warnings(): + # torch's MPS FFT emits an internal out-tensor resize notice + warnings.simplefilter("ignore", UserWarning) + im_fft = torch.fft.fft(im_stack, dim=-1) # (B, S, G) + # frequency ramp used to roll a template to an in-plane angle + k_ramp = torch.fft.fftfreq(G, d=1.0 / G).to(im_fft.dtype).to(device) + # orientations accepted so far in this batch, one (B, 4) per match + q_prev: list[torch.Tensor] = [] + + for m in range(M): + norms = torch.linalg.norm( + torch.fft.ifft(im_fft, dim=-1).real.reshape(B, -1), dim=1 + ).clamp_min(1e-12) + # contract shells: (B, Z, G) per channel + cc = torch.einsum("zsg,bsg->bzg", plan_fft, im_fft) + channels = [cc] + if include_mirror: + channels.append(torch.einsum("zsg,bsg->bzg", plan_fft, torch.conj(im_fft))) + corr_raw = torch.fft.ifft(torch.stack(channels, dim=1), dim=-1).real + # normalize: library slices are unit vectors, so dividing by the + # experimental norm makes corr a cosine similarity in [0, 1] + corr = corr_raw / norms[:, None, None, None] + if self.plan_norm_shift is not None: + # square-detector correction: renormalize by the on-detector + # template norm at each in-plane shift, and suppress + # rotations where most of the template is unmeasurable + n_ch = corr.shape[1] + corr = corr / self.plan_norm_shift[None, :n_ch].clamp_min(1e-3) + corr = corr.masked_fill( + self.plan_frac_shift[None, :n_ch] < min_detector_fraction, 0.0 + ) + # corr: (B, ch, Z, G) + if m == 0: + # the deflation below changes the image every match, so + # keep the correlation against the pattern as measured: + # selection uses the residual, the reported score does not + corr_full = corr + n_ch = corr.shape[1] + flat = corr.reshape(B, -1) + if M > 1: + # Rank the candidates and walk down until one is a + # genuinely different orientation from every earlier + # match. The test is the full misorientation, reduced by + # crystal symmetry: neither the zone axis nor the in-plane + # angle alone can tell a symmetry copy (same orientation, + # different library entry) from two grains sharing a zone + # axis but rotated in plane, and those must be treated + # oppositely. + K = min(flat.shape[1], top_k_matches) + top_v, top_i = flat.topk(K, dim=1) + q_top = self._grid_quats(top_i, Z, G, n_ch) # (B, K, 4) + keep = torch.zeros((B, K), dtype=torch.bool) + if m == 0: + keep[:, 0] = True + else: + ok = torch.ones((B, K), dtype=torch.bool) + for q_mm in q_prev: # (B, 4) each + ang = misorientation_angle_deg( + q_mm[:, None, :].expand(-1, K, -1).reshape(-1, 4), + q_top.reshape(-1, 4), + self.crystal.sym_quats_matching, + ).reshape(B, K) + ok &= ang >= min_angle_between_matches_deg + ok &= torch.isfinite(top_v.cpu()) + first = torch.where( + ok.any(dim=1), ok.double().argmax(dim=1), torch.full((B,), -1) + ) + for b in range(B): + if first[b] >= 0: + keep[b, int(first[b])] = True + sel = torch.where( + keep.any(dim=1), + keep.double().argmax(dim=1), + torch.zeros(B, dtype=torch.long), + ).to(flat.device) + flat_idx = top_i.gather(1, sel[:, None]).squeeze(1) + invalid = ~keep.any(dim=1).to(flat.device) + else: + flat_idx = flat.argmax(dim=1) + invalid = torch.zeros(B, dtype=torch.bool, device=flat.device) + ch_i = flat_idx // (Z * G) + z_i = (flat_idx // G) % Z + g_i = flat_idx % G + c_val = flat.gather(1, flat_idx[:, None]).squeeze(1) + c_val = c_val.masked_fill(invalid, -torch.inf) + + gamma = gamma_grid[g_i.cpu()].clone() + if subpixel_gamma: + b_ar = torch.arange(B, device=corr.device) + c1 = c_val + c0 = corr[b_ar, ch_i, z_i, (g_i - 1) % G] + c2 = corr[b_ar, ch_i, z_i, (g_i + 1) % G] + denom = 4 * c1 - 2 * c0 - 2 * c2 + dg = torch.where( + denom.abs() > 1e-12, + (c2 - c0) / denom, + torch.zeros_like(denom), + ) * (2 * np.pi / G) + gamma = gamma + dg.cpu().double() + + is_mirror = ch_i.cpu() == 1 + q_zone = self.zone_quats[z_i.cpu()] + + if subpixel_zone: + # sub-grid zone axis: correlation-weighted centroid over + # the symmetry-complete neighborhood of the best zone + # (see build_plan; images across the wedge boundary keep + # the centroid unbiased for boundary zones) + b_ar = torch.arange(B, device=corr.device) + corr_z = corr[b_ar, ch_i].amax(dim=-1).cpu().double() # (B, Z) + zi_cpu = z_i.cpu() + n_idx = self.zone_nbr_idx[zi_cpu] # (B, K) + n_pos = self.zone_nbr_pos[zi_cpu] # (B, K, 3) + n_ok = self.zone_nbr_valid[zi_cpu] # (B, K) + c_n = corr_z.gather(1, n_idx) # (B, K) + za_old = self.zone_axes[z_i.cpu()] + za_ref = _zone_peak_parabolic( + za_old, n_pos, c_n, n_ok, np.deg2rad(self.zone_step_deg) + ) + axis = torch.cross(za_ref, za_old, dim=-1) + sin_t = torch.linalg.norm(axis, dim=-1) + ang_t = torch.atan2(sin_t, (za_ref * za_old).sum(-1)) + ok_t = sin_t > 1e-12 + dq = torch.zeros((B, 4), dtype=torch.float64) + dq[:, 0] = 1.0 + if bool(ok_t.any()): + dq[ok_t] = quat_from_axis_angle( + axis[ok_t] / sin_t[ok_t, None], ang_t[ok_t] + ) + # rotate za_ref -> za_old in the crystal frame: R' = R S + q_zone = qmult(q_zone, dq) + + q_flip = torch.tensor([0.0, 1.0, 0.0, 0.0], dtype=torch.float64) + q_zone = torch.where(is_mirror[:, None], qmult(q_flip, q_zone), q_zone) + gamma = torch.where(is_mirror, -gamma - np.pi, gamma) + half = gamma / 2 + zeros = torch.zeros_like(half) + q_spin = torch.stack((torch.cos(half), zeros, zeros, torch.sin(half)), dim=-1) + q = qmult(q_spin, q_zone) + + # score every match against the pattern as measured, so the + # matches are comparable with each other and across positions + c_report = corr_full.reshape(B, -1).gather(1, flat_idx[:, None]).squeeze(1) + c_report = c_report.masked_fill(invalid, -torch.inf) + for b, (rx, ry) in enumerate(batch): + if torch.isfinite(c_val[b]): + quats[rx, ry, m] = q[b] + corr_out[rx, ry, m] = c_report[b].cpu().double() + corr_res[rx, ry, m] = c_val[b].cpu().double() + mirror_out[rx, ry, m] = bool(is_mirror[b]) + + if m == 0: + # second-best score at an orientation genuinely different + # from the best one -> reliability = corr - corr_second. + # Same misorientation test, so a symmetry copy of the + # winner never counts as the runner-up, while a real + # in-plane degeneracy does. + K2 = min(corr.reshape(B, -1).shape[1], top_k_matches) + tv, ti = corr.reshape(B, -1).topk(K2, dim=1) + q_t2 = self._grid_quats(ti, Z, G, n_ch) + q_best = self._grid_quats(flat_idx[:, None], Z, G, n_ch)[:, 0] + ang2 = misorientation_angle_deg( + q_best[:, None, :].expand(-1, K2, -1).reshape(-1, 4), + q_t2.reshape(-1, 4), + self.crystal.sym_quats_matching, + ).reshape(B, K2) + far2 = (ang2 >= min_angle_between_matches_deg) & torch.isfinite(tv.cpu()) + c2 = torch.where( + far2.any(dim=1), + tv.cpu().double().masked_fill(~far2, -torch.inf).amax(dim=1), + torch.full((B,), -torch.inf, dtype=torch.float64), + ) + for b, (rx, ry) in enumerate(batch): + if torch.isfinite(c2[b]): + corr_second[rx, ry] = c2[b] + + if M > 1: + q_prev.append(q.clone()) + + if M > 1 and m < M - 1 and suppress_matched > 0: + # Deflate the matched template out of the measured polar + # image, so the next match sees only what this one leaves + # unexplained. Without this, the exclusion ball keeps the + # next zone axis far away in orientation but nothing stops + # it from being fitted to the same peaks -- two grains in + # one probe then index as one, and the second match is a + # different view of the first. + # + # The templates are unit vectors, so the amount of this + # template present in the image is its raw inner product, + # read off the un-normalized correlation at the winning + # (zone, in-plane angle). Subtracting that multiple is the + # matching-pursuit step and removes it exactly. + b_ar = torch.arange(B, device=device) + alpha = suppress_matched * corr_raw[b_ar, ch_i, z_i, g_i].clamp_min(0).to( + im_fft.dtype + ) + # roll the template to the matched in-plane angle: a shift + # of g samples is a linear phase on its transform + phase = torch.exp( + -2j * np.pi * k_ramp[None, :] * g_i[:, None].to(k_ramp.dtype) / G + ) + t_fft = torch.conj(plan_fft[z_i]) * phase[:, None, :] # (B, S, G) + if include_mirror: + # the mirrored template is gamma -> -gamma, a conjugate + # in the transform, and its own roll + t_mir = torch.conj(t_fft) + t_fft = torch.where((ch_i == 1)[:, None, None].to(device), t_mir, t_fft) + im_fft = im_fft - alpha[:, None, None] * t_fft + # measured intensity is non-negative; keep it that way + im_real = torch.fft.ifft(im_fft, dim=-1).real.clamp_min(0) + with warnings.catch_warnings(): + warnings.simplefilter("ignore", UserWarning) + im_fft = torch.fft.fft(im_real.to(dtype), dim=-1) + + self.quats = quats + self.corr = corr_out + self.corr_residual = corr_res + self.corr_second = corr_second + self.reliability = corr_out[..., 0] - corr_second + self.mirror = mirror_out + return self + + # ------------------------------------------------------------------ + # sub-grid refinement + # ------------------------------------------------------------------ + + def smooth_orientations( + self, + match: int = 0, + sigma_px: float = 1.0, + sigma_deg: float = 1.0, + max_angle_deg: float = 5.0, + positions=None, + ) -> "OrientationMap": + """Average each orientation with its neighbours, keeping boundaries sharp. + + A bilateral filter on the orientation field: every position is + replaced by the weighted mean of the orientations around it, with + + w = exp(-r^2 / 2 sigma_px^2) * exp(-theta^2 / 2 sigma_deg^2) + + for a neighbour r probe positions away whose orientation differs by + theta, and with neighbours beyond `max_angle_deg` excluded outright. + The angular term is what keeps this from blurring across a grain + boundary or between two variants: those neighbours are tens of + degrees away and carry no weight. + + The point is the noise budget. Neighbouring probe positions inside a + grain measure the same orientation, so their scatter is measurement + error and averaging it down costs only spatial resolution, at the + scale of sigma_px probe steps. Running this before + `refine_orientations` starts the refinement from a cleaner field; + running it after smooths what the refinement leaves. + + This does not repair the ambiguities that make an orientation map + jump by tens of degrees, such as two variants with the same + zero-layer pattern: those differ by far more than `max_angle_deg` + and are excluded by design. Fold the in-plane angle by + `Crystal.projected_rotation_order` for those. + + Parameters + ---------- + match : int, default=0 + Which match index to smooth. + sigma_px : float, default=1.0 + Spatial width of the kernel in probe positions. The window is + three sigma wide. + sigma_deg : float, default=1.0 + Angular width: a neighbour misoriented by this much is weighted + down by 1/sqrt(e). + max_angle_deg : float, default=5.0 + Neighbours beyond this misorientation are excluded. + positions : list[tuple[int, int]] | np.ndarray | None + Positions to smooth; defaults to those carrying a match. + """ + assert self.quats is not None, "run match_orientations() first" + R, C = self.quats.shape[:2] + active = position_mask(positions, (R, C)) + if self.computed is not None: + active = active & self.computed + active = active & (self.corr[..., match] > 0) + self.quats[..., match, :] = smooth_quaternions( + self.quats[..., match, :], + active, + self.crystal.sym_quats_matching, + sigma_px=sigma_px, + sigma_deg=sigma_deg, + max_angle_deg=max_angle_deg, + ) + self.metadata["smooth"] = dict( + match=int(match), + sigma_px=float(sigma_px), + sigma_deg=float(sigma_deg), + max_angle_deg=float(max_angle_deg), + ) + return self + + def smoothed_quats(self, match: int = 0, **kwargs) -> torch.Tensor: + """Smoothed copy of the orientations, leaving the stored ones alone. + + Parameters + ---------- + match : int, default=0 + Which match index to smooth. + **kwargs + `sigma_px`, `sigma_deg` and `max_angle_deg` of + :func:`smooth_quaternions`. + + Returns + ------- + torch.Tensor + (R, C, 4) smoothed quaternions. + """ + assert self.quats is not None + R, C = self.quats.shape[:2] + active = torch.ones((R, C), dtype=torch.bool) + if self.computed is not None: + active = active & self.computed + active = active & (self.corr[..., match] > 0) + return smooth_quaternions( + self.quats[..., match, :], active, self.crystal.sym_quats_matching, **kwargs + ) + + def refine_orientations( + self, + num_iterations: int = 5, + positions=None, + pair_distance: float | None = None, + sigma_excitation: float | None = None, + min_pairs: int | None = None, + refine_tilt: bool = False, + refine_zone: bool = True, + zone_search_deg: float = 1.5, + sigma_envelope: float | None = None, + zone_max_total_deg: float | None = None, + power_intensity: float | None = None, + batched: bool = True, + neighbor_rescue: bool = True, + rescue_threshold_deg: float = 2.0, + rescue_passes: int = 3, + score_tol: float = 0.002, + consensus_tol: float = 0.01, + progress_bar: bool = True, + ) -> "OrientationMap": + """Refine matched orientations by least squares on paired peak positions. + + For each probe position and match, the simulated pattern is paired to + the measured peaks (nearest neighbor within `pair_distance`), and the + small rotation minimizing the weighted in-plane residuals is solved in + closed form and applied; repeated for `num_iterations` rounds with + re-pairing. This removes the in-plane quantization of the orientation + plan (typically to well below 0.1 degrees). + + By default only the in-plane rotation is refined. Zero-layer peak + positions carry almost no information about out-of-plane tilt (the + tilt terms in the position residual are proportional to g_z, which is + near zero for excited reflections), so fitting the full rotation from + positions is ill-conditioned and amplifies detection noise into + spurious tilts. Tilt is constrained by the diffracted *intensities* + (which reflections are excited) and belongs to the dynamical + refinement pass. + + Parameters + ---------- + num_iterations : int, default=5 + Pairing + rotation solve rounds. + positions : list[tuple[int, int]] | np.ndarray | None + Scan positions to refine, as for `match_orientations`. None + (default) refines every position that carries a match, so a + staged test run on a few positions is refined without repeating + the position list. + pair_distance : float | None + Maximum pairing distance (1/Angstroms); defaults to the plan's + corr_kernel_size. + sigma_excitation : float | None + Excitation error envelope used for simulation; defaults to the + plan's value. + min_pairs : int | None + Skip positions with fewer paired peaks; defaults to MIN_PAIRS (4). + refine_tilt : bool, default=False + Also solve the two tilt components from peak positions. Only + meaningful for noise-free simulated data. + refine_zone : bool, default=True + Refine the zone-axis tilt from the intensity envelope (the Laue + circle): the tilt that concentrates the measured intensity on + the Ewald sphere, searched over +/- zone_search_deg with + parabolic sub-stepping. Removes the zone-axis quantization of + the orientation plan. + zone_search_deg : float, default=1.5 + Half-range of the envelope tilt search, in degrees. + zone_max_total_deg : float | None + Trust region: cap on the cumulative envelope tilt applied to + each orientation, relative to its matched start. The coarse + match is grid-accurate to about half the zone-axis step, so tilt + corrections beyond that scale are noise walking the orientation + out of its basin. Defaults to 0.75 * the plan's zone step. + power_intensity : float | None + Power applied to the measured and predicted intensities in the + tilt envelope fit, inherited from the plan (0.25 by default). + Linear intensities let the strongest reflections dominate and, + on dynamical data, drive the fit to the edge of the search range. + sigma_envelope : float | None + Excitation-error width of the envelope objective; defaults to + half the plan's sigma_excitation (the plan value is widened for + grid robustness). + batched : bool, default=True + Refine positions in vectorized chunks instead of one at a time, + which is several times faster. `refine_tilt=True` disables it, + since only the per-position path solves the tilt from positions. + neighbor_rescue : bool, default=True + Retry every position that disagrees with a matched neighbour by + more than `rescue_threshold_deg`, from every distinct candidate + around it: all matches of the eight neighbours, this position's + own other matches, and its Friedel twin (the orientation rotated + 180 degrees about the beam). The one that best explains the + measured peaks is kept, judged by the same correlation matching + maximizes (see below); among candidates within `consensus_tol` + of the best, the one most neighbours agree with. Repairs wrong + local optima, near-degenerate variants, and probe positions + straddling two grains. + rescue_threshold_deg : float, default=2.0 + Misorientation to a neighbour that triggers a retry. + rescue_passes : int, default=3 + Rescue passes; each after the first revisits only positions next + to a change, since a corrected neighbour can offer a better + candidate. + score_tol : float, default=0.002 + Correlation margin. A refinement that moves an orientation by + more than `rescue_threshold_deg` is undone where it lowers the + correlation below the library match's by more than this, and a + rescue candidate replaces the current orientation only when it + beats it by more than this. + consensus_tol : float, default=0.01 + Correlation within which two candidates count as equally good. + Sparse patterns often cannot tell a few orientations apart: + pseudo-symmetric variants whose distinguishing reflections were + not recorded, and always the Friedel twin, as kinematic spot + positions are centrosymmetric and only the Ewald curvature + separates the two. Such ties are broken by agreement with the + eight neighbours; the Friedel twin is adopted only this way, + never on its score alone. 0 judges every position by its own + pattern alone. + progress_bar : bool, default=True + Show one progress bar covering refinement and neighbour rescue. + + Returns + ------- + OrientationMap + Self, with `quats` refined in place and `score` (R, C) holding + the correlation of match 0 with the measured peaks. + + Notes + ----- + Every orientation is judged by the correlation it gives with the + measured peaks -- the cosine similarity of the two patterns built + from the library's Gaussian pairing kernel -- and the result is + stored in :attr:`score`. Refinement itself works on paired peak + positions, a different objective; on sparse or ambiguous patterns it + can move an orientation downhill, and scoring every candidate the + same way is what keeps the stages consistent. The counts of reverted + refinements and rescued positions are in ``metadata['refine']``. + """ + assert self.quats is not None + plan_md = self.metadata.get("plan") + delta = resolve(pair_distance, "pair_distance", plan_md, default=self.corr_kernel_size) + sigma = resolve( + sigma_excitation, "sigma_excitation", plan_md, default=self.sigma_excitation + ) + min_pairs = resolve(min_pairs, "min_pairs", default=MIN_PAIRS) + self.metadata["refine"] = dict( + num_iterations=int(num_iterations), + positions=None if positions is None else "subset", + pair_distance=float(delta), + sigma_excitation=float(sigma), + min_pairs=int(min_pairs), + refine_tilt=bool(refine_tilt), + refine_zone=bool(refine_zone), + zone_search_deg=float(zone_search_deg), + zone_max_total_deg=zone_max_total_deg, + power_intensity=power_intensity, + sigma_envelope=sigma_envelope, + neighbor_rescue=bool(neighbor_rescue), + rescue_threshold_deg=float(rescue_threshold_deg), + rescue_passes=int(rescue_passes), + score_tol=float(score_tol), + consensus_tol=float(consensus_tol), + ) + peaks = self.peaks + R, C, M = self.quats.shape[:3] + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + g_all = self.crystal.g_vec + lam = self.wavelength + sigma_env = sigma_envelope if sigma_envelope is not None else sigma / 2 + power_env = resolve( + power_intensity, + "power_intensity", + self.metadata.get("plan", {}), + default=POWER_INTENSITY, + ) + if power_env <= 0: + # a positions-only plan (power 0) carries no intensity weighting; + # the envelope fit still needs one + power_env = POWER_INTENSITY + prec_ill = float(self.metadata.get("precession_deg", 0.0) or 0.0) + conv_ill = float(self.metadata.get("semiconv_mrad", 0.0) or 0.0) + n_rod = self._foil_normal_crystal() + from quantem.diffraction.illumination import relrod_factor + + def envelope(S, g_rows, f_rows=None): + # Laue-circle envelope of the paired reflections at shifted + # excitation errors S (P, T, T), averaged over the illumination + # recorded on this map (ring: Bessel series; disk: transform); + # f_rows scales the sweep onto the relrod for a foil normal + if prec_ill <= 0 and conv_ill <= 0: + return torch.exp(-(S**2) / (2 * sigma_env**2)) + from quantem.diffraction.illumination import ( + excitation_amplitudes, + gaussian_envelope, + gaussian_envelope_ring_torch, + ) + + a_r, b_r = excitation_amplitudes(g_rows, self.energy_ev, prec_ill, conv_ill) + if f_rows is not None: + a_r, b_r = a_r * f_rows.abs(), b_r * f_rows.abs() + if conv_ill <= 0: + return gaussian_envelope_ring_torch(S, a_r[:, None, None], sigma_env) + return torch.as_tensor( + gaussian_envelope( + S.numpy(), a_r[:, None, None].numpy(), b_r[:, None, None].numpy(), sigma_env + ), + dtype=torch.float64, + ) + + f_all = self.crystal.struct_factors_int.to(torch.float64) + tg = torch.deg2rad( + torch.linspace(-zone_search_deg, zone_search_deg, 17, dtype=torch.float64) + ) + eye3 = torch.eye(3, dtype=torch.float64) + tilt_cap = np.deg2rad( + zone_max_total_deg if zone_max_total_deg is not None else 0.75 * self.zone_step_deg + ) + + def refine_single(q, q_exp, w_exp): + """Refine one orientation; return (q, pairing score).""" + score = 0.0 + tilt_total = torch.zeros(2, dtype=torch.float64) + for _ in range(num_iterations): + g = qrotate(q, g_all) + gz, g2 = g[:, 2], (g**2).sum(dim=1) + s_g = (2 * gz - lam * g2) / (2 - 2 * lam * gz) + spot = g[:, :2] + f_rod = torch.ones_like(s_g) + if n_rod is not None: + # plate geometry: excitation along the rod, the spot + # where the rod meets the sphere + n_lab = qrotate(q, n_rod[None])[0] + f_rod = relrod_factor(g, n_lab, self.energy_ev, 0.0) + s_g = s_g * f_rod + spot = spot - s_g[:, None] * n_lab[None, :2] + sel = torch.abs(s_g) < 2 * sigma + g_sel = g[sel] + if g_sel.shape[0] == 0: + return q, score + d = torch.cdist(spot[sel], q_exp) + d_min, j_min = d.min(dim=1) + pair = d_min < delta + if int(pair.sum()) < min_pairs: + return q, score + gp = g_sel[pair] + tgt = q_exp[j_min[pair]] + w = w_exp[j_min[pair]] * (1 - d_min[pair] / delta) + score = float(w.sum()) + # solve min sum w | tgt - (g + omega x g)_xy |^2 for omega + r = tgt - spot[sel][pair] # (P, 2) + if refine_tilt: + A = torch.zeros((gp.shape[0], 2, 3), dtype=torch.float64) + A[:, 0, 1] = gp[:, 2] + A[:, 0, 2] = -gp[:, 1] + A[:, 1, 0] = -gp[:, 2] + A[:, 1, 2] = gp[:, 0] + Aw = A * w[:, None, None] + AtA = torch.einsum("pki,pkj->ij", Aw, A) + Atr = torch.einsum("pki,pk->i", Aw, r) + omega = torch.linalg.solve(AtA + 1e-12 * eye3, Atr) + else: + # in-plane only: residual model r = omega_z * (-g_y, g_x) + a = torch.stack((-gp[:, 1], gp[:, 0]), dim=1) # (P, 2) + num = (w[:, None] * a * r).sum() + den = (w[:, None] * a * a).sum().clamp_min(1e-12) + omega = torch.tensor([0.0, 0.0, float(num / den)], dtype=torch.float64) + angle = torch.linalg.norm(omega) + if angle > 1e-10: + dq = quat_from_axis_angle(omega / angle, angle) + q = qmult(dq, q) + + if refine_zone: + # continuous zone-axis tilt from the intensity envelope: + # a small lab-frame tilt (wx, wy) shifts every excitation + # error by s(w) = s0 + wx*gy - wy*gx; maximize the + # normalized cosine between the measured intensities and + # the predicted |F|^2 * envelope over a grid with + # parabolic sub-stepping (the Laue-circle fit -- peak + # positions carry no tilt information, the excitation + # pattern does) + s0 = s_g[sel][pair] + fr = f_rod[sel][pair] + a1 = gp[:, 1] * fr + a2 = -gp[:, 0] * fr + f_p = f_all[sel][pair] + S = ( + s0[:, None, None] + + tg[None, :, None] * a1[:, None, None] + + tg[None, None, :] * a2[:, None, None] + ) + pred = ( + f_p[:, None, None] + * envelope(S, g_sel[pair], None if n_rod is None else fr) + ).clamp_min(0) ** power_env + w_env = w**power_env + E = (w_env[:, None, None] * pred).sum(dim=0) / ( + (pred**2).sum(dim=0).sqrt().clamp_min(1e-12) + ) + ij = int(E.argmax()) + i0, j0 = ij // 17, ij % 17 + wx, wy = float(tg[i0]), float(tg[j0]) + step = float(tg[1] - tg[0]) + if 0 < i0 < 16: + c0, c1, c2 = ( + float(E[i0 - 1, j0]), + float(E[i0, j0]), + float(E[i0 + 1, j0]), + ) + den = 2 * c1 - c0 - c2 + if abs(den) > 1e-12: + wx += 0.5 * (c2 - c0) / den * step + if 0 < j0 < 16: + c0, c1, c2 = ( + float(E[i0, j0 - 1]), + float(E[i0, j0]), + float(E[i0, j0 + 1]), + ) + den = 2 * c1 - c0 - c2 + if abs(den) > 1e-12: + wy += 0.5 * (c2 - c0) / den * step + # trust region on the cumulative tilt from the start + prop = tilt_total + torch.tensor([wx, wy], dtype=torch.float64) + over = float(torch.linalg.norm(prop)) - tilt_cap + if over > 0: + prop = prop * tilt_cap / float(torch.linalg.norm(prop)) + step = prop - tilt_total + tilt_total = prop + tilt = torch.tensor( + [float(step[0]), float(step[1]), 0.0], + dtype=torch.float64, + ) + t_ang = torch.linalg.norm(tilt) + if t_ang > 1e-10: + dq = quat_from_axis_angle(tilt / t_ang, t_ang) + q = qmult(dq, q) + return q, score + + def get_exp(rx, ry): + data = peaks[rx, ry].numpy().astype(np.float64) + if data.shape[0] < min_pairs: + return None, None + q_exp = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + w_exp = torch.as_tensor(data[:, ix[2]], dtype=torch.float64).clamp_min(0) + w_exp = w_exp / w_exp.max().clamp_min(1e-12) + return q_exp, w_exp + + # positions to refine: those requested, or everything matched + active = position_mask(positions, (R, C)) + if self.computed is not None: + active = active & self.computed + # the library matches, kept so refinement can be undone where it + # made the fit worse + q_start = self.quats.clone() + # one bar per crystal covers refinement and neighbour rescue; each + # stage adds its own work to the total as it starts + bar = tqdm(total=0, desc=f"refining {self.crystal.name}") if progress_bar else None + if batched and not refine_tilt: + self._refine_batched( + active=active, + delta=delta, + sigma=sigma, + sigma_env=sigma_env, + tg=tg, + tilt_cap=tilt_cap, + num_iterations=num_iterations, + min_pairs=min_pairs, + refine_zone=refine_zone, + power_env=power_env, + progress_bar=bar if bar is not None else False, + ) + else: + iterator = [(rx, ry) for rx, ry in np.ndindex(R, C) if active[rx, ry]] + if bar is not None: + bar.total = (bar.total or 0) + len(iterator) + bar.refresh() + for rx, ry in iterator: + if bar is not None: + bar.update(1) + q_exp, w_exp = get_exp(rx, ry) + if q_exp is None: + continue + for m in range(M): + if self.corr[rx, ry, m] <= 0: + continue + q, _ = refine_single(self.quats[rx, ry, m], q_exp, w_exp) + self.quats[rx, ry, m] = q + + # Refinement polishes the orientation on paired peak positions, which + # is accurate for small corrections but can jump to another basin on + # sparse or ambiguous patterns. A polish within `rescue_threshold_deg` + # is trusted: the correlation depends on the excitation envelope, + # which is only approximately known, so it cannot referee sub-degree + # moves. A jump beyond that must explain the measured peaks better + # than the library match did, or it is undone. + act_list = [(rx, ry) for rx, ry in np.ndindex(R, C) if active[rx, ry]] + if bar is not None: + bar.set_description(f"{self.crystal.name} checking against the library match") + bar.total += len(act_list) + bar.refresh() + cscore = torch.zeros((R, C), dtype=torch.float64) + n_reverted = 0 + for rx, ry in act_list: + if bar is not None: + bar.update(1) + data = peaks[rx, ry].numpy().astype(np.float64) + meas = self._measured_term(data, ix) + for m in range(M): + if self.corr[rx, ry, m] <= 0: + continue + s_new = self._correlation_score(self.quats[rx, ry, m], data, ix, meas) + moved = float( + misorientation_angle_deg( + self.quats[rx, ry, m], q_start[rx, ry, m], self.crystal.sym_quats_matching + ) + ) + if moved > rescue_threshold_deg: + s_old = self._correlation_score(q_start[rx, ry, m], data, ix, meas) + if s_new < s_old - score_tol: + self.quats[rx, ry, m] = q_start[rx, ry, m] + s_new = s_old + n_reverted += int(m == 0) + if m == 0: + cscore[rx, ry] = s_new + self.metadata["refine"]["n_reverted"] = int(n_reverted) + + if neighbor_rescue: + # A wrong local optimum shows as a position disagreeing with a + # neighbour. Retry every such position from every distinct + # candidate around it -- all matches of the eight neighbours, + # this position's own other matches, and its Friedel twin -- and + # keep whichever explains the measured peaks best, by the same + # correlation; candidates within `consensus_tol` of the best are a + # tie, broken by how many neighbours agree. Only matched + # neighbours count, and the comparison is in the group the library + # was built with, where folded variants are one answer. Repeat + # while anything changes, up to `rescue_passes` times. + sym_m = self.crystal.sym_quats_matching + # 180 degrees about the beam: kinematic spot positions are + # centrosymmetric and the excitation errors nearly so, so only the + # Ewald curvature tells the two apart + q_twin = torch.tensor([0.0, 0.0, 0.0, 1.0], dtype=self.quats.dtype) + n_retried = n_rescued = 0 + changed = active.clone() + for _pass in range(max(int(rescue_passes), 0)): + q0 = self.quats[..., 0, :] + miso_max = torch.zeros((R, C), dtype=torch.float64) + for dr, dc in ((0, 1), (1, 0)): + both = active[: R - dr, : C - dc] & active[dr:, dc:] + mm = misorientation_angle_deg( + q0[: R - dr, : C - dc].reshape(-1, 4), q0[dr:, dc:].reshape(-1, 4), sym_m + ).reshape(R - dr, C - dc) + mm = torch.where(both, mm, torch.zeros_like(mm)) + miso_max[: R - dr, : C - dc] = torch.maximum(miso_max[: R - dr, : C - dc], mm) + miso_max[dr:, dc:] = torch.maximum(miso_max[dr:, dc:], mm) + # after the first pass, only where something nearby changed + near_change = changed.clone() + for dr in (-1, 0, 1): + for dc in (-1, 0, 1): + near_change |= torch.roll(torch.roll(changed, dr, 0), dc, 1) + retry = torch.nonzero((miso_max > rescue_threshold_deg) & active & near_change) + it2 = retry.tolist() + if not it2: + break + if bar is not None: + bar.set_description( + f"{self.crystal.name} neighbor rescue {_pass + 1}/{rescue_passes}" + ) + bar.set_postfix_str(f"{len(it2)} positions") + bar.total += len(it2) + bar.refresh() + changed = torch.zeros((R, C), dtype=torch.bool) + for rx, ry in it2: + if bar is not None: + bar.update(1) + n_retried += 1 + q_exp, w_exp = get_exp(rx, ry) + if q_exp is None: + continue + data = peaks[rx, ry].numpy().astype(np.float64) + cur_q = self.quats[rx, ry, 0].clone() + cur_s = float(cscore[rx, ry]) + # every candidate in order: this orientation, its Friedel + # twin, its own other matches, the neighbours' matches + raw = [cur_q, qmult(q_twin, cur_q)] + nbrs = [] + for m in range(1, M): + if self.corr[rx, ry, m] > 0: + raw.append(self.quats[rx, ry, m]) + for dr in (-1, 0, 1): + for dc in (-1, 0, 1): + nr, nc = rx + dr, ry + dc + if (dr == 0 and dc == 0) or not (0 <= nr < R and 0 <= nc < C): + continue + if not bool(active[nr, nc]): + continue + nbrs.append(self.quats[nr, nc, 0]) + for m in range(M): + if self.corr[nr, nc, m] > 0: + raw.append(self.quats[nr, nc, m]) + # drop repeats within 0.5 degrees, keeping the first, from + # one batched misorientation matrix + qr = torch.stack(raw) + dup = (misorientation_angle_deg(qr[:, None], qr[None], sym_m) <= 0.5).numpy() + keep: list[int] = [] + for i in range(len(raw)): + if not any(dup[i, j] for j in keep): + keep.append(i) + cands = [raw[i] for i in keep] + # none where the twin is a symmetry copy of this orientation + twin_ix = 1 if 1 in keep else None + # score every candidate as it stands -- the neighbours' + # were refined on a neighbouring pattern already + cq = torch.stack(cands) + meas = self._measured_term(data, ix) + cs = [cur_s] + [ + self._correlation_score(qc, data, ix, meas) for qc in cands[1:] + ] + # polish the best on this pattern, keeping it if it helps + k = int(np.argmax(cs)) + if k > 0: + q_ref, _ = refine_single(cq[k].clone(), q_exp, w_exp) + s_ref = self._correlation_score(q_ref, data, ix, meas) + if s_ref > cs[k]: + cq[k], cs[k] = q_ref, s_ref + cs = np.asarray(cs) + # ties within consensus_tol go to the candidate most + # neighbours agree with, then to the higher correlation + support = np.zeros(len(cs), dtype=int) + if nbrs: + nq = torch.stack(nbrs) + miso = misorientation_angle_deg(cq[:, None, :], nq[None], sym_m) + support = (miso <= rescue_threshold_deg).sum(dim=1).numpy() + tied = cs >= cs.max() - consensus_tol + pick = max(np.flatnonzero(tied), key=lambda j: (support[j], cs[j])) + agreed = ( + consensus_tol > 0 + and support[pick] > support[0] + and cs[pick] >= cur_s - consensus_tol + ) + if not agreed: + # by the pattern alone; the twin never wins here, as + # its pattern differs only by the Ewald curvature and + # a higher score is the forward model's error + own = [j for j in range(len(cs)) if j != twin_ix] + pick = max(own, key=lambda j: cs[j]) + if pick > 0 and (agreed or cs[pick] > cur_s + score_tol): + n_rescued += 1 + changed[rx, ry] = True + self.quats[rx, ry, 0] = cq[pick] + cscore[rx, ry] = float(cs[pick]) + self.metadata["refine"]["n_retried"] = int(n_retried) + self.metadata["refine"]["n_rescued"] = int(n_rescued) + if bar is not None: + bar.set_description(f"refined {self.crystal.name}") + bar.set_postfix_str( + f"kept the library match at {n_reverted}" + + ( + f", rescued {self.metadata['refine'].get('n_rescued', 0)}" + if neighbor_rescue + else "" + ) + ) + self.score = cscore + if bar is not None: + bar.close() + return self + + # ------------------------------------------------------------------ + # forward simulation of a match + # ------------------------------------------------------------------ + + def _measured_term(self, data: np.ndarray, ix: list[int]) -> tuple: + """Measured side of :meth:`_correlation_score`: positions, weights + and self-overlap, the same for every candidate at a position.""" + m_xy = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + m_i = torch.as_tensor(data[:, ix[2]], dtype=torch.float64).clamp_min(0) + w_m = ( + m_i**self.power_intensity_experiment + * torch.linalg.norm(m_xy, dim=1) ** self.power_radial + ) + inv = 1.0 / (4.0 * self.corr_kernel_size**2) + mm = float(w_m @ torch.exp(-(torch.cdist(m_xy, m_xy) ** 2) * inv) @ w_m) + return m_xy, w_m, max(mm, 0.0) + + def _correlation_score( + self, q: torch.Tensor, data: np.ndarray, ix: list[int], measured: tuple | None = None + ) -> float: + """How well orientation `q` explains the measured peaks `data`. + + The cosine similarity between the measured and the simulated + patterns, each a set of Gaussian spots of the library's pairing + width: the continuous form of the correlation matching maximizes, + with the same intensity and radial weights. Unlike a sum of paired + intensity it charges for predicted spots that were not measured, so + a denser pattern does not win by pairing more peaks. Simulated + reflections outside the detector are left out, as in the library. + """ + pd = float(self.metadata.get("precession_deg", 0.0) or 0.0) + sc = float(self.metadata.get("semiconv_mrad", 0.0) or 0.0) + sim = self.crystal.generate_pattern( + q, + energy_ev=self.energy_ev, + sigma_excitation=self.sigma_excitation, + precession_deg=pd, + semiconv_mrad=sc, + foil_normal=self.metadata.get("foil_normal"), + ) + s_xy = torch.stack((sim["qx"], sim["qy"]), dim=1).to(torch.float64) + s_i = sim["intensity"].to(torch.float64).clamp_min(0) + det = (self.metadata.get("plan") or {}).get("detector_q_max") + if det is not None and s_xy.shape[0]: + det = np.atleast_1d(det).astype(float) + qx_max, qy_max = (det[0], det[0]) if det.size == 1 else (det[0], det[1]) + rot = np.deg2rad(-float(self.peaks.metadata.get("rotation_ccw_deg", 0.0) or 0.0)) + c, s_ = np.cos(rot), np.sin(rot) + r_det = s_xy[:, 0] * c - s_xy[:, 1] * s_ + c_det = s_xy[:, 0] * s_ + s_xy[:, 1] * c + on = (r_det.abs() <= qx_max) & (c_det.abs() <= qy_max) + s_xy, s_i = s_xy[on], s_i[on] + p_rad = float(self.power_radial) + inv = 1.0 / (4.0 * self.corr_kernel_size**2) + + def overlap(a, wa, b, wb): + return float(wa @ torch.exp(-(torch.cdist(a, b) ** 2) * inv) @ wb) + + m_xy, w_m, mm = measured if measured is not None else self._measured_term(data, ix) + if s_xy.shape[0] == 0 or m_xy.shape[0] == 0: + return 0.0 + w_s = s_i**self.power_intensity * torch.linalg.norm(s_xy, dim=1) ** p_rad + norm = np.sqrt(mm * max(overlap(s_xy, w_s, s_xy, w_s), 0.0)) + return overlap(m_xy, w_m, s_xy, w_s) / norm if norm > 0 else 0.0 + + def _refine_batched( + self, + active: torch.Tensor, + delta: float, + sigma: float, + sigma_env: float, + tg: torch.Tensor, + tilt_cap: float, + num_iterations: int, + min_pairs: int, + refine_zone: bool, + progress_bar, + chunk: int = 64, + power_env: float = POWER_INTENSITY, + ) -> None: + """Chunk-vectorized in-plane + envelope refinement of the active positions.""" + from quantem.diffraction.rotations import quat_to_matrix + + peaks = self.peaks + R, C, M = self.quats.shape[:3] + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + g_all = self.crystal.g_vec # (G, 3) + f_all = self.crystal.struct_factors_int.to(torch.float64) + lam = self.wavelength + n_tg = tg.shape[0] + prec_ill = float(self.metadata.get("precession_deg", 0.0) or 0.0) + conv_ill = float(self.metadata.get("semiconv_mrad", 0.0) or 0.0) + n_rod = self._foil_normal_crystal() + from quantem.diffraction.illumination import relrod_factor + + def envelope(S, g_rows, f_rows=None): + # Laue-circle envelope of the paired reflections at shifted + # excitation errors S (P, T, T), averaged over the illumination + # recorded on this map (ring: Bessel series; disk: transform); + # f_rows scales the sweep onto the relrod for a foil normal + if prec_ill <= 0 and conv_ill <= 0: + return torch.exp(-(S**2) / (2 * sigma_env**2)) + from quantem.diffraction.illumination import ( + excitation_amplitudes, + gaussian_envelope, + gaussian_envelope_ring_torch, + ) + + a_r, b_r = excitation_amplitudes(g_rows, self.energy_ev, prec_ill, conv_ill) + if f_rows is not None: + a_r, b_r = a_r * f_rows.abs(), b_r * f_rows.abs() + if conv_ill <= 0: + return gaussian_envelope_ring_torch(S, a_r[:, None, None], sigma_env) + return torch.as_tensor( + gaussian_envelope( + S.numpy(), a_r[:, None, None].numpy(), b_r[:, None, None].numpy(), sigma_env + ), + dtype=torch.float64, + ) + + # flatten measured peaks once, padded per position + cells = [peaks[r, c].numpy().astype(np.float64) for r, c in np.ndindex(R, C)] + counts = np.array([c.shape[0] for c in cells]) + Pmax = max(1, counts.max()) + N = R * C + q_exp = torch.full((N, Pmax, 2), 1e6, dtype=torch.float64) + w_exp = torch.zeros((N, Pmax), dtype=torch.float64) + for i, arr in enumerate(cells): + n = arr.shape[0] + if n == 0: + continue + q_exp[i, :n] = torch.as_tensor(arr[:, ix[:2]], dtype=torch.float64) + wi = torch.as_tensor(arr[:, ix[2]], dtype=torch.float64).clamp_min(0) + w_exp[i, :n] = wi / wi.max().clamp_min(1e-12) + + quats = self.quats.reshape(N, M, 4) + corr = self.corr.reshape(N, M) + valid_pos = torch.as_tensor(counts >= min_pairs) & active.reshape(N) + + chunks = [i for i in range(0, N, chunk) if bool(valid_pos[i : i + chunk].any())] + # `progress_bar` is either a flag or the bar refine_orientations + # shares with the neighbour rescue, so one crystal shows one bar + bar = progress_bar if hasattr(progress_bar, "update") else None + if bar is not None: + bar.total = (bar.total or 0) + len(chunks) + bar.refresh() + elif progress_bar: + chunks = tqdm(chunks, desc=f"refining {self.crystal.name}") + for i0 in chunks: + if bar is not None: + bar.update(1) + i1 = min(i0 + chunk, N) + B = i1 - i0 + qe = q_exp[i0:i1] # (B, P, 2) + we = w_exp[i0:i1] # (B, P) + for m in range(M): + act = valid_pos[i0:i1] & (corr[i0:i1, m] > 0) + if not bool(act.any()): + continue + q = quats[i0:i1, m].clone() # (B, 4) + tilt_total = torch.zeros((B, 2), dtype=torch.float64) + for _ in range(num_iterations): + Rm = quat_to_matrix(q) # (B, 3, 3) + g = torch.einsum("bij,gj->bgi", Rm, g_all) # (B, G, 3) + gz, g2 = g[..., 2], (g**2).sum(dim=-1) + s_g = (2 * gz - lam * g2) / (2 - 2 * lam * gz) + spot = g[..., :2] + f_rod = None + if n_rod is not None: + # plate geometry: excitation along the rod, the spot + # where the rod meets the sphere + n_lab = torch.einsum("bij,j->bi", Rm, n_rod) # (B, 3) + f_rod = relrod_factor(g, n_lab[:, None, :], self.energy_ev, 0.0) + s_g = s_g * f_rod + spot = spot - s_g[..., None] * n_lab[:, None, :2] + sel = torch.abs(s_g) < 2 * sigma # (B, G) + d = torch.cdist(spot, qe) # (B, G, P) + d_min, j_min = d.min(dim=-1) # (B, G) + pair = sel & (d_min < delta) + w_g = torch.gather(we, 1, j_min) * (1 - d_min / delta).clamp_min(0) + w_g = w_g * pair # (B, G) + n_pair = pair.sum(dim=1) + ok = act & (n_pair >= min_pairs) + if not bool(ok.any()): + break + tgt = torch.gather(qe, 1, j_min[..., None].expand(-1, -1, 2)) # (B, G, 2) + r_vec = tgt - spot + # in-plane closed form + a_vec = torch.stack((-spot[..., 1], spot[..., 0]), dim=-1) + num = (w_g[..., None] * a_vec * r_vec).sum(dim=(1, 2)) + den = (w_g[..., None] * a_vec * a_vec).sum(dim=(1, 2)) + wz = torch.where(ok, num / den.clamp_min(1e-12), torch.zeros_like(num)) + half = wz / 2 + dq = torch.stack( + ( + torch.cos(half), + torch.zeros_like(half), + torch.zeros_like(half), + torch.sin(half), + ), + dim=-1, + ) + q = torch.where(ok[:, None], qmult(dq, q), q) + + if refine_zone: + # sparse over paired reflections only + idx_b, idx_g = torch.nonzero(pair, as_tuple=True) + s0f = s_g[idx_b, idx_g] + fr = None if f_rod is None else f_rod[idx_b, idx_g] + gyf = g[idx_b, idx_g, 1] * (1.0 if fr is None else fr) + gxf = g[idx_b, idx_g, 0] * (1.0 if fr is None else fr) + ff = f_all[idx_g] + wf = w_g[idx_b, idx_g] + S = ( + s0f[:, None, None] + + tg[None, :, None] * gyf[:, None, None] + - tg[None, None, :] * gxf[:, None, None] + ) # (Np, T, T) + pred = (ff[:, None, None] * envelope(S, g[idx_b, idx_g], fr)).clamp_min( + 0 + ) ** power_env + E_num = torch.zeros((B, n_tg, n_tg), dtype=torch.float64).index_add_( + 0, idx_b, (wf**power_env)[:, None, None] * pred + ) + E_den = torch.zeros((B, n_tg, n_tg), dtype=torch.float64).index_add_( + 0, idx_b, pred**2 + ) + E = E_num / E_den.sqrt().clamp_min(1e-12) # (B, T, T) + flat_ij = E.reshape(B, -1).argmax(dim=1) + i_b, j_b = flat_ij // n_tg, flat_ij % n_tg + step = float(tg[1] - tg[0]) + wx = tg[i_b].clone() + wy = tg[j_b].clone() + # parabolic sub-stepping where interior + b_ar = torch.arange(B) + for axis, idx, wv in ((0, i_b, wx), (1, j_b, wy)): + interior = (idx > 0) & (idx < n_tg - 1) + if not bool(interior.any()): + continue + if axis == 0: + c0 = E[b_ar, (idx - 1).clamp(0), j_b] + c1 = E[b_ar, idx, j_b] + c2 = E[b_ar, (idx + 1).clamp(max=n_tg - 1), j_b] + else: + c0 = E[b_ar, i_b, (idx - 1).clamp(0)] + c1 = E[b_ar, i_b, idx] + c2 = E[b_ar, i_b, (idx + 1).clamp(max=n_tg - 1)] + den2 = 2 * c1 - c0 - c2 + shift = torch.where( + interior & (den2.abs() > 1e-12), + 0.5 * (c2 - c0) / den2 * step, + torch.zeros_like(c1), + ) + wv += shift + prop = tilt_total + torch.stack((wx, wy), dim=-1) + norm = torch.linalg.norm(prop, dim=-1) + scale_f = torch.where( + norm > tilt_cap, + tilt_cap / norm.clamp_min(1e-12), + torch.ones_like(norm), + ) + prop = prop * scale_f[:, None] + step_t = torch.where( + ok[:, None], prop - tilt_total, torch.zeros_like(prop) + ) + tilt_total = torch.where(ok[:, None], prop, tilt_total) + t_ang = torch.linalg.norm(step_t, dim=-1) + axis_v = torch.zeros((B, 3), dtype=torch.float64) + nz = t_ang > 1e-10 + if bool(nz.any()): + axis_v[nz, 0] = step_t[nz, 0] / t_ang[nz] + axis_v[nz, 1] = step_t[nz, 1] / t_ang[nz] + dq_t = quat_from_axis_angle(axis_v[nz], t_ang[nz]) + q_nz = q[nz] + q[nz] = qmult(dq_t, q_nz) + quats[i0:i1, m] = torch.where(act[:, None], q, quats[i0:i1, m]) + self.quats = quats.reshape(R, C, M, 4) + + def generate_pattern(self, rx: int, ry: int, match: int = 0, **kwargs): + """Simulated pattern for the matched orientation at one probe position. + + The illumination (precession, convergence) and foil normal recorded + on this map are used unless overridden. + + Parameters + ---------- + rx, ry : int + Scan row and column. + match : int, default=0 + Which match index to simulate. + **kwargs + Passed to :meth:`Crystal.generate_pattern`, e.g. `k_max`. + + Returns + ------- + dict[str, torch.Tensor] + The simulated reflections, as returned by + :meth:`Crystal.generate_pattern` ("qx", "qy", "intensity", ...). + """ + assert self.quats is not None + kwargs.setdefault("precession_deg", self.metadata.get("precession_deg", 0.0)) + kwargs.setdefault("semiconv_mrad", self.metadata.get("semiconv_mrad", 0.0)) + kwargs.setdefault("foil_normal", self.metadata.get("foil_normal")) + return self.crystal.generate_pattern( + self.quats[rx, ry, match], + energy_ev=self.energy_ev, + sigma_excitation=self.sigma_excitation, + **kwargs, + ) + + def match_residual( + self, + other: "OrientationMap", + delete_radius: float = 0.04, + min_number_peaks: int = MIN_NUMBER_PEAKS, + min_corr_other: float = 0.0, + progress_bar: bool = True, + ) -> "OrientationMap": + """Re-match this crystal on the peaks another crystal cannot explain. + + For overlapping patterns (e.g. a thin lath on a matrix), the direct + match of the minority phase is poisoned by the majority phase's + peaks. Here the other crystal's simulated pattern is used to delete + its measured peaks at each position, and this crystal is matched and + refined against the remaining peaks only, with this map's plan. + Where the residual match scores above this map's stored second + match, it replaces it (match index 1), so the joint phase fit sees + one clean candidate per phase. A map with a single match is first + extended to two, the second empty (correlation zero). + + Parameters + ---------- + other : OrientationMap + The matched map of the (locally dominant) other crystal. + delete_radius : float, default=0.04 + Measured peaks within this distance (1/Angstroms) of one of the + other crystal's simulated peaks are removed. + min_number_peaks : int, default=MIN_NUMBER_PEAKS (5) + Positions with fewer measured peaks, or fewer residual peaks, + are not re-matched. + min_corr_other : float, default=0.0 + Positions where the other crystal's correlation is at or below + this are not re-matched (nothing trustworthy to delete). + progress_bar : bool, default=True + Show progress bars for the residual matching and refinement. + + Returns + ------- + OrientationMap + Self, with `quats`, `corr`, `corr_residual` and `mirror` holding + at least two matches. Where the residual match was taken, both + `corr[..., 1]` and `corr_residual[..., 1]` hold its correlation + with the residual peaks. When no position has enough residual + peaks, nothing is replaced. + """ + if self.quats is None or other.quats is None: + raise RuntimeError("Run match_orientations() on both maps first.") + if self.plan_fft is None: + raise RuntimeError("Run build_plan() first.") + peaks = self.peaks + R, C = peaks.shape[0], peaks.shape[1] + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + self.metadata["match_residual"] = dict( + other=other.crystal.name, + delete_radius=float(delete_radius), + min_number_peaks=int(min_number_peaks), + min_corr_other=float(min_corr_other), + ) + + cells = [] + for rx, ry in np.ndindex(R, C): + data = peaks[rx, ry].numpy().astype(np.float64) + if data.shape[0] < min_number_peaks or other.corr[rx, ry, 0] <= min_corr_other: + cells.append(np.zeros((0, 3))) + continue + sim = other.generate_pattern(rx, ry) + sq = torch.stack((sim["qx"], sim["qy"]), dim=1).to(torch.float64) + if sq.shape[0] == 0: + cells.append(data[:, ix]) + continue + qxy = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + d_min = torch.cdist(qxy, sq).min(dim=1).values + keep = (d_min > delete_radius).numpy() + cells.append(data[keep][:, ix]) + + # extend to two matches first, so the result has the same layout + # whether or not anything is replaced + if self.quats.shape[2] < 2: + pad_q = torch.zeros((R, C, 1, 4), dtype=self.quats.dtype) + pad_q[..., 0] = 1.0 + self.quats = torch.cat([self.quats, pad_q], dim=2) + pad = torch.zeros((R, C, 1), dtype=self.corr.dtype) + self.corr = torch.cat([self.corr, pad], dim=2) + if self.corr_residual is not None: + self.corr_residual = torch.cat([self.corr_residual, pad.clone()], dim=2) + self.mirror = torch.cat([self.mirror, torch.zeros((R, C, 1), dtype=torch.bool)], dim=2) + if self.corr_residual is None: + self.corr_residual = self.corr.clone() + + # positions with enough residual peaks, among those this map covers + wanted = torch.as_tensor( + np.array([c.shape[0] >= min_number_peaks for c in cells]).reshape(R, C) + ) + if self.computed is not None: + wanted &= self.computed + if not bool(wanted.any()): + return self + + residual = Vector.from_data( + [cells[r * C : (r + 1) * C] for r in range(R)], + fields=["qx", "qy", "intensity"], + units=["A^-1", "A^-1", "counts"], + name="residual_peaks", + metadata=dict(peaks.metadata or {}), + dtype=peaks.dtype, + ) + om_res = OrientationMap.from_vectors( + residual, + self.crystal, + self.energy_ev, + precession_deg=self.metadata.get("precession_deg", 0.0) or 0.0, + semiconv_mrad=self.metadata.get("semiconv_mrad", 0.0) or 0.0, + foil_normal=self.metadata.get("foil_normal"), + ) + # share this map's plan rather than rebuilding it + for attr in ( + "device", + "dtype", + "cdtype", + "corr_kernel_size", + "sigma_excitation", + "power_radial", + "power_intensity", + "power_intensity_experiment", + "zone_axis_range", + "zone_axes", + "zone_quats", + "zone_step_deg", + "zone_nbr_idx", + "zone_nbr_pos", + "zone_nbr_valid", + "plan_fft", + "shell_radii", + "num_gamma", + "gamma", + "detector_mask", + "plan_norm_shift", + "plan_frac_shift", + ): + setattr(om_res, attr, getattr(self, attr)) + om_res.metadata["plan"] = dict(self.metadata.get("plan") or {}) + om_res.match_orientations( + num_matches=1, + positions=wanted.numpy(), + min_number_peaks=min_number_peaks, + progress_bar=progress_bar, + ) + om_res.refine_orientations(progress_bar=progress_bar) + + # replace the stored second match where the residual match is better + better = om_res.computed & (om_res.corr[..., 0] > self.corr[..., 1]) + self.quats[..., 1, :] = torch.where( + better[..., None], om_res.quats[..., 0, :], self.quats[..., 1, :] + ) + self.corr[..., 1] = torch.where(better, om_res.corr[..., 0], self.corr[..., 1]) + self.corr_residual[..., 1] = torch.where( + better, om_res.corr[..., 0], self.corr_residual[..., 1] + ) + self.mirror[..., 1] = torch.where(better, om_res.mirror[..., 0], self.mirror[..., 1]) + return self + + def cluster_orientations( + self, + mask: np.ndarray | None = None, + threshold_deg: float = 5.0, + min_cluster_size: int = 10, + match: int = 0, + ) -> dict: + """Greedy clustering of the matched orientations into variants. + + Positions are visited in order of decreasing correlation; each seed + collects every unassigned position within `threshold_deg` + (symmetry-reduced misorientation) into a cluster. Follows the variant + analysis of MacLaren et al., J. Microscopy 295, 131 (2024). + + Parameters + ---------- + mask : np.ndarray | None + (R, C) boolean or weight mask of positions to include, e.g. a + phase mask. A position is included where the mask is above 0.5, + the same rule as :meth:`calculate_strain`. + threshold_deg : float, default=5.0 + Misorientation radius of a cluster, degrees. + min_cluster_size : int, default=10 + Smaller clusters are discarded (labels stay -1). + match : int, default=0 + Which match index to cluster. + + Returns + ------- + dict + 'labels' (R, C) int tensor, -1 = unassigned; 'mean_quats' + (K, 4) cluster mean orientations; 'sizes' (K,) member counts. + """ + assert self.quats is not None + R, C = self.quats.shape[:2] + q = self.quats[..., match, :].reshape(-1, 4) + corr = self.corr[..., match].reshape(-1) + ok = corr > 0 + if mask is not None: + ok &= torch.as_tensor(np.asarray(mask, dtype=float).reshape(-1)) > 0.5 + + labels = torch.full((R * C,), -1, dtype=torch.long) + # variants the library folds together belong to one grain + sym = self.crystal.sym_quats_matching + unassigned = ok.clone() + means, sizes = [], [] + k = 0 + while unassigned.any(): + seed = int(torch.where(unassigned, corr, torch.full_like(corr, -1)).argmax()) + miso = misorientation_angle_deg(q[seed][None], q, sym) + members = unassigned & (miso < threshold_deg) + unassigned &= ~members + if int(members.sum()) < min_cluster_size: + continue + labels[members] = k + # symmetry-align members to the seed, then average + qm = q[members] + dq = qmult(qconj(q[seed])[None], qm) + dq_sym = qmult(dq[:, None, :], sym) + best = dq_sym[..., 0].abs().argmax(dim=1) + dq_best = dq_sym[torch.arange(qm.shape[0]), best] + sign = torch.where(dq_best[:, :1] < 0, -1.0, 1.0) + q_aligned = qmult(q[seed][None], dq_best * sign) + means.append(qnormalize(q_aligned.mean(dim=0))) + sizes.append(int(members.sum())) + k += 1 + # order clusters by size, largest first + if means: + order = torch.argsort(torch.tensor(sizes), descending=True) + relabel = torch.full((len(sizes),), -1, dtype=torch.long) + relabel[order] = torch.arange(len(sizes)) + labels = torch.where(labels >= 0, relabel[labels.clamp_min(0)], labels) + means = [means[int(i)] for i in order] + sizes = [sizes[int(i)] for i in order] + return { + "labels": labels.reshape(R, C), + "mean_quats": torch.stack(means) if means else torch.zeros((0, 4)), + "sizes": torch.tensor(sizes), + } + + def calculate_strain( + self, + match: int = 0, + pair_distance: float | None = None, + min_pairs: int | None = None, + mask: np.ndarray | None = None, + ds_sampling: float | None = None, + ds_units: str | None = None, + progress_bar: bool = True, + ): + """Per-position strain from measured vs simulated peak positions. + + At each probe position the refined orientation's simulated pattern is + paired to the measured peaks and the in-plane deformation A + minimizing sum w |A q_sim - q_meas|^2 is solved in closed form. The + strain is referenced to the crystal's ideal lattice, so unlike + lattice-vector strain mapping it is absolute, not relative to a + reference region. + + Parameters + ---------- + match : int, default=0 + Which match index to measure. + pair_distance : float | None + Largest distance (1/Angstroms) between a simulated and a measured + peak that are paired; inherits the refinement's, then the plan's. + min_pairs : int | None + Positions with fewer paired peaks are left as NaN; inherits the + refinement's value, else MIN_PAIRS (4). + mask : np.ndarray | None + (R, C) boolean or weight mask of positions to measure. A position + is measured where the mask is above 0.5, the same rule as + :meth:`cluster_orientations`. None measures every matched + position. + ds_sampling : float | None + Scan step, passed to StrainMap for scale bars. + ds_units : str | None + Units of `ds_sampling`. + progress_bar : bool, default=True + Show a progress bar over the positions. + + Returns + ------- + StrainMap + The columns of A enter as per-position reciprocal lattice + vectors with the identity as the fixed reference, so all + StrainMap machinery applies: `plot_strain(rotation_angle=...)` + for user-chosen u/v directions, `rotate_strain`, masking, and + scale bars. `num_pairs` (R, C) is attached as an attribute. + """ + from quantem.diffraction.strain import StrainMap + + assert self.quats is not None + delta = resolve( + pair_distance, + "pair_distance", + self.metadata.get("refine"), + self.metadata.get("plan"), + default=self.corr_kernel_size, + ) + min_pairs = resolve(min_pairs, "min_pairs", self.metadata.get("refine"), default=MIN_PAIRS) + self.metadata["strain"] = dict( + match=int(match), pair_distance=float(delta), min_pairs=int(min_pairs) + ) + peaks = self.peaks + R, C = peaks.shape[0], peaks.shape[1] + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + + A_map = np.full((R, C, 2, 2), np.nan) + num_pairs = np.zeros((R, C), dtype=int) + include = None if mask is None else np.asarray(mask, dtype=float) > 0.5 + + iterator = list(np.ndindex(R, C)) + if progress_bar: + iterator = tqdm(iterator, desc=f"strain mapping {self.crystal.name}") + for rx, ry in iterator: + if include is not None and not include[rx, ry]: + continue + if self.corr[rx, ry, match] <= 0: + continue + data = peaks[rx, ry].numpy().astype(np.float64) + if data.shape[0] < min_pairs: + continue + q_exp = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + w_exp = torch.as_tensor(data[:, ix[2]], dtype=torch.float64).clamp_min(0) + sim = self.generate_pattern(rx, ry, match=match) + sq = torch.stack((sim["qx"], sim["qy"]), dim=1) + if sq.shape[0] == 0: + continue + d = torch.cdist(sq, q_exp) + d_min, j_min = d.min(dim=1) + pair = d_min < delta + n = int(pair.sum()) + if n < min_pairs: + continue + qs = sq[pair] + qm = q_exp[j_min[pair]] + w = w_exp[j_min[pair]] * (1 - d_min[pair] / delta) + # A = (sum w qm qs^T) (sum w qs qs^T)^-1 + M1 = torch.einsum("p,pi,pj->ij", w, qm, qs) + M2 = torch.einsum("p,pi,pj->ij", w, qs, qs) + A = M1 @ torch.linalg.inv(M2 + 1e-12 * torch.eye(2, dtype=torch.float64)) + A_map[rx, ry] = A.numpy() + num_pairs[rx, ry] = n + + # columns of A are the measured images of the reciprocal unit basis; + # StrainMap's reciprocal-space branch (U_ref @ inv(U) = F^T) then + # yields the real-space strain with the shared sign conventions + sm = StrainMap( + g1_array=A_map[..., :, 0], + g2_array=A_map[..., :, 1], + ds_shape=(R, C), + real_space=False, + g1_ref=np.array([1.0, 0.0]), + g2_ref=np.array([0.0, 1.0]), + mask=None if mask is None else np.asarray(mask, dtype=float), + ds_sampling=ds_sampling, + ds_units=ds_units, + ) + sm.num_pairs = num_pairs + return sm + + def in_plane_angle_deg( + self, match: int = 0, mod_deg: float | str | None = "auto" + ) -> torch.Tensor: + """In-plane angle of the crystal a-axis at every position (degrees). + + The angle of the projected crystal [100] Cartesian axis, measured + from the scan column axis toward the row axis. + + `mod_deg` wraps the angle, which is what makes the map continuous + where the in-plane orientation is not uniquely indexable. The default + "auto" wraps each position by 360 / n with n the apparent rotational + symmetry of its own zero-layer pattern + (`Crystal.projected_rotation_order`), so the map folds by exactly the + ambiguity the data carry and no more. For a body-centered cubic + crystal near <111> that is 60 degrees, where the crystal itself + repeats only every 120, and folding removes the 60 degree jumps + between two variants no zero-layer pattern can separate. A float + wraps everywhere by that value, and None returns the full range. + """ + from quantem.diffraction.rotations import quat_to_matrix + + assert self.quats is not None + R = quat_to_matrix(self.quats[..., match, :]) + a_lab = R[..., :, 0] # crystal x-axis in the lab frame + ang = torch.rad2deg(torch.atan2(a_lab[..., 0], a_lab[..., 1])) + if isinstance(mod_deg, str): + if mod_deg != "auto": + raise ValueError('mod_deg must be a number, None, or "auto"') + # beam direction in crystal coordinates, deduplicated on a coarse + # grid: the projected order is piecewise constant in the zone axis + zone_c = R[..., 2, :] + key = torch.round(zone_c.reshape(-1, 3) * 200) / 200 + uniq, inv = torch.unique(key, dim=0, return_inverse=True) + order = self.crystal.projected_rotation_order(uniq.numpy()) + n = torch.as_tensor(np.asarray(order), dtype=torch.float64)[inv] + return ang % (360.0 / n.reshape(ang.shape)) + if mod_deg is not None: + ang = ang % mod_deg + return ang + + def _default_mask(self, kwargs: dict) -> dict: + """After a staged run on a subset of positions, plot only those.""" + if kwargs.get("mask") is None and self.computed is not None: + if not bool(self.computed.all()): + kwargs["mask"] = self.computed.numpy().astype(float) + return kwargs + + @property + def scan_scalebar(self) -> dict | None: + """Real-space scale bar of the scan, carried from the dataset. + + `BraggVectors` stamps the scan sampling and units of the dataset onto + the detected peaks, and the calibration keeps them, so every map can + draw a scale bar without being told the step size. None when the + dataset was never calibrated, in which case set `dataset.sampling` + and `dataset.units` before detecting the disks. + """ + md = self.metadata.get("peaks", {}) or {} + return scan_scalebar(md) + + def plot_orientation(self, direction: str = "z", match: int = 0, **kwargs): + """IPF-colored orientation map with the color wedge beside it. + + Parameters + ---------- + direction : {"z", "r", "c"} | float | array-like, default="z" + Lab direction whose crystal-frame coordinates are colored; see + :func:`~quantem.diffraction.orientation_visualization.plot_orientation_map`. + match : int, default=0 + Which match index to plot. + **kwargs + Passed to + :func:`~quantem.diffraction.orientation_visualization.plot_orientation_map`. + The scale bar defaults to the scan calibration, and after a + staged run on a subset of positions the mask defaults to those + positions. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + from quantem.diffraction.orientation_visualization import plot_orientation_map + + kwargs.setdefault("scalebar", self.scan_scalebar) + return plot_orientation_map( + self, direction=direction, match=match, **self._default_mask(kwargs) + ) + + def plot_pole_figure(self, pole=(0, 0, 1), match: int = 0, **kwargs): + """Stereographic pole figure of a crystal direction family over the map. + + Parameters + ---------- + pole : array-like, default=(0, 0, 1) + Crystal direction in Miller indices, [uvw] or [uvtw]. + match : int, default=0 + Which match index to plot. + **kwargs + Passed to + :func:`~quantem.diffraction.orientation_visualization.plot_pole_figure`. + After a staged run on a subset of positions the mask defaults + to those positions. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + from quantem.diffraction.orientation_visualization import plot_pole_figure + + return plot_pole_figure(self, pole=pole, match=match, **self._default_mask(kwargs)) + + def plot_cluster_map(self, clusters: dict, **kwargs): + """Map of the orientation clusters, one color per cluster. + + Parameters + ---------- + clusters : dict + Output of :meth:`cluster_orientations`. + **kwargs + Passed to + :func:`~quantem.diffraction.orientation_visualization.plot_cluster_map`. + The scale bar defaults to the scan calibration. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + from quantem.diffraction.orientation_visualization import plot_cluster_map + + kwargs.setdefault("scalebar", self.scan_scalebar) + return plot_cluster_map(self, clusters, **kwargs) + + def plot_cluster_pole_figure(self, clusters: dict, pole=(0, 0, 1), **kwargs): + """Pole figure of the cluster mean orientations, one color per cluster. + + Parameters + ---------- + clusters : dict + Output of :meth:`cluster_orientations`. + pole : array-like, default=(0, 0, 1) + Crystal direction in Miller indices, [uvw] or [uvtw]. + **kwargs + Passed to + :func:`~quantem.diffraction.orientation_visualization.plot_cluster_pole_figure`, + e.g. `pole_label` and `overlay`. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + from quantem.diffraction.orientation_visualization import plot_cluster_pole_figure + + return plot_cluster_pole_figure(self, clusters, pole=pole, **kwargs) + + def misorientation_map( + self, reference: torch.Tensor | None = None, match: int = 0 + ) -> torch.Tensor: + """Misorientation angle of every position to a reference orientation. + + Parameters + ---------- + reference : torch.Tensor | None + (4,) reference quaternion; None is the identity. + match : int, default=0 + Which match index to compare. + + Returns + ------- + torch.Tensor + (R, C) misorientation angles in degrees, reduced by the matching + symmetry of the crystal. + """ + assert self.quats is not None + q = self.quats[..., match, :] + if reference is None: + reference = torch.tensor([1.0, 0, 0, 0], dtype=torch.float64) + return misorientation_angle_deg(reference, q, self.crystal.sym_quats_matching) diff --git a/src/quantem/diffraction/orientation_visualization.py b/src/quantem/diffraction/orientation_visualization.py new file mode 100644 index 000000000..6fcdf29bc --- /dev/null +++ b/src/quantem/diffraction/orientation_visualization.py @@ -0,0 +1,1371 @@ +"""Visualization of orientation maps: IPF maps, pattern overlays, pole figures.""" + +from __future__ import annotations + +import numpy as np +import torch + +from quantem.core.visualization.visualization_utils import add_scalebar_to_ax +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.rotations import quat_to_matrix + +# one color per candidate phase, used consistently across every plot; index +# with phase_color_cycle() so more crystals than colors cycle through them +DEFAULT_PHASE_COLORS = np.array( + [ + [1.00, 0.80, 0.25], # gold + [0.25, 0.80, 0.90], # cyan + [0.45, 0.80, 0.50], # green + [0.85, 0.50, 0.80], # purple + ] +) +# exponent on the distance from the wedge centre: >1 widens the white centre +# and softens the transition into it, <1 shrinks it (much below 0.5 leaves a +# bright point at the centre) +IPF_SATURATION_POWER = 0.5 +# colorfulness relative to the corner colors: 1 keeps them exact, lower values +# wash the whole wedge toward white +IPF_CHROMA = 1.0 +# corner colors: full red, green capped to avoid the fluorescent look, blue +# lifted off pure dark blue; pairwise blends give near-max-chroma orange, +# cyan and violet along the edges +IPF_CORNER_COLORS = np.array( + [ + [1.00, 0.00, 0.00], + [0.00, 0.70, 0.00], + [0.00, 0.30, 1.00], + ] +) +# cluster / grain label colors (tab10 cycle) +CLUSTER_COLORS = [ + (0.122, 0.467, 0.706), + (1.0, 0.498, 0.055), + (0.173, 0.627, 0.173), + (0.839, 0.153, 0.157), + (0.580, 0.404, 0.741), + (0.549, 0.337, 0.294), + (0.890, 0.467, 0.761), + (0.498, 0.498, 0.498), + (0.737, 0.741, 0.133), + (0.090, 0.745, 0.812), +] + + +def phase_color_cycle(n: int, colors=None) -> np.ndarray: + """One RGB color per phase, cycling through the palette. + + Parameters + ---------- + n : int + Number of phases. + colors : sequence | None + Palette of K matplotlib colors (RGB rows or names); None takes + `DEFAULT_PHASE_COLORS`. + + Returns + ------- + np.ndarray + (n, 3) RGB colors in [0, 1]; phase k takes palette entry k modulo K. + """ + from matplotlib.colors import to_rgb + + palette = DEFAULT_PHASE_COLORS if colors is None else colors + palette = np.array([to_rgb(c) for c in palette], dtype=float) + return palette[np.arange(n) % palette.shape[0]] + + +def _bary_to_rgb( + w: np.ndarray, + saturation_power: float | None = None, + chroma: float | None = None, +) -> np.ndarray: + """Barycentric wedge weights (..., 3) to RGB. + + The corner colors are mixed additively with the weights scaled by their + 4-norm, a smooth stand-in for dividing by the largest weight: the mix + stays bright between corners (orange, cyan and violet midway along the + edges) and the corners keep their own colors. The mix is then faded + toward white by 1 - 27 w0 w1 w2, which is zero at the wedge centre and + one on every edge. Both terms are smooth in the weights, so the colors + change gradually across the whole wedge, with no creases where one + corner takes over from another, and the mix never leaves the sRGB gamut. + + Parameters + ---------- + w : np.ndarray + Barycentric coordinates in the fundamental wedge, (..., 3). + saturation_power : float | None + Exponent on the distance from the wedge centre. Above 1 the color + builds up more slowly away from the centre, widening the white region + and softening the transition into it; below 1 shrinks it. None takes + `IPF_SATURATION_POWER`. + chroma : float | None + Colorfulness relative to the corner colors: 1 keeps them exact, lower + washes the wedge toward white. None takes `IPF_CHROMA`. + + Returns + ------- + np.ndarray + RGB array (..., 3) in [0, 1]. + """ + saturation_power = IPF_SATURATION_POWER if saturation_power is None else saturation_power + chroma = IPF_CHROMA if chroma is None else chroma + w = np.clip(np.asarray(w, dtype=float), 0, None) + w = w / np.clip(w.sum(axis=-1, keepdims=True), 1e-12, None) + u = w / np.clip((w**4).sum(axis=-1, keepdims=True) ** 0.25, 1e-12, None) + rgb = u @ IPF_CORNER_COLORS + # distance from the centre: 0 there, 1 on the edges, smooth in w + r = np.clip(1.0 - 27.0 * w[..., 0] * w[..., 1] * w[..., 2], 0, 1) ** saturation_power + return np.clip(1.0 - np.clip(chroma * r, 0, 1)[..., None] * (1.0 - rgb), 0, 1) + + +def _parse_direction(direction) -> torch.Tensor: + """Lab direction for IPF coloring: 'z' (beam), 'r' (scan row), 'c' (scan + col), an in-plane angle in degrees (measured from the column axis toward + the row axis), or an explicit [row, col] / [row, col, z] vector.""" + if isinstance(direction, str): + return { + "z": torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64), + "r": torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64), + "c": torch.tensor([0.0, 1.0, 0.0], dtype=torch.float64), + # back-compat aliases + "x": torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64), + "y": torch.tensor([0.0, 1.0, 0.0], dtype=torch.float64), + }[direction] + if isinstance(direction, (int, float)): + th = np.deg2rad(float(direction)) + return torch.tensor([np.sin(th), np.cos(th), 0.0], dtype=torch.float64) + v = torch.as_tensor(direction, dtype=torch.float64).reshape(-1) + if v.numel() == 2: + v = torch.cat([v, torch.zeros(1, dtype=torch.float64)]) + return v / torch.linalg.norm(v) + + +def _reduce_to_wedge(vectors: torch.Tensor, crystal: Crystal) -> torch.Tensor: + """Map crystal-frame directions into the fundamental zone-axis wedge. + + Applies all proper symmetry rotations to +/- v and returns, per input + vector, the orbit member inside the wedge (all barycentric coordinates + with respect to the wedge corners non-negative). + """ + corners = crystal.zone_axis_wedge() + if corners is None: + # hemisphere fallback: canonicalize to upper hemisphere only + v = vectors.clone() + v[v[..., 2] < 0] *= -1 + return v + Rs = quat_to_matrix(crystal.sym_quats_matching) # (S, 3, 3) + v = vectors.reshape(-1, 3) + orbit = torch.cat( + [torch.einsum("sij,nj->nsi", Rs, v), torch.einsum("sij,nj->nsi", Rs, -v)], + dim=1, + ) # (N, 2S, 3) + A_inv = torch.linalg.inv(corners.to(vectors.dtype).T) + w = torch.einsum("ij,nsj->nsi", A_inv, orbit) + inside = (w > -1e-6).all(dim=-1) + idx = inside.to(torch.float64).argmax(dim=1) + out = orbit[torch.arange(v.shape[0]), idx] + return out.reshape(vectors.shape) + + +def ipf_color( + orientations: torch.Tensor, + crystal: Crystal, + direction: str | torch.Tensor = "z", + saturation_power: float | None = None, + chroma: float | None = None, +) -> np.ndarray: + """Inverse pole figure RGB colors for orientations. + + Parameters + ---------- + orientations : torch.Tensor + Quaternions (..., 4). + crystal : Crystal + Provides symmetry and the fundamental wedge. + direction : {"z", "r", "c"} | float | array-like, default="z" + Lab direction whose crystal-frame coordinates are colored: "z" is + the beam direction (zone-axis map), "r" the scan row axis and "c" + the scan column axis ("x" and "y" are accepted as aliases of "r" + and "c"). A number is an in-plane angle in degrees from the column + axis toward the row axis; a 2 or 3 element vector is an explicit + (row, col[, z]) direction. + saturation_power : float | None + Width of the white centre; see :func:`_bary_to_rgb`. + chroma : float | None + Colorfulness relative to the corner colors; see :func:`_bary_to_rgb`. + + Returns + ------- + np.ndarray + RGB array (..., 3) in [0, 1]. + """ + direction = _parse_direction(direction) + R = quat_to_matrix(orientations) # v_lab = R v_crystal + v_crystal = torch.einsum("...ji,j->...i", R, direction) + v = _reduce_to_wedge(v_crystal, crystal) + + corners = crystal.zone_axis_wedge() + if corners is None: + # hemisphere: hue from azimuth, saturation from polar angle + from matplotlib.colors import hsv_to_rgb + + az = (torch.atan2(v[..., 1], v[..., 0]) / (2 * np.pi)) % 1.0 + pol = torch.acos(v[..., 2].clamp(-1, 1)) / (np.pi / 2) + sat = pol.clamp(0, 1) ** ( + IPF_SATURATION_POWER if saturation_power is None else saturation_power + ) + hsv = torch.stack((az, sat, torch.ones_like(az)), dim=-1) + return hsv_to_rgb(hsv.numpy()) + + A_inv = torch.linalg.inv(corners.to(v.dtype).T) + w = torch.einsum("ij,...j->...i", A_inv, v) + return _bary_to_rgb(w.numpy(), saturation_power, chroma) + + +def fold_in_plane(quats: torch.Tensor, crystal: Crystal, strict: bool = False) -> torch.Tensor: + """Fold the in-plane angle of each orientation by its projected symmetry. + + The zero-layer pattern of a zone axis can repeat more often under + rotation about the beam than the crystal does, and where it does, two + orientations produce the same measured pattern and the match returns one + of them arbitrarily. Rotating each orientation about the beam into the + first such sector makes those two identical, so any map colored from the + result is continuous across the ambiguity. + + Positions whose pattern is no more symmetric than the crystal itself are + returned unchanged, unless `strict`, which folds by the projected order + everywhere. + """ + from quantem.diffraction.rotations import ( + qmult, + qnormalize, + quat_from_axis_angle, + quat_to_matrix, + ) + + q = torch.as_tensor(quats, dtype=torch.float64) + shape = q.shape[:-1] + flat = q.reshape(-1, 4) + R = quat_to_matrix(flat) + zone = R[:, 2, :] # beam direction in crystal coordinates + # the projected order is piecewise constant in the zone axis: evaluate it + # once per distinct axis on a coarse grid + key = torch.round(zone * 200) / 200 + uniq, inv = torch.unique(key, dim=0, return_inverse=True) + n_proj = torch.as_tensor( + np.asarray(crystal.projected_rotation_order(uniq.numpy())), dtype=torch.float64 + )[inv] + if not strict: + # only fold where the pattern is more symmetric than the crystal is + # about that same axis, which is where the indexing is degenerate + n_cryst = _crystal_rotation_order(uniq, crystal)[inv].to(torch.float64) + n_proj = torch.where(n_proj > n_cryst, n_proj, torch.ones_like(n_proj)) + if bool((n_proj <= 1).all()): + return q + a_lab = R[:, :, 0] + ang = torch.rad2deg(torch.atan2(a_lab[:, 0], a_lab[:, 1])) + sector = 360.0 / n_proj + delta = ang - (ang % sector) + # the in-plane angle is measured from the column axis toward the row + # axis, which runs opposite to a right-handed rotation about the beam + beam = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + dq = quat_from_axis_angle(beam, torch.deg2rad(delta)) + return qnormalize(qmult(dq, flat)).reshape(*shape, 4) + + +def _crystal_rotation_order(axes: torch.Tensor, crystal: Crystal) -> torch.Tensor: + """Order of the crystal's own rotation axis along each direction, (N,).""" + from quantem.diffraction.rotations import quat_to_matrix + + Rs = quat_to_matrix(crystal.sym_quats) # (S, 3, 3) + u = axes / torch.linalg.norm(axes, dim=1, keepdim=True).clamp_min(1e-12) + # an operation is a rotation about u when it leaves u fixed + fixed = torch.einsum("sij,nj->nsi", Rs, u) + keeps = (fixed - u[:, None, :]).norm(dim=-1) < 1e-6 + return keeps.sum(dim=1) + + +def wedge_legend( + crystal: Crystal, + ax, + n: int = 120, + labels: bool = True, + orientation: str = "horizontal", + fontsize: int = 11, + saturation_power: float | None = None, + chroma: float | None = None, +) -> None: + """Draw the labeled IPF color triangle for the crystal's fundamental wedge. + + Corner direction labels use 4-index Miller-Bravais symbols for hexagonal + and trigonal crystals. orientation="vertical" rotates the wedge 90 + degrees to fill a tall side panel. + + `saturation_power` and `chroma` go to :func:`_bary_to_rgb` and must match + the map being labelled, which the plotting functions ensure. + """ + corners = crystal.zone_axis_wedge() + if corners is None: + ax.axis("off") + return + c = corners.numpy() + # vertical: rotate so the [001]/[0001] corner sits at the TOP of the + # tall panel with the wedge hanging straight down (the rotation aligns + # the wedge's angular bisector with the downward direction) + if orientation == "vertical": + # bisector from the summed directions, not the mean of two angles, + # which jumps by pi when a corner sits at +-180 degrees (-0.0 in y) + d = sum(c[k, :2] / (1 + c[k, 2]) / np.linalg.norm(c[k, :2]) for k in (1, 2)) + th = -np.pi / 2 - np.arctan2(d[1], d[0]) + else: + th = 0.0 + rot = np.array([[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]]) + cxy = np.stack([c[:, 0] / (1 + c[:, 2]), c[:, 1] / (1 + c[:, 2])], axis=1) @ rot.T + cx, cy = cxy[:, 0], cxy[:, 1] + + # wedge edges: stereographic great-circle arcs, which can bulge past the + # corners (the equator arc does once the wedge is rotated upright) + tt = np.linspace(0, 1, 60)[:, None] + edges = [] + for i0, i1 in ((0, 1), (1, 2), (2, 0)): + e = c[i0][None, :] * (1 - tt) + c[i1][None, :] * tt + e = e / np.linalg.norm(e, axis=1, keepdims=True) + edges.append(np.stack([e[:, 0] / (1 + e[:, 2]), e[:, 1] / (1 + e[:, 2])], axis=1) @ rot.T) + outline = np.concatenate(edges) + + # rasterize the wedge interior: invert the stereographic projection on a + # pixel grid covering the whole outline and alpha-mask outside the + # wedge, so no color spills past it and none is missing inside it + m = 8 + x0, x1 = outline[:, 0].min() - 0.02, outline[:, 0].max() + 0.02 + y0, y1 = outline[:, 1].min() - 0.02, outline[:, 1].max() + 0.02 + X, Y = np.meshgrid(np.linspace(x0, x1, n * m), np.linspace(y0, y1, n * m), indexing="xy") + Xu = np.cos(th) * X + np.sin(th) * Y + Yu = -np.sin(th) * X + np.cos(th) * Y + denom = 1 + Xu**2 + Yu**2 + V = np.stack([2 * Xu / denom, 2 * Yu / denom, (1 - Xu**2 - Yu**2) / denom], axis=-1) + A_inv = np.linalg.inv(c.T) + W = V @ A_inv.T + inside = (W > -1e-9).all(axis=-1) + rgba = np.zeros(X.shape + (4,)) + rgba[..., :3] = _bary_to_rgb(W, saturation_power, chroma) + rgba[..., 3] = inside + ax.imshow(rgba, extent=(x0, x1, y0, y1), origin="lower", interpolation="nearest") + for exy in edges: + ax.plot(exy[:, 0], exy[:, 1], color="k", lw=1.2) + if labels: + names = crystal.zone_axis_wedge_labels() or ["", "", ""] + center = np.array([cx.mean(), cy.mean()]) + for xi, yi, name in zip(cx, cy, names): + out = np.array([xi, yi]) - center + norm = np.linalg.norm(out) + off = out / norm * 0.08 if norm > 1e-6 else np.array([0, -0.08]) + ha = "left" if off[0] > 0.02 else ("right" if off[0] < -0.02 else "center") + va = "bottom" if off[1] > 0.02 else ("top" if off[1] < -0.02 else "center") + ax.text(xi + off[0], yi + off[1], name, fontsize=fontsize, ha=ha, va=va) + ox, oy = outline[:, 0], outline[:, 1] + pad = 0.45 * max(ox.max() - ox.min(), oy.max() - oy.min(), 0.2) + ax.set_xlim(ox.min() - pad, ox.max() + pad) + ax.set_ylim(oy.min() - pad, oy.max() + pad) + ax.set_aspect("equal") + ax.axis("off") + + +def plot_orientation_map( + om, + direction: str = "z", + match: int = 0, + mask: np.ndarray | None = None, + scalebar: dict | None = None, + figax=None, + legend: bool = True, + axsize: tuple[float, float] = (9.0, 4.5), + crop: tuple[int, int, int, int] | None = None, + title: str | None = None, + fold: bool | str = "auto", + smooth: dict | bool | None = None, + saturation_power: float | None = None, + chroma: float | None = None, +): + """IPF-colored orientation map with the wedge legend in an adjacent panel. + + Parameters + ---------- + om : OrientationMap + Matched orientation map. + direction : {"z", "r", "c"} | float | array-like, default="z" + Lab direction to color: "z" the beam (zone axis), "r" the scan row + axis, "c" the scan column axis ("x" and "y" are aliases of "r" and + "c"), a number for an in-plane angle in degrees, or an explicit + vector; see :func:`ipf_color`. + match : int, default=0 + Which match index to plot. + mask : np.ndarray | None + Multiplied into the RGB image (e.g. a phase or reliability mask). + figax : (fig, (ax_map, ax_legend)) | (fig, ax_map) | None + Existing axes; with a single axis the legend is skipped. + legend : bool, default=True + Draw the IPF color wedge in a panel beside the map. + axsize : tuple[float, float], default=(9.0, 4.5) + Figure size in inches of the map panel; the legend panel widens the + figure by 30%. Ignored when `figax` is given. + fold : bool | "auto", default="auto" + Fold the in-plane part of each orientation by the apparent + rotational symmetry of its own zero-layer pattern + (`Crystal.projected_rotation_order`) before coloring. Where that + symmetry exceeds the crystal's own, as for a cubic crystal near + <111>, two orientations give the same pattern and the match picks + between them at random; folding gives them the same color, which + removes jumps that no refinement can. It changes nothing for + direction="z", whose color depends only on the zone axis, and + nothing where the pattern is no more symmetric than the crystal. + "auto" folds only when some position needs it. + smooth : dict | bool | None + Smooth the orientations for display only, leaving the stored ones at + the fit to their own pattern. The average is bilateral and needs two + widths, not one: `sigma_px` over probe positions and `sigma_deg` over + misorientation, with `max_angle_deg` excluding anything further. The + angular pair is what stops the average at a grain boundary, so pass + a dict naming the values you want, such as + {"sigma_px": 1.0, "sigma_deg": 1.0, "max_angle_deg": 5.0}, which is + also what True uses. + saturation_power : float | None + Widens or narrows the white centre of the color wedge. Above 1 covers + a wider range of orientations near the wedge centre and softens the + transition into it. None takes `IPF_SATURATION_POWER`. + chroma : float | None + Colorfulness relative to the corner colors: 1 keeps them exact, lower + washes the map toward white. None takes `IPF_CHROMA`. + The legend is drawn with the same values. + scalebar : dict | None + Real-space scale bar, e.g. {"sampling": 30, "units": "A"}. + crop : (r0, r1, c0, c1) | None + Show only this window of the map (rows r0:r1, columns c0:c1). + title : str | None + Replaces the default title (crystal name and colored direction). + + Returns + ------- + tuple + ``(fig, ax)`` with `ax` the map panel. + """ + import matplotlib.pyplot as plt + + assert om.quats is not None + quats = om.quats[..., match, :] + if smooth is not None and smooth is not False: + if not isinstance(smooth, (dict, bool)): + raise TypeError( + "smooth must be a dict of widths or True; a bare number would set the " + "spatial width and leave the angular tolerance at its default, which " + "is the argument that keeps the average inside one grain" + ) + quats = om.smoothed_quats(match=match, **(smooth if isinstance(smooth, dict) else {})) + if fold and not (isinstance(direction, str) and direction == "z"): + # the zone-axis color does not depend on the in-plane angle + quats = fold_in_plane(quats, om.crystal, strict=fold != "auto") + rgb = ipf_color(quats, om.crystal, direction, saturation_power, chroma) + if mask is not None: + rgb = rgb * np.asarray(mask, dtype=float)[..., None] + if crop is not None: + r0, r1, c0, c1 = crop + rgb = rgb[r0:r1, c0:c1] + + ax_leg = None + if figax is None: + if legend: + fig, (ax, ax_leg) = plt.subplots( + 1, + 2, + figsize=(axsize[0] * 1.3, axsize[1]), + gridspec_kw={"width_ratios": [4, 1]}, + ) + else: + fig, ax = plt.subplots(figsize=axsize) + else: + fig, axs = figax + if isinstance(axs, (tuple, list, np.ndarray)) and len(axs) == 2: + ax, ax_leg = axs + else: + ax = axs + ax.imshow(rgb, interpolation="nearest") + ax.set_xticks([]) + ax.set_yticks([]) + if title is not None: + ax.set_title(title) + elif isinstance(direction, str) and direction == "z": + ax.set_title(f"{om.crystal.name} out-of-plane orientation") + else: + # arrow for the colored in-plane direction lives in the title, + # like the strain-map axis annotations + if isinstance(direction, str): + arrow = { + "r": r"$\downarrow$", + "c": r"$\rightarrow$", + "x": r"$\downarrow$", + "y": r"$\rightarrow$", + }.get(direction, "") + label = {"x": "r", "y": "c"}.get(direction, direction) + ax.set_title(f"{om.crystal.name} in-plane orientation {label} {arrow}") + elif isinstance(direction, (int, float)): + ax.set_title(f"{om.crystal.name} in-plane orientation ({direction:g}\u00b0)") + else: + ax.set_title(f"{om.crystal.name} in-plane orientation") + if scalebar is not None: + add_scalebar_to_ax( + ax, + array_size=rgb.shape[1], + sampling=scalebar.get("sampling", 1.0), + length_units=scalebar.get("length", None), + units=scalebar.get("units", "pixels"), + width_px=rgb.shape[0] / 40, + pad_px=rgb.shape[0] / 80, + color=scalebar.get("color", "white"), + loc="lower right", + ) + if legend and ax_leg is not None: + wedge_legend( + om.crystal, + ax_leg, + orientation="vertical", + saturation_power=saturation_power, + chroma=chroma, + ) + return fig, ax + + +def plot_pattern_matches( + orientation_maps, + positions, + dataset=None, + pixel_size: float | None = None, + origins: np.ndarray | None = None, + matches=(0, 1), + colors=None, + norm=None, + sigma_plot: float | None = 1.0, + q_max_plot: float | None = None, + q_max_quantile: float = 0.98, + scalebar: bool = True, + show_measured: bool = True, + marker_scale: float = 250.0, + measured_scale: float | None = None, + measured_power: float = 0.5, + marker: str | None = None, + transpose_plots: bool = False, + axsize: tuple[float, float] = (3.1, 3.1), +): + """Candidate matches side by side, py4DSTEM style. + + One row per probe position; one column per (crystal, match) candidate, + so alpha and beta fits sit next to each other for direct comparison. + Measured peaks are solid gray disks with area proportional to intensity; + each candidate's simulation is drawn as colored markers in its phase + color, also sized by intensity. With `dataset` given, the raw pattern is + shown behind the markers instead of the gray disks. + + Parameters + ---------- + orientation_maps : OrientationMap | list[OrientationMap] + Matched orientation maps sharing the same peaks. + positions : list[tuple[int, int]] + (row, col) probe positions to plot. + dataset : Dataset4dstem | None + If given, the diffraction pattern is shown behind the overlay and + the gray measured disks are omitted. + pixel_size : float | None + Reciprocal pixel size (1/Angstroms per pixel); required with + `dataset`. + origins : np.ndarray | None + (scan_r, scan_c, 2) fitted origins from measure_origins(); aligns + the background pattern with the origin-corrected peaks. + matches : tuple[int, ...], default=(0, 1) + Match indices per crystal. Indices a map does not hold (the second + match of a map matched with `num_matches=1`) are skipped. + norm : dict | str | None + `norm` of `show_2d`, which draws the recorded pattern, e.g. + {"power": 0.5, "upper_quantile": 0.98}. The default, + {"power": 0.4, "upper_quantile": 0.999}, keeps the direct beam from + flattening the disks. + sigma_plot : float | None, default=1.0 + Gaussian blur (pixels) of the displayed pattern only, which makes + the disks easier to see in low-dose data; None shows it raw. + q_max_plot : float | None + Half-width of every panel, 1/Angstroms. None fits it to the peaks + actually plotted, using `q_max_quantile`. + q_max_quantile : float, default=0.98 + Quantile of the measured peak radii that sets the automatic limit, + used only when `q_max_plot` is None and no `dataset` is given. A few + stray high-angle detections would otherwise set the scale for every + panel and leave the pattern in the middle of empty space, so the + default trims the furthest 2%. Pass 1.0 to enclose every peak. + colors : list | None + One color per crystal; defaults to `DEFAULT_PHASE_COLORS`, the + palette of the phase map, cycled when there are more crystals. + marker : str | None + Matplotlib marker for the simulated peaks. The default is an open + circle over a diffraction pattern, which leaves the measured disk + visible inside it, and a plus over the gray measured peaks. + measured_scale : float | None + Largest marker area of the gray measured peaks, in points^2. The + default fits the markers to the patterns shown: the largest disk is + about 0.6 of the median spacing between neighbouring peaks, so dense + patterns get small markers and sparse ones large, up to + ``1.5 * marker_scale``. The direct beam is far brighter than the + disks, so the areas are compressed by `measured_power` and floored, + which keeps the weak spots visible. + measured_power : float, default=0.5 + Compression applied to the measured intensities before sizing. + scalebar : bool, default=True + Draw a 0.5 1/Angstrom scale bar in the bottom-left panel. + show_measured : bool, default=True + Draw the measured peaks as gray disks. Ignored with `dataset`, + where the recorded pattern is shown instead. + marker_scale : float, default=250.0 + Marker area, in points^2, of the strongest simulated reflection; + the others scale with intensity. + transpose_plots : bool, default=False + Panel layout only, nothing in the data is transposed. By default + rows are probe positions and columns are candidates; True swaps + them, giving one row per candidate across the positions, which fits + a few candidates and many positions on a page. + axsize : tuple[float, float], default=(3.1, 3.1) + Size of one panel in inches. + + Returns + ------- + tuple + ``(fig, axs)`` with `axs` a 2D array of panels. + + Raises + ------ + ValueError + If `dataset` is given without `pixel_size`, or none of `matches` + exists in any map. + """ + import matplotlib.pyplot as plt + + from quantem.core.visualization import show_2d + from quantem.diffraction.bragg_vectors_visualization import _blur + + oms = ( + list(orientation_maps) + if isinstance(orientation_maps, (list, tuple)) + else [orientation_maps] + ) + if dataset is not None and pixel_size is None: + raise ValueError( + "plot_pattern_matches needs pixel_size (1/Angstroms per pixel) to place the " + "dataset behind the peaks" + ) + colors = [tuple(c) for c in phase_color_cycle(len(oms), colors)] + peaks = oms[0].peaks + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + # peaks may be rotated into the scan frame; the raw detector image is + # not, so rotate all overlay coordinates back to the detector frame + rot_deg = float(peaks.metadata.get("rotation_ccw_deg", 0.0) or 0.0) + th = np.deg2rad(-rot_deg) + rot_back = np.array([[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]]) + + # a map matched with fewer matches than requested has no panel for them + panels = [(i, m) for i, om in enumerate(oms) for m in matches if 0 <= m < om.quats.shape[2]] + if not panels: + raise ValueError(f"none of matches={tuple(matches)} exists in the orientation maps") + n_pos, n_pan = len(positions), len(panels) + n_r, n_c = (n_pan, n_pos) if transpose_plots else (n_pos, n_pan) + fig, axs = plt.subplots( + n_r, + n_c, + figsize=(axsize[0] * n_c, axsize[1] * n_r + 0.2), + squeeze=False, + ) + over_image = dataset is not None + if marker is None: + marker = "o" if over_image else "+" + ordinal = ["1st", "2nd", "3rd"] + [f"{k + 1}th" for k in range(3, 9)] + + # one limit for every panel, so positions are directly comparable. A + # position holding only the direct beam has q_max of zero, which would + # collapse its axes, so the limit is taken over all of them together. + if q_max_plot is not None: + q_lim = float(q_max_plot) + elif over_image: + q_lim = dataset.shape[-1] / 2 * pixel_size + else: + if not 0.0 < q_max_quantile <= 1.0: + raise ValueError(f"q_max_quantile must be in (0, 1], got {q_max_quantile}") + q_all = [ + np.hypot( + peaks[rx, ry].numpy().astype(np.float64)[:, ix[0]], + peaks[rx, ry].numpy().astype(np.float64)[:, ix[1]], + ) + for rx, ry in positions + ] + # the direct beam is at zero and every position has one, so it is + # dropped before taking the quantile over the diffracted peaks + q_flat = np.concatenate([q for q in q_all if q.size]) if q_all else np.empty(0) + q_flat = q_flat[q_flat > 0.05] + q_max = float(np.quantile(q_flat, q_max_quantile)) if q_flat.size else 0.0 + q_lim = 1.1 * q_max if q_max > 0 else 1.0 + + if measured_scale is None: + # size the measured disks to the spacing of the peaks on screen, so a + # dense pattern does not turn into overlapping blobs + spacings = [] + for rx, ry in positions: + xy = peaks[rx, ry].numpy().astype(np.float64)[:, [ix[0], ix[1]]] + xy = xy[(np.abs(xy) <= q_lim).all(axis=1)] + if xy.shape[0] > 2: + d = np.hypot(xy[:, None, 0] - xy[None, :, 0], xy[:, None, 1] - xy[None, :, 1]) + np.fill_diagonal(d, np.inf) + spacings.append(np.median(d.min(axis=1))) + if spacings: + spacing_pt = float(np.median(spacings)) / (2 * q_lim) * min(axsize) * 72 + measured_scale = min(1.5 * marker_scale, np.pi / 4 * (0.6 * spacing_pt) ** 2) + else: + measured_scale = 1.5 * marker_scale + + for pi, (rx, ry) in enumerate(positions): + data = peaks[rx, ry].numpy().astype(np.float64) + rc = data[:, [ix[0], ix[1]]] @ rot_back.T + data[:, ix[0]] = rc[:, 0] + data[:, ix[1]] = rc[:, 1] + # the direct beam outshines every disk, so scaling the areas by the + # brightest peak shrinks the real spots to nothing; normalize on the + # diffracted peaks instead, compress, and floor so none vanish + w_meas = data[:, ix[2]].clip(min=0) + q_meas = np.hypot(data[:, ix[0]], data[:, ix[1]]) + w_ref = w_meas[q_meas > 0.05] + hi = float(np.percentile(w_ref, 95)) if w_ref.size else float(w_meas.max(initial=0.0)) + w_meas = np.clip((w_meas / max(hi, 1e-12)) ** measured_power, 0.15, 1.0) + for ci, (i_om, m) in enumerate(panels): + om = oms[i_om] + ax = axs[ci, pi] if transpose_plots else axs[pi, ci] + if over_image: + H, W = dataset.shape[-2], dataset.shape[-1] + if origins is not None: + o_r, o_c = origins[rx, ry] + else: + o_r, o_c = H / 2, W / 2 + # the direct beam is orders of magnitude above the disks, so + # autoscaling to its peak flattens everything else + img = np.clip(np.asarray(dataset.array[rx, ry], dtype=float), 0, None) + show_2d( + _blur(img, sigma_plot), + norm=norm if norm is not None else {"power": 0.4, "upper_quantile": 0.999}, + cmap="gray_r", + figax=(fig, ax), + tight_layout=False, + ) + # pixel j has center (j - origin) * pixel_size; array edges + # sit half a pixel beyond the first/last centers + extent = ( + (-0.5 - o_c) * pixel_size, + (W - 0.5 - o_c) * pixel_size, + (H - 0.5 - o_r) * pixel_size, + (-0.5 - o_r) * pixel_size, + ) + ax.images[-1].set_extent(extent) + # the pattern is off-centre by the origin: show exactly the + # recorded area, so nothing is drawn beyond its edges + x_lim, y_lim = extent[:2], extent[2:] + else: + x_lim, y_lim = (-q_lim, q_lim), (q_lim, -q_lim) + if show_measured: + ax.scatter( + data[:, ix[1]], + data[:, ix[0]], + s=measured_scale * w_meas, + color="0.75", + lw=0, + ) + # a position with too few peaks was never matched, and its stored + # orientation is still the identity; drawing that [001] pattern + # would look like a fit where none was attempted + matched = float(om.corr[rx, ry, m]) > 0 + if om.computed is not None: + matched = matched and bool(om.computed[rx, ry]) + inten = np.zeros(0) + if matched: + sim = om.generate_pattern(rx, ry, match=m) + inten = sim["intensity"].numpy() + sim_rc = np.stack([sim["qx"].numpy(), sim["qy"].numpy()], axis=1) @ rot_back.T + if inten.size: + inside = ( + (sim_rc[:, 1] >= min(x_lim)) + & (sim_rc[:, 1] <= max(x_lim)) + & (sim_rc[:, 0] >= min(y_lim)) + & (sim_rc[:, 0] <= max(y_lim)) + ) + size = marker_scale * inten / inten.max() + sim_rc, size = sim_rc[inside], size[inside] + color = colors[i_om % len(colors)] + if marker == "o": + # open circles leave the measured disk visible inside + ax.scatter( + sim_rc[:, 1], + sim_rc[:, 0], + s=size, + marker="o", + facecolors="none", + edgecolors=color, + lw=1.4, + ) + else: + ax.scatter( + sim_rc[:, 1], + sim_rc[:, 0], + s=size, + marker=marker, + color=color, + lw=1.8, + ) + ax.set_xlim(*x_lim) + ax.set_ylim(*y_lim) + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_aspect("equal") + ax.set_title( + "%s %s\n(%d, %d) %s" + % ( + om.crystal.name, + ordinal[m], + rx, + ry, + ("corr = %.2f" % float(om.corr[rx, ry, m])) if matched else "no match", + ), + fontsize=9, + ) + last_row = (ci == n_r - 1) if transpose_plots else (pi == n_r - 1) + first_col = (pi == 0) if transpose_plots else (ci == 0) + if scalebar and last_row and first_col: + add_scalebar_to_ax( + ax, + array_size=abs(x_lim[1] - x_lim[0]), + sampling=1.0, + length_units=0.5, + units="A^-1", + width_px=q_lim / 45, + pad_px=q_lim / 60, + color="black", + loc="lower right", + fontsize=9, + ) + fig.tight_layout() + return fig, axs + + +def plot_cluster_map( + om, + clusters: dict, + colors: np.ndarray | None = None, + scalebar: dict | None = None, + figax=None, +): + """Map of orientation clusters (variants), one color per cluster. + + Also available as :meth:`OrientationMap.plot_cluster_map`. + + Parameters + ---------- + om : OrientationMap + The clustered map; gives the crystal name for the title. + clusters : dict + Output of :meth:`OrientationMap.cluster_orientations`. + colors : sequence | None + One color per cluster, cycled; defaults to `CLUSTER_COLORS`. + scalebar : dict | None + Real-space scale bar, e.g. {"sampling": 30, "units": "A"}. + figax : (fig, ax) | None + Existing axes to draw into. + + Returns + ------- + tuple + ``(fig, ax)``. Unassigned positions are black. + """ + import matplotlib.pyplot as plt + + labels = clusters["labels"].numpy() + if colors is None: + colors = CLUSTER_COLORS + K = int(labels.max()) + 1 + rgb = np.zeros(labels.shape + (3,)) + for k in range(K): + rgb[labels == k] = colors[k % len(colors)] + + if figax is None: + fig, ax = plt.subplots(figsize=(9, 4.5)) + else: + fig, ax = figax + ax.imshow(rgb, interpolation="nearest") + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_title(f"{om.crystal.name} orientation clusters") + handles = [ + plt.Line2D( + [0], + [0], + marker="s", + ls="", + color=colors[k % len(colors)], + label=f"{k + 1} ({int(clusters['sizes'][k])} px)", + ) + for k in range(K) + ] + ax.legend(handles=handles, loc="center left", bbox_to_anchor=(1.01, 0.5), fontsize=8) + if scalebar is not None: + add_scalebar_to_ax( + ax, + array_size=rgb.shape[1], + sampling=scalebar.get("sampling", 1.0), + length_units=scalebar.get("length", None), + units=scalebar.get("units", "pixels"), + width_px=rgb.shape[0] / 40, + pad_px=rgb.shape[0] / 80, + color=scalebar.get("color", "white"), + loc="lower right", + ) + return fig, ax + + +def _pole_family(crystal: Crystal, pole) -> torch.Tensor: + """Unit Cartesian vectors (F, 3) of every symmetry equivalent of a pole. + + Parameters + ---------- + crystal : Crystal + Gives the lattice and the symmetry rotations. + pole : array-like + Crystal direction in Miller indices, [uvw] or [uvtw]. + + Returns + ------- + torch.Tensor + The distinct symmetry images of the pole and their inverses. + """ + p = crystal.direction_vector(pole).to(torch.float64) + Rs = quat_to_matrix(crystal.sym_quats) + fam = torch.einsum("sij,j->si", Rs, p) + fam = torch.unique(torch.round(fam / 1e-6) * 1e-6, dim=0) + return torch.cat([fam, -fam]) + + +def _pole_points( + quats: torch.Tensor, crystal: Crystal, pole, mask=None +) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + """Stereographic projection of every symmetry-equivalent pole of a map. + + Parameters + ---------- + quats : torch.Tensor + (..., 4) orientations. + crystal : Crystal + Gives the lattice and the symmetry rotations. + pole : array-like + Crystal direction in Miller indices, [uvw] or [uvtw]. + mask : np.ndarray | None + Per-orientation weights; zero-weight orientations are dropped. + + Returns + ------- + tuple of np.ndarray + (x, y, weight, source) of the upper-hemisphere poles, with `source` + the flat index of the orientation each pole came from. + """ + q = quats.reshape(-1, 4) + fam = _pole_family(crystal, pole) + n_fam = fam.shape[0] + R = quat_to_matrix(q) + poles_lab = torch.einsum("nij,sj->nsi", R, fam) + if mask is not None: + w = torch.as_tensor(np.asarray(mask, dtype=float)).reshape(-1) + else: + w = torch.ones(q.shape[0], dtype=torch.float64) + w_all = w[:, None].expand(-1, n_fam).reshape(-1) + src_all = torch.arange(q.shape[0])[:, None].expand(-1, n_fam).reshape(-1) + v = poles_lab.reshape(-1, 3) + keep = (v[:, 2] > -1e-8) & (w_all > 0) + v, w_keep, src = v[keep], w_all[keep], src_all[keep] + x = (v[:, 0] / (1 + v[:, 2])).numpy() + y = (v[:, 1] / (1 + v[:, 2])).numpy() + return x, y, w_keep.numpy(), src.numpy() + + +def _pole_scatter_xy(quats: torch.Tensor, crystal: Crystal, pole) -> np.ndarray: + """Stereographic (x, y) of all symmetry-equivalent poles for orientations. + + Parameters + ---------- + quats : torch.Tensor + (4,) or (N, 4) orientations. + crystal : Crystal + Gives the lattice and the symmetry rotations. + pole : array-like + Crystal direction in Miller indices, [uvw] or [uvtw]. + + Returns + ------- + np.ndarray + (P, 2) projected upper-hemisphere poles. + """ + fam = _pole_family(crystal, pole) + R = quat_to_matrix(torch.atleast_2d(quats)) + v = torch.einsum("nij,sj->nsi", R, fam).reshape(-1, 3) + v = v[v[:, 2] > -1e-8] + x = (v[:, 0] / (1 + v[:, 2])).numpy() + y = (v[:, 1] / (1 + v[:, 2])).numpy() + return np.stack((x, y), axis=1) + + +def plot_cluster_pole_figure( + om, + clusters: dict, + pole, + pole_label: str = "", + overlay: dict | None = None, + colors: np.ndarray | None = None, + figax=None, +): + """Pole figure of the cluster mean orientations, one color per cluster. + + Also available as :meth:`OrientationMap.plot_cluster_pole_figure`. + + Parameters + ---------- + om : OrientationMap + Provides the crystal symmetry of the clustered phase. + clusters : dict + Output of :meth:`OrientationMap.cluster_orientations`. + pole : array-like + Crystal direction of the plotted family in Miller indices, [uvw] or + [uvtw]. + pole_label : str, default="" + Legend label prefix of the pole family, e.g. "[0001]". + overlay : dict | None + Second pole family drawn as open markers, e.g. + {"quats": q_beta_mean, "crystal": ti_beta, "pole": (1, 1, 0), + "label": "<110> beta"}, the standard Burgers relationship check. + Its "pole" is in Miller indices of its own crystal. + colors : sequence | None + One color per cluster, cycled; defaults to `CLUSTER_COLORS`. + figax : (fig, ax) | None + Existing axes to draw into. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + import matplotlib.pyplot as plt + + if colors is None: + colors = CLUSTER_COLORS + if figax is None: + fig, ax = plt.subplots(figsize=(6.5, 6.5)) + else: + fig, ax = figax + + theta = np.linspace(0, 2 * np.pi, 361) + for pol_deg in range(15, 91, 15): + r = np.tan(np.deg2rad(pol_deg) / 2) + lw = 1.0 if pol_deg == 90 else 0.4 + ax.plot(r * np.cos(theta), r * np.sin(theta), color="0.75", lw=lw) + for az in range(0, 180, 15): + ca, sa = np.cos(np.deg2rad(az)), np.sin(np.deg2rad(az)) + ax.plot([-ca, ca], [-sa, sa], color="0.85", lw=0.4) + + K = clusters["mean_quats"].shape[0] + for k in range(K): + xy = _pole_scatter_xy(clusters["mean_quats"][k], om.crystal, pole) + ax.scatter( + xy[:, 0], + xy[:, 1], + s=60, + marker="h", + color=colors[k % len(colors)], + edgecolors="k", + lw=0.4, + label=f"{pole_label} {k + 1}", + ) + if overlay is not None: + xy = _pole_scatter_xy(overlay["quats"], overlay["crystal"], overlay["pole"]) + ax.scatter( + xy[:, 0], + xy[:, 1], + s=70, + marker="D", + facecolors="none", + edgecolors="k", + lw=1.0, + label=overlay.get("label", "overlay"), + ) + ax.set_xlim(-1.15, 1.15) + ax.set_ylim(-1.15, 1.15) + ax.set_aspect("equal") + ax.set_xticks([]) + ax.set_yticks([]) + ax.legend(loc="center left", bbox_to_anchor=(1.02, 0.5), fontsize=8) + ax.set_title(f"{om.crystal.name} cluster pole figure") + return fig, ax + + +def plot_pole_figure( + om, + pole: list[float] | torch.Tensor = (0, 0, 1), + match: int = 0, + mask: np.ndarray | None = None, + bins: int = 181, + color_by: str = "density", + int_range: tuple[float, float] = (0.0, 1.0), + smooth_sigma: float = 1.5, + label: str | None = None, + grid: bool = True, + overlay: dict | None = None, + saturation_power: float | None = None, + chroma: float | None = None, + figax=None, +): + """Stereographic pole figure of a crystal direction family over the map. + + For every probe position, all symmetry equivalents of `pole` are rotated + into the lab frame; upper-hemisphere poles are projected + stereographically and accumulated into a 2D histogram. + + Parameters + ---------- + om : OrientationMap + Matched orientation map. + pole : array-like, default=(0, 0, 1) + Crystal direction of the pole family in Miller indices, [uvw] or + [uvtw]; converted to Cartesian with the crystal's lattice. + match : int, default=0 + Which match index to plot. + mask : np.ndarray | None + Per-position weights (e.g. phase mask). + bins : int, default=181 + Histogram bins across the stereographic disk. + color_by : {"density", "ipf"}, default="density" + "density": white through yellow and red to black with increasing + density. "ipf": each contribution is colored by the IPF (zone axis) + color of its probe position, blended from a white background as the + density rises, with the color wedge in a panel beside it. + int_range : tuple, default=(0.0, 1.0) + Density display range as fractions of the 98th percentile of the + occupied bins: values below the lower limit show as background, + above the upper limit at full strength. + smooth_sigma : float, default=1.5 + Gaussian blur of the histogram, in bins; 0 disables it. + label : str | None + Annotation for the pole family, e.g. "(0001)" or "{110}". + grid : bool, default=True + Draw polar-angle circles and azimuth spokes every 30 degrees. + overlay : dict | None + A second pole family drawn on top. With an "om" key, the density + of that map's poles is drawn as contours: + {"om": om_beta, "pole": (1, 1, 0), "match": 0, "mask": mask_beta, + "label": "<110> beta"}. Otherwise fixed orientations are drawn as + open markers: {"quats": q, "crystal": xtl, "pole": (1, 1, 0), + "label": ...}. Poles are Miller indices of the overlay's crystal. + saturation_power, chroma : float | None + Color wedge shape, used when `color_by` is "ipf"; see + :func:`plot_orientation_map`. + figax : (fig, ax) | (fig, (ax, ax_legend)) | None + Existing axes; the legend panel is used only with `color_by="ipf"`. + + Returns + ------- + tuple + ``(fig, ax)`` with `ax` the pole figure panel. + """ + import matplotlib.pyplot as plt + + assert om.quats is not None + qmap = om.quats[..., match, :] + x, y, wk, src = _pole_points(qmap, om.crystal, pole, mask) + + rng = [[-1.05, 1.05], [-1.05, 1.05]] + H, xe, ye = np.histogram2d(x, y, bins=bins, range=rng, weights=wk) + if smooth_sigma > 0: + from scipy.ndimage import gaussian_filter + + H = gaussian_filter(H, smooth_sigma) + lo, hi = int_range + # normalize against a high percentile of the occupied bins, not the single + # hottest bin -- one large uniform grain would otherwise black out the rest + occupied = H[H > 0] + h_ref = np.percentile(occupied, 98) if occupied.size else 1.0 + Hn = np.clip((H / max(h_ref, 1e-12) - lo) / max(hi - lo, 1e-12), 0, 1) + + if color_by == "ipf": + # white background: blend from white toward the per-position IPF + # color as the histogram density rises + rgb_pos = ipf_color( + om.quats[..., match, :], om.crystal, "z", saturation_power, chroma + ).reshape(-1, 3) + rgb_all = rgb_pos[src] + img = np.zeros((bins, bins, 3)) + cnt = np.zeros((bins, bins)) + ii = np.clip(((x - rng[0][0]) / (rng[0][1] - rng[0][0]) * bins).astype(int), 0, bins - 1) + jj = np.clip(((y - rng[1][0]) / (rng[1][1] - rng[1][0]) * bins).astype(int), 0, bins - 1) + for k in range(3): + np.add.at(img[..., k], (ii, jj), rgb_all[:, k] * wk) + np.add.at(cnt, (ii, jj), wk) + if smooth_sigma > 0: + from scipy.ndimage import gaussian_filter + + for k in range(3): + img[..., k] = gaussian_filter(img[..., k], smooth_sigma) + cnt = gaussian_filter(cnt, smooth_sigma) + img = img / np.maximum(cnt[..., None], 1e-12) + # white background blending toward the IPF color as density rises -- + # keeps dark corner colors (blue) legible + disp = 1.0 - Hn[..., None] * (1.0 - img) + else: + import matplotlib.cm as cm + + # white -> yellow -> red -> black with increasing density + disp = cm.hot_r(Hn)[..., :3] + + # display in the image frame: horizontal = c (col, rightward), vertical = + # r (row, downward), matching the orientation maps -- H is indexed + # [row-bin, col-bin] so no transpose, origin upper + yy, xx = np.meshgrid(0.5 * (ye[:-1] + ye[1:]), 0.5 * (xe[:-1] + xe[1:]), indexing="ij") + disp = disp.copy() + disp[(xx**2 + yy**2).T > 1.0] = 1.0 + + ax_leg = None + if figax is None: + if color_by == "ipf": + fig, (ax, ax_leg) = plt.subplots( + 1, 2, figsize=(7.2, 5.5), gridspec_kw={"width_ratios": [4, 1]} + ) + else: + fig, ax = plt.subplots(figsize=(5.5, 5.5)) + else: + fig, axs = figax + if isinstance(axs, (tuple, list, np.ndarray)) and len(np.atleast_1d(axs)) == 2: + ax, ax_leg = axs + else: + ax = axs + ax.imshow( + disp, + extent=(ye[0], ye[-1], xe[-1], xe[0]), + interpolation="nearest", + ) + if grid: + theta = np.linspace(0, 2 * np.pi, 361) + for pol_deg in (30, 60, 90): + r = np.tan(np.deg2rad(pol_deg) / 2) + lw = 1.0 if pol_deg == 90 else 0.5 + ax.plot(r * np.cos(theta), r * np.sin(theta), color="0.65", lw=lw) + if pol_deg < 90: + ax.text( + r * np.cos(np.deg2rad(45)), + r * np.sin(np.deg2rad(45)), + f"{pol_deg}°", + color="0.45", + fontsize=7, + ha="center", + va="center", + ) + for az in range(0, 180, 30): + ca, sa = np.cos(np.deg2rad(az)), np.sin(np.deg2rad(az)) + ax.plot([-ca, ca], [-sa, sa], color="0.85", lw=0.4) + # compact scan-axes glyph, top-left corner: the pole figure is in + # the scan (image) frame -- c rightward, r downward + gx, gy = -1.06, -1.06 + for dx, dy, lbl, ha, va in ( + (0.22, 0.0, "c", "left", "center"), + (0.0, 0.22, "r", "center", "top"), + ): + ax.annotate( + "", + xy=(gx + dx, gy + dy), + xytext=(gx, gy), + arrowprops=dict(arrowstyle="-|>", color="0.3", lw=1.2), + annotation_clip=False, + ) + ax.text( + gx + dx * 1.25, + gy + dy * 1.25, + lbl, + fontsize=9, + ha=ha, + va=va, + color="0.3", + ) + ax.text( + gx - 0.03, + gy - 0.06, + "scan axes", + fontsize=7, + ha="left", + va="bottom", + color="0.45", + ) + if overlay is not None: + if "om" in overlay: + # raw-histogram contour of another map's pole family + o_om = overlay["om"] + ox, oy, ow, _ = _pole_points( + o_om.quats[..., overlay.get("match", 0), :], + o_om.crystal, + overlay["pole"], + overlay.get("mask"), + ) + Ho, oxe, oye = np.histogram2d(ox, oy, bins=bins, range=rng, weights=ow) + if (Ho > 0).any(): + from scipy.ndimage import gaussian_filter + + Ho = gaussian_filter(Ho, max(smooth_sigma, 1.0)) + lev = np.percentile(Ho[Ho > 0], 99) * np.array([0.3, 0.7]) + xc = 0.5 * (oxe[:-1] + oxe[1:]) + yc = 0.5 * (oye[:-1] + oye[1:]) + # image frame: horizontal = col bins, vertical = row bins + ax.contour( + yc, + xc, + Ho, + levels=lev, + colors="k", + linewidths=[0.7, 1.3], + alpha=0.85, + ) + ax.plot([], [], color="k", lw=1.2, label=overlay.get("label", "overlay")) + else: + oxy = _pole_scatter_xy(overlay["quats"], overlay["crystal"], overlay["pole"]) + ax.scatter( + oxy[:, 1], + oxy[:, 0], + s=80, + marker="D", + facecolors="none", + edgecolors="k", + lw=1.2, + label=overlay.get("label", "overlay"), + ) + ax.legend(loc="upper right", fontsize=8) + ax.set_xlim(-1.12, 1.12) + ax.set_ylim(1.2, -1.2) + ax.set_xticks([]) + ax.set_yticks([]) + ax.set_frame_on(False) + title = f"{om.crystal.name} pole figure" + if label is not None: + title += f" {label}" + ax.set_title(title) + if ax_leg is not None: + if color_by == "ipf": + wedge_legend( + om.crystal, + ax_leg, + orientation="vertical", + saturation_power=saturation_power, + chroma=chroma, + ) + else: + ax_leg.axis("off") + return fig, ax diff --git a/src/quantem/diffraction/phase.py b/src/quantem/diffraction/phase.py new file mode 100644 index 000000000..b683dd1c2 --- /dev/null +++ b/src/quantem/diffraction/phase.py @@ -0,0 +1,690 @@ +"""Phase mapping from matched crystal orientations. + +PhaseMap compares candidate crystal patterns (each carrying orientations +matched and refined by OrientationMap) against the measured Bragg peaks at +every probe position. Every subset of candidates up to `max_patterns` is +scored as a joint model: the candidate pattern weights are solved by +non-negative least squares on the paired peak intensities, and the model cost +extends Diebold et al., Microsc. Microanal. 31, ozaf019 (2025): + + c(S) = sum_exp_peaks |I_m - sum_f w_f I_pred,f| + sum_f w_f I_unpaired,f + + penalty * (|S| - 1) + +normalized by the total measured intensity. Unpaired experimental intensity +appears in the first term (its prediction is zero); unpaired simulated +intensity is charged in full. The best subset answers orientation/phase +ambiguity directly: one orientation, two orientations of one phase, or two +phases, whichever explains the pattern best. Reliability is the cost gap +between the best models with and without the winning phase. +""" + +from __future__ import annotations + +from itertools import combinations + +import numpy as np +import torch +from tqdm import tqdm + +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.defaults import ( + MIN_NUMBER_PEAKS, + MIN_SIM_INTENSITY_REL, + PAIR_DISTANCE, + POWER_INTENSITY, + resolve, +) +from quantem.diffraction.orientation import OrientationMap, position_mask + +# Diffracted intensity is strongly skewed, so a linear brightness leaves most +# indexed positions nearly black. This exponent is applied wherever a map is +# shaded by signal, so the phase map and the orientation maps agree. +SHADE_GAMMA = 0.5 + + +def _majority_filter(phase: np.ndarray, radius: int) -> np.ndarray: + """Replace each position by the most common phase around it. + + Unindexed positions (-1) take part, so an isolated crystal pixel in + vacuum is removed rather than spreading. + + Parameters + ---------- + phase : np.ndarray + ``(scan_row, scan_col)`` phase indices, -1 where unindexed. + radius : int + Half-width of the square neighbourhood in probe positions. + + Returns + ------- + np.ndarray + Filtered phase indices, same shape and dtype. + """ + from scipy.ndimage import uniform_filter + + labels = np.unique(phase) + votes = np.stack( + [ + uniform_filter((phase == v).astype(float), size=2 * radius + 1, mode="nearest") + for v in labels + ] + ) + return labels[votes.argmax(axis=0)] + + +class PhaseMap(AutoSerialize): + """Assign best-fit phases to every probe position. + + Workflow:: + + pm = PhaseMap.from_orientation_maps([om_alpha, om_beta]) + pm.fit() + pm.plot_phase() + + Candidates are all (crystal, match) pairs of the input OrientationMaps, + so two matched orientations of one crystal compete on equal footing with + one orientation of each of two crystals. + """ + + _token = object() + + def __init__(self, orientation_maps: list[OrientationMap], _token=None): + """Private constructor; use :meth:`from_orientation_maps`. + + Enumerates the candidate list, one entry per (orientation map, match) + pair, so that two matched orientations of one crystal compete on equal + footing with one orientation of each of two crystals. The fit results + are left as None until :meth:`fit` runs. + + Parameters + ---------- + orientation_maps : list of OrientationMap + Per-crystal maps sharing one set of peaks. + _token : object + Guard against direct construction. + + Raises + ------ + RuntimeError + If called without the class token. + """ + if _token is not self._token: + raise RuntimeError("Use PhaseMap.from_orientation_maps().") + self.orientation_maps = orientation_maps + self.names = [om.crystal.name for om in orientation_maps] + # candidate list: (map index, match index) + self.candidates: list[tuple[int, int]] = [] + for i, om in enumerate(orientation_maps): + assert om.quats is not None + for m in range(om.quats.shape[2]): + self.candidates.append((i, m)) + + # hyperparameters inherited from the maps and recorded per stage + self.metadata: dict = {"orientation_maps": [dict(om.metadata) for om in orientation_maps]} + # fit results, (R, C, ...) over the scan; see fit() + self.phase_weights: torch.Tensor | None = None # (R, C, F) per candidate + self.crystal_weights: torch.Tensor | None = None # (R, C, n_crystals) + self.cost_best: torch.Tensor | None = None + self.phase_index: torch.Tensor | None = None + self.reliability: torch.Tensor | None = None + self.diffracted_intensity: torch.Tensor | None = None + self.num_diffracted: torch.Tensor | None = None + + @classmethod + def from_orientation_maps(cls, orientation_maps: list[OrientationMap]) -> "PhaseMap": + """Create from OrientationMaps that share the same peaks. + + Parameters + ---------- + orientation_maps : list of OrientationMap + Matched maps, one per candidate crystal. + + Returns + ------- + PhaseMap + + Raises + ------ + ValueError + If the maps have different scan shapes. + RuntimeError + If a map has not been matched. + """ + p0 = orientation_maps[0].peaks + for om in orientation_maps: + if om.peaks.shape != p0.shape: + raise ValueError("All OrientationMaps must share the same scan shape.") + if om.quats is None: + raise RuntimeError( + f"OrientationMap for {om.crystal.name}: run match_orientations() first." + ) + return cls(orientation_maps, _token=cls._token) + + def fit( + self, + positions=None, + pair_distance: float | None = None, + power_intensity: float | None = None, + max_patterns: int = 2, + complexity_penalty: float = 0.02, + weight_unmatched_sim: float = 0.5, + weight_overprediction: float = 1.0, + min_sim_intensity_rel: float | None = None, + k_max: float | None = None, + min_number_peaks: int | None = None, + min_diffracted_peaks: int = 2, + null_k_min: float = 0.05, + progress_bar: bool = True, + ) -> "PhaseMap": + """Score all candidate subsets at every probe position. + + Parameters left as None inherit the values the orientation + matching used (its plan kernel, intensity power and peak minimum), + so the phase decision compares the same peaks the same way; the + resolved values are recorded in `metadata['fit']`. + + Parameters + ---------- + positions : list[tuple[int, int]] | np.ndarray | None + Scan positions to fit: a list of (row, col) or an (R, C) + boolean mask. None (default) fits every position matched by + all of the orientation maps, so a staged test run on a few + positions carries through without repeating the list. + pair_distance : float | None + Pairing distance delta (1/Angstroms) between simulated and + measured peaks; inherits the plan's correlation kernel. + power_intensity : float | None + Intensities are raised to this power before comparison; + inherits the plan's value. + max_patterns : int, default=2 + Maximum number of candidate patterns fit simultaneously. + complexity_penalty : float, default=0.02 + Added cost per extra pattern in a model; sets how much better a + two-pattern fit must be to beat a single-pattern fit. + weight_unmatched_sim : float, default=0.5 + Cost weight of simulated intensity with no experimental partner. + weight_overprediction : float, default=1.0 + Cost weight of predicted intensity in excess of the measured + value on paired peaks. Unexplained measured intensity always + costs in full (coverage), but the symmetric intensity-mismatch + term assumes kinematical intensities are trustworthy; on + non-precession data with strong dynamical scattering, lowering + this weight makes the decision coverage-driven and removes the + bias toward sparse templates. + min_sim_intensity_rel : float | None + Simulated reflections weaker than this fraction of the pattern + maximum are dropped before comparison: kinematically weak spots + are frequently unobservable and should not penalize a phase whose + structure factors happen to include many of them. None takes + MIN_SIM_INTENSITY_REL (0.02). + k_max : float | None + Restrict the comparison below this scattering vector + (1/Angstroms). + min_number_peaks : int | None + Positions with fewer measured peaks (direct beam included) are + not fit and stay unindexed; inherits the matching's value. + min_diffracted_peaks : int, default=2 + Null hypothesis: a position needs at least this many measured + peaks beyond `null_k_min` before any phase is assigned. Vacuum + and amorphous support carry the direct beam and little else, and + a crystal fit to that is noise; those positions are left + unindexed (`phase_index` of -1) and plot black. + null_k_min : float, default=0.05 + Scattering vector (1/Angstroms) above which a measured peak + counts as diffracted. The default excludes the direct beam, + which sits at the origin after `correct_peak_origins`. + progress_bar : bool, default=True + Show a progress bar over the positions. + + Returns + ------- + PhaseMap + Self, with these (R, C, ...) results: + + - `phase_index`: winning crystal, -1 where unindexed (not fit, + rejected by the null hypothesis, or no candidate weight). + - `phase_weights`: (R, C, F) non-negative weight of every + candidate in the best model, F = len(`candidates`). + - `crystal_weights`: (R, C, n_crystals) those weights summed per + crystal and normalized to sum to one, zero where unindexed. + - `cost_best`: cost of the best model, NaN where not fit. + - `reliability`: cost gap to the best model without the winning + crystal, NaN where not fit or no such model exists. + - `diffracted_intensity`, `num_diffracted`: measured intensity + and number of peaks beyond `null_k_min`. + """ + from scipy.optimize import nnls + + oms = self.orientation_maps + plan_md = oms[0].metadata.get("plan") + match_md = oms[0].metadata.get("match") + pair_distance = resolve(pair_distance, "pair_distance", plan_md, default=PAIR_DISTANCE) + power_intensity = resolve( + power_intensity, "power_intensity", plan_md, default=POWER_INTENSITY + ) + min_sim_intensity_rel = resolve( + min_sim_intensity_rel, "min_sim_intensity_rel", default=MIN_SIM_INTENSITY_REL + ) + min_number_peaks = resolve( + min_number_peaks, "min_number_peaks", match_md, default=MIN_NUMBER_PEAKS + ) + self.metadata["fit"] = dict( + positions=None if positions is None else "subset", + pair_distance=float(pair_distance), + power_intensity=float(power_intensity), + max_patterns=int(max_patterns), + complexity_penalty=float(complexity_penalty), + weight_unmatched_sim=float(weight_unmatched_sim), + weight_overprediction=float(weight_overprediction), + min_sim_intensity_rel=float(min_sim_intensity_rel), + k_max=k_max, + min_number_peaks=int(min_number_peaks), + min_diffracted_peaks=int(min_diffracted_peaks), + null_k_min=float(null_k_min), + ) + peaks = oms[0].peaks + R, C = peaks.shape[0], peaks.shape[1] + cands = self.candidates + F = len(cands) + delta = pair_distance + + fields = peaks.fields + ix = [fields.index(f) for f in ("qx", "qy", "intensity")] + + subsets = [s for n in range(1, max_patterns + 1) for s in combinations(range(F), n)] + + cost_best = torch.full((R, C), torch.nan, dtype=torch.float64) + weights_out = torch.zeros((R, C, F), dtype=torch.float64) + reliability = torch.full((R, C), torch.nan, dtype=torch.float64) + diffracted = torch.zeros((R, C), dtype=torch.float64) + num_diffracted = torch.zeros((R, C), dtype=torch.long) + + active = position_mask(positions, (R, C)) + for om in oms: + if om.computed is not None: + active = active & om.computed + iterator = [(rx, ry) for rx, ry in np.ndindex(R, C) if active[rx, ry]] + if progress_bar: + iterator = tqdm(iterator, desc="phase mapping") + for rx, ry in iterator: + data = peaks[rx, ry].numpy().astype(np.float64) + if data.shape[0] < min_number_peaks: + continue + # null hypothesis: no diffracted signal, so no phase to decide. + # Vacuum and amorphous support carry the direct beam and nothing + # else, and the measured signal beyond it is the evidence that + # any crystal is present at all. + qr_meas = np.hypot(data[:, ix[0]], data[:, ix[1]]) + beyond = qr_meas > null_k_min + diffracted[rx, ry] = float(data[beyond, ix[2]].clip(min=0).sum()) + num_diffracted[rx, ry] = int(beyond.sum()) + if int(beyond.sum()) < min_diffracted_peaks: + continue + qxy = torch.as_tensor(data[:, ix[:2]], dtype=torch.float64) + im = torch.as_tensor(data[:, ix[2]], dtype=torch.float64).clamp_min(0) + if k_max is not None: + keep = torch.linalg.norm(qxy, dim=1) <= k_max + qxy, im = qxy[keep], im[keep] + im = im**power_intensity + n_exp = im.shape[0] + int_total = float(im.sum()) + + # per-candidate predicted intensity on each experimental peak, + # and unpaired simulated intensity + pred = np.zeros((n_exp, F)) + unpaired_sim = np.zeros(F) + for f, (i_om, m) in enumerate(cands): + om = oms[i_om] + if om.corr[rx, ry, m] <= 0: + continue + sim = om.generate_pattern(rx, ry, match=m, k_max=k_max) + sq = torch.stack((sim["qx"], sim["qy"]), dim=1) + s_raw = sim["intensity"] + if sq.shape[0] == 0: + continue + vis = s_raw > min_sim_intensity_rel * s_raw.max() + sq, s_raw = sq[vis], s_raw[vis] + si = s_raw**power_intensity + d = torch.cdist(sq, qxy) + d_min, j_min = d.min(dim=1) + pair = d_min < delta + frac = (d_min[pair] / delta).clamp(0, 1) + np.add.at( + pred[:, f], + j_min[pair].numpy(), + (si[pair] * (1 - frac)).numpy(), + ) + unpaired_sim[f] = weight_unmatched_sim * ( + float(si[~pair].sum()) + float((si[pair] * frac).sum()) + ) + + im_np = im.numpy() + results = [] + for s in subsets: + cols = [f for f in s if pred[:, f].any() or unpaired_sim[f] > 0] + if len(cols) == 0: + continue + B = pred[:, cols] + w, _ = nnls(B, im_np) + model = B @ w + under = np.maximum(im_np - model, 0).sum() # unexplained measured + over = np.maximum(model - im_np, 0).sum() # overpredicted paired + cost = (under + weight_overprediction * over + (w * unpaired_sim[cols]).sum()) / ( + int_total + 1e-12 + ) + complexity_penalty * (len(cols) - 1) + results.append((cost, s, cols, w)) + if not results: + continue + results.sort(key=lambda r: r[0]) + c_best, _, cols_best, w_best = results[0] + cost_best[rx, ry] = c_best + for f, w in zip(cols_best, w_best): + weights_out[rx, ry, f] = w + + # reliability: cost gap to the best model containing NO candidate + # of the dominant crystal (candidates of one crystal can be + # near-duplicates, e.g. after residual re-matching) + f_dom = cols_best[int(np.argmax(w_best))] + i_dom = cands[f_dom][0] + others = [c for c, s, _, _ in results if all(cands[f][0] != i_dom for f in s)] + reliability[rx, ry] = (min(others) - c_best) if others else torch.nan + + self.cost_best = cost_best + self.diffracted_intensity = diffracted + self.num_diffracted = num_diffracted + self.phase_weights = weights_out + self.reliability = reliability + + # dominant phase: candidate weights summed per crystal + n_maps = len(oms) + w_phase = torch.zeros((R, C, n_maps), dtype=torch.float64) + for f, (i_om, _) in enumerate(cands): + w_phase[..., i_om] += weights_out[..., f] + # argmax over all-zero weights returns 0, which would label every + # position that was never fit, or whose best model has no weight + # (NNLS returns all zeros when no candidate overlaps the peaks), as + # the first phase; mark them instead + w_sum = w_phase.sum(dim=-1) + self.phase_index = w_phase.argmax(dim=-1) + unindexed = torch.isnan(cost_best) | (w_sum <= 0) + self.phase_index[unindexed] = -1 + self.reliability[unindexed] = torch.nan + self.crystal_weights = w_phase / w_sum[..., None].clamp_min(1e-12) + return self + + def apply_dynamical(self, result: dict) -> "PhaseMap": + """Take the phase decision from a dynamical refinement. + + A phase map can be built from the orientation maps at any stage: + after matching, after refine_orientations, or after + bloch.refine_dynamical, whose per-candidate intensity costs decide + the phase here. The reliability becomes the cost gap between the + best candidates of the winning crystal and of the runner-up + crystal, and the kinematical result is kept under + `metadata['kinematical']`. Positions the refinement did not reach + (outside its `mask`) keep their current decision, so a refinement + of one region, or several in stages, updates only that region. + + Parameters + ---------- + result : dict + Output of :func:`~quantem.diffraction.bloch.refine_dynamical`: + "cost" (R, C, F) per-candidate cost, NaN where not refined, + "phase_index" (R, C) the winning crystal, and optionally + "metadata". + + Returns + ------- + PhaseMap + Self, with `phase_index`, `reliability` and `cost_best` updated + at the refined positions. With one crystal there is no + runner-up and the reliability is NaN, as in :meth:`fit`. + """ + cost = torch.nan_to_num(result["cost"], nan=torch.inf) + n_maps = len(self.orientation_maps) + R, C = cost.shape[:2] + cost_phase = torch.full((R, C, n_maps), torch.inf, dtype=cost.dtype) + for f, (i_om, _) in enumerate(self.candidates): + cost_phase[..., i_om] = torch.minimum(cost_phase[..., i_om], cost[..., f]) + order = cost_phase.sort(dim=-1).values + reliability = torch.where( + torch.isfinite(order[..., 0]), + (order[..., 1] - order[..., 0]).clamp_min(0) + if n_maps > 1 + else torch.full_like(order[..., 0], torch.nan), + torch.full_like(order[..., 0], torch.nan), + ) + done = torch.isfinite(order[..., 0]) + if self.phase_index is None: + self.phase_index = torch.full((R, C), -1, dtype=torch.long) + self.reliability = torch.full((R, C), torch.nan, dtype=cost.dtype) + self.cost_best = torch.full((R, C), torch.nan, dtype=cost.dtype) + # keep the kinematical decision from before the first dynamical pass + self.metadata.setdefault( + "kinematical", + { + "phase_index": self.phase_index.clone(), + "reliability": self.reliability.clone(), + }, + ) + self.phase_index = torch.where(done, result["phase_index"], self.phase_index) + self.reliability = torch.where(done, reliability, self.reliability) + self.cost_best = torch.where(done, order[..., 0], self.cost_best) + self.metadata["dynamical_applied"] = dict(result.get("metadata", {})) + return self + + def signal_confidence( + self, + signal_range: tuple[float, float] | str = "auto", + gamma: float = 1.0, + ) -> np.ndarray: + """Confidence in [0, 1] that a crystal is present, from the data alone. + + The measured intensity beyond the direct beam, scaled to [0, 1]. This + is the null-hypothesis test made visible: vacuum and amorphous support + diffract nothing, so they score zero however well some orientation + happens to correlate. Positions the null hypothesis left unindexed in + :meth:`fit` are forced to zero. + + Use it to fade the phase map and the orientation maps together:: + + conf = pm.signal_confidence() + mask_a = (pm.phase_index.numpy() == 0) * conf + + Parameters + ---------- + signal_range : tuple | "auto", default="auto" + Diffracted intensity mapped to 0 ... 1. "auto" spans zero to half + the median over the indexed positions, so a crystal shows at full + strength and only weakly diffracting positions fade. + gamma : float, default=1.0 + Exponent applied to the result. The default returns the raw + confidence, which is what a threshold should be taken on; + :attr:`SHADE_GAMMA` is the value used when shading a display, and + zero stays zero either way. + + Returns + ------- + np.ndarray + ``(scan_row, scan_col)`` confidence in [0, 1]. + """ + if gamma <= 0: + raise ValueError(f"gamma must be positive, got {gamma}") + if self.diffracted_intensity is None: + raise ValueError("Run fit() before signal_confidence().") + sig = np.nan_to_num(self.diffracted_intensity.numpy()) + indexed = self.phase_index.numpy() >= 0 + if isinstance(signal_range, str): + # full brightness from half the median signal of the indexed + # positions: crystals show at full strength and only positions + # that diffract well below typical fade + vals = sig[indexed] + hi = 0.5 * float(np.median(vals)) if vals.size else 1.0 + lo, hi = 0.0, max(hi, 1e-12) + else: + lo, hi = signal_range + conf = ((sig - lo) / max(hi - lo, 1e-12)).clip(0, 1) * indexed + return conf if gamma == 1.0 else np.power(conf, gamma) + + def plot_phase( + self, + phase_colors: np.ndarray | None = None, + shade_by: str = "signal", + shade_range: tuple[float, float] | str = "auto", + shade_gamma: float = SHADE_GAMMA, + majority_filter: int = 0, + reliability_range: tuple[float, float] | None = None, + scalebar: dict | str | None = "auto", + figax=None, + ): + """Dominant-phase map: color gives the crystal, brightness the evidence. + + By default the brightness is the diffracted signal, so vacuum and + unindexed positions are black; see `shade_by`. + + Parameters + ---------- + phase_colors : np.ndarray | None + One RGB color per phase, cycled when there are more phases; + defaults to `DEFAULT_PHASE_COLORS`, the palette shared with the + pattern overlay plots (gold, cyan, green, purple). + shade_by : {"signal", "reliability", "none"}, default="signal" + What the brightness means. "signal" fades each position by the + measured diffracted intensity (see :meth:`signal_confidence`), so + vacuum and amorphous support go black and the map shows where + crystals actually are. "reliability" uses the cost gap to the best + model without the winning crystal, which answers a different + question -- which phase, given that there is one -- and carries no + information about whether anything is there. "none" draws every + indexed position at full color. + shade_range : tuple | "auto", default="auto" + Values mapped to black ... full color. "auto" takes a high + percentile over the indexed positions, since the absolute scale + depends on the data. + shade_gamma : float, default=:attr:`SHADE_GAMMA` (0.5) + Exponent applied to the brightness, ``alpha ** shade_gamma``. + Diffracted intensity is strongly skewed, so a linear scale leaves + most indexed positions dark and only the brightest grains + readable. Values below 1 lift the faint ones: 0.5 is the default + and 1.0 restores the linear scale. Zero brightness is a fixed + point, so vacuum and unindexed positions stay black however low + this is set, and the colorbars carry the same curve. + majority_filter : int, default=0 + Radius in probe positions of a majority filter applied to the + phase decision for display only; the stored decision is + untouched. 1 replaces each position by the most common phase in + its 3x3 neighbourhood, which removes isolated single-pixel + phases without moving a real boundary. Orientation smoothing + does not do this: it averages orientations within one phase and + leaves the phase assignment alone. + reliability_range : tuple | None + Backwards-compatible shortcut: setting it selects + ``shade_by="reliability"`` with this range. + scalebar : dict | "auto" | None + Real-space scale bar. "auto" (the default) takes the scan step + and units carried from the dataset by the orientation maps; a + dict such as {"sampling": 30, "units": "A"} overrides it, and + None draws no bar. + figax : (fig, ax) | None + Existing axes to draw into. + + Returns + ------- + tuple + ``(fig, ax)``. + + Raises + ------ + ValueError + If `shade_by` is unknown or `shade_gamma` is not positive. + """ + if isinstance(scalebar, str): + scalebar = self.orientation_maps[0].scan_scalebar if scalebar == "auto" else None + import matplotlib.pyplot as plt + + from quantem.core.visualization.visualization_utils import add_scalebar_to_ax + from quantem.diffraction.orientation_visualization import phase_color_cycle + + assert self.phase_index is not None and self.reliability is not None + phase_colors = phase_color_cycle(len(self.names), phase_colors) + if reliability_range is not None: + shade_by, shade_range = "reliability", reliability_range + phase = self.phase_index.numpy() + indexed = phase >= 0 + if shade_by == "signal": + alpha = self.signal_confidence(shade_range) + lo, hi = 0.0, 1.0 + cbar_label = "diffracted signal" + elif shade_by == "reliability": + rel = np.nan_to_num(self.reliability.numpy(), nan=0.0) + if isinstance(shade_range, str): + vals = rel[indexed] + hi = float(np.percentile(vals, 98)) if vals.size else 1.0 + lo, hi = 0.0, max(hi, 1e-12) + else: + lo, hi = shade_range + alpha = ((rel - lo) / max(hi - lo, 1e-12)).clip(0, 1) * indexed + cbar_label = "reliability" + elif shade_by == "none": + alpha = indexed.astype(float) + lo, hi = 0.0, 1.0 + cbar_label = "indexed" + else: + raise ValueError( + f"shade_by must be 'signal', 'reliability' or 'none', got {shade_by!r}" + ) + if shade_gamma <= 0: + raise ValueError(f"shade_gamma must be positive, got {shade_gamma}") + # zero maps to zero under any positive exponent, so unindexed positions + # stay black and only the faint indexed ones are lifted + alpha = np.power(alpha, shade_gamma) + if majority_filter > 0: + # the filter can turn an indexed position unindexed (-1) and the + # reverse, so the colors and the black mask follow the filtered + # decision; a position it newly indexes has no brightness of its + # own and stays black + phase = _majority_filter(phase, int(majority_filter)) + alpha = alpha * (phase >= 0) + rgb = phase_colors[np.where(phase >= 0, phase, 0)] * alpha[..., None] + + if figax is None: + fig, ax = plt.subplots(figsize=(9, 4.5)) + else: + fig, ax = figax + ax.imshow(rgb, interpolation="nearest") + ax.set_xticks([]) + ax.set_yticks([]) + if scalebar is not None: + add_scalebar_to_ax( + ax, + array_size=rgb.shape[1], + sampling=scalebar.get("sampling", 1.0), + length_units=scalebar.get("length", None), + units=scalebar.get("units", "pixels"), + width_px=rgb.shape[0] / 40, + pad_px=rgb.shape[0] / 80, + color=scalebar.get("color", "white"), + loc="lower right", + ) + handles = [ + plt.Line2D([0], [0], marker="s", ls="", color=c, label=n) + for c, n in zip(phase_colors, self.names) + ] + ax.legend(handles=handles, loc="upper left", fontsize=8) + # stacked reliability colorbars, black -> phase color + from matplotlib.cm import ScalarMappable + from matplotlib.colors import LinearSegmentedColormap, Normalize + + n_ph = len(phase_colors) + for k, color in enumerate(phase_colors): + cmap_k = LinearSegmentedColormap.from_list( + f"rel{k}", [(0, 0, 0), tuple(color)], gamma=shade_gamma + ) + cax = ax.inset_axes([1.02 + 0.025 * k, 0.05, 0.025, 0.9]) + cb = fig.colorbar(ScalarMappable(norm=Normalize(lo, hi), cmap=cmap_k), cax=cax) + if k < n_ph - 1: + cb.set_ticks([]) + else: + cb.set_ticks([lo, hi]) + cb.set_label(cbar_label, fontsize=9) + return fig, ax diff --git a/src/quantem/diffraction/reverse_monte_carlo.py b/src/quantem/diffraction/reverse_monte_carlo.py new file mode 100644 index 000000000..ebda4cc25 --- /dev/null +++ b/src/quantem/diffraction/reverse_monte_carlo.py @@ -0,0 +1,2722 @@ +"""Reverse Monte Carlo fitting of diffuse electron scattering from several zone axes. + +One periodic supercell of a disordered crystal is fitted to every pattern at +once. Its amplitude ``F(q) = sum_j f_s(j)(q) exp(-2 pi i q.(r_j + u_j))`` +lives on the FFT grid of the (displaced) atom positions, with the Bragg +nodes of the average lattice removed; the diffuse intensity of each pattern +is ``|F|^2`` read where the pattern's Ewald sphere (with its fitted tilt) +cuts the grid, averaged over the cubic rotations so the model is as +symmetric as the (statistically cubic) foil. Each Monte Carlo move changes +``F`` by a few phase factors, so every move is scored exactly without +recomputing the supercell. The move types are species swaps, single-atom +displacements on the fine position grid and, optionally, omega embryos (see +`ReverseMonteCarlo.run`). + +Bragg peaks with their tails and the direct-beam bloom are masked by sigmoid +weights; a smooth background (constant, direct-beam Lorentzian, a wide +Gaussian about the zone-axis pole, Einstein thermal diffuse and powder rings) +and the diffuse envelope are solved in closed form. + +Electrons barely separate species of neighbouring atomic number (Nb and Zr +differ by a few percent in scattering factor), so their relative arrangement +is set by the moves' randomness unless something else tells them apart. +``set_size_effect`` / ``fit_size_effect`` add the linear size effect: each +species pushes its neighbours by its misfit through harmonic springs, and the +resulting displacement field (Huang and size-effect scattering, odd about +every Bragg peak) depends on which species sits where. + +Workflow: ``from_images`` -> ``set_crystal`` -> ``fit_geometry`` -> +``fit_thickness`` -> ``set_mask`` -> ``build_supercell`` -> ``fit_background`` +-> ``run`` -> analysis and plots. ``set_envelope``, ``set_size_effect`` and +``fit_size_effect`` are optional. See `ReverseMonteCarlo` for details. + +Limitations: cubic unit cells only; every mixed-occupancy site must share one +composition; static displacements (and ``displacement_correlations``) are +implemented for BCC site lattices only. +""" + +from __future__ import annotations + +from collections.abc import Sequence +from itertools import permutations, product + +import numpy as np +import torch +from scipy import ndimage, optimize, sparse +from scipy.spatial import cKDTree +from tqdm.auto import tqdm + +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.crystal import Crystal, electron_scattering_factor + + +def electron_wavelength(energy_ev: float) -> float: + """Relativistic electron wavelength in Angstroms.""" + return 12.2642598 / np.sqrt(energy_ev * (1.0 + 0.97847573e-6 * energy_ev)) + + +def cubic_rotations() -> np.ndarray: + """The 24 proper rotations of the cube as signed permutation matrices (24, 3, 3).""" + ops = [] + for perm in permutations(range(3)): + for signs in product((1, -1), repeat=3): + m = np.zeros((3, 3), dtype=int) + m[range(3), perm] = signs + if round(np.linalg.det(m)) == 1: + ops.append(m) + return np.stack(ops) + + +def _zone_frame(zone_axis) -> np.ndarray: + """Orthonormal crystal-frame basis (e1, e2, z) with z along the zone axis.""" + z = np.asarray(zone_axis, dtype=float) + z /= np.linalg.norm(z) + trial = np.eye(3)[np.argmin(np.abs(z))] + e1 = trial - (trial @ z) * z + e1 /= np.linalg.norm(e1) + return np.stack([e1, np.cross(z, e1), z]) + + +def _rot2(theta: float) -> np.ndarray: + c, s = np.cos(theta), np.sin(theta) + return np.array([[c, -s], [s, c]]) + + +def _sigmoid(x): + return 0.5 * (1.0 + np.tanh(0.5 * x)) + + +def _default_device() -> str: + """cuda when present, otherwise cpu (mps only on request: long runs heat a laptop).""" + return "cuda" if torch.cuda.is_available() else "cpu" + + +class ReverseMonteCarlo(AutoSerialize): + """Reverse Monte Carlo fit of one supercell to diffraction patterns along several zone axes. + + Notes + ----- + Workflow: + + 1. ``from_images``: patterns, their zone axes, beam energy and pixel size. + 2. ``set_crystal``: average structure with the mixed-occupancy sites. + 3. ``fit_geometry``: center, detector distortion and tilt of each pattern. + 4. ``fit_thickness``: Bloch-wave thickness and tilt from the Bragg + intensities (required for ``envelope="bloch"``; recommended otherwise + since it refines the tilts). + 5. ``set_mask``: diffuse weights, binned data and Bragg intensities. + 6. ``build_supercell``: random supercell at the crystal's composition. + 7. ``fit_background``: diffuse scale and smooth background per pattern. + 8. ``run``: Monte Carlo sweeps. + 9. Analysis (``r_factors``, ``warren_cowley``, ``displacement_correlations``, + ``diffuse_section``, ...) and ``plot_*`` methods. + + Optional steps after ``build_supercell``: ``set_envelope`` switches the + diffuse envelope; ``set_size_effect`` / ``fit_size_effect`` add the linear + size effect. + + Moves: ``run`` mixes three move types on the current supercell: + + - species swaps of two unlike atoms (composition conserved); + - single-atom random displacements by one step of the fine position grid + along each axis (``random_fraction``, default 0.5, when the supercell + was built with ``displacements=True``); + - omega embryos, three consecutive atoms of a <111> row with the last two + collapsed toward each other (``omega_fraction``, default 0). + + Limitations: cubic unit cells only. Every mixed-occupancy site must share one + composition. Static displacements and ``displacement_correlations`` are + implemented for BCC site lattices only. + + Saving and loading: save without the arrays that can be rebuilt, then load with + `quantem.core.io.load`, which calls ``_post_load`` to rebuild them:: + + rmc.save(path, mode="o", skip=rmc.DERIVED_ATTRIBUTES) + rmc = quantem.core.io.load(path) + """ + + _token = object() + + def __init__(self, images, zone_axes, sampling, energy, names, bin_factor, _token=None): + if _token is not self._token: + raise RuntimeError("Use ReverseMonteCarlo.from_images().") + self.images = [np.asarray(im, dtype=np.float32) for im in images] + self.zone_axes = [tuple(int(v) for v in z) for z in zone_axes] + self.sampling = float(sampling) + self.energy = float(energy) + self.wavelength = electron_wavelength(self.energy) + self.names = list(names) + self.bin_factor = int(bin_factor) + self.crystal: Crystal | None = None + self.geometry: dict | None = None + self.mask: dict | None = None + self.loss_history: list[float] = [] + + @classmethod + def from_images( + cls, + images: Sequence, + zone_axes: Sequence[Sequence[int]], + energy: float = 200e3, + bin_factor: int = 8, + sampling: float | None = None, + names: Sequence[str] | None = None, + ) -> "ReverseMonteCarlo": + """Patterns (Dataset2d or arrays) and the zone axis of each. + + Parameters + ---------- + images : sequence of Dataset2d or ndarray + One diffraction pattern per zone axis. + zone_axes : sequence of (u, v, w) + Nominal zone axis of each pattern, in the crystal's lattice + indices; the tilt off it is fitted. + energy : float + Beam energy in eV. + bin_factor : int + Detector binning for the diffuse fit (geometry uses full resolution). + sampling : float, optional + Detector pixel size in 1/Angstrom. Read from the first Dataset2d + (1/nm is converted) when not given. + names : sequence of str, optional + Panel titles; default the zone axes. + """ + arrays = [] + for im in images: + if hasattr(im, "array"): + if sampling is None: + s = float(np.asarray(im.sampling)[0]) + units = str(im.units[0]).lower() + sampling = s * 0.1 if "nm" in units else s + arrays.append(np.asarray(im.array)) + else: + arrays.append(np.asarray(im)) + if sampling is None: + raise ValueError("sampling (1/Angstrom per pixel) is required for plain arrays.") + if len(arrays) != len(zone_axes): + raise ValueError("one zone axis per image") + if names is None: + names = ["[" + "".join(str(v) for v in z) + "]" for z in zone_axes] + return cls(arrays, zone_axes, sampling, energy, names, bin_factor, _token=cls._token) + + def _binned(self, a: np.ndarray, reduce: str = "mean") -> np.ndarray: + b = self.bin_factor + ny, nx = (a.shape[0] // b) * b, (a.shape[1] // b) * b + out = a[:ny, :nx].reshape(ny // b, b, nx // b, b).sum(axis=(1, 3)) + return out / b**2 if reduce == "mean" else out + + # rebuilt from the saved state by _post_load: pass as ``skip`` to ``save`` to keep files small + DERIVED_ATTRIBUTES = ( + "_W_all", "_Wc", "_Wv", "_sym_index", "_needed", "_used", "_h_needed", "_cos", "_sin", + "_Fr", "_Fi", "_fs", "_keep", "_u", "_u_all", "_y", "_w", "_y_all", "_w_all", "_img", + "_fit", "_r", "_bg", "_scale", "_basis_all", "_pix_image", "_env_cols", "_site_lookup", + ) # fmt: skip + + def _post_load(self) -> None: + """AutoSerialize hook: rebuild the derived arrays skipped at save time.""" + if getattr(self, "site_x", None) is None or hasattr(self, "_W_all"): + return + self.device = torch.device(_default_device()) + nc = self.cells * self.grid_divisor + self._site_lookup = np.full((nc,) * 3, -1, dtype=np.int64) + xc = self.site_x // self.refine + self._site_lookup[xc[:, 0], xc[:, 1], xc[:, 2]] = np.arange(len(self.site_x)) + env_p = getattr(self, "_env_p", None) + self._setup_forward() + if self.envelope == "fitted" and env_p is not None: + self._env_p = list(env_p) + self._u_all = np.concatenate( + [self._env_cols[i] @ np.asarray(self._env_p[i]) for i in range(len(self.images))] + ) + self._u = torch.as_tensor( + self._u_all[self._fit], dtype=torch.float32, device=self.device + ) + self._update_residual() + + # ------------------------------------------------------------------ crystal + + def set_crystal( + self, + crystal: Crystal | None = None, + cif_file: str | None = None, + merge: dict[str, str] | None = None, + ) -> "ReverseMonteCarlo": + """Crystal whose mixed-occupancy sites are fitted. + + Parameters + ---------- + crystal, cif_file : Crystal or path + The average structure, with fractional occupancies on shared sites. + merge : dict, optional + Species to relabel before fitting, e.g. ``{"Zr": "Nb"}`` folds Zr + into Nb (their electron scattering factors differ by a few percent, + so the patterns barely tell them apart). + """ + from ase.data import atomic_numbers, chemical_symbols + + if crystal is None: + if cif_file is None: + raise ValueError("give crystal or cif_file") + crystal = Crystal.from_cif(cif_file, verbose=False) + self.crystal = crystal + cell = crystal.lat_real.numpy() + a = float(np.linalg.norm(cell[0])) + if not np.allclose(cell, a * np.eye(3), atol=1e-3 * a): + raise NotImplementedError("Only cubic cells are supported so far.") + merge = merge or {} + frac = np.mod(crystal.positions_frac.numpy(), 1.0) + symbols = [chemical_symbols[int(z)] for z in crystal.numbers] + symbols = [merge.get(s, s) for s in symbols] + occ = crystal.occupancy.numpy() + + # group species by site + sites: list[tuple[np.ndarray, dict[str, float]]] = [] + for f, s, o in zip(frac, symbols, occ): + for site in sites: + if np.allclose(site[0], f, atol=1e-4): + site[1][s] = site[1].get(s, 0.0) + float(o) + break + else: + sites.append((f, {s: float(o)})) + mixed = [(f, comp) for f, comp in sites if len(comp) > 1] + if not mixed: + raise ValueError("The crystal has no mixed-occupancy site to fit.") + species = sorted({s for _, comp in mixed for s in comp}) + comps = {tuple(sorted(comp.items())) for _, comp in mixed} + if len(comps) != 1: + raise NotImplementedError("All mixed sites must share one composition.") + comp = dict(next(iter(comps))) + total = sum(comp.values()) + + # smallest grid divisor that puts every mixed site on an integer grid + fr = np.stack([f for f, _ in mixed]) + for d in range(1, 13): + if np.allclose(fr * d, np.round(fr * d), atol=1e-4): + break + else: + raise ValueError("Mixed sites are not on a rational grid with denominator <= 12.") + + self._a_crystal = a + self.lattice_parameter = a + self.species = species + self.numbers = [atomic_numbers[s] for s in species] + self.concentrations = np.array([comp[s] / total for s in species]) + self.site_grid = np.round(fr * d).astype(int) + self.grid_divisor = d + print( + f"{len(mixed)} mixed site(s) per cell, " + + ", ".join(f"{s} {comp[s] / total:.3f}" for s in species) + + f", a = {a:.4f} A" + ) + return self + + def _zone_reflections(self, zone_axis, k_max: float): + """Allowed reflections in a zone: hkl (n, 3), zone-frame coords at the CIF a (n, 2), |F|^2.""" + crystal = self.crystal + crystal.calculate_structure_factors(k_max) + hkl = crystal.hkl.numpy() + inten = crystal.struct_factors_int.numpy() + keep = (hkl @ np.asarray(zone_axis) == 0) & (inten > 1e-4 * inten.max()) + hkl, inten = hkl[keep], inten[keep] + frame = _zone_frame(zone_axis) + g = hkl / self._a_crystal + return hkl, g @ frame[:2].T, inten + + # ----------------------------------------------------------------- geometry + + @staticmethod + def _find_peaks(im, n_peaks: int = 150): + smooth = ndimage.gaussian_filter(im, 2.0) + prom = smooth - ndimage.gaussian_filter(im, 25.0) + local = (prom == ndimage.maximum_filter(prom, 15)) & (prom > 0) + r, c = np.nonzero(local) + order = np.argsort(prom[r, c])[::-1][:n_peaks] + r, c = r[order], c[order] + pts = [] + for ri, ci in zip(r, c): + r0, r1 = max(ri - 4, 0), min(ri + 5, im.shape[0]) + c0, c1 = max(ci - 4, 0), min(ci + 5, im.shape[1]) + w = np.clip(prom[r0:r1, c0:c1] - 0.3 * prom[ri, ci], 0, None) + rr, cc = np.mgrid[r0:r1, c0:c1] + pts.append([(w * rr).sum() / w.sum(), (w * cc).sum() / w.sum()]) + return np.asarray(pts), prom[r, c], prom + + @staticmethod + def _halo_center(im) -> np.ndarray: + """Center of the broad inelastic halo, which sits on the direct beam.""" + small = ndimage.median_filter(im[::4, ::4].astype(np.float64), size=9) + b = ndimage.gaussian_filter(small, 10) + return np.asarray(np.unravel_index(np.argmax(b), b.shape), dtype=float) * 4 + 1.5 + + def fit_geometry( + self, + scale_range: tuple[float, float] = (0.85, 1.2), + k_max: float = 1.6, + centers: Sequence | None = None, + fit_tilt: bool = True, + verbose: bool = True, + ) -> "ReverseMonteCarlo": + """Index every pattern, fit its detector distortion and its tilt off the zone axis. + + The direct beam is the detected peak nearest the center of the broad + inelastic halo (a tilted pattern can have diffracted beams brighter + than the direct beam), or nearest ``centers[i]`` (row, col). Each + pattern then gets its own center and 2x2 detector matrix (rotation, + scale, ellipticity), fitted to the matched peaks. The lattice + parameter is the mean over patterns at the nominal pixel size. + + The tilt (beam direction off the zone axis, small-angle vector in the + zone frame) is fitted to the Bragg intensities: each reflection's + excitation error is ``s = -(|g|^2 / 2K + tilt . g)`` and its intensity + ``|F|^2 exp(-s^2 / 2 sigma^2)``. Zero tilt puts the Laue circle on the + direct beam. + + Parameters + ---------- + scale_range : (float, float), optional + Range of the pixel-size scale factor searched in the coarse + indexing step, relative to ``sampling``. Default (0.85, 1.2). + k_max : float, optional + Largest scattering vector (1/A) of the reflections used for + indexing and tilt fitting. Default 1.6. + centers : sequence of (row, col) or None, optional + Approximate direct-beam position of each pattern in pixels; None + entries (or ``centers=None``) use the halo center. + fit_tilt : bool, optional + Fit the tilt of each pattern. If False, the tilts are zero. + verbose : bool, optional + Print the fitted geometry of each pattern. + + Returns + ------- + ReverseMonteCarlo + self. The results are stored in ``self.geometry``, a dict of + per-pattern lists: "centers" (row, col) px, "matrices" (2x2, px + per 1/A at the CIF lattice parameter), "tilts" (zone-frame vector, + radians), "a" (lattice parameter, A), "rms_px", "n_matched", + "peaks", "bragg_hkl", "bragg_g", "bragg_px", "bragg_intensity" and + "excitation_width" (1/A). "excitation_width" is the width of the + Gaussian excitation-error profile fitted with the tilt; it is a + diagnostic only and is not used later. ``self.lattice_parameter`` + is set to the mean of "a". + """ + if self.crystal is None: + raise RuntimeError("set_crystal first") + pix = self.sampling + k_wave = 1.0 / self.wavelength + geo = dict(centers=[], matrices=[], tilts=[], a=[], rms_px=[], n_matched=[], peaks=[]) + geo.update(bragg_hkl=[], bragg_g=[], bragg_intensity=[], bragg_px=[], excitation_width=[]) + geo["_inten_kin"] = [] # kinematic |F|^2 of each pattern's reflections, for _fit_tilt + for i, (im, zone) in enumerate(zip(self.images, self.zone_axes)): + pts, heights, prom = self._find_peaks(im) + guess = ( + np.asarray(centers[i], dtype=float) + if centers is not None and centers[i] is not None + else self._halo_center(im) + ) + center = pts[np.argmin(np.linalg.norm(pts - guess, axis=1))] + hkl, g2, inten = self._zone_reflections(zone, k_max) + nz = np.linalg.norm(hkl, axis=1) > 0 + g2_nz, inten_nz = g2[nz], inten[nz] + + # coarse search: in-plane rotation and scale + score_img = ndimage.gaussian_filter(np.clip(prom, 0, None), 3.0) + wts = np.sqrt(inten_nz) + best = (-np.inf, 0.0, 1.0) + for scale in np.arange(scale_range[0], scale_range[1] + 1e-9, 0.004): + for th in np.deg2rad(np.arange(0.0, 360.0, 0.5)): + p = center + (g2_nz @ _rot2(th).T) / (pix * scale) + ok = ( + (p[:, 0] >= 0) + & (p[:, 0] < im.shape[0] - 1) + & (p[:, 1] >= 0) + & (p[:, 1] < im.shape[1] - 1) + ) + if ok.sum() < 4: + continue + pi = np.round(p[ok]).astype(int) + s = (wts[ok] * score_img[pi[:, 0], pi[:, 1]]).sum() / wts[ok].sum() + if s > best[0]: + best = (s, th, scale) + A = _rot2(best[1]) / (pix * best[2]) + c = center.copy() + + # refine center + 2x2 matrix on matched peaks, tightening the match + for tol in (12.0, 8.0, 5.0, 5.0): + p = c + g2_nz @ A.T + dist, j = cKDTree(pts).query(p) + ok = dist < tol + obs = pts[j[ok]] + X = np.column_stack([np.ones(ok.sum()), g2_nz[ok]]) + coef, *_ = np.linalg.lstsq(X, obs, rcond=None) + c, A = coef[0], coef[1:].T + p = c + g2_nz @ A.T + dist, j = cKDTree(pts).query(p) + ok = dist < 5.0 + rms = float(np.sqrt(np.mean(dist[ok] ** 2))) + a_i = self._a_crystal / (pix * np.sqrt(abs(np.linalg.det(A)))) + sv = np.linalg.svd(A, compute_uv=False) + + # Bragg intensities at the fitted positions + p_all = c + g2 @ A.T + r_core = 0.03 / pix + inten_meas = _integrate_spots(im, p_all, r_core) + + geo["centers"].append(c) + geo["matrices"].append(A) + geo["a"].append(float(a_i)) + geo["rms_px"].append(rms) + geo["n_matched"].append(int(ok.sum())) + geo["peaks"].append(pts) + geo["bragg_hkl"].append(hkl) + geo["bragg_g"].append(g2) # zone frame, at the CIF lattice parameter + geo["bragg_px"].append(p_all) + geo["bragg_intensity"].append(inten_meas) + geo["_inten_kin"].append(inten) + if verbose: + print( + f"{self.names[i]}: center ({c[0]:.1f}, {c[1]:.1f}), a = {a_i:.4f} A, " + f"anisotropy {100 * (sv[0] / sv[1] - 1):.2f}%, " + f"{int(ok.sum())} peaks, rms {rms:.2f} px" + ) + + self.lattice_parameter = float(np.mean(geo["a"])) + self.geometry = geo + for i in range(len(self.images)): + tilt, width = (np.zeros(2), np.nan) + if fit_tilt: + tilt, width = self._fit_tilt(i, k_wave) + geo["tilts"].append(tilt) + geo["excitation_width"].append(width) + if verbose and fit_tilt: + ang = np.rad2deg(np.linalg.norm(tilt)) + print(f"{self.names[i]}: tilt {ang:.2f} deg off the zone axis") + if verbose: + print(f"lattice parameter {self.lattice_parameter:.4f} A at {pix:.6f} 1/A per pixel") + return self + + def _fit_tilt(self, i: int, k_wave: float, max_tilt_deg: float = 3.5, prior_deg: float = 2.0): + """Tilt vector (zone frame, radians) from the Bragg intensities of pattern i, with a + Gaussian prior of ``prior_deg`` so patterns whose intensities barely constrain it stay + near the zone axis.""" + geo = self.geometry + hkl = geo["bragg_hkl"][i] + nz = np.linalg.norm(hkl, axis=1) > 0 + g = geo["bragg_g"][i][nz] * self._a_crystal / self.lattice_parameter + p_px = geo["bragg_px"][i][nz] + ny, nx = self.images[i].shape + inside = ( + (p_px[:, 0] > 20) & (p_px[:, 0] < ny - 20) & (p_px[:, 1] > 20) & (p_px[:, 1] < nx - 20) + ) + g = g[inside] + meas = geo["bragg_intensity"][i][nz][inside] + kin = geo["_inten_kin"][i][nz][inside] + ok = np.isfinite(meas) + g, meas, kin = g[ok], np.clip(meas[ok], 0, None), kin[ok] + y = np.sqrt(meas / meas.max()) + g2 = (g**2).sum(1) / (2 * k_wave) + + def model(x): + tilt, log_w, log_a = x[:2], x[2], x[3] + s = -(g2 + g @ tilt) + return ( + np.exp(log_a) * np.sqrt(kin / kin.max()) * np.exp(-0.25 * (s / np.exp(log_w)) ** 2) + ) + + best = None + for tx in np.linspace(-0.06, 0.06, 25): + for ty in np.linspace(-0.06, 0.06, 25): + for lw in (np.log(0.01), np.log(0.03)): + x = np.array([tx, ty, lw, 0.0]) + r = ((model(x) - y) ** 2).sum() + ((x[:2] / np.deg2rad(prior_deg)) ** 2).sum() + if best is None or r < best[0]: + best = (r, x) + lim = np.deg2rad(max_tilt_deg) + lo = np.array([-lim, -lim, np.log(0.003), -5.0]) + hi = np.array([lim, lim, np.log(0.05), 5.0]) + prior = np.deg2rad(prior_deg) + + def resid(x): + return np.concatenate([model(x) - y, x[:2] / prior]) + + sol = optimize.least_squares( + resid, np.clip(best[1], lo + 1e-9, hi - 1e-9), bounds=(lo, hi) + ) + return sol.x[:2], float(np.exp(sol.x[2])) + + def _orientation_quat(self, i: int, tilt: np.ndarray) -> torch.Tensor: + """Quaternion rotating crystal vectors into the lab frame of pattern i at a tilt. + + The zone axis is turned toward ``(tilt_x, tilt_y, 1)``, which puts the Bloch excitation + errors on the same Laue circle as the diffuse model's Ewald sphere.""" + from quantem.diffraction.rotations import quat_from_matrix + + n = np.array([tilt[0], tilt[1], 1.0]) + n /= np.linalg.norm(n) + z = np.array([0.0, 0.0, 1.0]) + axis = np.cross(z, n) + s_ang, c_ang = np.linalg.norm(axis), n[2] + if s_ang < 1e-12: + rot = np.eye(3) + else: + k = axis / s_ang + kx = np.array([[0, -k[2], k[1]], [k[2], 0, -k[0]], [-k[1], k[0], 0]]) + rot = np.eye(3) + s_ang * kx + (1 - c_ang) * kx @ kx + u = rot @ _zone_frame(self.zone_axes[i]) + return quat_from_matrix(torch.as_tensor(u, dtype=torch.float64)) + + def _measured_bragg(self, i: int): + """Measured integrated intensities of pattern i: hkl (n, 3) with the direct beam first.""" + hkl = np.asarray(self.geometry["bragg_hkl"][i]) + meas = np.asarray(self.geometry["bragg_intensity"][i], dtype=float) + nz = np.linalg.norm(hkl, axis=1) > 0 + hkl, meas = hkl[nz], meas[nz] + i000 = _integrate_spots( + self.images[i], self.geometry["centers"][i][None], 0.03 / self.sampling + )[0] + hkl = np.vstack([[0, 0, 0], hkl]) + meas = np.concatenate([[i000], meas]) + ok = np.isfinite(meas) + return hkl[ok].astype(int), np.clip(meas[ok], 0, None) + + def fit_thickness( + self, + thickness: tuple[float, float] = (20.0, 1000.0), + step: float = 10.0, + tilt_range_deg: float = 0.6, + tilt_step_deg: float = 0.1, + k_max: float = 1.6, + depth_samples: int = 24, + verbose: bool = True, + ) -> dict: + """Bloch-wave thickness and tilt of every pattern from its Bragg intensities. + + For each pattern, a grid of tilts about the current one and of thicknesses is scored by + comparing the Bloch exit intensities of the zone reflections with the measured + integrated intensities (each set normalized to unit sum, compared as square roots). + Absorptive (Weickenmeier-Kohl) structure factors of the average crystal are used. The + best thickness and tilt are stored, with each beam's intensity averaged over depth, + ``(1/t) int_0^t |phi_g(z)|^2 dz``: the beams that generate diffuse scattering inside + the foil, used by ``envelope="bloch"``. + + Parameters + ---------- + thickness : (float, float), optional + Thickness range searched, in A. Default (20, 1000). + step : float, optional + Thickness step, in A. Default 10. + tilt_range_deg : float, optional + Half width of the tilt search about the current tilt, along each + zone-frame axis, in degrees. Default 0.6. + tilt_step_deg : float, optional + Tilt search step, in degrees. Default 0.1. + k_max : float, optional + Largest scattering vector (1/A) of the Bloch-wave beams. Default + 1.6. Stored in ``geometry["thickness_k_max"]`` for + ``plot_thickness``. + depth_samples : int, optional + Depths at which the beam intensities are averaged. Default 24. + verbose : bool, optional + Print the result for each pattern. + + Returns + ------- + dict + ``{name: {"thickness": A, "tilt_deg": degrees off the zone axis}}`` + per pattern. ``self.geometry`` is updated in place: "tilts", + "thickness", "bloch_g", "bloch_p", "bloch_score" and + "thickness_k_max". + """ + if self.geometry is None: + raise RuntimeError("fit_geometry first") + from quantem.diffraction import bloch + + crystal = self.crystal + crystal.calculate_structure_factors(2 * k_max) + crystal.calculate_dynamical_structure_factors(self.energy, k_max=2 * k_max) + t_grid = np.arange(thickness[0], thickness[1] + 1e-9, step) + geo = self.geometry + geo.setdefault("thickness", [np.nan] * len(self.images)) + geo.setdefault("bloch_g", [None] * len(self.images)) + geo.setdefault("bloch_p", [None] * len(self.images)) + geo.setdefault("bloch_score", [None] * len(self.images)) + geo["thickness_k_max"] = float(k_max) + results = {} + for i in range(len(self.images)): + hkl_m, meas = self._measured_bragg(i) + a_meas = np.sqrt(meas / meas.sum()) + keys = {tuple(h): k for k, h in enumerate(hkl_m)} + tilt0 = np.asarray(geo["tilts"][i], dtype=float) + offs = np.deg2rad(np.arange(-tilt_range_deg, tilt_range_deg + 1e-9, tilt_step_deg)) + best = (np.inf, None, None, None) + curves = {} + for dx in offs: + for dy in offs: + tilt = tilt0 + np.array([dx, dy]) + out = bloch.dynamical_pattern( + crystal, self._orientation_quat(i, tilt), t_grid, self.energy, k_max=k_max + ) + calc = np.zeros((len(t_grid), len(hkl_m))) + calc[:, keys[(0, 0, 0)]] = out["intensity_000"].numpy() + for col, h in enumerate(out["hkl"].numpy().astype(int)): + k = keys.get(tuple(h)) + if k is not None: + calc[:, k] = out["intensity"][:, col].numpy() + a_calc = np.sqrt(calc / calc.sum(1, keepdims=True)) + score = ((a_calc - a_meas[None]) ** 2).sum(1) + curves[(dx, dy)] = score + j = int(np.argmin(score)) + if score[j] < best[0]: + best = (score[j], tilt, t_grid[j], (dx, dy)) + score, tilt, t_best, key = best + # depth-averaged beam intensities at the best thickness and tilt + z = (np.arange(depth_samples) + 0.5) / depth_samples * t_best + out = bloch.dynamical_pattern( + crystal, self._orientation_quat(i, tilt), z, self.energy, k_max=k_max + ) + frame = _zone_frame(self.zone_axes[i]) + hkl_b = np.vstack([[0, 0, 0], out["hkl"].numpy()]) + g_b = ( + (hkl_b / self._a_crystal) @ frame[:2].T * self._a_crystal / self.lattice_parameter + ) + p_b = np.concatenate( + [[out["intensity_000"].mean().item()], out["intensity"].mean(0).numpy()] + ) + geo["tilts"][i] = tilt + geo["thickness"][i] = float(t_best) + geo["bloch_g"][i] = g_b + geo["bloch_p"][i] = p_b + geo["bloch_score"][i] = dict(thickness=t_grid, score=curves[key], best=float(score)) + results[self.names[i]] = dict( + thickness=float(t_best), tilt_deg=float(np.rad2deg(np.linalg.norm(tilt))) + ) + if verbose: + print( + f"{self.names[i]}: thickness {t_best / 10:.0f} nm, tilt " + f"{np.rad2deg(np.linalg.norm(tilt)):.2f} deg, misfit {score:.3f}" + ) + return results + + def plot_thickness(self, **kwargs): + """Bloch thickness fit per pattern: misfit against thickness (top) and measured against + calculated Bragg intensities at the best fit, square-root scale (bottom). Uses the + ``k_max`` given to ``fit_thickness``.""" + import matplotlib.pyplot as plt + + from quantem.diffraction import bloch + + n = len(self.images) + fig, axs = plt.subplots(2, n, figsize=kwargs.pop("figsize", (4.2 * n, 7.5))) + for i in range(n): + sc = self.geometry["bloch_score"][i] + ax = axs[0, i] + ax.plot(sc["thickness"] / 10, sc["score"], "k-") + t_best = self.geometry["thickness"][i] + ax.axvline(t_best / 10, color="tab:red", lw=1) + ax.set_xlabel("thickness (nm)") + ax.set_ylabel("misfit") + ax.set_title(f"{self.names[i]}: {t_best / 10:.0f} nm") + hkl_m, meas = self._measured_bragg(i) + out = bloch.dynamical_pattern( + self.crystal, + self._orientation_quat(i, self.geometry["tilts"][i]), + [t_best], + self.energy, + k_max=self.geometry.get("thickness_k_max", 1.6), + ) + lookup = { + tuple(h): v + for h, v in zip(out["hkl"].numpy().astype(int), out["intensity"][0].numpy()) + } + lookup[(0, 0, 0)] = float(out["intensity_000"][0]) + calc = np.array([lookup.get(tuple(h), 0.0) for h in hkl_m]) + ax = axs[1, i] + x, y = np.sqrt(calc / calc.sum()), np.sqrt(meas / meas.sum()) + ax.plot(x, y, "o", ms=4, color="tab:blue") + lim = 1.05 * max(x.max(), y.max()) + ax.plot([0, lim], [0, lim], "k--", lw=0.8) + ax.set_xlim(0, lim) + ax.set_ylim(0, lim) + ax.set_aspect("equal") + ax.set_xlabel("Bloch sqrt(I)") + ax.set_ylabel("measured sqrt(I)") + fig.tight_layout() + return fig, axs + + def bragg_positions(self, i: int, k_max: float = 3.0) -> np.ndarray: + """Detector positions (row, col) of every zone reflection, direct beam included.""" + _, g2, _ = self._zone_reflections(self.zone_axes[i], k_max) + return self.geometry["centers"][i] + g2 @ self.geometry["matrices"][i].T + + def _q_zone(self, i: int, rows, cols) -> np.ndarray: + """Pixel coordinates -> in-plane scattering vector (..., 2) in the zone frame, 1/A.""" + p = np.stack([rows, cols], axis=-1) - self.geometry["centers"][i] + g = p @ np.linalg.inv(self.geometry["matrices"][i]).T + return g * self._a_crystal / self.lattice_parameter + + def _q_crystal(self, i: int, q2: np.ndarray) -> np.ndarray: + """In-plane zone-frame q (n, 2) -> crystal-frame q (n, 3) on the tilted Ewald sphere.""" + k_wave = 1.0 / self.wavelength + tilt = self.geometry["tilts"][i] + qz = -((q2**2).sum(1) / (2 * k_wave) + q2 @ tilt) + return np.column_stack([q2, qz]) @ _zone_frame(self.zone_axes[i]) + + # --------------------------------------------------------------------- mask + + def set_mask( + self, + bragg_radius: float = 0.12, + softness: float = 0.01, + q_max: float = 1.2, + center_radius: float = 0.25, + edge_px: int = 8, + ) -> "ReverseMonteCarlo": + """Diffuse-scattering weight, binned data and Bragg intensities. + + The weight is ``sigmoid((d - bragg_radius) / softness)``, with ``d`` + the distance (1/A) to the nearest reflection, the direct beam + included: 0 on every Bragg peak, 1 between them. Pixels beyond + ``q_max`` or within ``edge_px`` of the detector edge are dropped, and + a second sigmoid removes the direct beam's bloom out to + ``center_radius``. Each ``bin_factor`` square is reduced to its + weighted mean. + + The mask has to cover the peaks' tails (detector point spread and + near-peak scattering), which are not modelled: here they fall to a few + percent of the local diffuse level by 0.12 1/A. + + Each reflection's integrated intensity weights the diffuse envelope. + + Parameters + ---------- + bragg_radius : float, optional + Distance (1/A) from each reflection at which the weight reaches + 0.5. Default 0.12. + softness : float, optional + Width (1/A) of the sigmoid edges. Default 0.01. + q_max : float, optional + Largest scattering vector (1/A) fitted. Default 1.2. + center_radius : float, optional + Radius (1/A) of the direct-beam bloom removed. Default 0.25. + edge_px : int, optional + Detector-edge border (unbinned pixels) given zero weight. 0 keeps + the whole detector. Default 8. + + Returns + ------- + ReverseMonteCarlo + self. The results are stored in ``self.mask``: per-pattern lists + "y" (weighted mean of each binned pixel, normalized), "w" (binned + weight, 0 to 1), "k" (zone-frame q of each binned pixel, 1/A), + "data" (binned pattern, normalized), "scale" (normalization), + "bragg_k" and "bragg_intensity" (reflections fully on the + detector), plus the parameters above. + """ + if self.geometry is None: + raise RuntimeError("fit_geometry first") + b = self.bin_factor + out = dict( + bragg_radius=bragg_radius, softness=softness, q_max=q_max, center_radius=center_radius + ) + out.update(y=[], w=[], k=[], data=[], scale=[], bragg_k=[], bragg_intensity=[]) + for i, im in enumerate(self.images): + ny, nx = (im.shape[0] // b) * b, (im.shape[1] // b) * b + rows, cols = np.mgrid[0:ny, 0:nx].astype(np.float64) + k = self._q_zone(i, rows, cols) + bragg = self.bragg_positions(i) + g_q = self._q_zone(i, bragg[:, 0], bragg[:, 1]) + d, j = cKDTree(g_q).query(k.reshape(-1, 2)) + d = d.reshape(ny, nx) + j = j.reshape(ny, nx) + q = np.linalg.norm(k, axis=-1) + w = ( + _sigmoid((d - bragg_radius) / softness) + * _sigmoid((q - center_radius) / softness) + * (q < q_max) + ) + if edge_px > 0: + w[:edge_px] = w[-edge_px:] = 0 + w[:, :edge_px] = w[:, -edge_px:] = 0 + y = im[:ny, :nx].astype(np.float64) + + # integrated Bragg intensities over the local ring median + n_g = len(g_q) + core = d < bragg_radius + ring = (d >= bragg_radius) & (d < 1.6 * bragg_radius) + ring_med = np.zeros(n_g) + jr, yr = j[ring], y[ring] + order = np.argsort(jr, kind="stable") + jr, yr = jr[order], yr[order] + starts = np.searchsorted(jr, np.arange(n_g)) + ends = np.searchsorted(jr, np.arange(n_g), side="right") + for g in np.nonzero(ends > starts)[0]: + ring_med[g] = np.median(yr[starts[g] : ends[g]]) + core_sig = np.where(core, np.clip(y - ring_med[j], 0, None), 0.0) + inten = np.bincount(j[core], weights=core_sig[core], minlength=n_g) + n_core = np.bincount(j[core], minlength=n_g) + seen = n_core > 0.5 * n_core.max() # whole spot on the detector + + wb = self._binned(w, "sum") + yb = np.where(wb > 0, self._binned(w * y, "sum") / np.maximum(wb, 1e-12), 0.0) + rb, cb = np.mgrid[0 : ny // b, 0 : nx // b].astype(np.float64) * b + (b - 1) / 2 + norm = (wb * yb).sum() / wb.sum() + out["y"].append(yb / norm) + out["w"].append(wb / b**2) + out["k"].append(self._q_zone(i, rb, cb)) + out["data"].append(self._binned(y) / norm) + out["scale"].append(norm) + out["bragg_k"].append(g_q[seen]) + out["bragg_intensity"].append(inten[seen] / norm) + self.mask = out + return self + + # ---------------------------------------------------------------- supercell + + def build_supercell( + self, + cells: int = 16, + seed: int | None = 0, + displacements: bool = True, + displacement_grid: int = 24, + max_displacement: float = 0.3, + omega_amplitudes: Sequence[int] = (1, 2), + symmetrize: bool = True, + debye_waller: float = 0.5, + envelope: str = "measured", + resolution: float = 0.75, + shared_scale: bool = True, + max_beams: int = 40, + device: str | None = None, + ) -> "ReverseMonteCarlo": + """Random supercell of ``cells^3`` unit cells at the crystal's composition. + + Every species on the mixed sites is kept (use ``merge`` in + ``set_crystal`` to fold any together). The supercell amplitude is + ``F(q) = sum_j f_s(j)(q) exp(-2 pi i q.(r_j + u_j))`` with the + Bragg nodes of the average lattice removed. + + Parameters + ---------- + cells : int + Unit cells along each cube edge. The diffuse model is sampled + every ``1 / (cells a)`` in reciprocal space. + seed : int or None + Seed of the random generator (stored as ``self.rng``) used for the + initial species arrangement and for every later Monte Carlo move. + None seeds from fresh OS entropy. Default 0. + displacements : bool + Allow static displacements. Positions live on a grid + ``a / displacement_grid`` fine, so every move is still scored + exactly. Implemented for BCC site lattices only: the default + True raises NotImplementedError for any other site lattice, so + pass False there. + displacement_grid : int + Steps per lattice parameter of the displacement grid (a multiple + of 24): 24 gives 0.15 A steps, 48 gives 0.076 A for a = 3.66 A. + max_displacement : float + Largest displacement component (A) reached by random moves. + omega_amplitudes : sequence of int + Allowed displacements in units of a/24 along each axis: 2 is the + ideal omega collapse (a/12, 0.53 A along <111> for a = 3.66 A), 1 + a half collapse. + symmetrize : bool + Average every pattern over the 24 cubic rotations of the supercell. + debye_waller : float + Isotropic B (A^2) damping the diffuse intensity. + envelope : {"measured", "fitted", "bloch", "kinematic"} + "measured" redistributes the diffuse intensity over the Bragg + beams of each pattern, ``sum_g P_g fbar^2(q - g) / fbar^2(q)`` + with P_g the integrated Bragg intensities (exact for occupational + disorder, whose diffuse intensity is periodic in the reciprocal + lattice; approximate for displacements). "fitted" starts there + and re-solves the non-negative P_g of each pattern with its + background: diffuse scattering is generated by the beams' depth + averaged intensities, which dynamical diffraction makes differ + from their exit intensities. "bloch" takes the depth-averaged + Bloch-wave beam intensities at each pattern's fitted thickness and + tilt (requires ``fit_thickness``); use it with ``shared_scale`` + so one diffuse scale covers all patterns. "kinematic" keeps only + the direct beam. + resolution : float + Gaussian sigma, in supercell reciprocal-grid steps, with which + each pixel reads the diffuse grid (27 nearest points). A finite + supercell's intensity is speckle; reading it through a kernel of + about one step damps the speckle. 0 reads the 8 nearest points + trilinearly, about half as many grid points in total (faster). + shared_scale : bool + One diffuse scale for every pattern; the envelope carries each + pattern's absolute Bragg intensities. Ignored with + ``envelope="fitted"``, whose beam weights are solved per pattern. + max_beams : int + Strongest Bragg beams of each pattern in the diffuse envelope. + device : str, optional + torch device. Default cuda when available, otherwise cpu; mps is + used only when requested explicitly. + + Returns + ------- + ReverseMonteCarlo + self. + """ + if self.mask is None: + raise RuntimeError("set_mask first") + self._check_envelope(envelope) + rng = np.random.default_rng(seed) + self.rng = rng + d = self.grid_divisor + if displacements and not ( + d == 2 + and len(self.site_grid) == 2 + and np.array_equal(np.sort(self.site_grid.sum(1)), [0, 3]) + ): + raise NotImplementedError("Displacements are implemented for BCC sites only.") + if displacement_grid % 24: + raise ValueError("displacement_grid must be a multiple of 24.") + m = displacement_grid // d if displacements else 1 + unit = m * d // 24 # fine-grid steps per a/24 + self._omega_vectors = np.array( + [k * unit * np.array(v) for k in omega_amplitudes for v in product((1, -1), repeat=3)], + dtype=np.int64, + ) + + n = cells * d * m + idx = np.stack(np.meshgrid(*(np.arange(cells),) * 3, indexing="ij"), -1).reshape(-1, 3) + x0 = ((idx[:, None, :] * d + self.site_grid[None]) * m).reshape(-1, 3) + n_sites = len(x0) + counts = np.round(self.concentrations * n_sites).astype(int) + counts[-1] = n_sites - counts[:-1].sum() + spec = np.repeat(np.arange(len(counts)), counts) + rng.shuffle(spec) + self.cells, self.refine, self.grid_size = cells, m, n + self.site_x = x0 + self.species_index = spec + self.displacement = np.zeros((n_sites, 3), dtype=np.int64) # fine-grid steps + nc = cells * d + self._site_lookup = np.full((nc,) * 3, -1, dtype=np.int64) + xc = x0 // m + self._site_lookup[xc[:, 0], xc[:, 1], xc[:, 2]] = np.arange(n_sites) + self.displacements = bool(displacements) + self._max_steps = int(np.floor(max_displacement / (self._a_crystal / (d * m)) + 1e-9)) + self.size_eta = np.zeros(len(self.species)) + self.debye_waller = float(debye_waller) + self.symmetrize = bool(symmetrize) + self.envelope = envelope + self.resolution = float(resolution) + self.shared_scale = bool(shared_scale) + self.max_beams = int(max_beams) + self.device = torch.device(device or _default_device()) + self._setup_forward() + print( + f"{n_sites} sites (" + + ", ".join(f"{s} {c}" for s, c in zip(self.species, counts)) + + f"), {n}^3 grid, {len(self._needed)} grid points read, " + f"{int(self._fit.sum())} fitted pixels, {self.device}" + ) + return self + + def _positions(self, sites: np.ndarray, disp: np.ndarray | None = None) -> np.ndarray: + disp = self.displacement[sites] if disp is None else disp + return np.mod(self.site_x[sites] + disp, self.grid_size) + + # metallic (12-fold coordination) radii, A + _RADII = { + "V": 1.34, + "Nb": 1.46, + "Zr": 1.60, + "Ti": 1.47, + "Mo": 1.39, + "Ta": 1.46, + "Hf": 1.59, + "W": 1.39, + "Cr": 1.28, + "Fe": 1.26, + "Al": 1.43, + } + + def _pixel_grid(self, i: int, q2: np.ndarray): + """Grid indices and weights (n, 8) trilinear, or (n, 27) Gaussian of ``resolution`` + steps, for in-plane q of pattern i.""" + n = self.grid_size + h = self._q_crystal(i, q2) * self.lattice_parameter * self.cells + if self.resolution > 0: + stencil = np.array(list(product((-1, 0, 1), repeat=3))) + h0 = np.round(h).astype(np.int64) + f = h - h0 + wts = np.exp( + -0.5 * ((f[:, None, :] - stencil[None]) ** 2).sum(-1) / self.resolution**2 + ) + wts /= wts.sum(1, keepdims=True) + else: + stencil = np.array(list(product((0, 1), repeat=3))) + h0 = np.floor(h).astype(np.int64) + f = h - h0 + wts = np.prod(np.where(stencil[None], f[:, None, :], 1 - f[:, None, :]), axis=-1) + ijk = np.mod(h0[:, None, :] + stencil[None], n) + return (ijk[..., 0] * n + ijk[..., 1]) * n + ijk[..., 2], wts + + def _setup_forward(self): + """Ewald-sphere sampling of the supercell grid for every binned pixel.""" + n = self.grid_size + dev = self.device + cols_w, vals_w, u_all, basis_all, pix_image = [], [], [], [], [] + self._env_cols = [None] * len(self.images) + self._env_p = [None] * len(self.images) + for i, kk in enumerate(self.mask["k"]): + q2 = kk.reshape(-1, 2) + cols, wts = self._pixel_grid(i, q2) + cols_w.append(cols) + vals_w.append(wts) + u, tds = self._envelope(i, q2) + u_all.append(u) + qm = np.linalg.norm(q2, axis=1) + rings = self._powder_rings(qm) + basis_all.append(np.column_stack([np.ones_like(qm), qm, q2, tds, rings])) + pix_image.append(np.full(len(q2), i)) + cols = np.concatenate(cols_w) + vals = np.concatenate(vals_w) + # only pixels inside q_max are modelled: the fit never reads the rest, and every grid + # point read costs time in each move + inside = np.concatenate( + [ + np.linalg.norm(kk.reshape(-1, 2), axis=1) < self.mask["q_max"] + for kk in self.mask["k"] + ] + ) + used, inv = np.unique(cols[inside], return_inverse=True) + cols_used = np.zeros(cols.shape, dtype=np.int64) + cols_used[inside] = inv.reshape(-1, cols.shape[1]) + vals = np.where(inside[:, None], vals, 0.0) + n_pix = len(cols) + self._W_all = sparse.csr_matrix( + (vals.ravel(), (np.repeat(np.arange(n_pix), cols.shape[1]), cols_used.ravel())), + shape=(n_pix, len(used)), + ) + self._used = used + # symmetry: model reads mean_k I(S_k h) at each used h + ops = cubic_rotations() if self.symmetrize else np.eye(3, dtype=int)[None] + hh = np.stack(np.unravel_index(used, (n,) * 3), -1) + sym_flat = np.stack( + [(lambda s: (s[:, 0] * n + s[:, 1]) * n + s[:, 2])(np.mod(hh @ op.T, n)) for op in ops] + ) + needed, inv_s = np.unique(sym_flat, return_inverse=True) + self._needed = needed + self._sym_index = torch.as_tensor(inv_s.reshape(sym_flat.shape), device=dev) + hn = np.stack(np.unravel_index(needed, (n,) * 3), -1) + fs, keep = self._grid_factors(hn) + self._fs = torch.as_tensor(fs, dtype=torch.float32, device=dev) + self._keep = torch.as_tensor(keep, dtype=torch.float32, device=dev) + self._u_all = np.concatenate(u_all) + self._basis_all = np.concatenate(basis_all) + self._pix_image = np.concatenate(pix_image) + y = np.concatenate([a.ravel() for a in self.mask["y"]]) + w = np.concatenate([a.ravel() for a in self.mask["w"]]) + self._y_all, self._w_all = y, w + fit = w > 1e-3 + self._fit = fit + self._Wc = torch.as_tensor(cols_used[fit], device=dev) + self._Wv = torch.as_tensor(vals[fit], dtype=torch.float32, device=dev) + self._y = torch.as_tensor(y[fit], dtype=torch.float32, device=dev) + self._w = torch.as_tensor(w[fit], dtype=torch.float32, device=dev) + self._img = torch.as_tensor(self._pix_image[fit], device=dev) + self._u = torch.as_tensor(self._u_all[fit], dtype=torch.float32, device=dev) + self._h_needed = torch.as_tensor(hn.T.copy(), dtype=torch.int32, device=dev) # (3, n) + ang = 2 * np.pi * np.arange(n) / n + self._cos = torch.as_tensor(np.cos(ang), dtype=torch.float32, device=dev) + self._sin = torch.as_tensor(np.sin(ang), dtype=torch.float32, device=dev) + self._recompute_F() + if getattr(self, "background_sigmas", None) is None or len(self.background_sigmas) != len( + self.images + ): + # direct-beam Lorentzian half width, wide Gaussian sigma and its center, 1/A; the + # wide Gaussian's center floats because a tilted crystal centers its smooth + # background on the zone-axis pole rather than the direct beam + self.background_sigmas = [(0.1, 0.8, 0.0, 0.0) for _ in self.images] + self._sigma_bounds = (np.array([0.01, 0.3, -1.0, -1.0]), np.array([1.0, 3.0, 1.0, 1.0])) + self.coefficients = getattr(self, "coefficients", None) + if self.coefficients is not None: + n_coef = 1 + self._bg_basis(0, np.arange(1)).shape[1] + if self.coefficients.shape[1] != n_coef: # background terms changed: pad with zeros + c = np.zeros((len(self.images), n_coef)) + k = min(n_coef, self.coefficients.shape[1]) + c[:, :k] = self.coefficients[:, :k] + self.coefficients = c + self._update_residual() + + _ENVELOPES = ("measured", "fitted", "bloch", "kinematic") + + def _check_envelope(self, envelope: str) -> None: + """Raise if ``envelope`` is unknown, or is "bloch" before ``fit_thickness``.""" + if envelope not in self._ENVELOPES: + raise ValueError(f"unknown envelope {envelope!r}; use one of {self._ENVELOPES}") + if envelope == "bloch": + bloch_g = (self.geometry or {}).get("bloch_g") + if bloch_g is None or any(g is None for g in bloch_g): + raise RuntimeError("fit_thickness before envelope='bloch'") + + def set_envelope(self, envelope: str) -> "ReverseMonteCarlo": + """Switch the diffuse envelope and refit the scale and background. + + Parameters + ---------- + envelope : {"measured", "fitted", "bloch", "kinematic"} + See ``build_supercell``. "bloch" requires ``fit_thickness``. + + Returns + ------- + ReverseMonteCarlo + self. + """ + self._check_envelope(envelope) + self.envelope = envelope + self._setup_forward() + self._solve_linear(self._model_diffuse(), refit_sigmas=True) + self._update_residual() + return self + + def _grid_factors(self, h: np.ndarray, n: int | None = None): + """Scattering factors (K, n) of every species at grid points h (n, 3), and a 0/1 weight + removing the Bragg nodes of the average lattice.""" + n = self.grid_size if n is None else n + hm = np.where(h > n // 2, h - n, h) + q = np.linalg.norm(hm, axis=1) / (self.lattice_parameter * self.cells) + fs = electron_scattering_factor( + torch.tensor(self.numbers), torch.as_tensor(q, dtype=torch.float64) + ).numpy() + eta = getattr(self, "size_eta", None) + if eta is not None and np.any(eta != 0): + fs = ( + fs + + np.asarray(eta)[:, None] + * self._size_chi(hm / (self.lattice_parameter * self.cells))[None] + ) + on_node = np.all(np.mod(hm, self.cells) == 0, axis=1) + hkl = hm[on_node] // self.cells + frac = self.site_grid / self.grid_divisor + f_avg = np.exp(-2j * np.pi * hkl @ frac.T).sum(1) + keep = np.ones(len(h)) + keep[np.nonzero(on_node)[0][np.abs(f_avg) > 1e-6]] = 0.0 + return fs, keep + + def _size_chi(self, q: np.ndarray, k_ratio: float = 0.5, chunk: int = 200_000) -> np.ndarray: + """First-order size-effect factor chi(q) (A) at crystal-frame q (n, 3), 1/A. + + Each atom of species s pushes its 8 nearest and 6 next-nearest neighbours with Kanzaki + forces ``k_n eta_s |r_n| / 2`` along the bond; the lattice relaxes harmonically with the + same springs, u(q) = D(q)^-1 Phi(q) sum_s eta_s A_s(q). To first order in u the + amplitude gains ``-2 pi i fbar q.u``, i.e. each species' scattering factor becomes + ``f_s + eta_s chi(q)`` with ``chi = -2 pi i fbar q.D^-1 Phi`` (real). It is odd about + every Bragg node and grows as 1/|q - g| toward it: Huang and size-effect scattering. + """ + a = self.lattice_parameter + r_n = ( + np.array( + [list(v) for v in product((0.5, -0.5), repeat=3)] + + [[1, 0, 0], [-1, 0, 0], [0, 1, 0], [0, -1, 0], [0, 0, 1], [0, 0, -1]] + ) + * a + ) # A + length = np.linalg.norm(r_n, axis=1) + e_n = r_n / length[:, None] + k_n = np.where(np.arange(14) < 8, 1.0, k_ratio) + ee = e_n[:, :, None] * e_n[:, None, :] + out = np.zeros(len(q)) + for c0 in range(0, len(q), chunk): + qc = q[c0 : c0 + chunk] + theta = 2 * np.pi * qc @ r_n.T # (n, 14) + D = np.einsum("nk,kab->nab", k_n * (1 - np.cos(theta)), ee) + # Phi = sum_n k |r| / 2 e_n exp(-i theta) = -i sum_n k |r| / 2 e_n sin(theta) + psi = np.einsum("nk,ka->na", (k_n * length / 2) * np.sin(theta), e_n) + on_node = np.abs(np.linalg.det(D)) < 1e-12 + D[on_node] = np.eye(3) + x = np.linalg.solve(D, psi[..., None])[..., 0] # Phi = -i psi -> D^-1 Phi = -i x + fbar, _ = self._fbar(np.linalg.norm(qc, axis=1)) + chi = -2 * np.pi * fbar * (qc * x).sum(1) # -2 pi i fbar q.(-i x) + chi[on_node] = 0.0 + out[c0 : c0 + chunk] = chi + return out + + def set_size_effect(self, eta: dict[str, float] | str | None = "radii") -> "ReverseMonteCarlo": + """Set the linear size effect: species mismatch ``eta_s`` (dimensionless). + + Only differences between species matter (a common shift only moves the + Bragg peaks). The forward model is rebuilt; call ``fit_background`` + afterwards to refit the scale and background. + + Parameters + ---------- + eta : "radii", dict or None, optional + "radii" (default) takes ``(r_s - r_mean) / r_mean`` from metallic + radii (e.g. V 1.34, Nb 1.46, Zr 1.60 A; ASE covalent radii for + species without a tabulated metallic radius), with ``r_mean`` + the composition-weighted mean. A dict ``{species: eta}`` sets + them directly (missing species get 0). None switches the size + effect off. + + Returns + ------- + ReverseMonteCarlo + self. The mismatches are stored in ``self.size_eta``, in the + order of ``self.species``. + """ + if eta is None: + self.size_eta = np.zeros(len(self.species)) + elif isinstance(eta, str): + from ase.data import atomic_numbers, covalent_radii + + r = np.array( + [self._RADII.get(sp, covalent_radii[atomic_numbers[sp]]) for sp in self.species] + ) + r_mean = float((self.concentrations * r).sum()) + self.size_eta = (r - r_mean) / r_mean + else: + self.size_eta = np.array([float(eta.get(sp, 0.0)) for sp in self.species]) + self._setup_forward() + return self + + def _species_amplitudes(self) -> np.ndarray: + """Lattice sums A_s(h) = sum_{j in s} exp(-2 pi i h.x_j / N) on the needed points (K, n).""" + n = self.grid_size + hn = np.stack(np.unravel_index(self._needed, (n,) * 3), -1) + upper = hn[:, 2] > n // 2 + hr = np.where(upper[:, None], np.mod(-hn, n), hn) + flat = (hr[:, 0] * n + hr[:, 1]) * (n // 2 + 1) + hr[:, 2] + out = np.zeros((len(self.species), len(self._needed)), dtype=np.complex64) + for s_, A in self._species_fft(real=True): + v = A.ravel()[flat] + out[s_] = np.where(upper, np.conj(v), v) + return out + + def fit_size_effect(self, step: float = 0.005, verbose: bool = True) -> dict: + """Fit the species size mismatches eta_s to the diffuse scattering of the current supercell. + + Nelder-Mead from no size effect; the background and scale are + re-solved at each step, and one eta is fixed by ``sum c_s eta_s = 0``. + If the fit does not lower the loss, the size effect stays off. + + Parameters + ---------- + step : float, optional + Size of the initial Nelder-Mead simplex in eta. Default 0.005. + verbose : bool, optional + Print the fitted mismatches and the loss change. + + Returns + ------- + dict + "eta": ``{species: eta_s}``; "loss": (loss without size effect, + loss after the fit and a background refit). + """ + dev = self.device + A = torch.as_tensor(self._species_amplitudes(), device=dev) + n = self.grid_size + hn = np.stack(np.unravel_index(self._needed, (n,) * 3), -1) + hm = np.where(hn > n // 2, hn - n, hn) + chi = torch.as_tensor( + self._size_chi(hm / (self.lattice_parameter * self.cells)), + dtype=torch.float32, + device=dev, + ) + self.size_eta = np.zeros(len(self.species)) + fs0, _ = self._grid_factors(hn) + fs0 = torch.as_tensor(fs0, dtype=torch.float32, device=dev) + c = self.concentrations + + def full_eta(x): + e = np.concatenate([x, [0.0]]) + return e - (c * e).sum() + + def loss(x): + eta = torch.as_tensor(full_eta(x), dtype=torch.float32, device=dev) + F = ((fs0 + eta[:, None] * chi[None]) * A).sum(0) + self._Fr, self._Fi = F.real.contiguous(), F.imag.contiguous() + self._solve_linear(self._model_diffuse()) + return self._update_residual() + + # start from no size effect with a small simplex: at radius-sized mismatches the + # first-order displacements are no longer small and the loss is far from quadratic + x0 = np.zeros(len(c) - 1) + l0 = loss(x0) + simplex = np.vstack([x0, x0 + step * np.eye(len(x0))]) + sol = optimize.minimize( + loss, + x0, + method="Nelder-Mead", + options=dict(xatol=1e-5, fatol=1e-4, initial_simplex=simplex), + ) + if sol.fun > l0: + sol.x = x0 + self.size_eta = full_eta(sol.x) + self._setup_forward() + self._solve_linear(self._model_diffuse(), refit_sigmas=True) + l1 = self._update_residual() + out = dict(eta={sp: float(e) for sp, e in zip(self.species, self.size_eta)}, loss=(l0, l1)) + if verbose: + print( + "size mismatch eta: " + + ", ".join(f"{k} {v:+.4f}" for k, v in out["eta"].items()) + + f"; loss {l0:.2f} -> {l1:.2f}" + ) + return out + + def _read(self, i_used: torch.Tensor) -> torch.Tensor: + """Grid values on the used points (..., n_used) -> kernel-weighted values on fitted pixels.""" + return (i_used[..., self._Wc] * self._Wv).sum(-1) + + def _fbar(self, q: np.ndarray): + fe = electron_scattering_factor( + torch.tensor(self.numbers), torch.as_tensor(q, dtype=torch.float64) + ).numpy() + c = self.concentrations[:, None] + return (c * fe).sum(0), (c * fe**2).sum(0) + + def _envelope(self, i: int, q2: np.ndarray): + """Diffuse envelope (beam redistribution x Debye-Waller) and Einstein thermal diffuse. + + Also stores the per-beam envelope columns ``fbar^2(q - g) / fbar^2(q) x DW(q)`` and the + beam weights, which ``envelope="fitted"`` re-solves. + """ + if self.envelope == "bloch": + g = (self.geometry.get("bloch_g") or [None] * len(self.images))[i] + if g is None: + raise RuntimeError("fit_thickness before envelope='bloch'") + p = self.geometry["bloch_p"][i] + order = np.argsort(p)[::-1][: getattr(self, "max_beams", 40)] + g, p = g[order], p[order] + elif self.envelope in ("measured", "fitted"): + g = self.mask["bragg_k"][i] + p = self.mask["bragg_intensity"][i] + order = np.argsort(p)[::-1][: getattr(self, "max_beams", 40)] + g, p = g[order], p[order] + keep = p > 0.002 * p.max() + g, p = g[keep], p[keep] + elif self.envelope == "kinematic": + g, p = np.zeros((1, 2)), np.array([self.mask["bragg_intensity"][i].sum()]) + else: + raise ValueError(f"unknown envelope {self.envelope!r}") + q0 = np.linalg.norm(q2, axis=1) + fbar0, _ = self._fbar(q0) + dw0 = np.exp(-0.5 * self.debye_waller * q0**2) + cols = np.zeros((len(q2), len(g))) + tds = np.zeros(len(q2)) + for k, (gi, pi) in enumerate(zip(g, p)): + qm = np.linalg.norm(q2 - gi, axis=1) + fbar, f2 = self._fbar(qm) + dw = np.exp(-0.5 * self.debye_waller * qm**2) + cols[:, k] = fbar**2 / fbar0**2 * dw0 + tds += pi * f2 * (1 - dw) + self._env_cols[i] = cols + self._env_p[i] = p.astype(float) + return cols @ p, tds + + def _bg_lower(self, n: int) -> np.ndarray: + """Lower bounds of the background amplitudes: the constant may be negative, the others + are non-negative.""" + lo = np.zeros(n) + lo[0] = -np.inf + return lo + + def _powder_rings(self, q: np.ndarray, width: float = 0.025, k_max: float = 1.5) -> np.ndarray: + """Powder rings of the average crystal about the direct beam, ``sum_g |F_g|^2 / g^2`` + broadened by ``width`` (1/A): misoriented grains or a damaged surface layer.""" + self.crystal.calculate_structure_factors(k_max) + g = self.crystal.g_len.numpy() * self._a_crystal / self.lattice_parameter + f2 = self.crystal.struct_factors_int.numpy() + keep = g > 1e-6 + g, f2 = g[keep], f2[keep] + out = np.zeros_like(q) + for gi, fi in zip(g, f2): + out += fi / gi**2 * np.exp(-0.5 * ((q - gi) / width) ** 2) + return out / out.max() + + def _phases(self, pos: np.ndarray): + """cos and sin of 2 pi h.x / N for fine-grid positions (B, 3) on every needed point.""" + p = torch.as_tensor(pos, dtype=torch.int32, device=self.device) + h = self._h_needed + idx = (p[:, 0:1] * h[0] + p[:, 1:2] * h[1] + p[:, 2:3] * h[2]) % self.grid_size + idx = idx.long() + return self._cos[idx], self._sin[idx] + + def _species_fft(self, coarsen: int = 1, real: bool = False): + """Per-species FFT of the occupancy, one at a time. ``coarsen`` rounds positions onto a + grid that many times coarser; ``real`` returns the half spectrum (rfftn).""" + from scipy import fft as sfft + + n = self.grid_size // coarsen + pos = self._positions(np.arange(len(self.site_x))) + pos = np.mod((pos + coarsen // 2) // coarsen, n) + for s in range(len(self.species)): + occ = np.zeros((n,) * 3, dtype=np.float32) + sel = self.species_index == s + np.add.at(occ, (pos[sel, 0], pos[sel, 1], pos[sel, 2]), 1.0) + yield s, (sfft.rfftn if real else sfft.fftn)(occ, workers=-1) + + def _recompute_F(self): + n = self.grid_size + hn = np.stack(np.unravel_index(self._needed, (n,) * 3), -1) + upper = hn[:, 2] > n // 2 # read from the conjugate half + hr = np.where(upper[:, None], np.mod(-hn, n), hn) + flat = (hr[:, 0] * n + hr[:, 1]) * (n // 2 + 1) + hr[:, 2] + F = np.zeros(len(self._needed), dtype=np.complex128) + fs = self._fs.cpu().numpy() + for s, A in self._species_fft(real=True): + v = A.ravel()[flat] + F += fs[s] * np.where(upper, np.conj(v), v) + self._Fr = torch.as_tensor(F.real, dtype=torch.float32, device=self.device) + self._Fi = torch.as_tensor(F.imag, dtype=torch.float32, device=self.device) + + def _diffuse_used(self, Fr=None, Fi=None) -> torch.Tensor: + """Symmetrized diffuse intensity per site on the used grid points.""" + Fr = self._Fr if Fr is None else Fr + Fi = self._Fi if Fi is None else Fi + inten = (Fr**2 + Fi**2) * self._keep / len(self.site_x) + return inten[..., self._sym_index].mean(dim=-2) + + def diffuse_grid(self, max_size: int = 400) -> np.ndarray: + """Symmetrized diffuse intensity per site on the whole grid (Bragg nodes removed). + + Grids above ``max_size`` per edge are evaluated with positions rounded onto a coarser + grid (the same reciprocal sampling over a smaller q range), which slightly blurs the + displacement scattering.""" + coarsen = 1 + while self.grid_size // coarsen > max_size and self.refine % (2 * coarsen) == 0: + coarsen *= 2 + n = self.grid_size // coarsen + hh = np.stack(np.meshgrid(*(np.arange(n),) * 3, indexing="ij"), -1).reshape(-1, 3) + fs, keep = self._grid_factors(hh, n) + F = np.zeros(n**3, dtype=np.complex64) + for s, A in self._species_fft(coarsen): + F += (fs[s] * A.ravel()).astype(np.complex64) + S = (np.abs(F) ** 2 * keep).reshape((n,) * 3).astype(np.float32) / len(self.site_x) + if not self.symmetrize: + return S + og = np.ogrid[0:n, 0:n, 0:n] + out = np.zeros_like(S) + ops = cubic_rotations() + for op in ops: + perm = np.argmax(np.abs(op), axis=1) + sign = op[np.arange(3), perm] + out += S[tuple(np.mod(sign[r] * og[perm[r]], n) for r in range(3))] + return out / len(ops) + + # --------------------------------------------------------------- background + + def _bg_basis(self, i: int, sel: np.ndarray) -> np.ndarray: + q = self._basis_all[sel, 1] + q2 = self._basis_all[sel, 2:4] + s1, s2, cy, cx = self.background_sigmas[i] + return np.column_stack( + [ + np.ones_like(q), + 1.0 / (1.0 + (q / s1) ** 2), + np.exp(-0.5 * ((q2 - [cy, cx]) ** 2).sum(1) / s2**2), + self._basis_all[sel, 4:], + ] + ) + + def _solve_linear(self, diffuse_fit: np.ndarray, refit_sigmas: bool = False): + """Scale, background amplitudes (and optionally widths) per image, weighted least squares.""" + y = self._y_all[self._fit] + w = self._w_all[self._fit] + img = self._pix_image[self._fit] + sel_all = np.nonzero(self._fit)[0] + fitted = self.envelope == "fitted" + if fitted: + read = self._read(self._diffuse_used()).cpu().numpy().astype(np.float64) + coefs, loss = [], 0.0 + for i in range(len(self.images)): + m = img == i + sel = sel_all[m] + sw = np.sqrt(w[m]) + if fitted: + rows = sel - np.nonzero(self._pix_image == i)[0][0] + Xd = read[m][:, None] * self._env_cols[i][rows] + else: + Xd = diffuse_fit[m][:, None] + nd = Xd.shape[1] + + def solve(par): + par = np.clip(par, *self._sigma_bounds) + self.background_sigmas[i] = tuple(float(v) for v in par) + X_bg = self._bg_basis(i, sel) + X = np.column_stack([Xd, X_bg]) + lo = np.concatenate([np.zeros(nd), self._bg_lower(X_bg.shape[1])]) + x = _lsq_bounded(X, y[m], sw, lo) + return x, float(((X @ x - y[m]) ** 2 * w[m]).sum()) + + if refit_sigmas: + p0 = np.asarray(self.background_sigmas[i], dtype=float) + + def unpack(x): + return np.concatenate([np.exp(x[:2]), x[2:]]) + + r = optimize.minimize( + lambda x: solve(unpack(x))[1], + np.concatenate([np.log(p0[:2]), p0[2:]]), + method="Nelder-Mead", + options=dict(xatol=1e-3, fatol=1e-6, maxiter=300), + ) + solve(unpack(r.x)) + c, l_i = solve(self.background_sigmas[i]) + if fitted: + self._env_p[i] = c[:nd] + c = np.concatenate([[1.0], c[nd:]]) + coefs.append(c) + loss += l_i + self.coefficients = np.stack(coefs) + if fitted: + self._u_all = np.concatenate( + [self._env_cols[i] @ self._env_p[i] for i in range(len(self.images))] + ) + self._u = torch.as_tensor( + self._u_all[self._fit], dtype=torch.float32, device=self.device + ) + return loss + if not self.shared_scale: + return loss + + # one diffuse scale for all patterns, backgrounds per pattern + blocks = [self._bg_basis(i, sel_all[img == i]) for i in range(len(self.images))] + nb = blocks[0].shape[1] + X = np.zeros((len(y), 1 + nb * len(blocks))) + X[:, 0] = diffuse_fit + lo = np.zeros(X.shape[1]) + for i, bl in enumerate(blocks): + X[img == i, 1 + nb * i : 1 + nb * (i + 1)] = bl + lo[1 + nb * i : 1 + nb * (i + 1)] = self._bg_lower(nb) + sw = np.sqrt(w) + x = _lsq_bounded(X, y, sw, lo) + if x[0] <= 0: + # a random supercell explains nothing yet: start the scale from the background residual + r = y - X[:, 1:] @ x[1:] + x[0] = max((w * r * diffuse_fit).sum() / max((w * diffuse_fit**2).sum(), 1e-30), 0) + for i in range(len(blocks)): + self.coefficients[i, 0] = x[0] + self.coefficients[i, 1:] = x[1 + nb * i : 1 + nb * (i + 1)] + return float(((X @ x - y) ** 2 * w).sum()) + + def fit_background(self) -> float: + """Fit the diffuse scale and, per pattern, a constant, a direct-beam Lorentzian, a wide + Gaussian with a free center, Einstein thermal diffuse and powder rings of the + average crystal.""" + loss = self._solve_linear(self._model_diffuse(), refit_sigmas=True) + self._update_residual() + print( + "background widths (1/A): " + + ", ".join( + f"{n} {s[0]:.3f}/{s[1]:.3f} at ({s[2]:+.2f}, {s[3]:+.2f})" + for n, s in zip(self.names, self.background_sigmas) + ) + ) + return loss + + def _model_diffuse(self) -> np.ndarray: + """u(q) * (W I_sym) on the fitted pixels (before scale).""" + return (self._u * self._read(self._diffuse_used())).cpu().numpy().astype(np.float64) + + def _update_residual(self): + dev = self.device + img = self._pix_image[self._fit] + sel = np.nonzero(self._fit)[0] + bg = np.zeros(len(sel)) + for i in range(len(self.images)): + m = img == i + bg[m] = self._bg_basis(i, sel[m]) @ self.coefficients[i, 1:] + self._bg = torch.as_tensor(bg, dtype=torch.float32, device=dev) + c = torch.as_tensor(self.coefficients[:, 0], dtype=torch.float32, device=dev) + self._scale = c[self._img] + model = self._scale * self._u * self._read(self._diffuse_used()) + self._bg + self._r = self._y - model + return float((self._w * self._r**2).sum()) + + # ---------------------------------------------------------------------- RMC + + def run( + self, + n_sweeps: int = 20, + batch: int = 32, + temperature: float = 0.05, + random_fraction: float = 0.5, + omega_fraction: float = 0.0, + max_static_b: float | None = 0.5, + refit_every: int = 2, + progress: bool = True, + ) -> "ReverseMonteCarlo": + """Species swaps and displacements until the diffuse fit converges. + + Each batch proposes ``batch`` moves on distinct sites against the + current supercell and scores each exactly. With probability + ``random_fraction`` the batch moves single atoms by one grid step + (-1, 0 or +1 along each axis, up to ``max_displacement``), with no + assumed pattern, so any displacement correlation has to come from the + data. With probability ``omega_fraction`` it proposes omega embryos + instead: three consecutive atoms of a <111> row, the first fixed and + the next two collapsed toward each other, (0, +v, -v), or cleared + where one already stands. Otherwise it swaps two unlike atoms + (composition conserved). ``max_static_b`` (A^2) caps the + static Debye-Waller factor of all displacements, 8 pi^2 / 3, so + they stay consistent with how slowly the Bragg intensities fall off + (a displacement field this strong would damp them; the diffuse scale + alone does not fix how many atoms are displaced). Metropolis acceptance + at ``temperature`` (a fraction of the median score change, falling + linearly to 0); the accepted moves are applied together, capped at a + number that adapts so the joint step never raises the loss. A sweep + is one proposal per site. Scale and background are refit every + ``refit_every`` sweeps. + + Parameters + ---------- + n_sweeps : int, optional + Number of sweeps. Default 20. + batch : int, optional + Moves proposed per batch. Default 32. + temperature : float, optional + Initial Metropolis temperature as a fraction of the median + absolute score change of the first batch. Default 0.05. + random_fraction : float, optional + Probability that a batch proposes single-atom displacements. + Ignored (0) if the supercell has no displacements. Default 0.5. + omega_fraction : float, optional + Probability that a batch proposes omega embryos. Ignored (0) if + the supercell has no displacements. Default 0. + max_static_b : float or None, optional + Cap on the static Debye-Waller B (A^2) of all displacements; None + for no cap. Default 0.5. + refit_every : int, optional + Sweeps between refits of the scale and background. Default 2. + progress : bool, optional + Show a progress bar. + + Returns + ------- + ReverseMonteCarlo + self. The weighted loss after each sweep is appended to + ``self.loss_history``. + """ + if self.coefficients is None: + self.fit_background() + n_sites = len(self.site_x) + loss = self._update_residual() + if not self.loss_history: + self.loss_history.append(loss) + cap = max(batch // 8, 1) + t0 = None + n_batches = max(n_sites // batch, 1) + p_omega = omega_fraction if self.displacements else 0.0 + p_random = random_fraction if self.displacements else 0.0 + budget = ( + 3 * max_static_b / (8 * np.pi**2) * n_sites if max_static_b is not None else np.inf + ) + keep = self._keep / n_sites + sweeps = tqdm(range(n_sweeps), desc="RMC sweeps", disable=not progress) + su = None + for sweep in sweeps: + accepted = {"swap": 0, "random": 0, "omega": 0} + for _ in range(n_batches): + if su is None: + su = self._scale * self._u + spec = self.species_index + roll = self.rng.random() + if roll < p_random: + kind = "random" + sites, new = self._random_proposals(batch) + elif roll < p_random + p_omega: + kind = "omega" + sites, new = self._omega_proposals(batch) + else: + kind = "swap" + j = self.rng.choice(n_sites, 2 * batch, replace=False) + j1, j2 = j[:batch], j[batch:] + ok = spec[j1] != spec[j2] + sites = np.stack([j1[ok], j2[ok]], 1) + new = None + if len(sites) == 0: + continue + if kind == "swap": + c1, s1 = self._phases(self._positions(sites[:, 0])) + c2, s2 = self._phases(self._positions(sites[:, 1])) + df = ( + self._fs[torch.as_tensor(spec[sites[:, 1]], device=self.device)] + - self._fs[torch.as_tensor(spec[sites[:, 0]], device=self.device)] + ) + dFr = df * (c1 - c2) + dFi = df * (s2 - s1) + else: + dFr = torch.zeros((len(sites), len(self._needed)), device=self.device) + dFi = torch.zeros_like(dFr) + for k in range(sites.shape[1]): + j = sites[:, k] + changed = np.any(new[:, k] != self.displacement[j], axis=1) + if not changed.any(): + continue + co, so = self._phases(self._positions(j)) + cn, sn = self._phases(self._positions(j, new[:, k])) + f = self._fs[torch.as_tensor(spec[j], device=self.device)] + f = ( + f + * torch.as_tensor(changed, dtype=torch.float32, device=self.device)[ + :, None + ] + ) + dFr += f * (cn - co) + dFi += f * (so - sn) + d_int = (2 * (self._Fr * dFr + self._Fi * dFi) + dFr**2 + dFi**2) * keep + d_used = d_int[:, self._sym_index].mean(dim=1) + dm = su * self._read(d_used) + dL = (self._w * (dm**2 - 2 * self._r * dm)).sum(-1) + dl = dL.cpu().numpy() + nb = len(dl) + if not t0: + t0 = float(np.median(np.abs(dl))) + temp = temperature * t0 * (1 - sweep / n_sweeps) + if temp > 0: + acc = (dl < 0) | (self.rng.random(nb) < np.exp(-np.clip(dl / temp, 0, 50))) + else: + acc = dl < 0 + pick = np.nonzero(acc)[0] + pick = pick[np.argsort(dl[pick])][:cap] + if kind != "swap" and max_static_b is not None: + du2 = (self._u2(new) - self._u2(self.displacement[sites])).sum(1) + total = self._u2(self.displacement).sum() + chosen = [] + for b in pick: + if du2[b] <= 0 or total + du2[b] <= budget: + chosen.append(b) + total += du2[b] + pick = np.asarray(chosen, dtype=int) + if len(pick) == 0: + continue + pk = torch.as_tensor(pick, device=self.device) + Fr_new = self._Fr + dFr[pk].sum(0) + Fi_new = self._Fi + dFi[pk].sum(0) + r_new = self._y - su * self._read(self._diffuse_used(Fr_new, Fi_new)) - self._bg + loss_new = float((self._w * r_new**2).sum()) + spec_new, disp_new = spec.copy(), self.displacement.copy() + if kind == "swap": + a, b = sites[pick, 0], sites[pick, 1] + spec_new[a], spec_new[b] = spec[b], spec[a] + else: + disp_new[sites[pick]] = new[pick] + if loss_new > loss + max(temp, 0.0) * len(pick) and len(pick) > 1: + cap = max(cap // 2, 1) + continue + self._Fr, self._Fi, self._r = Fr_new, Fi_new, r_new + self.species_index[:] = spec_new + self.displacement[:] = disp_new + loss = loss_new + accepted[kind] += len(pick) + cap = min(int(cap * 1.25) + 1, batch) + if (sweep + 1) % refit_every == 0: + self._recompute_F() + self._solve_linear(self._model_diffuse()) + loss = self._update_residual() + su = None + self.loss_history.append(loss) + sweeps.set_postfix(loss=f"{loss:.4g}", **accepted) + self._recompute_F() + self._solve_linear(self._model_diffuse()) + self.loss_history[-1] = self._update_residual() + return self + + def _random_proposals(self, batch: int): + """Single atoms moved by -1, 0 or +1 grid steps along each axis (not all zero), within + ``max_displacement``: sites (B, 1) and new displacement vectors (B, 1, 3).""" + n_sites = len(self.site_x) + j = self.rng.choice(n_sites, batch, replace=False) + delta = self.rng.integers(-1, 2, (batch, 3)) + zero = ~np.any(delta, axis=1) + delta[zero, self.rng.integers(0, 3, zero.sum())] = self.rng.choice((-1, 1), zero.sum()) + new = self.displacement[j] + delta + ok = np.all(np.abs(new) <= self._max_steps, axis=1) + return j[ok][:, None], new[ok][:, None, :] + + def _omega_proposals(self, batch: int): + """Disjoint omega embryos along <111> rows: sites (B, 3) and their new displacement + vectors (B, 3, 3), the (0, +v, -v) collapse.""" + n_sites = len(self.site_x) + nc = self._site_lookup.shape[0] + j0 = self.rng.choice(n_sites, batch, replace=False) + vec = self._omega_vectors[self.rng.integers(0, len(self._omega_vectors), batch)] + step = np.sign(vec) # nearest neighbour along that <111>, in site units + xc = self.site_x[j0] // self.refine + sites = np.stack( + [self._site_lookup[tuple(np.mod(xc + t * step, nc).T)] for t in range(3)], 1 + ) + pattern = np.stack([np.zeros_like(vec), vec, -vec], 1) # (B, 3, 3) + current = self.displacement[sites] + standing = np.all(current == pattern, axis=(1, 2)) + new = np.where(standing[:, None, None], 0, pattern) + keep = _disjoint_rows(sites) + return sites[keep], new[keep] + + # ---------------------------------------------------------------- analysis + + def to_atoms(self, lattice_parameter: float | None = None): + """The fitted supercell as an ``ase.Atoms`` object, one atom per site. + + The mixed sites carry the species and static displacements of the fit; any ordered + sites of the average crystal are repeated over the supercell without displacements. + Atoms are grouped by species. + Species are labelled as fitted, so with ``set_crystal(merge={"Zr": "Nb"})`` the Nb + sites stand for both Nb and Zr. + + Parameters + ---------- + lattice_parameter : float, optional + Edge of the cubic unit cell in A. Default None: the lattice parameter of the + crystal passed to :meth:`set_crystal`. The fractional coordinates do not depend + on it; the lattice parameter fitted by :meth:`fit_geometry` also absorbs any error + in the detector pixel size. + + Returns + ------- + ase.Atoms + Periodic cubic supercell of ``cells`` unit cells along each axis. + """ + from ase import Atoms + from ase.data import chemical_symbols + + if getattr(self, "site_x", None) is None: + raise RuntimeError("build_supercell first") + a = self._a_crystal if lattice_parameter is None else float(lattice_parameter) + if not np.isfinite(a) or a <= 0: + raise ValueError("lattice_parameter must be positive") + frac = [self._positions(np.arange(len(self.site_x))) / self.grid_size] + symbols = [self.species[k] for k in self.species_index] + + # ordered sites of the average crystal (any site that is not a fitted mixed site) + mixed = self.site_grid / self.grid_divisor + pos0 = np.mod(self.crystal.positions_frac.numpy(), 1.0) + numbers = self.crystal.numbers.numpy() + ordered = {} + for f, z in zip(pos0, numbers): + if np.any(np.all(np.abs(np.mod(f - mixed + 0.5, 1.0) - 0.5) < 1e-4, axis=1)): + continue + key = tuple(np.round(f, 6)) + ordered.setdefault(key, chemical_symbols[int(z)]) + if ordered: + idx = np.stack(np.meshgrid(*(np.arange(self.cells),) * 3, indexing="ij"), -1).reshape( + -1, 3 + ) + f0 = np.array(list(ordered)) + frac.append(((idx[:, None, :] + f0[None]) / self.cells).reshape(-1, 3)) + symbols += list(ordered.values()) * len(idx) + # atoms grouped by species (stable within each), so the formula reads e.g. Nb12480V3520 + order = np.argsort(np.array(symbols), kind="stable") + return Atoms( + [symbols[i] for i in order], + scaled_positions=np.mod(np.concatenate(frac), 1.0)[order], + cell=np.eye(3) * a * self.cells, + pbc=True, + ) + + def to_cif(self, path, lattice_parameter: float | None = None): + """Write the fitted supercell (see :meth:`to_atoms`) to a CIF file with ASE. + + Parameters + ---------- + path : str or Path + Output file. + lattice_parameter : float, optional + Edge of the cubic unit cell in A. Default None: the crystal's lattice parameter. + + Returns + ------- + pathlib.Path + The file written. + """ + from pathlib import Path + + from ase.io import write + + path = Path(path) + write(path, self.to_atoms(lattice_parameter), format="cif") + return path + + def model_images(self, diffuse_only: bool = False) -> list[np.ndarray]: + """Model on the binned grid of every pattern (data units): scaled diffuse + background, + or the scaled diffuse term alone. The diffuse term is zero beyond ``q_max``.""" + i_used = self._diffuse_used().cpu().numpy().astype(np.float64) + diffuse = self._u_all * (self._W_all @ i_used) + out = [] + for i, k in enumerate(self.mask["k"]): + m = self._pix_image == i + sel = np.nonzero(m)[0] + c = self.coefficients[i] + img = c[0] * diffuse[m] + if not diffuse_only: + img = img + self._bg_basis(i, sel) @ c[1:] + out.append(img.reshape(k.shape[:2])) + return out + + def r_factors(self) -> dict: + """Weighted R of the diffuse fit per pattern. + + ``R = sqrt(sum w (y - model)^2 / sum w (y - background)^2)``: the + fraction of the diffuse signal (experiment minus fitted background) + the supercell leaves unexplained. + + Returns + ------- + dict + ``{pattern name: R}``. + """ + full = self.model_images() + bg = self.background_images() + out = {} + for i, name in enumerate(self.names): + w, y = self.mask["w"][i], self.mask["y"][i] + out[name] = float( + np.sqrt((w * (y - full[i]) ** 2).sum() / max((w * (y - bg[i]) ** 2).sum(), 1e-30)) + ) + return out + + def background_images(self) -> list[np.ndarray]: + """Fitted smooth background per pattern.""" + full = self.model_images() + diffuse = self.model_images(diffuse_only=True) + return [f - d for f, d in zip(full, diffuse)] + + def warren_cowley(self, n_shells: int | None = 6, max_radius: float | None = None) -> dict: + """Warren-Cowley alpha for every species pair over the first neighbour shells. + + ``alpha[shell, s, t] = 1 - P(t | s) / c_t`` for unlike pairs and + ``(P(s | s) - c_s) / (1 - c_s)`` for like pairs, with ``P(t | s)`` the + fraction of the shell around an ``s`` atom occupied by ``t``. + Negative: unlike neighbours preferred; positive: like. + + Parameters + ---------- + n_shells : int, optional + Largest number of neighbour shells. Default 6; None keeps every shell within + `max_radius`. + max_radius : float, optional + Largest neighbour distance in A. Default None (no limit beyond half the + supercell). + + Returns + ------- + dict + "radius": (n_shells,) shell radii in A; "alpha": (n_shells, K, K) + Warren-Cowley parameters; "species": the K species in index order. + """ + radii, prob = self._pair_probabilities(n_shells=n_shells, max_radius=max_radius) + c = self.concentrations + K = len(self.species) + alpha = np.zeros_like(prob) + for s in range(K): + for t in range(K): + p = prob[:, s, t] + alpha[:, s, t] = (p - c[t]) / (1 - c[t]) if s == t else 1 - p / c[t] + return dict(radius=radii, alpha=alpha, species=list(self.species)) + + def _pair_probabilities( + self, n_shells: int | None = None, max_radius: float | None = None + ) -> tuple[np.ndarray, np.ndarray]: + """Shell radii (A) and ``P(t | s)`` (n_shells, K, K): the fraction of the sites in each + neighbour shell of an ``s`` atom that hold a ``t`` atom, averaged over every ``s`` atom + of the supercell (periodic boundaries). Shells are kept up to `n_shells` and up to + `max_radius` (A), whichever ends first, and never beyond half the supercell.""" + d = self.grid_divisor + n = self.cells * d + xs = self.site_x // self.refine + K = len(self.species) + site = np.zeros((n,) * 3) + site[xs[:, 0], xs[:, 1], xs[:, 2]] = 1 + fsite = np.fft.fftn(site) + focc = [] + for s in range(K): + occ = np.zeros((n,) * 3) + sel = self.species_index == s + occ[xs[sel, 0], xs[sel, 1], xs[sel, 2]] = 1 + focc.append(np.fft.fftn(occ)) + ss = np.real(np.fft.ifftn(fsite * np.conj(fsite))) + v = np.stack(np.meshgrid(*(np.fft.fftfreq(n, 1 / n),) * 3, indexing="ij"), -1) + r = np.linalg.norm(v, axis=-1) / d * self.lattice_parameter + half_box = self.cells * self.lattice_parameter / 2 + r_max = half_box if max_radius is None else min(max_radius, half_box) + valid = (ss > 0.5) & (r <= r_max + 1e-6) + radii = np.unique(np.round(r[valid], 4))[1:] + if n_shells is not None: + radii = radii[:n_shells] + shell = np.searchsorted(radii, np.round(r, 4)) + inside = valid & (shell < len(radii)) + inside[inside] = radii[shell[inside]] == np.round(r[inside], 4) + index = shell[inside] + + def shell_sums(a): + return np.bincount(index, weights=a[inside], minlength=len(radii)) + + prob = np.zeros((len(radii), K, K)) + for s in range(K): + cs_site = shell_sums(np.real(np.fft.ifftn(focc[s] * np.conj(fsite)))) + for t in range(K): + cst = shell_sums(np.real(np.fft.ifftn(focc[s] * np.conj(focc[t])))) + prob[:, s, t] = cst / cs_site + return radii, prob + + def shell_correlations( + self, n_shells: int | None = None, max_radius: float | None = None + ) -> dict: + """Pair correlation of every species pair per neighbour shell, relative to a random + alloy. + + ``ratio[shell, s, t] = P(t | s) / c_t``, the number of ``s``-``t`` pairs in the shell + divided by the number expected for a random arrangement at the same composition. It + is symmetric in ``s`` and ``t``, equals 1 for a random alloy, is above 1 for pairs that + are favoured and below 1 for pairs that are avoided. ``c_t`` is the composition of + the supercell itself. For an unlike pair, ``ratio = 1 - alpha`` with the Warren-Cowley + parameter of :meth:`warren_cowley`, up to the rounding of the composition to whole + atoms. + + Parameters + ---------- + n_shells : int, optional + Largest number of neighbour shells. Default None (every shell within + `max_radius`). + max_radius : float, optional + Largest neighbour distance in A. Default None: five lattice parameters, far + enough for the short-range order to decay. Never beyond half the supercell. + + Returns + ------- + dict + "radius": (n_shells,) shell radii in A; "shell": bond-vector labels in units of a + (e.g. "1/2<111>"), or None for site lattices other than BCC; "ratio": + (n_shells, K, K) pair correlations relative to random; "species": the K species + in index order. + """ + if max_radius is None: + max_radius = 5 * self.lattice_parameter + radii, prob = self._pair_probabilities(n_shells=n_shells, max_radius=max_radius) + # the supercell's own composition, so that the ratio is exactly symmetric + counts = np.bincount(self.species_index, minlength=len(self.species)) + ratio = prob / (counts / counts.sum())[None, None, :] + shell = None + if self.grid_divisor == 2: + reach = int(np.ceil(radii[-1] / (self.lattice_parameter / 2))) + 1 + offsets = self._shells(len(radii), reach=reach) + dist = [np.linalg.norm(o[0]) * self.lattice_parameter / 2 for o in offsets] + if len(offsets) == len(radii) and np.allclose(dist, radii, atol=1e-3): + shell = [self._shell_label(o[0]) for o in offsets] + return dict(radius=radii, shell=shell, ratio=ratio, species=list(self.species)) + + def _u2(self, disp: np.ndarray) -> np.ndarray: + """Squared displacement (A^2) of each displacement vector (..., 3).""" + step = self.lattice_parameter / (self.grid_divisor * self.refine) + return (disp**2).sum(-1) * step**2 + + def static_b(self) -> float: + """Static Debye-Waller B (A^2) of the displacements, 8 pi^2 / 3.""" + return float(8 * np.pi**2 * self._u2(self.displacement).mean() / 3) + + def _omega_like(self) -> np.ndarray: + """Atoms displaced along a <111> by at least a half omega collapse.""" + a = np.abs(self.displacement) + unit = self.refine * self.grid_divisor // 24 + return (a[:, 0] >= unit) & (a[:, 0] == a[:, 1]) & (a[:, 1] == a[:, 2]) + + def displacement_summary(self) -> dict: + """Static Debye-Waller B, mean displacement and omega fraction per species.""" + step = self.lattice_parameter / (self.grid_divisor * self.refine) + mag = np.linalg.norm(self.displacement, axis=1) * step + omega = self._omega_like() + out = dict( + static_b=self.static_b(), + mean_displacement={ + s: float(mag[self.species_index == k].mean()) for k, s in enumerate(self.species) + }, + omega_fraction={ + s: float(omega[self.species_index == k].mean()) for k, s in enumerate(self.species) + }, + ) + return out + + def diffuse_section( + self, + normal=(0, 0, 1), + extent: float = 2.0, + smooth: bool = True, + grid: np.ndarray | None = None, + ): + """Symmetrized supercell diffuse intensity on a reciprocal-lattice plane, in Laue units. + + Parameters + ---------- + normal : (h, k, l) + Plane normal; the plane passes through the origin. + extent : float + Half width in units of the cubic reciprocal lattice vector 1/a. + smooth : bool + Blur by the fit's ``resolution`` kernel. + grid : ndarray, optional + Output of ``diffuse_grid()``, to reuse. + + Returns + ------- + image, (u, v) in-plane axes (crystal frame, unit vectors), distance of each pixel from + the nearest reciprocal-lattice node (1/a units) + """ + grid = self.diffuse_grid() if grid is None else grid + if smooth and self.resolution > 0: + grid = ndimage.gaussian_filter(grid, self.resolution, mode="wrap") + frame = _zone_frame(normal) + u, v = frame[0], frame[1] + steps = np.arange(-extent * self.cells, extent * self.cells + 1) + su, sv = np.meshgrid(steps, steps, indexing="ij") + pts = su[..., None] * u + sv[..., None] * v # grid units (cells per 1/a) + img = ndimage.map_coordinates( + grid, np.moveaxis(pts, -1, 0).reshape(3, -1), order=1, mode="grid-wrap" + ).reshape(su.shape) + q = np.linalg.norm(pts, axis=-1) / (self.cells * self.lattice_parameter) + fbar, f2 = self._fbar(q.ravel()) + laue = (f2 - fbar**2).reshape(q.shape) + frac = pts / self.cells + node_dist = np.linalg.norm(frac - np.round(frac), axis=-1) # 1/a units + return img / laue, (u, v), node_dist + + # ----------------------------------------------------------------- plotting + + def plot_images( + self, quantiles: tuple[float, float] = (0.86, 0.98), cmap: str = "turbo_black", **kwargs + ): + """Binned patterns on a linear scale. Once a mask is set, the scale spans min to max of the + diffuse region of each pattern; before that, ``quantiles`` of the whole pattern.""" + from quantem.core.visualization import show_2d + + arrays, norms = [], [] + for i, im in enumerate(self.images): + b = self._binned(im) + if self.mask is not None: + q = np.linalg.norm(self.mask["k"][i], axis=-1) + sel = (self.mask["w"][i] > 0.5) & (q < self.mask["q_max"]) + vals = b[sel] + norms.append(dict(interval_type="manual", vmin=vals.min(), vmax=vals.max())) + else: + lo, hi = np.quantile(b, quantiles) + norms.append(dict(interval_type="manual", vmin=lo, vmax=hi)) + arrays.append(b) + return show_2d( + arrays, + title=self.names, + norm=norms, + cmap=cmap, + axsize=kwargs.pop("axsize", (6, 4)), + **kwargs, + ) + + def plot_geometry(self, power: float = 0.3, q_view: float = 1.4, **kwargs): + """Each pattern with its fitted reflections (red), direct beam (cyan) and Laue circle (yellow).""" + import matplotlib.patches as mpatches + + from quantem.core.visualization import show_2d + + fig, axs = show_2d( + self.images, + title=self.names, + norm={ + "stretch_type": "power", + "power": power, + "lower_quantile": 0.05, + "upper_quantile": 0.999, + }, + axsize=kwargs.pop("axsize", (5, 5)), + **kwargs, + ) + axs = np.atleast_1d(axs).ravel() + r = 0.03 / self.sampling + k_wave = 1.0 / self.wavelength + for i, ax in enumerate(axs): + p = self.bragg_positions(i) + c = self.geometry["centers"][i] + for pr, pc in p: + ax.add_patch(mpatches.Circle((pc, pr), r, fill=False, color="tab:red", lw=1.0)) + ax.add_patch(mpatches.Circle((c[1], c[0]), 1.5 * r, fill=False, color="cyan", lw=1.5)) + tilt = self.geometry["tilts"][i] + if np.linalg.norm(tilt) > 0: + th = np.linspace(0, 2 * np.pi, 361) + circ = -k_wave * tilt + k_wave * np.linalg.norm(tilt) * np.column_stack( + [np.cos(th), np.sin(th)] + ) + px = ( + c + + (circ * self.lattice_parameter / self._a_crystal) + @ self.geometry["matrices"][i].T + ) + ax.plot(px[:, 1], px[:, 0], "--", color="yellow", lw=1.0) + half = q_view / self.sampling + ax.set_xlim(c[1] - half, c[1] + half) + ax.set_ylim(c[0] + half, c[0] - half) + return fig, axs + + def plot_mask(self, **kwargs): + """Binned diffuse weight of every pattern.""" + from quantem.core.visualization import show_2d + + return show_2d( + self.mask["w"], + title=self.names, + cmap="gray", + axsize=kwargs.pop("axsize", (4, 2.7)), + **kwargs, + ) + + def plot_fit( + self, + columns: Sequence[str] = ("experiment", "background", "difference", "model", "residual"), + diffuse_only: bool = False, + sigma: float = 0.7, + quantiles: tuple[float, float] = (0.01, 0.99), + cmap: str = "turbo_black", + **kwargs, + ): + """Experiment, fitted background, experiment - background, supercell diffuse model and + residual, one row per zone, cropped to ``q_max``. + + Experiment and background share a linear scale spanning ``quantiles`` of the + experiment inside the diffuse mask (weight > 0.5); experiment - background and the + model share one spanning ``quantiles`` of the difference; the residual uses a diverging + map at half that range (white = no difference). Masked pixels are black in the last + three columns. The experiment is blurred by ``sigma`` binned pixels. ``diffuse_only`` + keeps the last three columns. + """ + import matplotlib + + from quantem.core.visualization import show_2d + + if diffuse_only: + columns = ("difference", "model", "residual") + full = self.model_images() + diffuse = self.model_images(diffuse_only=True) + cmap_obj = matplotlib.colormaps[cmap].with_extremes(bad="black") + diverging = matplotlib.colormaps["RdBu_r"].with_extremes(bad="black") + titles_of = { + "experiment": "experiment", + "background": "background", + "difference": "experiment - background", + "model": "model", + "residual": "residual", + } + rows, titles, norms, cmaps = [], [], [], [] + for i in range(len(self.images)): + q = np.linalg.norm(self.mask["k"][i], axis=-1) + inside = q < self.mask["q_max"] + w = self.mask["w"][i] + sel = (w > 0.5) & inside + rr, cc = np.nonzero(inside) + crop = (slice(rr.min(), rr.max() + 1), slice(cc.min(), cc.max() + 1)) + background = full[i] - diffuse[i] + exp = self.mask["data"][i] + if sigma: + exp = ndimage.gaussian_filter(exp, sigma) + diff = self.mask["y"][i] - background + if sigma: + diff = _nan_blur(np.where(w > 0.2, diff, np.nan), sigma) + diff = np.where(sel, diff, np.nan) + panels = { + "experiment": np.where(inside, exp, np.nan), + "background": np.where(inside, background, np.nan), + "difference": diff, + "model": np.where(sel, diffuse[i], np.nan), + "residual": diff - np.where(sel, diffuse[i], np.nan), + } + lo, hi = np.quantile(exp[sel], quantiles) + dlo, dhi = np.nanquantile(diff, quantiles) + half = 0.5 * (dhi - dlo) + scales = { + "experiment": (lo, hi), + "background": (lo, hi), + "difference": (dlo, dhi), + "model": (dlo, dhi), + "residual": (-half, half), + } + rows.append([panels[c][crop] for c in columns]) + titles.append([f"{self.names[i]} {titles_of[c]}" for c in columns]) + norms.append( + [ + dict(interval_type="manual", vmin=scales[c][0], vmax=scales[c][1]) + for c in columns + ] + ) + cmaps.append([diverging if c == "residual" else cmap_obj for c in columns]) + return show_2d( + rows, + title=titles, + norm=norms, + cmap=cmaps, + axsize=kwargs.pop("axsize", (3.4, 3.4)), + **kwargs, + ) + + def plot_warren_cowley(self, max_radius: float | None = None): + """Warren-Cowley alpha against neighbour distance (0 = random, < 0 unlike neighbours + preferred, > 0 like); one curve for a binary site, one per species pair otherwise. + + Parameters + ---------- + max_radius : float, optional + Largest neighbour distance in A. Default None: five lattice parameters. + + Returns + ------- + fig, ax + """ + import matplotlib.pyplot as plt + + max_radius = self._default_shell_radius(max_radius) + sro = self.warren_cowley(n_shells=None, max_radius=max_radius) + sc = self.shell_correlations(max_radius=max_radius) + labelled = self._direction_shells(sc["radius"], sc["shell"]) + K = len(self.species) + fig, ax = plt.subplots(figsize=(8.0, 3.6)) + ax.axhline(0, color="0.6", lw=0.8) + for radius, _ in labelled: + ax.axvline(radius, color="0.9", lw=0.6, zorder=0) + if K == 2: # a binary site has one alpha for every pair + ax.plot(sro["radius"], sro["alpha"][:, 0, 0], "o-", ms=3, lw=1.2, color="tab:blue") + ax.set_ylabel(f"Warren-Cowley alpha, {self.species[0]}-{self.species[1]}") + else: + for s_ in range(K): + for t in range(s_, K): + ax.plot( + sro["radius"], + sro["alpha"][:, s_, t], + "o-" if s_ == t else "s--", + ms=3, + lw=1.2, + label=f"{self.species[s_]}-{self.species[t]}", + ) + ax.legend(fontsize=8) + ax.set_ylabel("Warren-Cowley alpha") + ax.set_xlabel("neighbour distance (A)") + ax.set_xlim(0, sro["radius"][-1] * 1.02) + fig.tight_layout() + self._label_direction_shells(fig, ax, labelled, sro["radius"][-1] * 1.02) + return fig, ax + + def _default_shell_radius(self, max_radius: float | None) -> float: + """Largest neighbour distance of the shell plots: `max_radius`, or five lattice parameters.""" + return 5 * self.lattice_parameter if max_radius is None else max_radius + + @staticmethod + def _direction_shells(radii, labels) -> list[tuple[float, str]]: + """(radius, label) of the shells along <100>, <110> and <111>, from the bond-vector + labels of :meth:`shell_correlations`; empty when there are no labels.""" + if labels is None: + return [] + out = [] + for radius, label in zip(radii, labels): + digits = [int(c) for c in label.split("<")[1].rstrip(">")] + if ( + digits[1:] == [0, 0] + or (digits[0] == digits[1] and digits[2] == 0) + or digits[0] == digits[1] == digits[2] + ): + out.append((radius, label)) + return out + + @staticmethod + def _label_direction_shells(fig, ax, labelled, span: float) -> None: + """Write shell labels above `ax`, stacked in rows so that none overlap, and make room + for them at the top of the figure. Call after ``tight_layout``.""" + last: list[float] = [] + for radius, label in labelled: + width = 0.012 * span * len(label) + row = next((k for k, x in enumerate(last) if radius - width / 2 > x), len(last)) + if row == len(last): + last.append(0.0) + last[row] = radius + width / 2 + ax.annotate( + label, + (radius, 1.0), + xycoords=("data", "axes fraction"), + xytext=(0, 3 + 11 * row), + textcoords="offset points", + ha="center", + va="bottom", + fontsize=8, + annotation_clip=False, + ) + if last: + height = fig.get_size_inches()[1] + fig.subplots_adjust(top=fig.subplotpars.top - (0.03 + 0.15 * len(last)) / height) + + def plot_shell_correlations( + self, + max_radius: float | None = None, + panel_size: tuple[float, float] = (8.0, 2.2), + ): + """Pair correlation of every species pair against neighbour distance, relative to a + random alloy (1 = random), one panel per pair stacked on shared distance and ratio axes. The + shells along <100>, <110> and <111> are labelled above the top panel (BCC only). + + Parameters + ---------- + max_radius : float, optional + Largest neighbour distance in A. Default None: five lattice parameters. + panel_size : tuple of float, optional + Width and height of each panel in inches. Default (8.0, 2.2). + + Returns + ------- + fig, axs + The figure and the array of panels, top to bottom. + """ + import matplotlib.pyplot as plt + + sc = self.shell_correlations(max_radius=self._default_shell_radius(max_radius)) + K = len(self.species) + pairs = [(s, t) for s in range(K) for t in range(s, K)] + fig, axs = plt.subplots( + len(pairs), + 1, + sharex=True, + sharey=True, + figsize=(panel_size[0], panel_size[1] * len(pairs)), + ) + axs = np.atleast_1d(axs) + labelled = self._direction_shells(sc["radius"], sc["shell"]) + for ax, (s, t) in zip(axs, pairs): + ax.axhline(1, color="0.6", lw=0.8) + for radius, _ in labelled: + ax.axvline(radius, color="0.9", lw=0.6, zorder=0) + ax.plot(sc["radius"], sc["ratio"][:, s, t], "o-", ms=3, lw=1.2, color="tab:blue") + ax.set_ylabel(f"{self.species[s]}-{self.species[t]}\n/ random") + axs[-1].set_xlabel("neighbour distance (A)") + axs[-1].set_xlim(0, sc["radius"][-1] * 1.02) + fig.tight_layout() + self._label_direction_shells(fig, axs[0], labelled, sc["radius"][-1] * 1.02) + return fig, axs + + def plot_diffuse_sections(self, extent: float = 2.0, normals=((0, 0, 1), (1, -1, 0))): + """Symmetrized diffuse intensity of the supercell in Laue units (1 = random alloy) on + reciprocal-lattice planes through the origin. Maxima at special points name the order + (100: B2-type, 1/2 1/2 1/2: D0_3, 2/3 2/3 2/3: omega). The color scale is set away + from the reciprocal-lattice nodes.""" + import matplotlib.pyplot as plt + + grid = self.diffuse_grid() + secs = [self.diffuse_section(n, extent, grid=grid) for n in normals] + between = np.concatenate([sec[d > 0.2] for sec, _, d in secs]) + vmax = np.quantile(between, 0.995) + fig, axs = plt.subplots(1, len(normals), figsize=(5.2 * len(normals), 4.4)) + axs = np.atleast_1d(axs) + ext = [-extent, extent, -extent, extent] + for ax, n, (sec, (u, v), _) in zip(axs, normals, secs): + im = ax.imshow( + sec.T, origin="lower", extent=ext, cmap="turbo_black", vmin=0, vmax=vmax + ) + ax.set_title("(" + "".join(f"{int(x)}" for x in n) + ") section, Laue units") + ax.set_xlabel("[" + " ".join(f"{x:.2f}" for x in u) + "] (1/a)") + ax.set_ylabel("[" + " ".join(f"{x:.2f}" for x in v) + "] (1/a)") + fig.colorbar(im, ax=ax, fraction=0.046) + fig.tight_layout() + return fig, axs + + def _shells(self, n_shells: int, reach: int = 4): + """Neighbour offsets of the BCC site lattice (site units, a / 2) grouped by distance, + complete for every shell within `reach` site units.""" + r = np.arange(-reach, reach + 1) + v = np.stack(np.meshgrid(r, r, r, indexing="ij"), -1).reshape(-1, 3) + bcc = np.all(v % 2 == 0, axis=1) | np.all(v % 2 == 1, axis=1) + v = v[bcc & np.any(v != 0, axis=1)] + d2 = (v**2).sum(1) + return [v[d2 == d] for d in np.unique(d2[d2 <= reach**2])[:n_shells]] + + @staticmethod + def _shell_label(offset) -> str: + """Bond vector of a BCC neighbour shell in units of a, e.g. "1/2<111>" or "<100>", from + one offset in site units (a / 2).""" + v = np.sort(np.abs(np.asarray(offset)))[::-1] + if np.all(v % 2 == 0): + return "<" + "".join(str(int(x)) for x in v // 2) + ">" + return "1/2<" + "".join(str(int(x)) for x in v) + ">" + + def displacement_correlations(self, n_shells: int = 6) -> dict: + """Displacement short-range order per neighbour shell. + + ``longitudinal``: <(u_i.r)(u_j.r)> / <(u.r)^2> with r the bond direction; + ``transverse``: the same for the components normal to the bond. Omega embryos give a + strong negative longitudinal correlation on the nearest-neighbour <111> bond (the + collapsing pair moves together). + + Parameters + ---------- + n_shells : int, optional + Number of neighbour shells. Default 6. + + Returns + ------- + dict + "radius": (n_shells,) shell radii in A; "shell": bond-vector + labels (e.g. "1/2<111>", "<100>"); "longitudinal" and + "transverse": lists of correlations, dimensionless. + + Raises + ------ + NotImplementedError + For site lattices other than BCC. + """ + if self.grid_divisor != 2: + raise NotImplementedError("Displacement correlations are implemented for BCC sites.") + step = self.lattice_parameter / (self.grid_divisor * self.refine) + u = self.displacement * step + nc = self._site_lookup.shape[0] + xc = self.site_x // self.refine + out = dict(radius=[], shell=[], longitudinal=[], transverse=[]) + for offs in self._shells(n_shells): + out["shell"].append(self._shell_label(offs[0])) + rhat = offs / np.linalg.norm(offs, axis=1, keepdims=True) + nb = np.stack([self._site_lookup[tuple(np.mod(xc + o, nc).T)] for o in offs], 1) + ui = u[:, None, :] + uj = u[nb] + li = (ui * rhat[None]).sum(-1) + lj = (uj * rhat[None]).sum(-1) + ti = ui - li[..., None] * rhat[None] + tj = uj - lj[..., None] * rhat[None] + ll = (li**2).mean() + tt = (ti**2).sum(-1).mean() + out["radius"].append(np.linalg.norm(offs[0]) * self.lattice_parameter / 2) + out["longitudinal"].append(float((li * lj).mean() / ll) if ll > 0 else 0.0) + out["transverse"].append(float((ti * tj).sum(-1).mean() / tt) if tt > 0 else 0.0) + out["radius"] = np.asarray(out["radius"]) + return out + + def plot_displacement_correlations(self, n_shells: int = 8): + """Displacement short-range order: longitudinal and transverse displacement + correlations for each neighbour shell, labelled by the shell's bond vector in units of + a. BCC has no shell between a (second neighbours) and a sqrt(2) (third).""" + import matplotlib.pyplot as plt + + c = self.displacement_correlations(n_shells) + fig, ax = plt.subplots(figsize=(9, 4.5)) + ax.axhline(0, color="0.6", lw=0.8) + ax.plot(c["radius"], c["longitudinal"], "o-", ms=5, label="longitudinal (along bond)") + ax.plot(c["radius"], c["transverse"], "s--", ms=5, label="transverse") + ax.set_xticks(c["radius"]) + ax.set_xticklabels(c["shell"], rotation=45, fontsize=8) + ax.set_xlabel("neighbour shell (bond vector / a)") + ax.set_ylabel("displacement correlation") + top = ax.secondary_xaxis("top") + top.set_xticks(c["radius"]) + top.set_xticklabels([f"{r:.2f}" for r in c["radius"]], fontsize=7) + top.set_xlabel("distance (A)") + ax.legend(fontsize=8) + fig.tight_layout() + return fig, ax + + def displacement_distributions(self) -> dict: + """Probability of each projected displacement u.n (A) per species, pooled over the + symmetry-equivalent directions n of <100>, <110> and <111> (both senses).""" + step = self.lattice_parameter / (self.grid_divisor * self.refine) + families = { + "<100>": np.eye(3), + "<110>": np.array( + [[1, 1, 0], [1, -1, 0], [1, 0, 1], [1, 0, -1], [0, 1, 1], [0, 1, -1]] + ), + "<111>": np.array([[1, 1, 1], [1, 1, -1], [1, -1, 1], [-1, 1, 1]]), + } + out = {} + for name, dirs in families.items(): + n = dirs / np.linalg.norm(dirs, axis=1, keepdims=True) + proj = self.displacement @ n.T * step # (n_sites, n_dirs) + proj = np.concatenate([proj, -proj], axis=1) + out[name] = {} + for k, sp in enumerate(self.species): + v = np.round(proj[self.species_index == k].ravel(), 4) + vals, counts = np.unique(v, return_counts=True) + out[name][sp] = (vals, counts / counts.sum()) + return out + + def plot_displacements(self, **kwargs): + """Probability distribution of the static displacement of each species projected on + <100>, <110> and <111> (pooled over equivalent directions and both senses), on a log + scale. Displacements live on a grid, so each curve connects the grid values.""" + import matplotlib.pyplot as plt + + dist = self.displacement_distributions() + fig, axs = plt.subplots(1, 3, figsize=kwargs.pop("figsize", (14, 3.8)), sharey=True) + colors = plt.get_cmap("tab10")(np.arange(len(self.species))) + for ax, (name, per_species) in zip(axs, dist.items()): + for k, sp in enumerate(self.species): + vals, prob = per_species[sp] + ax.plot(vals, prob, "o-", ms=4, lw=1.2, color=colors[k], label=sp) + ax.set_yscale("log") + ax.set_xlabel(f"u . n, n along {name} (A)") + ax.set_title(name) + axs[0].set_ylabel("probability") + axs[0].legend() + summary = self.displacement_summary() + fig.suptitle( + f"static B = {summary['static_b']:.3f} A^2; mean |u| " + + ", ".join(f"{k} {v:.3f} A" for k, v in summary["mean_displacement"].items()) + ) + fig.tight_layout() + return fig, axs + + def plot_loss(self): + """Weighted loss after each sweep (``loss_history``) on a log scale. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + import matplotlib.pyplot as plt + + fig, ax = plt.subplots(figsize=(5, 3)) + ax.plot(self.loss_history, "k.-") + ax.set_xlabel("sweep") + ax.set_ylabel("weighted loss") + ax.set_yscale("log") + fig.tight_layout() + return fig, ax + + +def _lsq_bounded(X: np.ndarray, y: np.ndarray, sw: np.ndarray, lo: np.ndarray) -> np.ndarray: + """Weighted least squares with lower bounds, columns normalized for conditioning.""" + Xw = X * sw[:, None] + norm = np.linalg.norm(Xw, axis=0) + norm[norm == 0] = 1.0 + res = optimize.lsq_linear(Xw / norm, y * sw, bounds=(lo * norm, np.inf)) + return res.x / norm + + +def _disjoint_rows(sites: np.ndarray) -> list[int]: + """Rows of ``sites`` sharing no site with an earlier kept row.""" + seen, keep = set(), [] + for b, row in enumerate(sites): + r = row.tolist() + if seen.isdisjoint(r): + seen.update(r) + keep.append(b) + return keep + + +def _integrate_spots(im: np.ndarray, positions: np.ndarray, radius: float) -> np.ndarray: + """Integrated intensity inside ``radius`` px of each position, over the median of the ring + out to 1.6 radius; NaN for spots off the detector.""" + out = np.full(len(positions), np.nan) + r_out = int(np.ceil(1.6 * radius)) + 1 + yy, xx = np.mgrid[-r_out : r_out + 1, -r_out : r_out + 1] + for n, (pr, pc) in enumerate(positions): + r0, c0 = int(round(pr)), int(round(pc)) + if ( + r0 - r_out < 0 + or c0 - r_out < 0 + or r0 + r_out >= im.shape[0] + or c0 + r_out >= im.shape[1] + ): + continue + win = im[r0 - r_out : r0 + r_out + 1, c0 - r_out : c0 + r_out + 1] + dist = np.hypot(yy + r0 - pr, xx + c0 - pc) + bg = np.median(win[(dist >= radius) & (dist < 1.6 * radius)]) + out[n] = (win[dist < radius] - bg).sum() + return out + + +def _nan_blur(a: np.ndarray, sigma: float) -> np.ndarray: + """Gaussian blur that ignores NaNs.""" + ok = np.isfinite(a) + num = ndimage.gaussian_filter(np.where(ok, a, 0.0), sigma) + den = ndimage.gaussian_filter(ok.astype(float), sigma) + return np.where(ok, num / np.maximum(den, 1e-12), np.nan) diff --git a/src/quantem/diffraction/rotations.py b/src/quantem/diffraction/rotations.py new file mode 100644 index 000000000..60fd668b1 --- /dev/null +++ b/src/quantem/diffraction/rotations.py @@ -0,0 +1,777 @@ +"""Quaternion rotation utilities for orientation mapping. + +All orientations in quantem.diffraction are represented as unit quaternions, +stored as torch tensors of shape (..., 4) in scalar-first order (w, x, y, z). +Rotation matrices, Euler angles, and axis-angle forms are provided only as +conversions at the boundaries. + +Convention +---------- +A quaternion q represents the rotation of crystal-frame vectors into the +laboratory (beam) frame:: + + v_lab = R(q) @ v_crystal + +The electron beam travels along -z in the lab frame. The zone axis is the +crystal direction that points from the specimen back toward the source, +lab +z, expressed in crystal Cartesian coordinates. It is therefore the +third row of R(q):: + + zone_axis = R(q).T @ [0, 0, 1] + +Euler angles use the Z-X-Z convention (Rowenhorst et al., 2015). +""" + +from __future__ import annotations + +import numpy as np +import torch + + +def qmult(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """Hamilton product a * b of quaternions, broadcasting over leading dims. + + The product applies `b` first, then `a`: R(a * b) = R(a) @ R(b). + + Parameters + ---------- + a, b : torch.Tensor + Scalar-first quaternions (..., 4), broadcastable. + + Returns + ------- + torch.Tensor + Product quaternions (..., 4), not renormalized. + """ + aw, ax, ay, az = a.unbind(-1) + bw, bx, by, bz = b.unbind(-1) + return torch.stack( + ( + aw * bw - ax * bx - ay * by - az * bz, + aw * bx + ax * bw + ay * bz - az * by, + aw * by - ax * bz + ay * bw + az * bx, + aw * bz + ax * by - ay * bx + az * bw, + ), + dim=-1, + ) + + +def qconj(q: torch.Tensor) -> torch.Tensor: + """Quaternion conjugate (inverse for unit quaternions). + + Parameters + ---------- + q : torch.Tensor + Scalar-first quaternions (..., 4). + + Returns + ------- + torch.Tensor + (w, -x, -y, -z), shape (..., 4). For a crystal-to-lab orientation + this is the lab-to-crystal rotation. + """ + w, x, y, z = q.unbind(-1) + return torch.stack((w, -x, -y, -z), dim=-1) + + +def qnormalize(q: torch.Tensor) -> torch.Tensor: + """Normalize to unit length, with w >= 0 canonicalization. + + Parameters + ---------- + q : torch.Tensor + Scalar-first quaternions (..., 4), nonzero. + + Returns + ------- + torch.Tensor + Unit quaternions (..., 4) with w >= 0; q and -q describe the same + rotation, so this picks one of the two. + """ + q = q / torch.linalg.norm(q, dim=-1, keepdim=True) + return torch.where(q[..., :1] < 0, -q, q) + + +def qrotate(q: torch.Tensor, v: torch.Tensor) -> torch.Tensor: + """Rotate vectors by quaternions, v' = R(q) @ v. + + Parameters + ---------- + q : torch.Tensor + Unit scalar-first quaternions (..., 4). For an orientation, crystal + frame vectors are rotated into the lab frame. + v : torch.Tensor + Vectors (..., 3), broadcastable against `q`. + + Returns + ------- + torch.Tensor + Rotated vectors (..., 3). + """ + qv = torch.cat((torch.zeros_like(v[..., :1]), v), dim=-1) + return qmult(qmult(q, qv), qconj(q))[..., 1:] + + +def quat_to_matrix(q: torch.Tensor) -> torch.Tensor: + """Convert quaternions to rotation matrices. + + Parameters + ---------- + q : torch.Tensor + Unit scalar-first quaternions (..., 4). + + Returns + ------- + torch.Tensor + Rotation matrices (..., 3, 3) with v_lab = R @ v_crystal for an + orientation. + """ + w, x, y, z = q.unbind(-1) + two = 2.0 + R = torch.stack( + ( + 1 - two * (y * y + z * z), + two * (x * y - w * z), + two * (x * z + w * y), + two * (x * y + w * z), + 1 - two * (x * x + z * z), + two * (y * z - w * x), + two * (x * z - w * y), + two * (y * z + w * x), + 1 - two * (x * x + y * y), + ), + dim=-1, + ) + return R.reshape(q.shape[:-1] + (3, 3)) + + +def quat_from_matrix(R: torch.Tensor) -> torch.Tensor: + """Convert rotation matrices to unit quaternions. + + Uses the numerically stable branch selection of Shepperd's method, + vectorized over leading dimensions. + + Parameters + ---------- + R : torch.Tensor + Proper rotation matrices (..., 3, 3). + + Returns + ------- + torch.Tensor + Unit scalar-first quaternions (..., 4) with w >= 0, the inverse of + :func:`quat_to_matrix`. + """ + batch_shape = R.shape[:-2] + R = R.reshape(-1, 3, 3) + m00, m01, m02 = R[:, 0, 0], R[:, 0, 1], R[:, 0, 2] + m10, m11, m12 = R[:, 1, 0], R[:, 1, 1], R[:, 1, 2] + m20, m21, m22 = R[:, 2, 0], R[:, 2, 1], R[:, 2, 2] + + # four candidate solutions, one per branch + q_w = torch.stack((1 + m00 + m11 + m22, m21 - m12, m02 - m20, m10 - m01), dim=-1) + q_x = torch.stack((m21 - m12, 1 + m00 - m11 - m22, m01 + m10, m02 + m20), dim=-1) + q_y = torch.stack((m02 - m20, m01 + m10, 1 - m00 + m11 - m22, m12 + m21), dim=-1) + q_z = torch.stack((m10 - m01, m02 + m20, m12 + m21, 1 - m00 - m11 + m22), dim=-1) + q_all = torch.stack((q_w, q_x, q_y, q_z), dim=1) # (N, 4, 4) + + trace_terms = torch.stack( + (1 + m00 + m11 + m22, 1 + m00 - m11 - m22, 1 - m00 + m11 - m22, 1 - m00 - m11 + m22), + dim=-1, + ) + branch = trace_terms.argmax(dim=-1) + q = q_all[torch.arange(R.shape[0], device=R.device), branch] + return qnormalize(q).reshape(batch_shape + (4,)) + + +def quat_from_axis_angle(axis: torch.Tensor, angle: torch.Tensor) -> torch.Tensor: + """Quaternion for a right-handed rotation about an axis. + + Parameters + ---------- + axis : torch.Tensor + Rotation axes (..., 3); need not be normalized. + angle : torch.Tensor + Rotation angles (...,), radians. + + Returns + ------- + torch.Tensor + Unit scalar-first quaternions (..., 4). + """ + axis = axis / torch.linalg.norm(axis, dim=-1, keepdim=True) + angle = torch.as_tensor(angle, dtype=axis.dtype, device=axis.device) + half = angle[..., None] / 2 + return torch.cat((torch.cos(half), torch.sin(half) * axis), dim=-1) + + +def quat_to_axis_angle(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Rotation axis and angle of unit quaternions. + + Parameters + ---------- + q : torch.Tensor + Scalar-first quaternions (..., 4). + + Returns + ------- + axis : torch.Tensor + Unit rotation axes (..., 3). Undefined (near zero) for the identity. + angle : torch.Tensor + Rotation angles (...,) in [0, pi], radians, after w >= 0 + canonicalization. + """ + q = qnormalize(q) + angle = 2 * torch.acos(q[..., 0].clamp(-1, 1)) + sin_half = torch.sqrt((1 - q[..., 0] ** 2).clamp_min(1e-24)) + axis = q[..., 1:] / sin_half[..., None] + return axis, angle + + +def quat_from_euler_zxz(angles: torch.Tensor) -> torch.Tensor: + """Quaternion from Z-X-Z Euler angles. + + Parameters + ---------- + angles : torch.Tensor + (phi1, Phi, phi2) Euler angles (..., 3), radians, applied as + R = Rz(phi1) @ Rx(Phi) @ Rz(phi2). + + Returns + ------- + torch.Tensor + Scalar-first quaternions (..., 4). + """ + a, b, c = angles.unbind(-1) + z = torch.zeros_like(a) + qa = torch.stack((torch.cos(a / 2), z, z, torch.sin(a / 2)), dim=-1) + qb = torch.stack((torch.cos(b / 2), torch.sin(b / 2), z, z), dim=-1) + qc = torch.stack((torch.cos(c / 2), z, z, torch.sin(c / 2)), dim=-1) + return qmult(qmult(qa, qb), qc) + + +def quat_to_euler_zxz(q: torch.Tensor) -> torch.Tensor: + """Z-X-Z Euler angles from unit quaternions. + + Parameters + ---------- + q : torch.Tensor + Unit scalar-first quaternions (..., 4). + + Returns + ------- + torch.Tensor + (phi1, Phi, phi2) Euler angles (..., 3), radians, the inverse of + :func:`quat_from_euler_zxz`. In the gimbal-locked case (Phi = 0 or + pi) the whole rotation about z is put in phi1 and phi2 is 0. + """ + R = quat_to_matrix(q) + beta = torch.acos(R[..., 2, 2].clamp(-1, 1)) + alpha = torch.atan2(R[..., 0, 2], -R[..., 1, 2]) + gamma = torch.atan2(R[..., 2, 0], R[..., 2, 1]) + # gimbal-locked cases: fold everything into alpha + locked = torch.sin(beta).abs() < 1e-8 + alpha_locked = torch.atan2(R[..., 1, 0], R[..., 0, 0]) + alpha = torch.where(locked, alpha_locked, alpha) + gamma = torch.where(locked, torch.zeros_like(gamma), gamma) + return torch.stack((alpha, beta, gamma), dim=-1) + + +def quat_from_zone_axis( + zone_axis: torch.Tensor, + in_plane_deg: torch.Tensor | float = 0.0, +) -> torch.Tensor: + """Orientation with the given crystal direction along the beam. + + Parameters + ---------- + zone_axis : torch.Tensor + Crystal-frame Cartesian direction(s) (..., 3) to place along lab +z, + pointing toward the source; need not be normalized. + in_plane_deg : torch.Tensor | float, default=0.0 + Additional rotation about lab z, degrees. + + Returns + ------- + torch.Tensor + Unit quaternions (..., 4), crystal to lab, such that + quat_to_matrix(q).T @ [0, 0, 1] == zone_axis. + """ + v = zone_axis / torch.linalg.norm(zone_axis, dim=-1, keepdim=True) + zhat = torch.zeros_like(v) + zhat[..., 2] = 1.0 + # minimal rotation taking zone axis to z + axis = torch.cross(v, zhat, dim=-1) + sin_t = torch.linalg.norm(axis, dim=-1) + cos_t = v[..., 2] + angle = torch.atan2(sin_t, cos_t) + # antiparallel / parallel cases: rotate about x + fallback = torch.zeros_like(v) + fallback[..., 0] = 1.0 + axis = torch.where(sin_t[..., None] < 1e-12, fallback, axis) + q_tilt = quat_from_axis_angle(axis, angle) + in_plane = torch.deg2rad( + torch.as_tensor(in_plane_deg, dtype=v.dtype, device=v.device) + ).broadcast_to(v.shape[:-1]) + z3 = torch.zeros_like(in_plane) + q_spin = torch.stack((torch.cos(in_plane / 2), z3, z3, torch.sin(in_plane / 2)), dim=-1) + return qnormalize(qmult(q_spin, q_tilt)) + + +def zone_axis_from_quat(q: torch.Tensor) -> torch.Tensor: + """Zone axis of orientations, in crystal Cartesian coordinates. + + Parameters + ---------- + q : torch.Tensor + Unit scalar-first quaternions (..., 4), crystal to lab. + + Returns + ------- + torch.Tensor + Unit crystal directions (..., 3) along lab +z, toward the source: + the third row of R(q). + """ + return quat_to_matrix(q)[..., 2, :] + + +def misorientation_angle_deg( + qa: torch.Tensor, + qb: torch.Tensor, + sym_ops: torch.Tensor | None = None, +) -> torch.Tensor: + """Misorientation angle between orientations, minimized over symmetry. + + Parameters + ---------- + qa, qb : torch.Tensor + Quaternions (..., 4), broadcastable against each other. + sym_ops : torch.Tensor | None + Proper rotation symmetry quaternions (S, 4) of the crystal. If None, + the raw rotation angle between qa and qb is returned. + + Returns + ------- + torch.Tensor + Misorientation angles in degrees (...,). + """ + dq = qmult(qconj(qa), qb) + if sym_ops is None: + w = dq[..., 0].abs().clamp(-1, 1) + else: + dq_sym = qmult(dq[..., None, :], sym_ops) # (..., S, 4) + w = dq_sym[..., 0].abs().amax(dim=-1).clamp(-1, 1) + return torch.rad2deg(2 * torch.acos(w)) + + +def misorientation_axis_angle( + qa: torch.Tensor, + qb: torch.Tensor, + sym_ops: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Symmetry-reduced misorientation axis and angle. + + The misorientation dq = conj(qa) * qb takes the crystal frame of `qa` + onto that of `qb`, R(qb) = R(qa) @ R(dq). Among the equivalent + dq * s over the symmetry operators s, the one with the smallest angle is + returned. + + Parameters + ---------- + qa, qb : torch.Tensor + Unit scalar-first quaternions (..., 4), crystal to lab, + broadcastable against each other. + sym_ops : torch.Tensor | None + Proper rotation symmetry quaternions (S, 4) of the crystal. If None, + the raw misorientation is returned. + + Returns + ------- + axis : torch.Tensor + Unit rotation axes (..., 3) in the crystal Cartesian frame of `qa`. + Undefined (near zero) when the angle is zero. + angle : torch.Tensor + Misorientation angles (...,) in degrees, in [0, 180]; equal to + :func:`misorientation_angle_deg` for the same inputs. + """ + dq = qmult(qconj(qa), qb) + if sym_ops is not None: + dq_sym = qmult(dq[..., None, :], sym_ops) # (..., S, 4) + best = dq_sym[..., 0].abs().argmax(dim=-1) + dq = torch.gather(dq_sym, -2, best[..., None, None].expand(*best.shape, 1, 4)).squeeze(-2) + dq = qnormalize(dq) + axis, angle = quat_to_axis_angle(dq) + return axis, torch.rad2deg(angle) + + +def slerp(v0: torch.Tensor, v1: torch.Tensor, t: torch.Tensor) -> torch.Tensor: + """Spherical linear interpolation between directions. + + Parameters + ---------- + v0, v1 : torch.Tensor + End directions (..., 3); normalized internally. + t : torch.Tensor + Interpolation fractions (...,), 0 at `v0` and 1 at `v1`. + + Returns + ------- + torch.Tensor + Unit vectors (..., 3) along the great circle from `v0` to `v1`. + """ + v0 = v0 / torch.linalg.norm(v0, dim=-1, keepdim=True) + v1 = v1 / torch.linalg.norm(v1, dim=-1, keepdim=True) + omega = torch.acos((v0 * v1).sum(-1, keepdim=True).clamp(-1, 1)) + so = torch.sin(omega) + t = t[..., None] + small = so.abs() < 1e-12 + w0 = torch.where(small, 1 - t, torch.sin((1 - t) * omega) / so) + w1 = torch.where(small, t, torch.sin(t * omega) / so) + return w0 * v0 + w1 * v1 + + +def sample_zone_axes( + corners: torch.Tensor, + step_deg: float, +) -> tuple[torch.Tensor, torch.Tensor]: + """Isotropic SLERP grid of unit zone-axis vectors inside a spherical triangle. + + Rows run from corner 0 toward the opposite edge; the number of points + in each row follows the row's own arc length, so the spacing is close + to `step_deg` in every direction whatever the apex angle of the wedge + (a 30 degree hexagonal wedge and a 120 degree trigonal wedge get the + same density). + + Parameters + ---------- + corners : torch.Tensor + (3, 3) rows are the Cartesian corner directions of the fundamental + zone-axis wedge, e.g. [001], [011], [111] for m-3m. + step_deg : float + Angular step between neighboring zone axes, degrees. + + Returns + ------- + vectors : torch.Tensor + (N, 3) unit vectors sampling the wedge. + inds : torch.Tensor + (N, 2) integer (row, col) indices in the triangular grid. + """ + c = corners / torch.linalg.norm(corners, dim=-1, keepdim=True) + a01 = torch.rad2deg(torch.acos((c[0] * c[1]).sum().clamp(-1, 1))) + a02 = torch.rad2deg(torch.acos((c[0] * c[2]).sum().clamp(-1, 1))) + n_steps = int(torch.ceil(torch.maximum(a01, a02) / step_deg).item()) + n_steps = max(n_steps, 1) + + vecs, inds = [], [] + for i in range(n_steps + 1): + t = torch.tensor(i / n_steps, dtype=c.dtype, device=c.device) + pv = slerp(c[0], c[1], t) + pw = slerp(c[0], c[2], t) + if i == 0: + vecs.append(pv[None]) + inds.append(torch.tensor([[0, 0]])) + continue + arc = torch.rad2deg(torch.acos((pv * pw).sum().clamp(-1, 1))) + n_i = max(1, int(torch.ceil(arc / step_deg).item())) + s = torch.linspace(0, 1, n_i + 1, dtype=c.dtype, device=c.device) + row = slerp(pv.expand(n_i + 1, 3), pw.expand(n_i + 1, 3), s) + row = row / torch.linalg.norm(row, dim=-1, keepdim=True) + vecs.append(row) + inds.append(torch.stack((torch.full((n_i + 1,), i), torch.arange(n_i + 1)), dim=-1)) + return torch.cat(vecs), torch.cat(inds).to(torch.long) + + +def symmetry_axes(sym_quats: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Distinct rotation axes of a proper point group and their orders. + + Parameters + ---------- + sym_quats : torch.Tensor + Proper rotation quaternions (S, 4) of the group, crystal Cartesian + frame. + + Returns + ------- + axes : torch.Tensor + Unit axes (A, 3), each listed once with the sign that makes it + point into the upper hemisphere (+z, then +y, then +x on ties). + orders : torch.Tensor + Long (A,), the highest rotation order about each axis (a 4-fold + axis is listed as order 4, not also as 2). + """ + axis, angle = quat_to_axis_angle(sym_quats) + keep = angle > 1e-6 + axis, angle = axis[keep], angle[keep] + order = torch.round(2 * np.pi / angle).to(torch.long) + # orient each axis to a canonical hemisphere so +-axis merge + sign = torch.sign(axis[:, 2] + 1e-3 * axis[:, 1] + 1e-6 * axis[:, 0]) + sign[sign == 0] = 1 + axis = axis * sign[:, None] + out_axes, out_orders = [], [] + for a, n in zip(axis, order): + for k, b in enumerate(out_axes): + if float(torch.abs(a @ b)) > 1 - 1e-6: + out_orders[k] = max(out_orders[k], int(n)) + break + else: + out_axes.append(a) + out_orders.append(int(n)) + if not out_axes: + return torch.zeros((0, 3), dtype=torch.float64), torch.zeros(0, dtype=torch.long) + return torch.stack(out_axes), torch.tensor(out_orders, dtype=torch.long) + + +def _rotate_about(v: torch.Tensor, axis: torch.Tensor, angle_rad: float) -> torch.Tensor: + q = quat_from_axis_angle(axis, torch.tensor(angle_rad, dtype=v.dtype)) + return qrotate(q, v) + + +def _closest(cands: torch.Tensor, prefer: torch.Tensor) -> torch.Tensor: + """The candidate (sign chosen freely) closest to the preferred direction, + with deterministic tie-breaking toward +z, then +y, then +x.""" + signed = torch.cat([cands, -cands]) + key = signed @ prefer + 1e-3 * signed[:, 2] + 1e-6 * signed[:, 1] + 1e-9 * signed[:, 0] + return signed[int(torch.argmax(key))] + + +def symmetry_aligned( + reference: torch.Tensor, + quats: torch.Tensor, + sym_quats: torch.Tensor, +) -> torch.Tensor: + """Symmetry images of `quats` that lie nearest to `reference`, (N, 4). + + Two quaternions can describe the same crystal orientation while being far + apart as quaternions, so any average over orientations has to bring them + into a common symmetry branch first. For each input this returns the + symmetry-equivalent quaternion whose misorientation to the reference is + smallest, which makes a weighted quaternion mean well defined. + + Parameters + ---------- + reference : torch.Tensor + Unit scalar-first quaternion (4,), crystal to lab. + quats : torch.Tensor + Unit quaternions, reshaped to (N, 4). + sym_quats : torch.Tensor + Proper rotation symmetry quaternions (S, 4) of the crystal, crystal + Cartesian frame. Each candidate is q * s, the same lab orientation + of the symmetric crystal. + + Returns + ------- + torch.Tensor + (N, 4) float64 quaternions with w >= 0, each describing the same + orientation as the corresponding input. + """ + ref = torch.as_tensor(reference, dtype=torch.float64).reshape(4) + q = torch.as_tensor(quats, dtype=torch.float64).reshape(-1, 4) + sym = torch.as_tensor(sym_quats, dtype=torch.float64).reshape(-1, 4) + cand = qmult(q[:, None, :], sym[None, :, :]) # (N, S, 4) + dots = torch.abs(torch.einsum("nsi,i->ns", cand, ref)) + best = dots.argmax(dim=1) + return qnormalize(cand[torch.arange(q.shape[0]), best]) + + +def sample_zone_axis_cap( + axis: torch.Tensor, + half_angle_deg: float, + step_deg: float, +) -> torch.Tensor: + """Near-uniform sampling of a spherical cap of directions, (N, 3). + + The fiber-texture case: the zone axis is known to lie within + `half_angle_deg` of `axis` (a fiber axis normal to a 2D material, or a + textured film), and only that cap needs a library. A half angle of zero + returns the axis itself, so the match is over the in-plane angle alone. + + Points are placed on a Fibonacci spiral restricted to the cap, which + gives an equal-area covering; the count follows the cap area divided by + `step_deg` squared. + + Parameters + ---------- + axis : torch.Tensor + Center of the cap (3,), crystal Cartesian; need not be normalized. + half_angle_deg : float + Angular radius of the cap, degrees. + step_deg : float + Approximate spacing between neighboring directions, degrees. + + Returns + ------- + torch.Tensor + Unit directions (N, 3), float64, all within `half_angle_deg` of + `axis`; (1, 3) holding `axis` itself when `half_angle_deg` is zero. + """ + axis = torch.as_tensor(axis, dtype=torch.float64) + axis = axis / torch.linalg.norm(axis).clamp_min(1e-12) + half = np.deg2rad(float(half_angle_deg)) + if half <= 0: + return axis[None, :] + step = np.deg2rad(float(step_deg)) + n = max(1, int(np.ceil(2 * np.pi * (1 - np.cos(half)) / step**2))) + i = torch.arange(n, dtype=torch.float64) + 0.5 + z = 1.0 - (1.0 - np.cos(half)) * i / n + r = torch.sqrt((1 - z**2).clamp_min(0)) + phi = i * (np.pi * (3 - np.sqrt(5))) + pts = torch.stack([r * torch.cos(phi), r * torch.sin(phi), z], dim=1) + # rotate the +z pole onto `axis` + zhat = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + v = torch.linalg.cross(zhat, axis) + c = float(torch.dot(zhat, axis)) + if float(torch.linalg.norm(v)) < 1e-12: + return pts if c > 0 else -pts + vx = torch.tensor( + [[0.0, -v[2], v[1]], [v[2], 0.0, -v[0]], [-v[1], v[0], 0.0]], dtype=torch.float64 + ) + R = torch.eye(3, dtype=torch.float64) + vx + vx @ vx / (1 + c) + return pts @ R.T + + +def fundamental_zone_axis_wedge(sym_quats: torch.Tensor) -> torch.Tensor | None: + """Fundamental zone-axis wedge of a Laue group from its proper rotations. + + Zone axes are directions modulo inversion (Friedel), so the wedge is a + fundamental domain of the Laue group on the projective hemisphere, + built from the actual symmetry axes in the crystal's Cartesian frame + rather than from a table keyed on the Laue class. This makes it correct + for every setting: for -3m the wedge is bounded by the mirror planes + (perpendicular to the in-plane 2-fold axes), which is 30 degrees away + from a wedge bounded by the 2-fold axes themselves; for a cell in a + non-standard Cartesian setting the corners follow the axes wherever + they point. + + Parameters + ---------- + sym_quats : torch.Tensor + Proper rotation quaternions (S, 4) of the group, crystal Cartesian + frame. + + Returns + ------- + torch.Tensor | None + Corner directions as rows (3, 3), or None for Laue classes -1 and + 2/m whose fundamental domain is not a spherical triangle (sample the + hemisphere instead). + """ + axes, orders = symmetry_axes(sym_quats) + n_ops = sym_quats.shape[0] + x, y, z = (torch.eye(3, dtype=torch.float64)[i] for i in range(3)) + if n_ops <= 2: + return None + + three = axes[orders == 3] + if three.shape[0] >= 4: # cubic + four = axes[orders == 4] + two = axes[orders == 2] + if four.shape[0] > 0: # m-3m: 4-fold, <110> 2-fold, 3-fold + c0 = _closest(four, z) + c2 = _closest(three, c0) + c1 = _closest(two, c0 + c2) + else: # m-3: two cubic 2-fold axes and the 3-fold between them + c0 = _closest(two, z) + rest = two[torch.abs(two @ c0) < 0.5] + c1 = _closest(rest, x) + c2 = _closest(three, c0 + c1) + return torch.stack([c0, c1, c2]) + + n_max = int(orders.max()) + if n_max in (3, 4, 6): # uniaxial classes + c0 = _closest(axes[orders == n_max], z) + in_plane = axes[(orders == 2) & (torch.abs(axes @ c0) < 1e-6)] + # a direction perpendicular to the axis, nearest +x + ref = x - (x @ c0) * c0 + if torch.linalg.norm(ref) < 1e-6: + ref = y - (y @ c0) * c0 + ref = ref / torch.linalg.norm(ref) + if in_plane.shape[0] == 0: # 6/m, 4/m, -3: any 360/n sector + c2 = ref + c1 = _rotate_about(c2, c0, 2 * np.pi / n_max) + elif n_max == 3: # -3m: the sector between adjacent mirror planes, + # which are perpendicular to the in-plane 2-fold axes + a = _closest(in_plane, ref) + c2 = _rotate_about(a, c0, np.pi / 2) + c1 = _rotate_about(a, c0, 5 * np.pi / 6) + else: # 6/mmm, 4/mmm: between adjacent in-plane 2-fold axes + c2 = _closest(in_plane, ref) + c1 = _rotate_about(c2, c0, np.pi / n_max) + return torch.stack([c0, c1, c2]) + + if n_max == 2 and axes.shape[0] == 3: # mmm: the three 2-fold axes + c0 = _closest(axes, z) + rest = axes[torch.abs(axes @ c0) < 0.5] + c1 = _closest(rest, x) + c2 = _closest(rest[torch.abs(rest @ c1) < 0.5], y) + return torch.stack([c0, c1, c2]) + return None + + +def symmetry_reduced_zone_angles( + zone_axes: torch.Tensor, sym_quats: torch.Tensor, chunk: int = 8 +) -> torch.Tensor: + """Angular distances between zone axes, minimized over symmetry. + + The minimum is over the symmetry operations and the inversion (zone axes + are directions modulo sign). Symmetry-equivalent zones are at distance + zero, so an exclusion ball around a match also excludes its symmetry + copies. + + Parameters + ---------- + zone_axes : torch.Tensor + Unit directions (Z, 3), crystal Cartesian. + sym_quats : torch.Tensor + Proper rotation quaternions (S, 4) of the crystal. + chunk : int, default=8 + Symmetry operations processed at once, which bounds the memory to + chunk * Z * Z. + + Returns + ------- + torch.Tensor + (Z, Z) angles in degrees. + """ + Rs = quat_to_matrix(sym_quats).to(zone_axes.dtype) + best = torch.full((zone_axes.shape[0],) * 2, -1.0, dtype=zone_axes.dtype) + for s0 in range(0, Rs.shape[0], chunk): + imgs = torch.einsum("sij,zj->szi", Rs[s0 : s0 + chunk], zone_axes) + dots = torch.einsum("szi,wi->szw", imgs, zone_axes).abs().amax(dim=0) + best = torch.maximum(best, dots) + return torch.rad2deg(torch.acos(best.clamp(-1, 1))) + + +def symmetry_quaternions( + rotations: np.ndarray, + lat_real: np.ndarray, +) -> torch.Tensor: + """Convert spglib integer rotation matrices to Cartesian quaternions. + + Parameters + ---------- + rotations : np.ndarray + (S, 3, 3) integer rotation matrices in the lattice basis, as returned + by spglib (improper operations are discarded). + lat_real : np.ndarray + (3, 3) real-space lattice vectors as rows. + + Returns + ------- + torch.Tensor + (S', 4) unique proper-rotation quaternions in Cartesian coordinates, + float64. + """ + A = torch.as_tensor(lat_real, dtype=torch.float64).T # columns are a, b, c + W = torch.as_tensor(np.array(rotations), dtype=torch.float64) + R_cart = A @ W @ torch.linalg.inv(A) + proper = torch.linalg.det(R_cart) > 0 + # for a pseudo-symmetry group the lattice is slightly distorted from the + # ideal one, so A W A^-1 is only approximately orthogonal: take the + # nearest rotation (polar decomposition) so the operators are exact + # rotations and the wedge and misorientation math stay consistent + U, _, Vh = torch.linalg.svd(R_cart[proper]) + q = quat_from_matrix(U @ Vh) + # deduplicate (q and -q are the same rotation; qnormalize fixed the sign) + q_unique = torch.unique(torch.round(q / 1e-6) * 1e-6, dim=0) + return qnormalize(q_unique) diff --git a/src/quantem/diffraction/strain.py b/src/quantem/diffraction/strain.py new file mode 100644 index 000000000..cde29b50b --- /dev/null +++ b/src/quantem/diffraction/strain.py @@ -0,0 +1,1076 @@ +from __future__ import annotations + +import warnings + +import numpy as np +from numpy.lib.stride_tricks import sliding_window_view +from numpy.typing import NDArray + +from quantem.core.datastructures.dataset2d import Dataset2d +from quantem.core.io.serialize import AutoSerialize +from quantem.diffraction.strain_visualization import ( + plot_strain_panels, + plot_strain_precision_histogram, +) + + +class StrainMap(AutoSerialize): + """Strain tensor maps fit from per-position lattice vectors. + + Stores the reference-frame strain components ``e_rr`` (row), ``e_cc`` (col), + ``e_rc`` (shear), and ``phi`` (infinitesimal rotation, radians). The reference + lattice is the weighted median (or mean) of the fitted ``g1``/``g2`` over a + mask/ROI; the strain tensor is recomputed by :meth:`update_reference`. + + The lattice vectors are first mapped from the detector frame into the scan + frame (optional row/col transpose, then a counter-clockwise rotation by + ``q_to_r_rotation_ccw_deg``), so ``e_rr``/``e_cc`` refer to the scan rows and + columns. :attr:`g1_array`, :attr:`g2_array`, :attr:`g1_ref` and :attr:`g2_ref` + are stored in that rotated frame. + + Two measurement modalities are supported and give identical strain for the same + deformation, so correlation and cepstral maps can be compared directly: + reciprocal-space Bragg vectors (``real_space=False``, nanobeam correlation) and + real-space cepstral/autocorrelation vectors (``real_space=True``). + + Parameters + ---------- + g1_array : np.ndarray + Per-position first lattice vector, shape ``(scan_row, scan_col, 2)``. + g2_array : np.ndarray + Per-position second lattice vector, shape ``(scan_row, scan_col, 2)``. + ds_shape : tuple of int + Shape of the parent scan grid, used to size the strain maps. + real_space : bool + ``False`` for reciprocal-space (Bragg/correlation) lattice vectors; ``True`` + for real-space (cepstral autocorrelation / DPC) vectors. Both modalities are + arranged to yield matching strain (see :func:`_strain_tensor`). + g1_ref : np.ndarray, optional + Fixed reference for ``g1``, in the same (unrotated) frame as ``g1_array``; + if omitted the ``calculation_metric`` over the mask/ROI is used. A value + supplied here persists across re-fits. + g2_ref : np.ndarray, optional + Fixed reference for ``g2``, as ``g1_ref``. + mask : np.ndarray, optional + ``(scan_row, scan_col)`` weighting/ROI mask; defaults to all ones (the full + scan). Normalized to ``[0, 1]`` on assignment. + ds_sampling : float, optional + Real-space scan sampling (step size); defaults to ``1.0``. + ds_units : str, optional + Units for ``ds_sampling``; defaults to ``"pixels"``. + q_to_r_rotation_ccw_deg : float, default=0.0 + Counter-clockwise rotation in degrees from the detector (``q``) frame to + the scan (``r``) frame, applied to the lattice vectors and to ``g1_ref`` / + ``g2_ref``. + q_transpose : bool, default=False + If ``True``, swap the detector row/col axes before the rotation. + calculation_metric : {"median", "mean"}, default="median" + Statistic used for the automatic reference lattice. Stored and reused by + :meth:`update_reference` unless overridden there. + """ + + mask: np.ndarray | None = None + real_space: bool = False + + e_rr: Dataset2d + e_cc: Dataset2d + e_rc: Dataset2d + phi: Dataset2d + + g1_ref: np.ndarray | None = None + g2_ref: np.ndarray | None = None + g1_array: np.ndarray + g2_array: np.ndarray + + ds_sampling: float = 1.0 + ds_units: str = "pixels" + ds_shape: tuple[int, ...] + + calculation_metric: str = "median" + + def __init__( + self, + g1_array: np.ndarray, + g2_array: np.ndarray, + ds_shape: tuple[int, ...], + real_space: bool, + g1_ref: np.ndarray | None = None, + g2_ref: np.ndarray | None = None, + mask: np.ndarray | None = None, + ds_sampling: float | None = None, + ds_units: str | None = None, + q_to_r_rotation_ccw_deg: float = 0.0, + q_transpose: bool = False, + calculation_metric: str = "median", + ): + super().__init__() + self.g1_array = g1_array + self.g2_array = g2_array + + self.q_to_r_rotation_ccw_deg = q_to_r_rotation_ccw_deg + self.q_transpose = q_transpose + + self.g1_array = _raw_vec_to_display( + self.g1_array, rotation_ccw_deg=q_to_r_rotation_ccw_deg, transpose=q_transpose + ) + self.g2_array = _raw_vec_to_display( + self.g2_array, rotation_ccw_deg=q_to_r_rotation_ccw_deg, transpose=q_transpose + ) + + self.ds_shape = ds_shape + self.real_space = real_space + + self.ds_sampling = 1.0 if ds_sampling is None else ds_sampling + self.ds_units = "pixels" if ds_units is None else ds_units + + m = np.ones(ds_shape[:2], dtype=float) if mask is None else np.asarray(mask, dtype=float) + m_lo = np.nanmin(m) + m_hi = np.nanmax(m) + if not (np.isfinite(m_lo) and np.isfinite(m_hi)) or m_hi <= m_lo: + m = np.ones_like(m) + elif m_lo < 0.0 or m_hi > 1.0: + m = (m - m_lo) / (m_hi - m_lo) + self.mask = m + + # user-supplied reference vectors persist across re-fits (None = automatic) + self.g1_ref_fixed = ( + None + if g1_ref is None + else _raw_vec_to_display( + np.asarray(g1_ref, dtype=float), + rotation_ccw_deg=q_to_r_rotation_ccw_deg, + transpose=q_transpose, + ) + ) + self.g2_ref_fixed = ( + None + if g2_ref is None + else _raw_vec_to_display( + np.asarray(g2_ref, dtype=float), + rotation_ccw_deg=q_to_r_rotation_ccw_deg, + transpose=q_transpose, + ) + ) + self.g1_ref = None + self.g2_ref = None + self.calculation_metric = calculation_metric + self.update_reference(calculation_metric=calculation_metric) + + # ---- main methods ---- + + def update_reference( + self, + strain_mask: np.ndarray | None = None, + g1_ref: np.ndarray | None = None, + g2_ref: np.ndarray | None = None, + plot_strain_roi: bool = False, + define_in_rotated_frame: bool = False, + calculation_metric: str | None = None, + **plot_kwargs, + ) -> "StrainMap": + """(Re)compute the reference lattice and strain tensor maps. + + Reference precedence: explicit ``g1_ref``/``g2_ref`` argument > vectors fixed at + construction > ``calculation_metric`` (weighted median or mean) over + ``strain_mask`` (if given) else weighted by ``self.mask``. + + Parameters + ---------- + strain_mask : np.ndarray, optional + ``(scan_row, scan_col)`` ROI or weights selecting the positions used to + compute the automatic reference lattice. If omitted, ``self.mask`` is + used. + g1_ref : np.ndarray, optional + Explicit reference for ``g1``; overrides both the construction-time fixed + value and the median. + g2_ref : np.ndarray, optional + Explicit reference for ``g2``; overrides both the construction-time fixed + value and the median. + plot_strain_roi : bool, default=False + If ``True``, show the recomputed strain via :meth:`plot_strain_roi` + (color-scaled to the ROI) so the chosen reference region can be checked + for flatness. + define_in_rotated_frame : bool, default=False + If ``True``, the ``g1_ref`` and ``g2_ref`` passed here are already in + the rotated (scan) frame; otherwise they are in the detector frame and + are rotated like the lattice vectors. + calculation_metric : {"median", "mean"}, optional + Statistic for the automatic reference lattice. If given, it is stored on the + object and used by later calls; if ``None`` (default), the stored value is + used. + **plot_kwargs + Forwarded to :meth:`plot_strain_roi` when ``plot_strain_roi=True``. + + Returns + ------- + StrainMap + ``self``, with the reference lattice and strain maps recomputed. + """ + if calculation_metric is not None: + self.calculation_metric = calculation_metric + g1_med, g2_med = _reference_lattice( + self.g1_array, + self.g2_array, + self.mask, + strain_mask, + calculation_metric=self.calculation_metric, + ) + + if g1_ref is not None: + if define_in_rotated_frame: + self.g1_ref = np.asarray(g1_ref, dtype=float) + else: + self.g1_ref = _raw_vec_to_display( + np.asarray(g1_ref, dtype=float), + rotation_ccw_deg=self.q_to_r_rotation_ccw_deg, + transpose=self.q_transpose, + ) + elif self.g1_ref_fixed is not None: + self.g1_ref = self.g1_ref_fixed + else: + self.g1_ref = g1_med + + if g2_ref is not None: + if define_in_rotated_frame: + self.g2_ref = np.asarray(g2_ref, dtype=float) + else: + self.g2_ref = _raw_vec_to_display( + np.asarray(g2_ref, dtype=float), + rotation_ccw_deg=self.q_to_r_rotation_ccw_deg, + transpose=self.q_transpose, + ) + elif self.g2_ref_fixed is not None: + self.g2_ref = self.g2_ref_fixed + else: + self.g2_ref = g2_med + + e_rr, e_cc, e_rc, phi = _strain_tensor( + self.g1_array, self.g2_array, self.g1_ref, self.g2_ref, self.real_space + ) + self.e_rr = Dataset2d.from_array(e_rr, name="strain e_rr", signal_units="fractional") + self.e_cc = Dataset2d.from_array(e_cc, name="strain e_cc", signal_units="fractional") + self.e_rc = Dataset2d.from_array(e_rc, name="strain e_rc", signal_units="fractional") + self.phi = Dataset2d.from_array(phi, name="strain rotation", signal_units="radians") + + if plot_strain_roi: + self.plot_strain_roi(strain_mask=strain_mask, **plot_kwargs) + return self + + def rotate_strain( + self, rotation_angle: float = 0.0 + ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Tensor-rotate the strain into a frame rotated by ``rotation_angle`` (degrees). + + The rotation field ``phi`` is invariant under frame rotation and is not + transformed. + + Parameters + ---------- + rotation_angle : float, default=0.0 + Frame rotation angle, in degrees. + + Returns + ------- + tuple of np.ndarray + ``(e_uu, e_vv, e_uv)`` strain components in the rotated frame. + """ + return _rotate_strain_tensor( + self.e_rr.array, self.e_cc.array, self.e_rc.array, rotation_angle + ) + + def plot_strain_roi( + self, + strain_mask: np.ndarray | None = None, + plot_rotation: bool = True, + cmap_strain: str = "RdBu_r", + cmap_rotation: str = "PiYG", + strain_range_percent: tuple[float, float] | None = None, + rotation_range_degrees: tuple[float, float] | None = None, + transpose_image: bool = False, + rotate_title: bool = False, + plot_dilation: bool = False, + layout: str = "horizontal", + arrow_style: str = "title", + figsize: tuple[float, float] | None = None, + **kwargs, + ): + """Plot the strain in the raw row/col reference frame, color-scaled to the ROI. + + The color range is symmetric about zero and set by the largest absolute + strain (and rotation) *inside the reference ROI* — ``strain_mask`` if given, + else ``self.mask`` — so a well-chosen, strain-free reference region reads as + flat (near mid-color) and any residual gradient or tilt stands out. The ROI + itself is drawn in color while everything outside it is shown in greyscale, + so the chosen reference region is obvious at a glance. Unlike + :meth:`plot_strain`, no display rotation is applied: the panels show the raw + ``e_rr``/``e_cc``/``e_rc`` that :meth:`update_reference` just computed. + + Parameters + ---------- + strain_mask : np.ndarray, optional + ROI defining the color range (and the reference region). If omitted, + ``self.mask`` is used. + plot_rotation : bool, default=True + Whether to include the rotation (``phi``) panel. + cmap_strain : str, default="RdBu_r" + Colormap for the strain panels. + cmap_rotation : str, default="PiYG" + Colormap for the rotation panel. + strain_range_percent : tuple of float, optional + Color range for the strain panels, in percent. ``None`` (default) uses + a symmetric range set by the largest absolute strain inside the ROI. + rotation_range_degrees : tuple of float, optional + Color range for the rotation panel, in degrees. ``None`` (default) uses + a symmetric range set by the largest absolute rotation inside the ROI. + transpose_image : bool, default=False + If ``True``, transpose the real-space images before plotting. + rotate_title : bool, default=False + If ``True``, rotate the panel titles by 90 degrees. + plot_dilation : bool, default=False + If ``True``, plot ``e_rr + e_cc`` and ``e_rc`` instead of ``e_rr``, + ``e_cc`` and ``e_rc``. + layout : {"horizontal", "vertical"}, default="horizontal" + Panel arrangement. + arrow_style : {"title", "legend"}, default="title" + Accepted for symmetry with :meth:`plot_strain`; the row/col titles + carry their own arrows, so no extra arrows are drawn here. + figsize : tuple of float, optional + Figure size in inches; if omitted it is derived from the layout. + **kwargs + Forwarded to + :func:`~quantem.diffraction.strain_visualization.plot_strain_panels`. + + Returns + ------- + tuple + ``(fig, ax)`` from :func:`plot_strain_panels`. + """ + + if arrow_style not in ("title", "legend"): + raise ValueError("arrow_style must be 'title' or 'legend'") + + if plot_dilation: + panel_titles = ( + r"$\epsilon_{rr} + \epsilon_{cc}$", + r"$\epsilon_{rc}$ $\nwarrow\!\!\!\!\!\!\!\!\!\:\searrow$", + ) + else: + panel_titles = ( + r"$\epsilon_{rr}$ $\updownarrow$", + r"$\epsilon_{cc}$ $\leftrightarrow$", + r"$\epsilon_{rc}$ $\nwarrow\!\!\!\!\!\!\!\!\!\:\searrow$", + ) + + roi_src = self.mask if strain_mask is None else strain_mask + e_rr, e_cc, e_rc, phi = ( + self.e_rr.array, + self.e_cc.array, + self.e_rc.array, + self.phi.array, + ) + + inside = np.asarray(roi_src) > 0 if roi_src is not None else np.ones(e_rr.shape, bool) + if not inside.any(): + inside = np.ones(e_rr.shape, bool) + + strain_stack = np.stack([e_rr[inside], e_cc[inside], e_rc[inside]]) + smax = float(np.nanmax(np.abs(strain_stack))) * 100.0 + rmax = float(np.rad2deg(np.nanmax(np.abs(phi[inside])))) + smax = smax if smax > 0 else 1e-6 + rmax = rmax if rmax > 0 else 1e-6 + + return plot_strain_panels( + e_rr, + e_cc, + e_rc, + phi, + self.mask, + self.g1_ref, + self.g2_ref, + self.ds_shape, + ds_sampling=self.ds_sampling, + ds_units=self.ds_units, + strain_range_percent=(-smax, smax) + if strain_range_percent is None + else strain_range_percent, + rotation_range_degrees=(-rmax, rmax) + if rotation_range_degrees is None + else rotation_range_degrees, + roi=inside, + plot_rotation=plot_rotation, + cmap_strain=cmap_strain, + cmap_rotation=cmap_rotation, + layout=layout, + transpose_image=transpose_image, + rotate_title=rotate_title, + plot_dilation=plot_dilation, + figsize=figsize, + panel_titles=panel_titles, + arrow_style=arrow_style, + **kwargs, + ) + + def plot_strain( + self, + rotation_angle: float = 0.0, + strain_range_percent: tuple[float, float] = (-3.0, 3.0), + rotation_range_degrees: tuple[float, float] = (-2.0, 2.0), + mask_range: tuple[float, float] = (0.0, 1.0), + plot_rotation: bool = True, + plot_gvecs: bool = False, + plot_scalebar: bool = False, + cmap_strain: str = "RdBu_r", + cmap_rotation: str = "PiYG", + transpose_image: bool = False, + transpose_strain: bool = False, + rotate_title: bool = False, + plot_dilation: bool = False, + layout: str = "horizontal", + arrow_style: str = "title", + figsize: tuple[float, float] | None = None, + **kwargs, + ): + """Plot the strain (rotated into the display frame) and rotation panels. + + Parameters + ---------- + rotation_angle : float, default=0.0 + Angle (degrees) by which the strain tensor is rotated into the display + frame before plotting. + strain_range_percent : tuple of float, default=(-3.0, 3.0) + Symmetric color range for the strain panels, in percent. + rotation_range_degrees : tuple of float, default=(-2.0, 2.0) + Symmetric color range for the rotation panel, in degrees. + mask_range : tuple of float, default=(0.0, 1.0) + ``(low, high)`` window remapping the mask brightness: positions with + mask ``>= high`` are shown at full color, ``<= low`` are black, and + values between ramp linearly from black to full. The default leaves the + normalized mask unchanged. + plot_rotation : bool, default=True + Whether to include the rotation (``phi``) panel. + plot_gvecs : bool, default=False + Whether to overlay the reference lattice vectors. + plot_scalebar : bool, default=False + Whether to draw a real-space scale bar. + cmap_strain : str, default="RdBu_r" + Colormap for the strain panels. + cmap_rotation : str, default="PiYG" + Colormap for the rotation panel. + transpose_image : bool, default=False + If ``True``, transpose the real-space images before plotting. + transpose_strain : bool, default=False + If ``True``, transpose the (row/col) axes of the strain tensor before + rotating, matching the DPC convention of :func:`_raw_vec_to_display` + (transpose first, then rotate). This swaps the normal strain + components, leaves the shear unchanged, and reverses the sign of the + rotation field. + rotate_title : bool, default=False + If ``True``, rotate the panel titles by 90 degrees. + plot_dilation : bool, default=False + If ``True``, plot ``e_uu + e_vv`` and ``e_uv`` instead of ``e_uu``, + ``e_vv`` and ``e_uv``. + layout : {"horizontal", "vertical"}, default="horizontal" + Panel arrangement. + arrow_style : {"title", "legend"}, default="title" + Draw the strain direction arrows next to the panel titles, or in a + legend at the side of the figure. + figsize : tuple of float, optional + Figure size in inches; if omitted it is derived from the layout. + **kwargs + Forwarded to + :func:`~quantem.diffraction.strain_visualization.plot_strain_panels`. + + Returns + ------- + tuple + ``(fig, ax)`` from :func:`plot_strain_panels`. + """ + if arrow_style not in ("title", "legend"): + raise ValueError("arrow_style must be 'title' or 'legend'") + + e_rr = self.e_rr.array + e_cc = self.e_cc.array + e_rc = self.e_rc.array + phi = self.phi.array + if transpose_strain: + # Detector-axis transpose, applied BEFORE the rotation to match the DPC + # convention shared across quantem (see _raw_vec_to_display): swapping the + # (row, col) axes swaps the normal strains, keeps the shear unchanged, and + # reverses the sense of the rotation field. + e_rr, e_cc = e_cc, e_rr + phi = -phi + e_uu, e_vv, e_uv = _rotate_strain_tensor(e_rr, e_cc, e_rc, rotation_angle) + return plot_strain_panels( + e_uu, + e_vv, + e_uv, + phi, + self.mask, + self.g1_ref, + self.g2_ref, + self.ds_shape, + ds_sampling=self.ds_sampling, + ds_units=self.ds_units, + strain_range_percent=strain_range_percent, + rotation_range_degrees=rotation_range_degrees, + mask_range=mask_range, + plot_rotation=plot_rotation, + plot_gvecs=plot_gvecs, + plot_scalebar=plot_scalebar, + cmap_strain=cmap_strain, + cmap_rotation=cmap_rotation, + layout=layout, + transpose_image=transpose_image, + rotate_title=rotate_title, + plot_dilation=plot_dilation, + figsize=figsize, + strain_rotation_angle=rotation_angle, + arrow_style=arrow_style, + **kwargs, + ) + + def estimate_strain_precision( + self, + mask_range: tuple[float, float] = (0.0, 1.0), + rotation_angle: float = 0.0, + window: int = 5, + mask_threshold: float = 0.5, + min_neighbors: int = 3, + require_full_neighborhood: bool = True, + component: str = "combined", + bins: int = 50, + bounds: tuple[float, float] | None = None, + plot: bool = True, + returnfig: bool = False, + verbose: bool = False, + ): + """Estimate strain *precision* (random scatter) from local median deviations. + + This measures repeatability, not accuracy. Without a ground truth (e.g. a + simulation) it cannot detect systematic error — only how far each position + scatters from its local neighborhood. For every position the deviation from + the median of its surrounding well-indexed neighbors is + + ``error(r, c) = | strain(r, c) - median( strain over neighbors with + scaled mask > mask_threshold ) |`` + + computed for each tensor component (the center position is excluded from its + own median). The three strain components are reduced to one rotation-invariant + number via the Frobenius norm of the symmetric strain-tensor deviation, + + ``combined = sqrt(d_uu**2 + d_vv**2 + 2*d_uv**2)``, + + (equivalently the root-sum-square of the principal-strain deviations) so a + single strain precision can be quoted and compared between datasets. Rotation + precision is reported separately, not folded into ``combined``. + + Each component's precision is summarized by the mask-weighted **median** of its + per-position deviations — the center of the histogram bulk. The median is used + (not the mean or RMS) because a handful of bad-fit pixels form a heavy tail that + would drag a second moment far to the right of where the distribution actually + sits, leaving the reported number disconnected from the histogram; the median + ignores that tail. A weighted histogram of the chosen component is shown, marked + with its median. + + Parameters + ---------- + mask_range : tuple of float, default=(0.0, 1.0) + ``(low, high)`` window remapping :attr:`mask` to ``[0, 1]`` (same + convention as :meth:`plot_strain`); the remapped mask both selects which + positions are trusted (``> mask_threshold`` -- used as neighbors *and* as + the set the precision is computed over) and weights the histogram and the + median. + rotation_angle : float, default=0.0 + Frame rotation (degrees) applied before measuring per-component precision, + matching :meth:`plot_strain`. ``0`` reports the raw row/col frame + (``e_uu == e_rr`` ...). The combined number is rotation-invariant. + window : int, default=5 + Odd edge length (px) of the neighborhood bounding box; the footprint is + the inscribed disk of radius ``window / 2`` (3 -> 8 neighbors, 5 -> 20, + 7 -> 36). A pure linear strain ramp cancels in the (symmetric) median, so + larger windows mostly just steady the median — at the cost of blurring + *curved* strain and biasing the masked edges. ``5`` roughly halves the + noise-floor over-estimate of ``3`` (~9% -> ~4%) while staying local. + mask_threshold : float, default=0.5 + A position is trusted only if its scaled mask exceeds this value. Trusted + positions are the ones used as local-median neighbors *and* the ones whose + deviations enter the reported median and histogram; sub-threshold positions + are excluded from both (not merely down-weighted), so a poorly-indexed + pixel cannot leak its scatter into the precision. + min_neighbors : int, default=3 + Minimum number of valid neighbors required; positions with fewer get no + precision estimate (``nan``, dropped from the statistics). Ignored when + ``require_full_neighborhood=True``. + require_full_neighborhood : bool, default=True + If ``True``, require every neighbor in the footprint to be valid (sets + ``min_neighbors`` to the footprint size), so positions next to masked + or missing data get no estimate. + component : {"combined","e_uu","e_vv","e_uv","rotation"}, default="combined" + Which error distribution to histogram. + bins : int, default=50 + Number of histogram bins, or a sequence of explicit bin edges. With a + bin *count* and no ``bounds``, the range defaults to ``[0, weighted 99th + percentile]`` of the trusted deviations -- robust to the heavy outlier + tail, which otherwise sets the range to its max and crushes the bulk into + the first bin. Passing explicit edges (or ``bounds``) overrides this. + bounds : tuple of float, optional + ``(low, high)`` histogram range in display units (percent for strain, + degrees for rotation). Fix it to compare datasets on the same axis, or to + see the full tail. Values outside the range are left out of the bars (no + overflow spike); the median is computed from all trusted positions + regardless, and ``out_of_range_fraction`` records how much was off-range. + plot : bool, default=True + If ``True``, draw the weighted precision histogram. + returnfig : bool, default=False + If ``True``, return ``(fig, ax)`` instead of the results dict. + verbose : bool, default=False + If ``True``, print a summary of the precision per component. + + Returns + ------- + dict or tuple + A results dict with the ``precision`` (mask-weighted median local + deviation) per component and ``combined`` (strain in percent, rotation in + degrees), the normalized ``counts`` and ``edges`` of the histogrammed + ``component`` (and ``counts_raw``, the weighted bin sums), + ``out_of_range_fraction`` (weighted mass outside the histogram range, + excluded from the bars), and the chosen settings; or ``(fig, ax)`` when + ``returnfig=True``. + """ + if window < 3 or window % 2 == 0: + raise ValueError("window must be an odd integer >= 3.") + valid_components = ("combined", "e_uu", "e_vv", "e_uv", "rotation") + if component not in valid_components: + raise ValueError(f"component must be one of {valid_components}.") + + # number of neighbors in the circular footprint (matches _local_masked_median) + p = window // 2 + oy, ox = np.ogrid[-p : p + 1, -p : p + 1] + n_neighbors = int(np.sum((oy**2 + ox**2) <= (window / 2.0) ** 2) - 1) + if require_full_neighborhood: + min_neighbors = n_neighbors + + # per-component fields in the (optionally rotated) display frame; phi is + # rotation-invariant and is carried through unchanged + e_uu, e_vv, e_uv = self.rotate_strain(rotation_angle) + fields = {"e_uu": e_uu, "e_vv": e_vv, "e_uv": e_uv, "rotation": self.phi.array} + + # remap the mask exactly as plot_strain does, then use it both to select + # neighbors (> mask_threshold) and to weight the histogram / mean + low, high = float(mask_range[0]), float(mask_range[1]) + m = np.asarray(self.mask, dtype=float) + if high > low: + scaled = np.clip((m - low) / (high - low), 0.0, 1.0) + else: + scaled = (m >= high).astype(float) + valid = scaled > float(mask_threshold) + + # per-component local-median deviation, native units (fractional / radians) + dev = { + name: np.abs(field - _local_masked_median(field, valid, window, min_neighbors)) + for name, field in fields.items() + } + # single rotation-invariant number: Frobenius norm of the symmetric + # strain-tensor deviation (== root-sum-square of the principal-strain + # deviations). Rotation is reported separately, not folded in: in nanobeam + # data it is partly a systematic (tilt/descan) and would mix radians into a + # percent figure. + dev["combined"] = np.sqrt(dev["e_uu"] ** 2 + dev["e_vv"] ** 2 + 2.0 * dev["e_uv"] ** 2) + + # display-unit scaling: strain -> percent, rotation -> degrees + scale = { + "e_uu": 100.0, + "e_vv": 100.0, + "e_uv": 100.0, + "rotation": float(np.rad2deg(1.0)), + "combined": 100.0, + } + + # Precision = the weighted MEDIAN of each per-position deviation distribution, + # in display units. Restricted to trusted positions (valid == scaled > + # mask_threshold, the SAME set used to pick neighbors) and mask-weighted within + # it -- otherwise sub-threshold junk pixels, already excluded as neighbors, + # would leak in. The median sits at the center of the histogram bulk and is + # immune to the heavy outlier tail that a mean / RMS would chase out to the + # right (a few bad-fit pixels dominate a second moment but not the median). + def _weighted_median(err_native: np.ndarray, factor: float) -> float: + e = err_native * factor + use = np.isfinite(e) & valid + return _weighted_quantile(e[use], scaled[use], 0.5) + + precision = {name: _weighted_median(dev[name], scale[name]) for name in scale} + + # weighted histogram of the chosen component, over the same trusted positions + # as the median above. This is purely a picture of the common error values, so + # anything beyond the bin range is left OUT of the bars -- no overflow spike at + # the edge to crush the bulk. Nothing is lost: the median is computed from all + # trusted positions regardless. Bars are normalized by the total trusted + # weight, so each bar is the true fraction of all trusted positions and the + # off-range mass simply isn't drawn (the bars sum to 1 - out_of_range_fraction). + e = dev[component] * scale[component] + use = np.isfinite(e) & valid # trusted positions only, consistent with median + e_f = e[use] + w_f = scaled[use] + # Default histogram range: a robust weighted upper percentile, NOT the raw + # max. A handful of bad-fit positions can reach tens of percent; used as the + # range they crush the entire bulk into the first bin and leave the rest of + # the axis empty (a spurious "spike at 0" plus a far outlier spike). Capping + # at the weighted 99th percentile keeps the common error values readable; the + # few positions past it spill into out_of_range_fraction (reported, not + # drawn). An explicit `bounds`, or passing bin EDGES as `bins`, overrides it. + if bounds is None and np.ndim(bins) == 0 and e_f.size and float(w_f.sum()) > 0: + hi_default = _weighted_quantile(e_f, w_f, 0.99) + if np.isfinite(hi_default) and hi_default > 0: + bounds = (0.0, hi_default) + edges = np.histogram_bin_edges(e_f, bins=bins, range=bounds) + lo, hi = float(edges[0]), float(edges[-1]) + wtot = float(w_f.sum()) + frac_below = float(w_f[e_f < lo].sum()) / wtot if wtot > 0 else 0.0 + frac_above = float(w_f[e_f > hi].sum()) / wtot if wtot > 0 else 0.0 + out_of_range_fraction = frac_below + frac_above + counts_raw, edges = np.histogram(e_f, bins=edges, weights=w_f) + counts = counts_raw / wtot if wtot > 0 else counts_raw + + unit = "°" if component == "rotation" else "%" + result = { + "precision": precision, + "component": component, + "unit": unit, + "counts": counts, + "counts_raw": counts_raw, + "edges": edges, + "out_of_range_fraction": out_of_range_fraction, + "window": int(window), + "n_neighbors": n_neighbors, + "mask_threshold": float(mask_threshold), + "mask_range": (low, high), + "rotation_angle": float(rotation_angle), + "min_neighbors": int(min_neighbors), + "require_full_neighborhood": bool(require_full_neighborhood), + } + + if verbose: + print("Strain precision (median local deviation, mask-weighted)") + print( + f" reference={n_neighbors} neighbors (disk, window={window}) " + f"mask>{mask_threshold:g} min_neighbors={min_neighbors} " + f"rotation_angle={rotation_angle:g} deg" + ) + for name in ("e_uu", "e_vv", "e_uv"): + print(f" {name:<9}: {precision[name]:7.4f} %") + print(f" {'rotation':<9}: {precision['rotation']:7.4f} deg") + print( + f" {'combined':<9}: {precision['combined']:7.4f} % " + "(strain-only Frobenius norm; rotation excluded)" + ) + + if not (plot or returnfig): + return result + + fig, ax = plot_strain_precision_histogram(edges, counts, precision, component, unit) + if returnfig: + return fig, ax + return result + + +# ---- module-level fitting functions ---- + + +def _weighted_quantile(values: np.ndarray, weights: np.ndarray, q: float) -> float: + """Weighted ``q``-quantile of ``values`` (``q`` in ``[0, 1]``); ``nan`` if no weight. + + Uses cumulative-weight interpolation with weights centered on each sorted sample, + so with uniform weights it tracks ``np.quantile``'s linear interpolation and is + robust to a heavy upper tail (the median ignores how far the outliers reach). + """ + values = np.asarray(values, dtype=float) + weights = np.asarray(weights, dtype=float) + total = float(weights.sum()) + if values.size == 0 or total <= 0: + return float("nan") + order = np.argsort(values) + v = values[order] + w = weights[order] + cw = np.cumsum(w) - 0.5 * w + return float(np.interp(q * total, cw, v)) + + +def _local_masked_median( + field: np.ndarray, + valid: np.ndarray, + window: int, + min_neighbors: int, +) -> np.ndarray: + """Median of each position's surrounding neighbors over valid (masked) pixels. + + The center position is excluded ("surrounding" only); a neighbor contributes + only where ``valid`` is True and the field is finite. Neighbors are taken over a + circular (isotropic) footprint of radius ``window / 2`` inscribed in the + ``window`` x ``window`` box — a disk avoids the square's far corners, which + over-weight the diagonals and sample the most strain-different points. Positions + left with fewer than ``min_neighbors`` contributing neighbors return ``nan``. + + Parameters + ---------- + field : np.ndarray + ``(scan_row, scan_col)`` field to take local medians of. + valid : np.ndarray + ``(scan_row, scan_col)`` boolean mask of usable neighbor positions. + window : int + Odd edge length of the bounding box; the footprint is the disk of radius + ``window / 2`` within it (3 -> 8 neighbors, 5 -> 20, 7 -> 36). + min_neighbors : int + Minimum contributing neighbors required, else ``nan``. + + Returns + ------- + np.ndarray + ``(scan_row, scan_col)`` local masked median (``nan`` where undefined). + """ + p = window // 2 + fpad = np.pad(np.asarray(field, dtype=float), p, mode="constant", constant_values=np.nan) + vpad = np.pad(np.asarray(valid, dtype=bool), p, mode="constant", constant_values=False) + + # writable per-position (window, window) neighborhoods + fw = sliding_window_view(fpad, (window, window)).copy() + vw = sliding_window_view(vpad, (window, window)) + fw[~vw] = np.nan + fw[:, :, p, p] = np.nan # exclude the center position from its own median + # restrict the square box to a circular footprint of radius window/2 + oy, ox = np.ogrid[-p : p + 1, -p : p + 1] + outside = (oy**2 + ox**2) > (window / 2.0) ** 2 + fw[:, :, outside] = np.nan + + flat = fw.reshape(fw.shape[0], fw.shape[1], -1) + count = np.sum(np.isfinite(flat), axis=-1) + with warnings.catch_warnings(): + warnings.simplefilter("ignore", category=RuntimeWarning) + med = np.nanmedian(flat, axis=-1) + med[count < min_neighbors] = np.nan + return med + + +def _reference_lattice( + g1_array: np.ndarray, + g2_array: np.ndarray, + mask: np.ndarray | None = None, + strain_mask: np.ndarray | None = None, + calculation_metric: str = "median", +) -> tuple[np.ndarray, np.ndarray]: + """Reference lattice vectors: weighted median or mean over a mask/ROI. + + Each component of the reference is the weighted median + (``calculation_metric="median"``) or weighted mean (``"mean"``) of the lattice + vectors over finite positions with positive weight. Weights come from + ``strain_mask`` if given, else from the continuous ``mask`` (the ``[0, 1]`` + per-position weight, e.g. ``BraggVectors.mask_weight`` from + ``BraggVectors.fit_lattice``). A boolean ROI + (weights in ``{0, 1}``) reduces to the plain median/mean over the selected + positions. If no position has positive weight, the unweighted statistic over all + finite positions is used. A continuous weight (rather than ``mask == 1``) is used + because a min-max normalized mask rarely hits exactly 1. + + Parameters + ---------- + g1_array : np.ndarray + Per-position first lattice vector, shape ``(scan_row, scan_col, 2)``. + g2_array : np.ndarray + Per-position second lattice vector, shape ``(scan_row, scan_col, 2)``. + mask : np.ndarray, optional + ``(scan_row, scan_col)`` per-position weight in ``[0, 1]``. Used as the + weights when ``strain_mask`` is not given. + strain_mask : np.ndarray, optional + ``(scan_row, scan_col)`` ROI / weight taking precedence over ``mask``. + calculation_metric : {"median", "mean"}, default="median" + Statistic used for the reference: weighted median or weighted mean. + + Returns + ------- + tuple of np.ndarray + ``(g1_ref, g2_ref)``, each a length-2 reference vector. + """ + if calculation_metric not in ("mean", "median"): + raise ValueError("calculation metric must be mean or median") + + def reduce(v: np.ndarray, ww: np.ndarray) -> float: + if calculation_metric == "mean": + return float(np.average(v, weights=ww)) + return _weighted_quantile(v, ww, 0.5) + + if strain_mask is not None: + w = np.asarray(strain_mask, dtype=float).reshape(-1) + elif mask is not None: + w = np.asarray(mask, dtype=float).reshape(-1) + else: + w = None + + g1_flat = g1_array.reshape(-1, 2) + g2_flat = g2_array.reshape(-1, 2) + + def _wmed(vals: np.ndarray) -> float: + # weighted statistic over finite, positively-weighted positions; positions + # fit_lattice could not fit are NaN and must be dropped, else the reference + # (and the whole strain map) collapses to NaN. Falls back to the unweighted + # statistic over finite positions when no weight survives. + finite = np.isfinite(vals) + ww = np.ones_like(vals) if w is None else w + use = finite & (ww > 0) + if not use.any(): + return reduce(vals[finite], np.ones(finite.sum())) if finite.any() else float("nan") + return reduce(vals[use], ww[use]) + + g1_ref = np.array((_wmed(g1_flat[:, 0]), _wmed(g1_flat[:, 1])), dtype=float) + g2_ref = np.array((_wmed(g2_flat[:, 0]), _wmed(g2_flat[:, 1])), dtype=float) + return g1_ref, g2_ref + + +def _strain_tensor( + g1_array: np.ndarray, + g2_array: np.ndarray, + g1_ref: np.ndarray, + g2_ref: np.ndarray, + real_space: bool, +) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + """Per-position strain tensor from lattice vectors relative to a reference. + + Two measurement modalities are supported and are arranged to give *identical* + strain for the same physical deformation, so correlation (Bragg) and cepstral + (autocorrelation) maps can be compared directly: + + * ``real_space=False`` -- reciprocal-space lattice vectors (nanobeam Bragg + disks), which contract under tension. The per-position transform is + ``strain_trans = U_ref @ inv(U)``. + * ``real_space=True`` -- real-space lattice vectors (cepstral / Patterson + autocorrelation peaks, or DPC), which expand under tension. The transform is + ``strain_trans = (U @ inv(U_ref)).T``. + + Both expressions evaluate to ``F.T`` (the transpose of the real-space deformation + gradient ``F``), so the normal strains, shear, and rotation come out the same + regardless of modality: ``e_rr = F[0, 0] - 1``, ``e_cc = F[1, 1] - 1``, + ``e_rc = (F[0, 1] + F[1, 0]) / 2`` and ``phi = (F[1, 0] - F[0, 1]) / 2``, so a + lattice rotated counter-clockwise in the (row, col) frame by a small angle + ``theta`` gives ``phi = theta``. + + Parameters + ---------- + g1_array : np.ndarray + Per-position first lattice vector, shape ``(scan_row, scan_col, 2)``. + g2_array : np.ndarray + Per-position second lattice vector, shape ``(scan_row, scan_col, 2)``. + g1_ref : np.ndarray + Reference first lattice vector (length 2). + g2_ref : np.ndarray + Reference second lattice vector (length 2). + real_space : bool + ``False`` for reciprocal-space (Bragg/correlation) vectors; ``True`` for + real-space (cepstral autocorrelation / DPC) vectors. Selects the per-position + transform above; both yield matching strain. + + Returns + ------- + tuple of np.ndarray + ``(e_rr, e_cc, e_rc, phi)``, each of shape ``(scan_row, scan_col)``. + """ + scan_r, scan_c = g1_array.shape[0], g1_array.shape[1] + Uref = np.stack((g1_ref, g2_ref), axis=1).astype(float) + strain_trans = np.zeros((scan_r, scan_c, 2, 2)) + + # For real-space vectors the reference is inverted once (it is shared by every + # position); a non-finite or singular reference leaves the whole map undefined. + Uref_inv = None + if real_space and np.all(np.isfinite(Uref)) and abs(np.linalg.det(Uref)) >= 1e-12: + Uref_inv = np.linalg.inv(Uref) + + for r in range(scan_r): + for c in range(scan_c): + U = np.stack((g1_array[r, c, :], g2_array[r, c, :]), axis=1) + # Positions fit_lattice could not fit are NaN; a degenerate (collinear) + # fit is singular. Either way there is no meaningful inverse -- leave the + # strain NaN (masked out downstream) rather than feeding NaN into pinv, + # whose SVD does not converge and raises LinAlgError. + if not np.all(np.isfinite(U)) or abs(np.linalg.det(U)) < 1e-12: + strain_trans[r, c, :, :] = np.nan + continue + if real_space: + # real-space vectors expand under tension: (U @ U_ref^-1).T == F.T + if Uref_inv is None: + strain_trans[r, c, :, :] = np.nan + else: + strain_trans[r, c, :, :] = (U @ Uref_inv).T + else: + # reciprocal-space vectors contract under tension: U_ref @ U^-1 == F.T + strain_trans[r, c, :, :] = Uref @ np.linalg.inv(U) + + e_rr = strain_trans[:, :, 0, 0] - 1 + e_cc = strain_trans[:, :, 1, 1] - 1 + e_rc = strain_trans[:, :, 1, 0] * 0.5 + strain_trans[:, :, 0, 1] * 0.5 + phi = strain_trans[:, :, 1, 0] * -0.5 + strain_trans[:, :, 0, 1] * 0.5 + return e_rr, e_cc, e_rc, phi + + +def _rotate_strain_tensor( + e_rr: np.ndarray, + e_cc: np.ndarray, + e_rc: np.ndarray, + rotation_angle: float, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Rotate a 2D strain tensor by ``rotation_angle`` (degrees). + + Parameters + ---------- + e_rr : np.ndarray + Row-row (normal) strain component. + e_cc : np.ndarray + Column-column (normal) strain component. + e_rc : np.ndarray + Row-column (shear) strain component. + rotation_angle : float + Frame rotation angle, in degrees. + + Returns + ------- + tuple of np.ndarray + ``(e_uu, e_vv, e_uv)`` in the rotated frame. + """ + angle = np.deg2rad(rotation_angle) + c = np.cos(angle) + s = np.sin(angle) + e_uu = e_rr * (c * c) + 2.0 * e_rc * (c * s) + e_cc * (s * s) + e_vv = e_rr * (s * s) - 2.0 * e_rc * (c * s) + e_cc * (c * c) + e_uv = (e_cc - e_rr) * (c * s) + e_rc * (c * c - s * s) + return e_uu, e_vv, e_uv + + +def _raw_vec_to_display(vec_rc: NDArray, *, rotation_ccw_deg: float, transpose: bool) -> NDArray: + """Map a raw-detector ``(row, col)`` vector into the rotated display frame. + + Applies the optional axis transpose, then a counter-clockwise rotation of + ``rotation_ccw_deg``. + + Parameters + ---------- + vec_rc : np.ndarray + Vector(s) with ``(row, col)`` in the last axis, shape ``(..., 2)``. + rotation_ccw_deg : float + Counter-clockwise rotation in degrees. + transpose : bool + If ``True``, swap row and col before rotating. + + Returns + ------- + np.ndarray + Rotated vector(s), same shape as ``vec_rc``. + """ + v = np.asarray(vec_rc, dtype=float) + dr, dc = v[..., 0], v[..., 1] + + if transpose: + dr, dc = dc, dr + + theta = np.deg2rad(rotation_ccw_deg) + ct = np.cos(theta) + st = np.sin(theta) + + dr2 = ct * dr - st * dc + dc2 = st * dr + ct * dc + return np.stack((dr2, dc2), axis=-1) diff --git a/src/quantem/diffraction/strain_visualization.py b/src/quantem/diffraction/strain_visualization.py new file mode 100644 index 000000000..f899e1685 --- /dev/null +++ b/src/quantem/diffraction/strain_visualization.py @@ -0,0 +1,672 @@ +from __future__ import annotations + +import warnings + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.cm import ScalarMappable +from matplotlib.colors import Normalize +from matplotlib.patches import FancyArrowPatch +from matplotlib.ticker import FuncFormatter, MaxNLocator + +from quantem.core.visualization.visualization_utils import ScalebarConfig, add_scalebar_to_ax + + +def plot_strain_panels( + e_uu: np.ndarray, + e_vv: np.ndarray, + e_uv: np.ndarray, + rotation: np.ndarray, + mask: np.ndarray | None, + g1_ref: np.ndarray | None, + g2_ref: np.ndarray | None, + ds_shape: tuple[int, ...], + ds_sampling: float = 1.0, + ds_units: str = "pixels", + strain_range_percent: tuple[float, float] = (-3.0, 3.0), + rotation_range_degrees: tuple[float, float] = (-2.0, 2.0), + mask_range: tuple[float, float] = (0.0, 1.0), + roi: np.ndarray | None = None, + plot_rotation: bool = True, + plot_gvecs: bool = False, + plot_scalebar: bool = False, + cmap_strain: str = "RdBu_r", + cmap_rotation: str = "PiYG", + layout: str = "horizontal", + transpose_image: bool = False, + rotate_title: bool = False, + plot_dilation: bool = False, + figsize: tuple[float, float] | None = None, + panel_titles: tuple[str, str, str] | None = None, + strain_rotation_angle: float = 0.0, + arrow_style: str = "title", + **kwargs, +): + """Render strain (e_uu, e_vv, e_uv) and rotation panels. + + Strain arrays are fractional (multiplied by 100 for display); ``rotation`` is + in radians (converted to degrees for display). ``panel_titles`` overrides the + three strain-panel titles (e.g. to label the raw row/col reference frame). + + The mask modulates panel brightness (black where masked out). ``mask_range`` + ``(low, high)`` remaps it linearly before display: mask values ``>= high`` show + full color, ``<= low`` go black, and values between ramp from black to full. + The default ``(0.0, 1.0)`` leaves the already-normalized mask unchanged. + + When ``roi`` (a boolean ``(scan_row, scan_col)`` array) is given, positions + inside it are drawn in color and positions outside it in greyscale (the same + field, desaturated), so a chosen reference region stands out from its context. + + Parameters + ---------- + e_uu, e_vv, e_uv : np.ndarray + ``(scan_row, scan_col)`` fractional strain components in the display frame. + rotation : np.ndarray + ``(scan_row, scan_col)`` infinitesimal rotation in radians. + mask : np.ndarray or None + ``(scan_row, scan_col)`` brightness weights in ``[0, 1]``; ``None`` shows + every position at full brightness. + g1_ref, g2_ref : np.ndarray or None + Reference lattice vectors ``(row, col)``; only their directions are drawn, + and only when ``plot_gvecs=True``. + ds_shape : tuple of int + Scan shape; the first two entries size the default mask and the scale bar. + ds_sampling : float, default=1.0 + Scan step size used by the scale bar, in ``ds_units``. + ds_units : str, default="pixels" + Units of ``ds_sampling``. + strain_range_percent : tuple of float, default=(-3.0, 3.0) + Color range of the strain panels, in percent. + rotation_range_degrees : tuple of float, default=(-2.0, 2.0) + Color range of the rotation panel, in degrees. + mask_range : tuple of float, default=(0.0, 1.0) + ``(low, high)`` window remapping ``mask`` before display (see above). + roi : np.ndarray, optional + Boolean ``(scan_row, scan_col)`` region drawn in color; the rest is grey. + plot_rotation : bool, default=True + Whether to add the rotation panel. + plot_gvecs : bool, default=False + Whether to draw the directions of ``g1_ref`` and ``g2_ref`` beside the panels. + plot_scalebar : bool, default=False + Whether to draw a scale bar on the first panel. + cmap_strain : str, default="RdBu_r" + Colormap of the strain panels. + cmap_rotation : str, default="PiYG" + Colormap of the rotation panel; ``None`` uses ``cmap_strain``. + layout : {"horizontal", "vertical"}, default="horizontal" + Panel arrangement. + transpose_image : bool, default=False + If ``True``, transpose every panel (swap scan rows and columns) for display. + rotate_title : bool, default=False + If ``True``, draw the panel titles vertically. + plot_dilation : bool, default=False + If ``True``, show ``e_uu + e_vv`` and ``e_uv`` (plus rotation) instead of + the three strain components. Forces the rotation panel on. + figsize : tuple of float, optional + Figure size in inches; derived from ``layout`` if omitted. + panel_titles : tuple of str, optional + Titles of the strain panels (three, or two with ``plot_dilation``). When + given, no direction arrows are drawn. + strain_rotation_angle : float, default=0.0 + Angle in degrees by which the default direction arrows are rotated, to + match a strain tensor rotated by this angle. + arrow_style : {"title", "legend"}, default="title" + Draw the direction arrows next to the panel titles, or in a legend at the + side of the figure. + **kwargs + Keys starting with ``scalebar_`` are passed to the scale bar with the prefix + removed (for example ``scalebar_length``, ``scalebar_color``, + ``scalebar_box``, ``scalebar_box_color``, ``scalebar_box_alpha``). Other + keys are ignored. + + Returns + ------- + tuple + ``(fig, ax)`` with ``ax`` the array of panel axes. + """ + if mask is None: + mask = np.ones(ds_shape[:2]) + + # remap the mask brightness onto the [low, high] window: <= low -> black, + # >= high -> full color, linear between. default (0, 1) is a no-op. + low, high = float(mask_range[0]), float(mask_range[1]) + if high > low: + mask = np.clip((np.asarray(mask, dtype=float) - low) / (high - low), 0.0, 1.0) + else: + mask = (np.asarray(mask, dtype=float) >= high).astype(float) + + if cmap_rotation is None: + cmap_rotation = cmap_strain + + if layout not in ["horizontal", "vertical"]: + raise ValueError("layout must be 'horizontal' or 'vertical'") + + ncols = 4 if plot_rotation else 3 + is_horizontal = layout == "horizontal" + if plot_dilation: + ncols = 3 + plot_rotation = True + + n_strain = 2 if plot_dilation else 3 + + if figsize is None: + figsize = (8, 3) if is_horizontal else (6, 6) + + if is_horizontal: + fig, ax = plt.subplots(1, ncols, figsize=figsize) + else: + fig, ax = plt.subplots(ncols, 1, figsize=figsize) + + cm_strain = plt.get_cmap(cmap_strain).copy() + cm_strain.set_bad(color="black") + cm_rot = plt.get_cmap(cmap_rotation).copy() + cm_rot.set_bad(color="black") + + euu_pct = e_uu * 100 + evv_pct = e_vv * 100 + euv_pct = e_uv * 100 + rot_deg = np.rad2deg(rotation) + + roi_bool = None if roi is None else np.asarray(roi).astype(bool) + gray_cm = plt.get_cmap("gray").copy() + gray_cm.set_bad(color="black") + + def _roi_compose(norm_vals, color_cm): + """Color the field inside the ROI; show it in greyscale outside the ROI.""" + rgb = color_cm(norm_vals)[:, :, :3] + if roi_bool is None: + return rgb + rgb_gray = gray_cm(norm_vals)[:, :, :3] + return np.where(roi_bool[:, :, np.newaxis], rgb, rgb_gray) + + norm_strain = Normalize(vmin=strain_range_percent[0], vmax=strain_range_percent[1]) + euu_disp = _roi_compose(norm_strain(euu_pct), cm_strain) + evv_disp = _roi_compose(norm_strain(evv_pct), cm_strain) + euv_disp = _roi_compose(norm_strain(euv_pct), cm_strain) + + if transpose_image: + euu_disp = euu_disp.transpose(1, 0, 2) + evv_disp = evv_disp.transpose(1, 0, 2) + euv_disp = euv_disp.transpose(1, 0, 2) + mask = mask.T + + if plot_dilation: + etot_pct = (e_uu + e_vv) * 100 + etot_disp = _roi_compose(norm_strain(etot_pct), cm_strain) + if transpose_image: + etot_disp = etot_disp.transpose(1, 0, 2) + ax[0].imshow(etot_disp * mask[:, :, np.newaxis]) + ax[1].imshow(euv_disp * mask[:, :, np.newaxis]) + else: + ax[0].imshow(euu_disp * mask[:, :, np.newaxis]) + ax[1].imshow(evv_disp * mask[:, :, np.newaxis]) + ax[2].imshow(euv_disp * mask[:, :, np.newaxis]) + + ref_dim = figsize[1] if is_horizontal else figsize[0] + fs_threshold = 3.0 + fs_scale = min(1.0, max(0.5, ref_dim / fs_threshold)) + title_fs = 16 * fs_scale + tick_fs = 12 * fs_scale + title_val = "vertical" if rotate_title else "horizontal" + if panel_titles is None: + if plot_dilation: + panel_titles = ( + r"$\epsilon_{uu} + \epsilon_{vv}$", + r"$\epsilon_{uv}$", + "", + ) + title_arrow_angles = (None, -45 + strain_rotation_angle, None) + else: + panel_titles = ( + r"$\epsilon_{uu}$", + r"$\epsilon_{vv}$", + r"$\epsilon_{uv}$", + ) + title_arrow_angles = ( + 90 + strain_rotation_angle, + 0 + strain_rotation_angle, + -45 + strain_rotation_angle, + ) + if transpose_image: + title_arrow_angles = ( + 0 + strain_rotation_angle, + 90 + strain_rotation_angle, + 45 + strain_rotation_angle, + ) + else: + title_arrow_angles = (None, None, None) + + if plot_rotation: + norm_rot = Normalize(vmin=rotation_range_degrees[0], vmax=rotation_range_degrees[1]) + rot_disp = _roi_compose(norm_rot(rot_deg), cm_rot) + if transpose_image: + rot_disp = rot_disp.transpose(1, 0, 2) + ax[-1].imshow(rot_disp * mask[:, :, np.newaxis]) + if arrow_style == "title": + ax[-1].set_title(r"$\phi$ $\circlearrowleft$", fontsize=title_fs, rotation=title_val) + else: + ax[-1].set_title(r"$\phi$", fontsize=title_fs, rotation=title_val) + + for a in ax: + a.set_xticks([]) + a.set_yticks([]) + a.set_facecolor("black") + a.set_aspect("equal") + a.set_anchor("W" if not is_horizontal else "C") + + if plot_scalebar: + scalebar_kwargs = {} + for key, value in kwargs.items(): + if key.startswith("scalebar_"): + scalebar_key = key[len("scalebar_") :] + scalebar_kwargs[scalebar_key] = value + + # default: white bar on a translucent black box, readable on the + # diverging strain colormaps; override with scalebar_color / + # scalebar_box / scalebar_box_color / scalebar_box_alpha + box = scalebar_kwargs.pop("box", True) + box_color = scalebar_kwargs.pop("box_color", "black") + box_alpha = scalebar_kwargs.pop("box_alpha", 0.45) + scalebar_defaults = { + "sampling": ds_sampling, + "units": ds_units, + "length": None, + "width_px": 1, + "pad_px": 0.5, + "color": "white" if box and box_color == "black" else "black", + "loc": "lower left", + "fontsize": 12, + "bold": True, + } + scalebar_defaults.update(scalebar_kwargs) + scalebar_config = ScalebarConfig(**scalebar_defaults) + add_scalebar_to_ax( + ax[0], + array_size=int(ds_shape[0]), + sampling=scalebar_config.sampling, + length_units=scalebar_config.length, + units=scalebar_config.units, + width_px=scalebar_config.width_px, + pad_px=scalebar_config.pad_px, + color=scalebar_config.color, + loc=scalebar_config.loc, + fontsize=scalebar_config.fontsize, + bold=scalebar_config.bold, + box=box, + box_color=box_color, + box_alpha=box_alpha, + ) + + cb_size = 0.02 + cb_pad = 0.03 + cb_min_len = 0.16 + + def _finalize_layout(): + # set_aspect("equal") only resizes/recenters each panel at draw time, so + # get_position() before a draw returns stale boxes -- placing the colorbars + # and g-vector compass off those boxes then spills them off the figure. + # Settle the layout cheaply first so every box read below is the real one. + try: + fig.draw_without_rendering() + except AttributeError: # matplotlib < 3.5 + fig.canvas.draw() + + need_side_panel = plot_gvecs or arrow_style == "legend" + if is_horizontal: + # Reserve a bottom band wide enough for the colorbar + its tick labels and + # title (fontsize 16) and a right band for the rotation-panel gap; widen the + # right band when the g-vector compass is drawn in it. These keep the figure + # usable when saved "as is" (no bbox_inches='tight'). + right = 0.72 if need_side_panel else 0.93 + fig.subplots_adjust(left=0.04, right=right, top=0.88, bottom=0.24, wspace=0.05) + if plot_rotation: + # nudge the rotation panel right for a visual gap from the strain panels; + # 0.03 stays inside the reserved right band so nothing is clipped. + pos3 = ax[-1].get_position() + ax[-1].set_position([pos3.x0 + 0.03, pos3.y0, pos3.width, pos3.height]) + _finalize_layout() + + cb_orientation = "horizontal" + b0 = ax[0].get_position() + b2 = ax[n_strain - 1].get_position() + cb_y = b2.y0 - cb_pad - cb_size + strain_cb_pos = [b0.x0, cb_y, b2.x1 - b0.x0, cb_size] + + if plot_rotation: + b3 = ax[-1].get_position() + rot_cb_w = max(b3.x1 - b3.x0, cb_min_len) + rot_cb_cx = 0.5 * (b3.x0 + b3.x1) + rot_cb_x0 = min(max(rot_cb_cx - 0.5 * rot_cb_w, 0.0), 0.99 - rot_cb_w) + rot_cb_pos = [rot_cb_x0, cb_y, rot_cb_w, cb_size] + last_pos = b3 + else: + rot_cb_pos = None + last_pos = b2 + + else: + # Top band for the panel titles, right band for the vertical colorbars + labels. + right = 0.55 if need_side_panel else 0.80 + fig.subplots_adjust(left=0.04, right=right, top=0.92, bottom=0.06, hspace=0.15) + _finalize_layout() + + cb_orientation = "vertical" + b0 = ax[0].get_position() + b2 = ax[n_strain - 1].get_position() + title_gap = 0.15 if arrow_style == "title" else cb_pad + cb_x0 = b0.x1 + title_gap + strain_cb_pos = [cb_x0, b2.y0, cb_size, b0.y1 - b2.y0] + + if plot_rotation: + b3 = ax[-1].get_position() + rot_cb_h = max(b3.y1 - b3.y0, cb_min_len) + rot_cb_cy = 0.5 * (b3.y0 + b3.y1) + rot_cb_y0 = min(max(rot_cb_cy - 0.5 * rot_cb_h, 0.0), 0.99 - rot_cb_h) + rot_cb_pos = [cb_x0, rot_cb_y0, cb_size, rot_cb_h] + last_pos = b3 + else: + rot_cb_pos = None + last_pos = b2 + + cax1 = fig.add_axes(strain_cb_pos) + sm_strain = ScalarMappable(norm=norm_strain, cmap=cm_strain) + cbar1 = fig.colorbar(sm_strain, cax=cax1, orientation=cb_orientation) + cbar1.set_label("Strain", fontsize=title_fs) + cbar1.formatter = FuncFormatter(lambda v, _pos: f"{v:g}%") + cbar1.update_ticks() + cbar1.ax.tick_params(labelsize=tick_fs) + + if plot_rotation and rot_cb_pos is not None: + cax2 = fig.add_axes(rot_cb_pos) + sm_rot = ScalarMappable(norm=norm_rot, cmap=cm_rot) + cbar2 = fig.colorbar(sm_rot, cax=cax2, orientation=cb_orientation) + cbar2.set_label("Rotation", fontsize=title_fs) + cbar2.formatter = FuncFormatter(lambda v, _pos: f"{v:g}°") + cbar2.locator = MaxNLocator(nbins=2) + cbar2.update_ticks() + cbar2.ax.tick_params(labelsize=tick_fs) + + def _add_title_arrow(ax, angle_deg, gap_pt=4.0, color="black", fontsize=None): + fs = fontsize if fontsize is not None else title_fs + try: + ax.figure.draw_without_rendering() + except AttributeError: # matplotlib < 3.5 + ax.figure.canvas.draw() + renderer = ax.figure.canvas.get_renderer() + bbox_ax = ax.title.get_window_extent(renderer=renderer).transformed( + ax.transAxes.inverted() + ) + y = 0.5 * (bbox_ax.y0 + bbox_ax.y1) + ax.annotate( + "\u2194", + xy=(bbox_ax.x1, y), + xycoords=ax.transAxes, + xytext=(gap_pt + fs / 2.0, 0), + textcoords="offset points", + ha="center", + va="center", + rotation=angle_deg, + rotation_mode="anchor", + fontsize=fs, + color=color, + annotation_clip=False, + ) + + def _add_arrow_legend(fig, x0, y_top, entries, plot_rotation, fontsize, color="black"): + fig_w_in, fig_h_in = figsize + row_h_in = fontsize * 1.6 / 72.0 + box_w_in = 1.6 + n_rows = len(entries) + 1 + (2 if plot_rotation else 0) + row_h = row_h_in / fig_h_in + box_h = row_h * n_rows + box_w = min(box_w_in / fig_w_in, 0.99 - x0) + leg_ax = fig.add_axes([x0, y_top - box_h, box_w, box_h]) + leg_ax.set_xlim(0, 1) + leg_ax.set_ylim(0, 1) + leg_ax.axis("off") + + dy = 1.0 / n_rows + y = 1.0 - dy / 2 + leg_ax.text(0.0, y, "Strain", fontsize=fontsize, fontweight="bold", ha="left", va="center") + for label, angle_deg in entries: + y -= dy + leg_ax.text( + 0.15, + y, + "\u2194", + rotation=angle_deg, + rotation_mode="anchor", + ha="center", + va="center", + fontsize=fontsize, + color=color, + ) + leg_ax.text(0.32, y, label, fontsize=fontsize, ha="left", va="center") + + if plot_rotation: + y -= dy + leg_ax.text( + 0.0, y, "Rotation", fontsize=fontsize, fontweight="bold", ha="left", va="center" + ) + y -= dy + leg_ax.text(0.15, y, "\u21ba", fontsize=fontsize, ha="center", va="center") + leg_ax.text(0.32, y, r"$\phi$", fontsize=fontsize, ha="left", va="center") + return box_h + + for i in range(n_strain): + ax[i].set_title(panel_titles[i], fontsize=title_fs, rotation=title_val) + angle = title_arrow_angles[i] + if arrow_style == "title" and angle is not None: + _add_title_arrow(ax[i], angle, color="black") + + if is_horizontal: + _finalize_layout() + renderer = fig.canvas.get_renderer() + panel_edge = last_pos.x1 + if plot_rotation: + title_edge = ( + ax[-1] + .title.get_window_extent(renderer=renderer) + .transformed(fig.transFigure.inverted()) + .x1 + ) + panel_edge = max(panel_edge, title_edge) + margin_x0 = panel_edge + 0.03 + else: + _finalize_layout() + renderer = fig.canvas.get_renderer() + margin_x0 = cax1.get_tightbbox(renderer).transformed(fig.transFigure.inverted()).x1 + 0.02 + if plot_rotation and rot_cb_pos is not None: + rot_edge = cax2.get_tightbbox(renderer).transformed(fig.transFigure.inverted()).x1 + margin_x0 = max(margin_x0, rot_edge + 0.02) + top_bound = 0.88 if is_horizontal else 0.92 + bottom_bound = 0.24 if is_horizontal else 0.06 + center_y = 0.5 * (top_bound + bottom_bound) + + entries = [] + leg_h = 0.0 + if arrow_style == "legend": + entries = [ + (panel_titles[i], title_arrow_angles[i]) + for i in range(n_strain) + if title_arrow_angles[i] is not None + ] + n_rows = len(entries) + 1 + (2 if plot_rotation else 0) + leg_h = (title_fs * 1.6 / 72.0 / figsize[1]) * n_rows + + show_gvecs = plot_gvecs and g1_ref is not None and g2_ref is not None + if plot_gvecs and not show_gvecs: + warnings.warn( + "plot_gvecs=True but g1_ref and g2_ref are not set; run " + "StrainMap.update_reference() first.", + UserWarning, + ) + fig_aspect = figsize[0] / figsize[1] + gvec_w = min(0.99 - margin_x0, 0.15) if show_gvecs else 0.0 + gvec_h = gvec_w * fig_aspect if show_gvecs else 0.0 + + gap = 0.03 if (leg_h > 0 and gvec_h > 0) else 0.0 + total_needed = leg_h + gap + gvec_h + available_span = top_bound - bottom_bound + side_scale = min(1.0, available_span / total_needed) if total_needed > 0 else 1.0 + leg_h *= side_scale + gvec_w *= side_scale + gvec_h *= side_scale + legend_fontsize = title_fs * side_scale + + y_top = center_y + (leg_h + gap * side_scale + gvec_h) / 2.0 + + if leg_h > 0: + _add_arrow_legend( + fig, + margin_x0, + y_top, + entries, + plot_rotation=plot_rotation, + fontsize=legend_fontsize, + color="black", + ) + y_top -= leg_h + gap * side_scale + + if show_gvecs: + ref_ax = fig.add_axes([margin_x0, y_top - gvec_h, gvec_w, gvec_h]) + ref_ax.set_xlim(-1.5, 1.5) + ref_ax.set_ylim(-1.5, 1.5) + ref_ax.set_aspect("equal") + ref_ax.axis("off") + g1_norm = g1_ref / np.linalg.norm(g1_ref) + g2_norm = g2_ref / np.linalg.norm(g2_ref) + g1_row, g1_col = g1_norm + g2_row, g2_col = g2_norm + arrow_props_ref = dict(arrowstyle="->", lw=3, mutation_scale=25) + ref_ax.add_patch( + FancyArrowPatch((0, 0), (g1_col, -g1_row), color="darkred", **arrow_props_ref) + ) + ref_ax.add_patch( + FancyArrowPatch((0, 0), (g2_col, -g2_row), color="darkblue", **arrow_props_ref) + ) + ref_ax.text( + g1_col * 1.3, + -g1_row * 1.3, + r"$\mathbf{g}_{1}$", + fontsize=14, + fontweight="bold", + color="darkred", + ha="center", + va="center", + ) + ref_ax.text( + g2_col * 1.3, + -g2_row * 1.3, + r"$\mathbf{g}_{2}$", + fontsize=14, + fontweight="bold", + color="darkblue", + ha="center", + va="center", + ) + + return fig, ax + + +def plot_strain_precision_histogram( + edges: np.ndarray, + counts: np.ndarray, + precision: dict[str, float], + component: str, + unit: str, + *, + figsize: tuple[float, float] = (6.0, 4.0), +): + """Weighted histogram of the local-deviation strain precision. + + ``edges``/``counts`` describe the (mask-weighted, normalized) distribution of the + chosen ``component`` deviation in display units (``unit``). ``precision`` is the + weighted-median local deviation per component (used for the annotation box); the + plotted component's median is marked with a solid line. + + Parameters + ---------- + edges : np.ndarray + ``(bins + 1,)`` histogram bin edges, in ``unit``. + counts : np.ndarray + ``(bins,)`` weighted fraction of positions in each bin. + precision : dict of str to float + Median local deviation per component; must contain ``"e_uu"``, ``"e_vv"``, + ``"e_uv"`` (percent), ``"rotation"`` (degrees) and ``"combined"`` (percent). + component : str + Key of ``precision`` that was histogrammed; its median is marked. + unit : str + Display unit of ``component`` (``"%"`` or ``"°"``). + figsize : tuple of float, default=(6.0, 4.0) + Figure size in inches. + + Returns + ------- + tuple + ``(fig, ax)``. + """ + fig, ax = plt.subplots(figsize=figsize) + edges = np.asarray(edges, dtype=float) + counts = np.asarray(counts, dtype=float) + centers = 0.5 * (edges[:-1] + edges[1:]) + widths = np.diff(edges) + + ax.bar( + centers, + counts, + width=widths, + align="center", + color="#4C72B0", + edgecolor="white", + linewidth=0.3, + ) + + median_value = precision[component] + if np.isfinite(median_value): + ax.axvline(median_value, color="crimson", ls="-", lw=2) + # label the line inline -- a legend box here would sit on top of the info box. + # Put the text on whichever side of the line keeps it clear of the right box. + span = float(edges[-1] - edges[0]) + on_right = span > 0 and (median_value - edges[0]) / span > 0.5 + ax.annotate( + f"median = {median_value:.3g} {unit}", + xy=(median_value, 0.96), + xycoords=("data", "axes fraction"), + xytext=(-6 if on_right else 6, 0), + textcoords="offset points", + ha="right" if on_right else "left", + va="top", + color="crimson", + fontsize=9, + ) + + label = "combined" if component == "combined" else component + ax.set_xlabel(f"{label} deviation ({unit})", fontsize=12) + ax.set_ylabel("weighted fraction", fontsize=12) + ax.set_title("Strain precision (median local deviation)", fontsize=13) + ax.tick_params(labelsize=10) + + annotation = "\n".join( + [ + r"median:", + rf" $\epsilon_{{uu}}$: {precision['e_uu']:.3g} %", + rf" $\epsilon_{{vv}}$: {precision['e_vv']:.3g} %", + rf" $\epsilon_{{uv}}$: {precision['e_uv']:.3g} %", + rf" rotation: {precision['rotation']:.3g} °", + rf" combined: {precision['combined']:.3g} %", + ] + ) + ax.text( + 0.97, + 0.97, + annotation, + transform=ax.transAxes, + ha="right", + va="top", + fontsize=9, + family="monospace", + bbox=dict(boxstyle="round", fc="white", ec="0.7", alpha=0.9), + ) + + fig.tight_layout() + return fig, ax diff --git a/src/quantem/diffraction/wk_scattering_factors.py b/src/quantem/diffraction/wk_scattering_factors.py new file mode 100644 index 000000000..744acefd8 --- /dev/null +++ b/src/quantem/diffraction/wk_scattering_factors.py @@ -0,0 +1,548 @@ +"""Weickenmeier-Kohl absorptive electron scattering factors. + +Elastic form factors use the 8-parameter fit of Weickenmeier & Kohl, +Acta Cryst. A47, 590 (1991); the absorptive (core-loss and phonon/TDS) +parts are computed analytically from the same fit. The implementation +follows Weickenmeier's original F77 code as adapted by Marc De Graef in +EMsoftLib/others.f90, translated to Python by SE Zeltmann for py4DSTEM and +vectorized over g by Colin Ophus. It was vendored from py4DSTEM with one +correction: the middle branch of WEKO now starts at argu >= 0.1, as in +others.f90 (py4DSTEM used argu >= 1.0, which dropped every term with argu +in [0.1, 1) and made the form factors non-monotonic at small g). Comments +in quotation marks are carried over from the original code. +""" + +import numpy as np +from scipy.special import expi + +from quantem.core.utils.utils import electron_wavelength_angstrom + + +def compute_WK_factor( + g: np.ndarray, + Z: int, + accelerating_voltage: float, + thermal_sigma: float | None = None, + include_core: bool = True, + include_phonon: bool = True, +) -> np.ndarray: + """Absorptive, relativistically corrected Weickenmeier-Kohl scattering factors. + + The elastic part uses the 8-parameter fit of the elastic form factors; + the absorptive part is computed analytically from the same fitting + function, following EMsoftLib/others.f90. Vectorized over `g` only. + + Parameters + ---------- + g : array-like + Scattering vector magnitudes 1/d_hkl (crystallographic convention, + no 2 pi), 1/Angstroms. + Z : int + Atomic number, 1 (H) to 98 (Cf). + accelerating_voltage : float + Beam energy, eV. + thermal_sigma : float | None + RMS atomic displacement for the Debye-Waller factor and the TDS + absorption, Angstroms (often written in papers). None or 0 + means no thermal motion: no Debye-Waller damping and no phonon + absorption. + include_core : bool, default=True + Include the core-loss contribution to the absorptive form factor. + include_phonon : bool, default=True + Include the phonon/TDS contribution to the absorptive form factor. + + Returns + ------- + np.ndarray + Complex128 form factors, same shape as `g`, Angstroms. The real part + is elastic, the imaginary part absorptive. + """ + g = np.atleast_1d(np.asarray(g, dtype=float)) + + # the WK Fortran code works in its own units: + # lowercase "g", our input, is the standard crystallographic quantity, in A^-1 + # uppercase "G" is the "G" in others.f90:FSCATT, g * 2 pi + # uppercase "S" is the "S" in others.f90:FSCATT, G / 4 pi = g / 2 + G = g * 2.0 * np.pi + S = g / 2.0 + + accelerating_voltage_kV = accelerating_voltage / 1.0e3 + + if thermal_sigma is not None: + UL = float(thermal_sigma) + DWF = np.exp(-0.5 * UL**2 * G**2) + else: + UL = 0.0 + DWF = 1.0 + + A = WK_A_param[int(Z) - 1] + B = WK_B_param[int(Z) - 1] + + # WEKO(A,B,S) + # NOTE: the py4DSTEM version this was vendored from used `argu >= 1.0` + # for the middle branch, silently dropping every term with argu in + # [0.1, 1) and producing non-monotonic form factors at small g. The + # original EMsoftLib/others.f90 WEKO uses ARGU >= 0.1, restored here. + WK = np.zeros_like(S) + for i in range(4): + argu = B[i] * S**2 + sub = argu < 0.1 + WK[sub] += A[i] * B[i] * (1.0 - 0.5 * argu[sub]) + sub = np.logical_and(argu >= 0.1, argu <= 20.0) + WK[sub] += A[i] * (1.0 - np.exp(-argu[sub])) / S[sub] ** 2 + sub = argu > 20.0 + WK[sub] += A[i] / S[sub] ** 2 + + Freal = 4.0 * np.pi * DWF * WK + + ################################################# + # calculate "core" contribution, following FCORE: + k0 = ( + 2.0 * np.pi / electron_wavelength_angstrom(accelerating_voltage) + ) # remember, physicist units here + + if include_core: + # "CALCULATE CHARACTERISTIC ENERGY LOSS AND ANGLE" + DE = 6.0e-3 * Z + theta_e = ( + DE + / (2.0 * accelerating_voltage_kV) + * (2.0 * accelerating_voltage_kV + 1022.0) + / (accelerating_voltage_kV + 1022.0) + ) + + # "SCREENING PARAMETER OF YUKAWA POTENTIAL" + R = 0.885 * 0.5289 / Z ** (1.0 / 3.0) + + # "CALCULATE NORMALISING ANGLE" + TA = 1.0 / (k0 * R) + + # "CALCULATE BRAGG ANGLE" + TB = G / (2.0 * k0) + + # "NORMALIZE" + OMEGA = 2.0 * TB / TA + KAPPA = theta_e / TA + + K2 = KAPPA * KAPPA + O2 = OMEGA * OMEGA + + X1 = ( + OMEGA + / ((1.0 + O2) * np.sqrt(O2 + 4.0 * K2)) + * np.log((OMEGA + np.sqrt(O2 + 4.0 * K2)) / (2.0 * KAPPA)) + ) + X2 = ( + 1.0 + / np.sqrt((1.0 + O2) * (1.0 + O2) + 4.0 * K2 * O2) + * np.log( + (1.0 + 2.0 * K2 + O2 + np.sqrt((1.0 + O2) * (1.0 + O2) + 4.0 * K2 * O2)) + / (2.0 * KAPPA * np.sqrt(1.0 + K2)) + ) + ) + + X3 = np.zeros_like(OMEGA) + sub = OMEGA > 1e-2 + X3[sub] = ( + 1.0 + / (OMEGA[sub] * np.sqrt(O2[sub] + 4.0 * (1.0 + K2))) + * np.log( + (OMEGA[sub] + np.sqrt(O2[sub] + 4.0 * (1.0 + K2))) / (2.0 * np.sqrt(1.0 + K2)) + ) + ) + sub = np.logical_not(sub) + X3[sub] = 1.0 / (4.0 * (1.0 + K2)) + + HI = 2 * Z / (TA * TA) * (-X1 + X2 - X3) + + A0 = 0.5289 + Fcore = 4.0 / (A0 * A0) * 2.0 * np.pi / (k0 * k0) * HI + else: + Fcore = 0.0 + + ########################################################## + # calculate phonon contribution, following FPHON(G,UL,A,B) + # without thermal motion RI2 equals RI1 and the phonon term vanishes; + # evaluating it at U = 0 instead gives inf - inf in the asymptotic branch + Fphon = 0.0 + if include_phonon and UL > 0.0: + A1 = A * (4.0 * np.pi) ** 2 + B1 = B / (4.0 * np.pi) ** 2 + + for jj in range(4): + for ii in range(jj + 1): + Fphon += ( + (2.0 if jj != ii else 1.0) + * A1[jj] + * A1[ii] + * (DWF * RI1(B1[ii], B1[jj], G) - RI2(B1[ii], B1[jj], G, UL)) + ) + + Fimag = (Fcore * DWF) + Fphon + + # perform relativistic correction + gamma = (accelerating_voltage_kV + 511.0) / (511.0) + + Fscatt = np.asarray((Freal * gamma) + (1.0j * (Fimag * gamma**2 / k0)), dtype=np.complex128) + + # convert to Angstroms and remove the extra physicist factors, as in + # diffraction.f90:427,576,630 + return Fscatt * 0.4787801 * 0.664840340614319 / (4.0 * np.pi) + + +############################################## +# Helper integral functions for DW calculation + + +def RI1(BI, BJ, G): + # "ERSTES INTEGRAL FUER DIE ABSORPTIONSPOTENTIALE" + eps = np.max([BI, BJ]) * G**2 + + ri1 = np.zeros_like(G) + + sub = eps <= 0.1 + ri1[sub] = np.pi * (BI * np.log((BI + BJ) / BI) + BJ * np.log((BI + BJ) / BJ)) + + sub = np.logical_and(eps <= 0.1, G > 0.0) + temp = 0.5 * BI**2 * np.log(BI / (BI + BJ)) + 0.5 * BJ**2 * np.log(BJ / (BI + BJ)) + temp += 0.75 * (BI**2 + BJ**2) - 0.25 * (BI + BJ) ** 2 + temp -= 0.5 * (BI - BJ) ** 2 + ri1[sub] += np.pi * G[sub] ** 2 * temp + + sub = eps > 0.1 + ri1[sub] = ( + 2.0 * 0.5772157 + + np.log(BI * G[sub] ** 2) + + np.log(BJ * G[sub] ** 2) + - 2.0 * expi(-BI * BJ * G[sub] ** 2 / (BI + BJ)) + ) + + ri1[sub] += RIH1(BI * G[sub] ** 2, BI * G[sub] ** 2 * BI / (BI + BJ), BI * G[sub] ** 2) + + ri1[sub] += RIH1(BJ * G[sub] ** 2, BJ * G[sub] ** 2 * BJ / (BI + BJ), BJ * G[sub] ** 2) + ri1[sub] *= np.pi / G[sub] ** 2 + + return ri1 + + +def RI2(BI, BJ, G, U): + # "ZWEITES INTEGRAL FUER DIE ABSORPTIONSPOTENTIALE" + U2 = U**2 + U22 = 0.5 * U2 + G2 = G**2 + BIUH = BI + 0.5 * U2 + BJUH = BJ + 0.5 * U2 + BIU = BI + U2 + BJU = BJ + U2 + + # "IST DIE ASYMPTOTISCHE ENTWICKLUNG ANWENDBAR?"" + EPS = np.max([BI, BJ, U2]) + EPS = EPS * G2 + + ri2 = np.zeros_like(G) + + sub = EPS <= 0.1 + ri2[sub] = (BI + U2) * np.log((BI + BJ + U2) / (BI + U2)) + BJ * np.log( + (BI + BJ + U2) / (BJ + U2) + ) + if U2 > 0.0: + ri2[sub] += U2 * np.log(U2 / (BJ + U2)) + ri2[sub] *= np.pi + + if U2 > 0.0: + TEMP = 0.5 * U22 * U22 * np.log(BIU * BJU / (U2 * U2)) + else: + TEMP = 0.0 + TEMP = TEMP + 0.5 * BIUH * BIUH * np.log(BIU / (BIUH + BJUH)) + TEMP = TEMP + 0.5 * BJUH * BJUH * np.log(BJU / (BIUH + BJUH)) + TEMP = TEMP + 0.25 * BIU * BIU + 0.5 * BI * BI + TEMP = TEMP + 0.25 * BJU * BJU + 0.5 * BJ * BJ + TEMP = TEMP - 0.25 * (BIUH + BJUH) * (BIUH + BJUH) + TEMP = TEMP - 0.5 * ((BI * BIU - BJ * BJU) / (BIUH + BJUH)) ** 2 + TEMP = TEMP - U22 * U22 + ri2[sub] += np.pi * G2[sub] * TEMP + + sub = EPS > 0.1 + ri2[sub] = expi(-0.5 * U2 * G2[sub] * BIUH / BIU) + expi(-0.5 * U2 * G2[sub] * BJUH / BJU) + ri2[sub] -= expi(-BIUH * BJUH * G2[sub] / (BIUH + BJUH)) + expi(-0.25 * U2 * G2[sub]) + ri2[sub] *= 2.0 + X1 = 0.5 * U2 * G2[sub] + X2 = 0.25 * U2 * G2[sub] + X3 = 0.25 * U2 * U2 * G2[sub] / BIU + ri2[sub] += RIH1(X1, X2, X3) + + X1 = 0.5 * U2 * G2[sub] + X2 = 0.25 * U2 * G2[sub] + X3 = 0.25 * U2 * U2 * G2[sub] / BJU + ri2[sub] += RIH1(X1, X2, X3) + + X1 = BIUH * G2[sub] + X2 = BIUH * BIUH * G2[sub] / (BIUH + BJUH) + X3 = BIUH * BIUH * G2[sub] / BIU + ri2[sub] += RIH1(X1, X2, X3) + + X1 = BJUH * G2[sub] + X2 = BJUH * BJUH * G2[sub] / (BIUH + BJUH) + X3 = BJUH * BJUH * G2[sub] / BJU + ri2[sub] += RIH1(X1, X2, X3) + + ri2[sub] *= np.pi / G2[sub] + + return ri2 + + +def RIH1(X1, X2, X3): + # "WERTET DEN AUSDRUCK EXP(-X1) * ( EI(X2)-EI(X3) ) AUS" + rih1 = np.zeros(X1.shape) + + sub = np.logical_and(X2 <= 20.0, X3 <= 20.0) + rih1[sub] = np.exp(-X1[sub]) * (expi(X2[sub]) - expi(X3[sub])) + + sub = np.logical_and(X2 > 20.0, X3 <= 20.0) + rih1[sub] = np.exp(X2[sub] - X1[sub]) * RIH2(X2[sub]) / X2[sub] - np.exp(-X1[sub]) * expi( + X3[sub] + ) + + sub = np.logical_and(X2 <= 20.0, X3 > 20.0) + rih1[sub] = ( + np.exp(-X1[sub]) * expi(X2[sub]) - np.exp(X3[sub] - X1[sub]) * RIH2(X3[sub]) / X3[sub] + ) + + sub = np.logical_and(X2 > 20.0, X3 > 20.0) + rih1[sub] = ( + np.exp(X2[sub] - X1[sub]) * RIH2(X2[sub]) / X2[sub] + - np.exp(X3[sub] - X1[sub]) * RIH2(X3[sub]) / X3[sub] + ) + + return rih1 + + +def RIH2(X): + """ + WERTET X*EXP(-X)*EI(X) AUS FUER GROSSE X + DURCH INTERPOLATION DER TABELLE ... AUS ABRAMOWITZ + """ + idx = np.floor(200.0 / X).astype("int") + + sig = RIH2_tabulated_data[idx] + 200.0 * ( + RIH2_tabulated_data[idx + 1] - RIH2_tabulated_data[idx] + ) * ((1.0 / X) - 0.5e-3 * idx) + + return sig + + +################## +# TABULATED DATA # +################## + +# fmt:off + +RIH2_tabulated_data = np.array([1.000000,1.005051,1.010206,1.015472,1.020852, + 1.026355,1.031985,1.037751,1.043662,1.049726, + 1.055956,1.062364,1.068965,1.075780,1.082830, + 1.090140,1.097737,1.105647,1.113894,1.122497, + 1.131470]) + + +WK_A_param = np.array([ + 0.00427, 0.00957, 0.00802, 0.00209, + 0.01217, 0.02616,-0.00884, 0.01841, + 0.00251, 0.03576, 0.00988, 0.02370, + 0.01596, 0.02959, 0.04024, 0.01001, + 0.03652, 0.01140, 0.05677, 0.01506, + 0.04102, 0.04911, 0.05296, 0.00061, + 0.04123, 0.05740, 0.06529, 0.00373, + 0.03547, 0.03133, 0.10865, 0.01615, + 0.03957, 0.07225, 0.09581, 0.00792, + 0.02597, 0.02197, 0.13762, 0.05394, + 0.03283, 0.08858, 0.11688, 0.02516, + 0.03833, 0.17124, 0.03649, 0.04134, + 0.04388, 0.17743, 0.05047, 0.03957, + 0.03812, 0.17833, 0.06280, 0.05605, + 0.04166, 0.17817, 0.09479, 0.04463, + 0.04003, 0.18346, 0.12218, 0.03753, + 0.04245, 0.17645, 0.15814, 0.03011, + 0.05011, 0.16667, 0.17074, 0.04358, + 0.04058, 0.17582, 0.20943, 0.02922, + 0.04001, 0.17416, 0.20986, 0.05497, + 0.09685, 0.14777, 0.20981, 0.04852, + 0.06667, 0.17356, 0.22710, 0.05957, + 0.05118, 0.16791, 0.26700, 0.06476, + 0.03204, 0.18460, 0.30764, 0.05052, + 0.03866, 0.17782, 0.31329, 0.06898, + 0.05455, 0.16660, 0.33208, 0.06947, + 0.05942, 0.17472, 0.34423, 0.06828, + 0.06049, 0.16600, 0.37302, 0.07109, + 0.08034, 0.15838, 0.40116, 0.05467, + 0.02948, 0.19200, 0.42222, 0.07480, + 0.16157, 0.32976, 0.18964, 0.06148, + 0.16184, 0.35705, 0.17618, 0.07133, + 0.06190, 0.18452, 0.41600, 0.12793, + 0.15913, 0.41583, 0.13385, 0.10549, + 0.16514, 0.41202, 0.12900, 0.13209, + 0.15798, 0.41181, 0.14254, 0.14987, + 0.16535, 0.44674, 0.24245, 0.03161, + 0.16039, 0.44470, 0.24661, 0.05840, + 0.16619, 0.44376, 0.25613, 0.06797, + 0.16794, 0.44505, 0.27188, 0.07313, + 0.16552, 0.45008, 0.30474, 0.06161, + 0.17327, 0.44679, 0.32441, 0.06143, + 0.16424, 0.45046, 0.33749, 0.07766, + 0.18750, 0.44919, 0.36323, 0.05388, + 0.16081, 0.45211, 0.40343, 0.06140, + 0.16599, 0.43951, 0.41478, 0.08142, + 0.16547, 0.44658, 0.45401, 0.05959, + 0.17154, 0.43689, 0.46392, 0.07725, + 0.15752, 0.44821, 0.48186, 0.08596, + 0.15732, 0.44563, 0.48507, 0.10948, + 0.16971, 0.42742, 0.48779, 0.13653, + 0.14927, 0.43729, 0.49444, 0.16440, + 0.18053, 0.44724, 0.48163, 0.15995, + 0.13141, 0.43855, 0.50035, 0.22299, + 0.31397, 0.55648, 0.39828, 0.04852, + 0.32756, 0.53927, 0.39830, 0.07607, + 0.30887, 0.53804, 0.42265, 0.09559, + 0.28398, 0.53568, 0.46662, 0.10282, + 0.35160, 0.56889, 0.42010, 0.07246, + 0.33810, 0.58035, 0.44442, 0.07413, + 0.35449, 0.59626, 0.43868, 0.07152, + 0.35559, 0.60598, 0.45165, 0.07168, + 0.38379, 0.64088, 0.41710, 0.06708, + 0.40352, 0.64303, 0.40488, 0.08137, + 0.36838, 0.64761, 0.47222, 0.06854, + 0.38514, 0.68422, 0.44359, 0.06775, + 0.37280, 0.67528, 0.47337, 0.08320, + 0.39335, 0.70093, 0.46774, 0.06658, + 0.40587, 0.71223, 0.46598, 0.06847, + 0.39728, 0.73368, 0.47795, 0.06759, + 0.40697, 0.73576, 0.47481, 0.08291, + 0.40122, 0.78861, 0.44658, 0.08799, + 0.41127, 0.76965, 0.46563, 0.10180, + 0.39978, 0.77171, 0.48541, 0.11540, + 0.39130, 0.80752, 0.48702, 0.11041, + 0.40436, 0.80701, 0.48445, 0.12438, + 0.38816, 0.80163, 0.51922, 0.13514, + 0.39551, 0.80409, 0.53365, 0.13485, + 0.40850, 0.83052, 0.53325, 0.11978, + 0.40092, 0.85415, 0.53346, 0.12747, + 0.41872, 0.88168, 0.54551, 0.09404, + 0.43358, 0.88007, 0.52966, 0.12059, + 0.40858, 0.87837, 0.56392, 0.13698, + 0.41637, 0.85094, 0.57749, 0.16700, + 0.38951, 0.83297, 0.60557, 0.20770, + 0.41677, 0.88094, 0.55170, 0.21029, + 0.50089, 1.00860, 0.51420, 0.05996, + 0.47470, 0.99363, 0.54721, 0.09206, + 0.47810, 0.98385, 0.54905, 0.12055, + 0.47903, 0.97455, 0.55883, 0.14309, + 0.48351, 0.98292, 0.58877, 0.12425, + 0.48664, 0.98057, 0.61483, 0.12136, + 0.46078, 0.97139, 0.66506, 0.13012, + 0.49148, 0.98583, 0.67674, 0.09725, + 0.50865, 0.98574, 0.68109, 0.09977, + 0.46259, 0.97882, 0.73056, 0.12723, + 0.46221, 0.95749, 0.76259, 0.14086, + 0.48500, 0.95602, 0.77234, 0.13374, + ]).reshape(98,4) + + +WK_B_param = np.array([ + 4.17218, 16.05892, 26.78365, 69.45643, + 1.83008, 7.20225, 16.13585, 18.75551, + 0.02620, 2.00907, 10.80597,130.49226, + 0.38968, 1.99268, 46.86913,108.84167, + 0.50627, 3.68297, 27.90586, 74.98296, + 0.41335, 10.98289, 34.80286,177.19113, + 0.29792, 7.84094, 22.58809, 72.59254, + 0.17964, 2.60856, 11.79972, 38.02912, + 0.16403, 3.96612, 12.43903, 40.05053, + 0.09101, 0.41253, 5.02463, 17.52954, + 0.06008, 2.07182, 7.64444,146.00952, + 0.07424, 2.87177, 18.06729, 97.00854, + 0.09086, 2.53252, 30.43883, 98.26737, + 0.05396, 1.86461, 22.54263, 72.43144, + 0.05564, 1.62500, 24.45354, 64.38264, + 0.05214, 1.40793, 23.35691, 53.59676, + 0.04643, 1.15677, 19.34091, 52.88785, + 0.07991, 1.01436, 15.67109, 39.60819, + 0.03352, 0.82984, 14.13679,200.97722, + 0.02289, 0.71288, 11.18914,135.02390, + 0.12527, 1.34248, 12.43524,131.71112, + 0.05198, 0.86467, 10.59984,103.56776, + 0.03786, 0.57160, 8.30305, 91.78068, + 0.00240, 0.44931, 7.92251, 86.64058, + 0.01836, 0.41203, 6.73736, 76.30466, + 0.03947, 0.43294, 6.26864, 71.29470, + 0.03962, 0.43253, 6.05175, 68.72437, + 0.03558, 0.39976, 5.36660, 62.46894, + 0.05475, 0.45736, 5.38252, 60.43276, + 0.00137, 0.26535, 4.48040, 54.26088, + 0.10455, 2.18391, 9.04125, 75.16958, + 0.09890, 2.06856, 9.89926, 68.13783, + 0.01642, 0.32542, 3.51888, 44.50604, + 0.07669, 1.89297, 11.31554, 46.32082, + 0.08199, 1.76568, 9.87254, 38.10640, + 0.06939, 1.53446, 8.98025, 33.04365, + 0.07044, 1.59236, 17.53592,215.26198, + 0.06199, 1.41265, 14.33812,152.80257, + 0.06364, 1.34205, 13.66551,125.72522, + 0.06565, 1.25292, 13.09355,109.50252, + 0.05921, 1.15624, 13.24924, 98.69958, + 0.06162, 1.11236, 12.76149, 90.92026, + 0.05081, 0.99771, 11.28925, 84.28943, + 0.05120, 1.08672, 12.23172, 85.27316, + 0.04662, 0.85252, 10.51121, 74.53949, + 0.04933, 0.79381, 9.30944, 41.17414, + 0.04481, 0.75608, 9.34354, 67.91975, + 0.04867, 0.71518, 8.40595, 64.24400, + 0.03672, 0.64379, 7.83687, 73.37281, + 0.03308, 0.60931, 7.04977, 64.83582, + 0.04023, 0.58192, 6.29247, 55.57061, + 0.02842, 0.50687, 5.60835, 48.28004, + 0.03830, 0.58340, 6.47550, 47.08820, + 0.02097, 0.41007, 4.52105, 37.18178, + 0.07813, 1.45053, 15.05933,199.48830, + 0.08444, 1.40227, 13.12939,160.56676, + 0.07206, 1.19585, 11.55866,127.31371, + 0.05717, 0.98756, 9.95556,117.31874, + 0.08249, 1.43427, 12.37363,150.55968, + 0.07081, 1.31033, 11.44403,144.17706, + 0.07442, 1.38680, 11.54391,143.72185, + 0.07155, 1.34703, 11.00432,140.09138, + 0.07794, 1.55042, 11.89283,142.79585, + 0.08508, 1.60712, 11.45367,116.64063, + 0.06520, 1.32571, 10.16884,134.69034, + 0.06850, 1.43566, 10.57719,131.88972, + 0.06264, 1.26756, 9.46411,107.50194, + 0.06750, 1.35829, 9.76480,127.40374, + 0.06958, 1.38750, 9.41888,122.10940, + 0.06574, 1.31578, 9.13448,120.98209, + 0.06517, 1.29452, 8.67569,100.34878, + 0.06213, 1.30860, 9.18871, 91.20213, + 0.06292, 1.23499, 8.42904, 77.59815, + 0.05693, 1.15762, 7.83077, 67.14066, + 0.05145, 1.11240, 8.33441, 65.71782, + 0.05573, 1.11159, 8.00221, 57.35021, + 0.04855, 0.99356, 7.38693, 51.75829, + 0.04981, 0.97669, 7.38024, 44.52068, + 0.05151, 1.00803, 8.03707, 45.01758, + 0.04693, 0.98398, 7.83562, 46.51474, + 0.05161, 1.02127, 9.18455, 64.88177, + 0.05154, 1.03252, 8.49678, 58.79463, + 0.04200, 0.90939, 7.71158, 57.79178, + 0.04661, 0.87289, 6.84038, 51.36000, + 0.04168, 0.73697, 5.86112, 43.78613, + 0.04488, 0.83871, 6.44020, 43.51940, + 0.05786, 1.20028, 13.85073,172.15909, + 0.05239, 1.03225, 11.49796,143.12303, + 0.05167, 0.98867, 10.52682,112.18267, + 0.04931, 0.95698, 9.61135, 95.44649, + 0.04748, 0.93369, 9.89867,102.06961, + 0.04660, 0.89912, 9.69785,100.23434, + 0.04323, 0.78798, 8.71624, 92.30811, + 0.04641, 0.85867, 9.51157,111.02754, + 0.04918, 0.87026, 9.41105,104.98576, + 0.03904, 0.72797, 8.00506, 86.41747, + 0.03969, 0.68167, 7.29607, 75.72682, + 0.04291, 0.69956, 7.38554, 77.18528, + ]).reshape(98,4) diff --git a/tests/core/test_clustering.py b/tests/core/test_clustering.py new file mode 100644 index 000000000..d2bfa6ca2 --- /dev/null +++ b/tests/core/test_clustering.py @@ -0,0 +1,65 @@ +"""Torch DBSCAN against unambiguous synthetic ground truth.""" + +import numpy as np + +from quantem.core.utils.clustering import cluster_vector, dbscan, filter_rows + + +def test_dbscan_blobs(): + rng = np.random.default_rng(0) + centers = np.array([[0, 0], [10, 0], [0, 10], [30, 30]], dtype=float) + pts = np.concatenate( + [c + rng.normal(0, 0.5, (200, 2)) for c in centers] + + [rng.uniform(-5, 40, (40, 2))] # sparse background + ) + labels = dbscan(pts, eps=1.5, min_samples=10) + # four clusters, sorted by size; blob members agree internally + assert labels.max() + 1 == 4 + for k in range(4): + blob = labels[200 * k : 200 * (k + 1)] + vals, cnt = np.unique(blob[blob >= 0], return_counts=True) + assert cnt.max() > 190 # one dominant label per blob + assert (labels[800:] == -1).mean() > 0.8 # background mostly noise + + +def test_cluster_vector_and_filter(): + from quantem.core.datastructures.vector import Vector + + rng = np.random.default_rng(1) + R, C = 4, 5 + nested = [] + for r in range(R): + row = [] + for c in range(C): + # one tight cluster at (1,2) plus background + n = 30 if (r, c) == (1, 2) else 3 + q = ( + rng.normal(0, 0.01, (n, 2)) + [0.5, -0.3] + if (r, c) == (1, 2) + else rng.uniform(-1, 1, (n, 2)) + ) + inten = rng.uniform(1, 2, (n, 1)) + row.append(np.concatenate([q, inten], axis=1)) + nested.append(row) + vec = Vector.from_data(nested, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + labeled, labels = cluster_vector(vec, fields=("qx", "qy"), eps=0.05, min_samples=10) + assert "cluster" in labeled.fields + assert labels.max() == 0 # exactly one cluster found + got = labeled[1, 2].numpy() + assert (got[:, -1] == 0).sum() >= 28 + + kept = filter_rows(vec, labels == 0) + assert kept.total_rows == int((labels == 0).sum()) + assert kept.shape[:2] == vec.shape[:2] + + +def test_filter_rows_checks_mask_length(): + import pytest + + from quantem.core.datastructures.vector import Vector + + nested = [[np.ones((2, 1)), np.ones((3, 1))]] + vec = Vector.from_data(nested, fields=["intensity"], units=["counts"], name="t") + assert filter_rows(vec, [True, False, True, True, False]).total_rows == 3 + with pytest.raises(ValueError, match="rows"): + filter_rows(vec, [True, False]) diff --git a/tests/core/test_file_readers.py b/tests/core/test_file_readers.py new file mode 100644 index 000000000..3d2c5123e --- /dev/null +++ b/tests/core/test_file_readers.py @@ -0,0 +1,93 @@ +import numpy as np +import pytest + +from quantem.core.io import file_readers +from quantem.core.io.file_readers import _resolve_rsciio_plugin, read_4dstem + + +class TestResolveRsciioPlugin: + def test_plugin_name(self): + assert _resolve_rsciio_plugin("x.dat", "digitalmicrograph") == "rsciio.digitalmicrograph" + assert _resolve_rsciio_plugin("x.dat", "DigitalMicrograph") == "rsciio.digitalmicrograph" + + def test_unique_extension(self): + assert _resolve_rsciio_plugin("scan.dm4") == "rsciio.digitalmicrograph" + assert _resolve_rsciio_plugin("scan.DM3") == "rsciio.digitalmicrograph" + assert _resolve_rsciio_plugin("scan.mib") == "rsciio.quantumdetector" + assert _resolve_rsciio_plugin("x.dat", file_type="mib") == "rsciio.quantumdetector" + + def test_extension_equal_to_plugin_name_wins(self): + # "mrc" is listed by both the MRC and MRCZ plugins + assert _resolve_rsciio_plugin("stack.mrc") == "rsciio.mrc" + + def test_ambiguous_extension_raises(self): + with pytest.raises(ValueError, match="file_type") as err: + _resolve_rsciio_plugin("data.h5") + assert "arina" in str(err.value) + + def test_unknown_or_missing_type_raises(self): + with pytest.raises(ValueError, match="No RosettaSciIO reader"): + _resolve_rsciio_plugin("data.notaformat") + with pytest.raises(ValueError, match="Cannot infer"): + _resolve_rsciio_plugin("data_without_extension") + + +def _fake_reader(monkeypatch, entries): + def file_reader(path, **kwargs): + return entries + + monkeypatch.setattr( + file_readers, "_rsciio_reader", lambda path, file_type=None: ("rsciio.fake", file_reader) + ) + + +def _axis(scale, offset, units, name): + return {"scale": scale, "offset": offset, "units": units, "name": name} + + +@pytest.mark.parametrize("scan_axis", [0, 1]) +def test_read_4dstem_reshapes_3d_stack(monkeypatch, scan_axis): + n_frames, ny, nx, scan_length = 12, 5, 7, 4 + frames = np.arange(n_frames * ny * nx, dtype=np.float32).reshape(n_frames, ny, nx) + det_y = _axis(0.1, -0.25, "1/nm", "ky") + det_x = _axis(0.2, -0.7, "1/nm", "kx") + scan = _axis(1.0, 0.0, "1", "frame") + if scan_axis == 0: + data, axes = frames, [scan, det_y, det_x] + else: + data, axes = np.moveaxis(frames, 0, 1), [det_y, scan, det_x] + _fake_reader(monkeypatch, [{"data": data, "axes": axes}]) + + ds = read_4dstem("stack.fake", scan_length=scan_length, scan_axis=scan_axis) + + assert ds.shape == (n_frames // scan_length, scan_length, ny, nx) + assert np.array_equal(ds.array[1, 2], frames[1 * scan_length + 2]) + assert np.allclose(ds.sampling, [1.0, 1.0, 0.1, 0.2]) + assert np.allclose(ds.origin, [0.0, 0.0, -0.25, -0.7]) + assert list(ds.units[2:]) == ["1/nm", "1/nm"] + + +def test_read_4dstem_transpose_scan_axes(monkeypatch): + frames = np.random.default_rng(0).random((6, 3, 3)) + axes = [_axis(1.0, 0.0, "1", "f"), _axis(1.0, 0.0, "1", "y"), _axis(1.0, 0.0, "1", "x")] + _fake_reader(monkeypatch, [{"data": frames, "axes": axes}]) + ds = read_4dstem("stack.fake", scan_length=3, transpose_scan_axes=True) + assert ds.shape == (3, 2, 3, 3) + assert np.array_equal(ds.array[2, 1], frames[1 * 3 + 2]) + + +def test_read_4dstem_3d_needs_scan_length(monkeypatch): + axes = [_axis(1.0, 0.0, "1", n) for n in "fyx"] + _fake_reader(monkeypatch, [{"data": np.zeros((4, 3, 3)), "axes": axes}]) + with pytest.raises(ValueError, match="scan_length"): + read_4dstem("stack.fake") + with pytest.raises(ValueError, match="divisible"): + read_4dstem("stack.fake", scan_length=3) + + +def test_read_4dstem_rejects_bad_scan_axis(monkeypatch): + axes = [_axis(1.0, 0.0, "1", n) for n in "fyx"] + _fake_reader(monkeypatch, [{"data": np.zeros((4, 3, 3)), "axes": axes}]) + for scan_axis in (2, -1): + with pytest.raises(ValueError, match="scan_axis must be 0 or 1"): + read_4dstem("stack.fake", scan_length=2, scan_axis=scan_axis) diff --git a/tests/core/test_polar4dstem.py b/tests/core/test_polar4dstem.py new file mode 100644 index 000000000..0c1311a6f --- /dev/null +++ b/tests/core/test_polar4dstem.py @@ -0,0 +1,53 @@ +import numpy as np +import pytest +import torch + +from quantem.core.datastructures import Dataset4dstem +from quantem.core.datastructures.polar4dstem import Polar4dstem + + +def _ring_dataset(radius=9.0, center=(16.0, 15.0), shape=(2, 3, 33, 33)): + rr, cc = np.mgrid[0 : shape[2], 0 : shape[3]].astype(float) + r = np.hypot(rr - center[0], cc - center[1]) + ring = np.exp(-0.5 * ((r - radius) / 1.0) ** 2) + array = np.broadcast_to(ring, shape).astype(np.float32).copy() + return array, Dataset4dstem.from_array( + array, sampling=[1.0, 1.0, 0.05, 0.05], units=["A", "A", "1/A", "1/A"] + ) + + +def test_polar_transform_ring_peaks_at_radius(): + _, ds = _ring_dataset() + polar = ds.polar_transform(origin_row=16.0, origin_col=15.0, num_annular_bins=36) + assert isinstance(polar, Polar4dstem) + assert polar.shape[:2] == (2, 3) + assert polar.n_phi == 36 + profile = polar.array.mean(axis=(0, 1, 2)) + assert int(np.argmax(profile)) == 9 # radial_step 1 px from radial_min 0 + # the ring is isotropic: every azimuthal bin peaks at the same radius + assert np.all(np.argmax(polar.array[0, 0], axis=1) == 9) + assert polar.sampling[3] == pytest.approx(0.05) + assert polar.sampling[2] == pytest.approx(10.0) + assert polar.units[2] == "deg" + assert polar.metadata["polar_origin_row"] == 16.0 + assert polar.n_r == len(profile) + + +def test_polar_transform_two_fold_and_radial_range(): + _, ds = _ring_dataset() + polar = ds.polar_transform( + 16.0, 15.0, num_annular_bins=18, radial_min=4.0, radial_max=12.0, + radial_step=0.5, two_fold_rotation_symmetry=True, + ) # fmt: skip + assert polar.n_r == 16 + assert polar.sampling[2] == pytest.approx(10.0) # 180 deg over 18 bins + profile = polar.array.mean(axis=(0, 1, 2)) + assert 4.0 + 0.5 * int(np.argmax(profile)) == pytest.approx(9.0) + + +def test_polar_transform_tensor_backed(): + array, ds = _ring_dataset() + ds_t = Dataset4dstem.from_tensor(torch.as_tensor(array)) + a = ds.polar_transform(16.0, 15.0, num_annular_bins=12).array + b = ds_t.polar_transform(16.0, 15.0, num_annular_bins=12).array + assert np.allclose(a, b) diff --git a/tests/core/test_serialize_bundle.py b/tests/core/test_serialize_bundle.py new file mode 100644 index 000000000..8f3af7e0f --- /dev/null +++ b/tests/core/test_serialize_bundle.py @@ -0,0 +1,74 @@ +import numpy as np +import pytest +import torch +from ase import Atoms + +from quantem.core.datastructures import Dataset2d +from quantem.core.io import load +from quantem.core.io.serialize import AutoSerialize, Bundle + + +class _Holder(AutoSerialize): + def __init__(self, **kwargs): + for k, v in kwargs.items(): + setattr(self, k, v) + + +def _partial_atoms() -> Atoms: + atoms = Atoms( + "NbVZr", + scaled_positions=[[0, 0, 0], [0.5, 0.5, 0.5], [0, 0, 0]], + cell=[3.2, 3.2, 3.2], + pbc=[True, True, False], + ) + atoms.set_array("occupancy", np.array([0.4, 1.0, 0.6])) + atoms.set_tags([1, 2, 1]) + atoms.set_masses([92.9, 50.9, 91.2]) + atoms.info["spacegroup_number"] = 229 + atoms.info["not_json"] = object() + return atoms + + +def test_ase_atoms_round_trip_keeps_occupancy(tmp_path): + atoms = _partial_atoms() + path = tmp_path / "atoms.zip" + _Holder(atoms=atoms).save(path, mode="o") + back = load(path).atoms + + assert isinstance(back, Atoms) + assert back.get_chemical_symbols() == ["Nb", "V", "Zr"] + assert np.allclose(back.get_positions(), atoms.get_positions()) + assert np.allclose(back.get_cell(), atoms.get_cell()) + assert np.array_equal(back.get_pbc(), atoms.get_pbc()) + assert np.allclose(back.arrays["occupancy"], [0.4, 1.0, 0.6]) + assert np.array_equal(back.get_tags(), [1, 2, 1]) + assert np.allclose(back.get_masses(), [92.9, 50.9, 91.2]) + assert back.info == {"spacegroup_number": 229} + + +def test_torch_device_round_trip(tmp_path): + path = tmp_path / "device.zip" + _Holder(device=torch.device("cpu"), name="x").save(path, mode="o") + back = load(path) + assert isinstance(back.device, torch.device) + assert back.device == torch.device("cpu") + assert back.name == "x" + + +def test_bundle_round_trip(tmp_path): + image = Dataset2d.from_array(np.arange(12.0).reshape(3, 4), name="adf") + bundle = Bundle(adf=image, atoms=_partial_atoms(), note="IM689") + assert "adf: Dataset2d" in repr(bundle) + path = tmp_path / "bundle.zip" + bundle.save(path, mode="o") + back = load(path) + assert isinstance(back, Bundle) + assert np.array_equal(back.adf.array, image.array) + assert back.note == "IM689" + assert np.allclose(back.atoms.arrays["occupancy"], [0.4, 1.0, 0.6]) + + +@pytest.mark.parametrize("name", ["save", "print_tree", "_recursive_save"]) +def test_bundle_rejects_reserved_names(name): + with pytest.raises(ValueError, match="shadow"): + Bundle(**{name: 1}) diff --git a/tests/diffraction/test_bloch.py b/tests/diffraction/test_bloch.py new file mode 100644 index 000000000..5650f0375 --- /dev/null +++ b/tests/diffraction/test_bloch.py @@ -0,0 +1,1256 @@ +"""Tests for quantem.diffraction.bloch.""" + +from functools import lru_cache + +import numpy as np +import pytest +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction import bloch +from quantem.diffraction.bloch import dynamical_pattern, refine_thickness +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.phase import PhaseMap +from quantem.diffraction.rotations import quat_from_zone_axis + + +@pytest.fixture(scope="module") +def ti_beta(): + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), name="Ti beta") + # 2x coverage so all coupling vectors g - h have structure factors + xtl.calculate_structure_factors(k_max=3.0, tol_structure_factor=1e-6) + return xtl + + +def test_flux_conservation(ti_beta): + q = quat_from_zone_axis(torch.tensor([0.0, 1.0, 1.0], dtype=torch.float64)) + p = dynamical_pattern( + ti_beta, q, np.arange(50, 1500, 50.0), energy_ev=200e3, sg_max=0.08, k_max=1.5 + ) + total = p["intensity"].sum(dim=1) + assert float(total.max()) <= 1.0 + 1e-6 + + +def test_thin_limit_matches_kinematical(ti_beta): + q = quat_from_zone_axis(torch.tensor([0.0, 1.0, 1.0], dtype=torch.float64)) + p = dynamical_pattern(ti_beta, q, 25.0, energy_ev=200e3, sg_max=0.08, k_max=1.5) + kin = ti_beta.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + top_dyn = set(map(tuple, p["hkl"][p["intensity"][0].argsort(descending=True)[:4]].tolist())) + top_kin = set(map(tuple, kin["hkl"][kin["intensity"].argsort(descending=True)[:4]].tolist())) + assert top_dyn == top_kin + + +def test_thickness_recovery(ti_beta): + """Simulate dynamical peaks at a known thickness, recover it.""" + t_true = 600.0 + torch.manual_seed(0) + zones = torch.tensor([[0.1, 0.9, 1.0], [0.3, 0.5, 1.0], [0.05, 1.0, 1.1]], dtype=torch.float64) + N = zones.shape[0] + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + q_true = quat_from_zone_axis(zones) + for i in range(N): + p = dynamical_pattern(ti_beta, q_true[i], t_true, energy_ev=200e3, sg_max=0.06, k_max=1.5) + keep = p["intensity"][0] > 1e-4 + peaks[0, i] = np.stack( + [ + p["qx"][keep].numpy(), + p["qy"][keep].numpy(), + p["intensity"][0][keep].numpy(), + ], + axis=1, + ) + + om = OrientationMap.from_vectors(peaks, ti_beta, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0) + om.match_orientations(progress_bar=False) + # thickness oscillations are sensitive to ~1 degree tilt errors, beyond + # what kinematical matching provides for dynamical patterns; test the + # thickness scan itself with the true orientations (the joint tilt and + # thickness search is refine_dynamical, tested below) + om.quats[0, :, 0] = q_true + + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(max_patterns=1, progress_bar=False) + res = refine_thickness( + pm, + thicknesses_A=np.arange(100, 1200, 50.0), + sg_max=0.06, + progress_bar=False, + ) + t_fit = res["thickness"][0].numpy() + assert (np.abs(t_fit - t_true) <= 50.0).all() + + +# ---------------------------------------------------------------------- +# CBED / LACBED / Kossel / Kossel reference pattern +# ---------------------------------------------------------------------- + + +@lru_cache(maxsize=None) +def _si(absorptive: bool) -> Crystal: + """Silicon, built once per session; the tests only read it (the Bloch + code caches lattice data on it, which every test shares safely).""" + si = Crystal.from_ase(bulk("Si", "diamond", a=5.431, cubic=True), name="Si", verbose=False) + si.calculate_structure_factors(k_max=3.0) + if absorptive: + si.calculate_dynamical_structure_factors(energy_ev=200e3, k_max=3.0) + return si + + +def _zone_110() -> torch.Tensor: + return quat_from_zone_axis(torch.tensor([[1.0, 1.0, 0.0]]) / np.sqrt(2))[0] + + +def test_zero_tilt_matches_dynamical_pattern(): + si = _si(absorptive=True) + q = _zone_110() + t = torch.tensor([800.0]) + tilts = torch.zeros((1, 2), dtype=torch.float64) + inten, g_xy, hkl = bloch._cbed_amplitudes(si, q, tilts, t, 200e3, sg_max=0.08, k_max=1.3) + ref = bloch.dynamical_pattern(si, q, t, energy_ev=200e3, sg_max=0.08, k_max=1.3) + # same beams (000 first in CBED) and identical intensities + assert hkl.shape[0] == ref["hkl"].shape[0] + 1 + assert torch.allclose(inten[0, 0, 1:], ref["intensity"][0], rtol=1e-10, atol=1e-12) + + +def test_unitarity_without_absorption(): + si = _si(absorptive=False) + q = _zone_110() + tilts = bloch.tilt_grid(2.0, 200e3, n_rings=2) + inten, _, _ = bloch._cbed_amplitudes( + si, q, tilts, torch.tensor([500.0, 1500.0]), 200e3, sg_max=0.08, k_max=1.3 + ) + # Hermitian structure matrix: evolution is unitary in the beam space + total = inten.sum(dim=-1) + assert torch.allclose(total, torch.ones_like(total), atol=1e-8) + + +def test_lacbed_centrosymmetric_disk(): + si = _si(absorptive=True) + q = _zone_110() + res = bloch.calculate_lacbed( + si, + q, + 800.0, + hkl=(0, 0, 0), + energy_ev=200e3, + semiconv_mrad=6.0, + n_pixels=24, + sg_max=0.08, + k_max=1.3, + ) + disk = res["disk"] + # Si is centrosymmetric: the bright field rocking surface at a zone axis + # is inversion symmetric, I(t) = I(-t). Small residuals come from beam + # truncation at the s_g cutoff (|s_g| differs slightly for +g and -g), + # so the tolerance is physical rather than numerical. + flipped = disk[::-1, ::-1] + m = np.isfinite(disk) & np.isfinite(flipped) + assert m.sum() > 100 + assert np.allclose(disk[m], flipped[m], rtol=2e-3, atol=1e-5) + + +def test_cbed_library_common_grid(): + si = _si(absorptive=True) + q0 = _zone_110() + quats = torch.stack([q0, q0]) + lib = bloch.calculate_cbed_library( + si, + quats, + thickness_A=600.0, + energy_ev=200e3, + semiconv_mrad=3.0, + k_max=1.0, + progress_bar=False, + ) + assert lib["patterns"].shape[0] == 2 + assert lib["patterns"].shape[1] == lib["patterns"].shape[2] + assert np.allclose(lib["patterns"][0], lib["patterns"][1]) + assert lib["patterns"][0].max() > 0 + + +def test_cbed_geometry_and_normalization(): + """Without absorption the evolution is unitary, so a detector holding + every disk carries the full intensity (the pattern is the tilt + average). In a thin crystal the direct beam disk is uniform and centered + on the center pixel; diffracted disks sit at their g (rows qx, columns + qy). Checked on and off a zone axis.""" + si = _si(absorptive=False) + for q in (_zone_110(), _tilted_110()): + res = bloch.calculate_cbed( + si, + q, + [10.0, 500.0], + energy_ev=200e3, + semiconv_mrad=2.0, + n_rings=3, + sg_max=0.06, + k_max=0.8, + ) + thin, thick = res["pattern"] + H = thin.shape[0] + c = (H - 1) // 2 + assert np.allclose(res["pattern"].sum(axis=(1, 2)), 1.0, rtol=1e-9) + assert np.allclose(res["g_xy"][0], 0.0) + r_px = res["disk_radius"] / res["sampling"] + yy, xx = np.mgrid[0:H, 0:H] + w = thin * (np.hypot(yy - c, xx - c) <= r_px + 1.5) + assert w.sum() > 0.95 + assert abs((w * yy).sum() / w.sum() - c) < 0.05 + assert abs((w * xx).sum() / w.sum() - c) < 0.05 + # every disk sits at its g: the pattern is nonzero only inside the + # disks centered at (row, col) = (qx, qy) / sampling + center + centers = res["g_xy"] / res["sampling"] + c + d = np.hypot(yy[..., None] - centers[:, 0], xx[..., None] - centers[:, 1]).min(-1) + assert np.all(thick[d > r_px + 1.5] == 0) + + +def test_cbed_detector_crop_drops_outside_samples(): + """A smaller detector is a crop of the larger one: samples beyond its + edge are dropped, not piled onto the border pixels.""" + si = _si(absorptive=True) + q = _tilted_110() + kw = dict(energy_ev=200e3, semiconv_mrad=2.0, n_rings=3, sg_max=0.06, k_max=0.8) + full = bloch.calculate_cbed(si, q, 500.0, **kw) + s = full["sampling"] + small = bloch.calculate_cbed(si, q, 500.0, pixel_size=s, q_max_plot=0.3, **kw) + h_full = (full["pattern"].shape[0] - 1) // 2 + h_small = (small["pattern"].shape[0] - 1) // 2 + crop = full["pattern"][ + h_full - h_small : h_full + h_small + 1, h_full - h_small : h_full + h_small + 1 + ] + assert small["pattern"].sum() < 0.99 * full["pattern"].sum() # disks were cut + assert np.allclose(small["pattern"], crop, rtol=1e-12, atol=1e-15) + + +def test_kossel_bright_field_matches_lacbed(): + si = _si(absorptive=True) + q = _zone_110() + kw = dict(energy_ev=200e3, semiconv_mrad=15.0, sg_max=0.06, k_max=0.8) + kos = bloch.calculate_kossel(si, q, 900.0, n_pixels=32, progress_bar=False, **kw) + lac = bloch.calculate_lacbed(si, q, 900.0, hkl=(0, 0, 0), n_pixels=32, **kw) + a, b = kos["bright_field"], lac["disk"] + m = np.isfinite(a) & np.isfinite(b) + assert m.sum() > 300 + assert np.allclose(a[m], b[m], rtol=1e-10, atol=1e-12) + + # the summed pattern includes the direct beam plus every diffracted + # cone, so inside the aperture it can only exceed the bright field + pat = kos["pattern"] + assert (pat[m] >= a[m] - 1e-9).mean() > 0.99 + assert np.all(pat[~np.isfinite(a)] == 0) + + +def test_kossel_pattern_orientation_matches_bright_field(): + """The full Kossel pattern is stored on the bright field's axes, (row, + col) = (theta_y, theta_x): off a zone axis, where no symmetry hides a + transposition, the deficiency lines of the direct beam make the two + correlate, and the transposed pattern does not.""" + si = _si(absorptive=True) + kos = bloch.calculate_kossel( + si, + _tilted_110(), + 800.0, + energy_ev=200e3, + semiconv_mrad=15.0, + n_pixels=32, + sg_max=0.06, + k_max=0.7, + progress_bar=False, + ) + a, p = kos["bright_field"], kos["pattern"] + m = np.isfinite(a) + cc = np.corrcoef(a[m], p[m])[0, 1] + cc_t = np.corrcoef(a[m], p.T[m])[0, 1] + assert cc > 0.4 + assert cc - cc_t > 0.3 + assert np.isclose(kos["mrad_per_pixel"], 2 * 15.0 / 31) + + +# The reference-pattern tests share one coarse reference: 3 mrad sampling and +# beams to 0.7 1/A cost ~3 s, against ~30 s for 2 mrad and 1.0 1/A, and keep +# the comparisons meaningful (the lookup and the line model are compared at +# the reference's resolution). +REF_KW = dict(energy_ev=200e3, sg_max=0.06, k_max=0.7) +REF_STEP_MRAD = 3.0 + + +def _tilted_110(): + from quantem.diffraction.rotations import qmult, quat_from_axis_angle + + tilt = quat_from_axis_angle( + torch.tensor([1.0, 0.3, 0.0], dtype=torch.float64) / np.hypot(1, 0.3), + torch.tensor(np.deg2rad(5.0), dtype=torch.float64), + ) + return qmult(tilt, _zone_110()) + + +@pytest.fixture(scope="module") +def si_reference(): + return bloch.calculate_kossel_reference( + _si(absorptive=True), + [800.0], + angle_step_mrad=REF_STEP_MRAD, + progress_bar=False, + **REF_KW, + ) + + +@pytest.fixture(scope="module") +def si_lines(): + return bloch.kossel_lines(_si(absorptive=True), 800.0, energy_ev=200e3, k_max=REF_KW["k_max"]) + + +def _direct_bright_field(q, semiconv_mrad=25.0, n_pixels=48): + return bloch.calculate_kossel( + _si(absorptive=True), + q, + 800.0, + semiconv_mrad=semiconv_mrad, + n_pixels=n_pixels, + progress_bar=False, + **REF_KW, + )["bright_field"] + + +def _blurred_cc(a, b, n_pixels=48, semiconv_mrad=25.0): + """Correlation of two bright fields inside the aperture after blurring + both to the reference's resolution. Bilinear splatting onto the Lambert + grid and bilinear lookup are two triangle kernels of one grid step, + together a blur of standard deviation step / sqrt(3).""" + from scipy.ndimage import gaussian_filter + + px_mrad = 2 * semiconv_mrad / (n_pixels - 1) + sigma = REF_STEP_MRAD / np.sqrt(3) / px_mrad + m = np.isfinite(a) & np.isfinite(b) + assert m.sum() > 1000 + af = gaussian_filter(np.nan_to_num(a), sigma) + bf = gaussian_filter(np.nan_to_num(b), sigma) + return np.corrcoef(af[m], bf[m])[0, 1] + + +def test_reference_pattern_lookup(si_reference): + q = _zone_110() + assert si_reference["k_max"] == REF_KW["k_max"] + fast = bloch.kossel_from_reference(si_reference, q, semiconv_mrad=25.0, n_pixels=48) + assert np.isclose(fast["mrad_per_pixel"], 2 * 25.0 / 47) + assert _blurred_cc(fast["bright_field"], _direct_bright_field(q)) > 0.9 + + # off-zone orientation: catches in-plane sign errors that zone-axis + # symmetry hides (the reference stores the ANTI-propagation direction) + q2 = _tilted_110() + fast2 = bloch.kossel_from_reference(si_reference, q2, semiconv_mrad=25.0, n_pixels=48) + assert _blurred_cc(fast2["bright_field"], _direct_bright_field(q2)) > 0.9 + + +def test_kossel_polar_from_reference_matches_cartesian(si_reference): + """The polar lookup samples the same function: on a Cartesian grid with + pixels at the polar radii, the azimuth 0 and 90 degree rows coincide + with the center row and column.""" + q = _tilted_110() + n_r = 8 + pol = bloch.kossel_polar_from_reference( + si_reference, q, semiconv_mrad=20.0, n_radial=n_r, n_azimuthal=16 + ) + assert pol["polar"].shape == (16, n_r) + assert np.allclose(pol["radii_mrad"], 20.0 * np.arange(1, n_r + 1) / n_r) + cart = bloch.kossel_from_reference(si_reference, q, semiconv_mrad=20.0, n_pixels=2 * n_r + 1) + bf = cart["bright_field"] + # (row, col) = (theta_y, theta_x): azimuth 0 runs along +col, 90 along +row + assert np.allclose(pol["polar"][0], bf[n_r, n_r + 1 :]) + assert np.allclose(pol["polar"][4], bf[n_r + 1 :, n_r]) + + +def test_plot_kossel_reference(si_reference, si_lines, monkeypatch): + import matplotlib.pyplot as plt + + fig, ax = bloch.plot_kossel_reference(si_reference, _si(True), lines=si_lines, upsample=1) + assert len(ax.images) == 1 and len(ax.texts) > 0 + plt.close(fig) + + # a reference computed without a beam cutoff stores k_max=None; the + # default line set then falls back to 1.2 1/A instead of failing + seen = {} + + def fake_lines(crystal, thicknesses_A, energy_ev, k_max): + seen["k_max"] = k_max + return si_lines + + monkeypatch.setattr(bloch, "kossel_lines", fake_lines) + fig, ax = bloch.plot_kossel_reference({**si_reference, "k_max": None}, _si(True), upsample=1) + assert seen["k_max"] == 1.2 + plt.close(fig) + + +def test_kossel_lines_rejects_missing_cutoff_and_empty_set(): + si = _si(absorptive=True) + with pytest.raises(ValueError, match="k_max"): + bloch.kossel_lines(si, 800.0, energy_ev=200e3, k_max=None) + with pytest.raises(ValueError, match="min_depth"): + bloch.kossel_lines(si, 800.0, energy_ev=200e3, k_max=0.4, min_depth=2.0) + + +def test_kossel_lines_render_matches_direct(si_lines): + q = _tilted_110() + lines = si_lines + direct = _direct_bright_field(q, semiconv_mrad=40.0) + # every line is a band edge: the +g and -g cones of a row sit at + # +-theta_B, never on the zone plane + first = lines["line_order"] == 1 + assert torch.all(lines["line_u"][first] > 0) + assert torch.all(lines["line_u"][lines["line_order"] == -1] < 0) + r = bloch.render_kossel_lines(lines, q, semiconv_mrad=40.0, n_pixels=48) + a = r["bright_field"] + m = np.isfinite(a) & np.isfinite(direct) + assert m.sum() > 1000 + assert np.corrcoef(a[m], direct[m])[0, 1] > 0.98 + + # polar rendering samples the same function: its first ring must + # agree with the Cartesian pattern evaluated at those angles + pol = bloch.render_kossel_lines( + lines, q, semiconv_mrad=40.0, polar=True, n_radial=10, n_azimuthal=12 + )["polar"] + assert pol.shape == (12, 10) + assert np.all(np.isfinite(pol)) + assert pol.min() > 0 and pol.max() < 1.5 * float(lines["background"][0]) + + +def test_kossel_line_segments_on_cones(si_lines): + from quantem.diffraction.rotations import quat_to_matrix + + lines = si_lines + q = _tilted_110() + alpha = 40.0 + seg = bloch.kossel_line_segments(lines, q, semiconv_mrad=alpha) + n = seg["depth"].shape[0] + assert n >= 3 + R = quat_to_matrix(q).numpy() + g_c = lines["g_hat"].numpy() + hkl_row = lines["hkl_row"].numpy() + for k in range(n): + # row of this line from its hkl (an integer multiple of the row vector) + h = seg["hkl"][k] + ri = next(i for i in range(g_c.shape[0]) if np.all(np.cross(hkl_row[i], h) == 0)) + n_ord = int(np.round(np.dot(h, hkl_row[ri]) / np.dot(hkl_row[ri], hkl_row[ri]))) + u = n_ord * bloch.electron_wavelength_angstrom(200e3) * float(lines["g_len"][ri]) / 2 + g_lab = R @ g_c[ri] + for key in ("start_mrad", "stop_mrad"): + row, col = seg[key][k] * 1e-3 + assert np.isclose(np.hypot(row, col), alpha * 1e-3) + d_lab = np.array([-col, -row, np.sqrt(1 - row**2 - col**2)]) + assert abs(d_lab @ g_lab - u) < 1e-9 + # polar end points are the same points + phi, r = seg["start_polar"][k] + assert np.isclose(r, alpha) + assert np.allclose([alpha * np.sin(phi), alpha * np.cos(phi)], seg["start_mrad"][k]) + assert seg["width_mrad"][k] > 0 and 0 < seg["depth"][k] <= 1 + + +def test_overlay_kossel_segments_registration(si_lines): + """Segment end points lie on the aperture edge, which is the pixel + circle of radius (n - 1) / 2 about the center pixel (Cartesian) and the + last column (polar).""" + import matplotlib.pyplot as plt + + q = _tilted_110() + alpha, n = 40.0, 33 + seg = bloch.kossel_line_segments(si_lines, q, semiconv_mrad=alpha) + n_draw = int((seg["depth"] >= 0.01).sum()) + assert n_draw >= 3 + fig, ax = plt.subplots() + bloch.overlay_kossel_segments(ax, seg, alpha, n_pixels=n, min_depth=0.01) + assert len(ax.lines) == n_draw + c = (n - 1) / 2 + for ln in ax.lines: + x, y = ln.get_xdata(), ln.get_ydata() + assert np.allclose(np.hypot(np.asarray(x) - c, np.asarray(y) - c), c) + plt.close(fig) + + n_r, n_az = 16, 90 + fig, ax = plt.subplots() + bloch.overlay_kossel_segments( + ax, seg, alpha, polar=True, n_radial=n_r, n_azimuthal=n_az, min_depth=0.01 + ) + assert len(ax.lines) == n_draw + for ln in ax.lines: + cols = np.asarray(ln.get_xdata()) + # radius semiconv sits in the last column, n_radial - 1 + assert np.isclose(cols.max(), n_r - 1) + assert np.all(cols <= n_r - 1 + 1e-9) + plt.close(fig) + + +def test_reference_residual_hybrid(si_reference, si_lines): + q = _zone_110() + si = _si(absorptive=True) + reference = dict(si_reference) # the residual is added in place + bloch.kossel_reference_residual(reference, si_lines, si) + assert reference["residual"].shape == reference["lambert"].shape + assert np.all(np.isfinite(reference["residual"])) + assert "residual" not in si_reference + direct = _direct_bright_field(q) + plain = bloch.render_kossel_lines(si_lines, q, semiconv_mrad=25.0, n_pixels=48) + hybrid = bloch.render_kossel_lines( + si_lines, q, semiconv_mrad=25.0, n_pixels=48, reference=reference + ) + cc_plain = _blurred_cc(plain["bright_field"], direct) + cc_hybrid = _blurred_cc(hybrid["bright_field"], direct) + # on the zone axis the many-beam residual must improve the line model + assert cc_hybrid > cc_plain + 0.05 + assert cc_hybrid > 0.9 + + +def test_refine_dynamical_recovery(): + """Bragg-vector dynamical refinement: thickness, tilt, in-plane strain and + rotation recovered from noise-free patterns of a strained, tilted cell.""" + from quantem.core.datastructures.vector import Vector + from quantem.diffraction.orientation import OrientationMap + from quantem.diffraction.phase import PhaseMap + from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qmult, + qnormalize, + quat_from_axis_angle, + ) + + energy_ev = 200e3 + xtl = _si(absorptive=True) + torch.manual_seed(0) + rng = np.random.default_rng(0) + N = 6 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + t_true = torch.tensor([300.0, 450.0, 600.0, 300.0, 450.0, 600.0]) + A_true = torch.tensor([[1.010, 0.003], [0.003, 0.995]], dtype=torch.float64) + rot = np.deg2rad(0.3) + qz = torch.tensor([np.cos(rot / 2), 0.0, 0.0, np.sin(rot / 2)], dtype=torch.float64) + q_expect = torch.stack([qmult(qz, q_true[i]) for i in range(N)]) + deform3 = torch.eye(3, dtype=torch.float64) + deform3[:2, :2] = A_true + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + inten, g_xy, _ = bloch._cbed_amplitudes( + xtl, + q_expect[i], + torch.zeros((1, 2), dtype=torch.float64), + t_true[i : i + 1], + energy_ev, + 0.06, + 1.0, + progress_bar=False, + deform=deform3, + ) + inten_np = inten[0, 0, 1:].numpy() + keep = inten_np > 1e-3 * inten_np.max() + peaks[0, i] = np.column_stack([g_xy[1:].numpy()[keep], inten_np[keep]]) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=energy_ev) + om.build_plan(angle_step_zone_axis_deg=3.0, angle_step_in_plane_deg=5.0, verbose=False) + om.match_orientations(progress_bar=False) + # start 0.15 degrees off the truth about random in-plane axes, unstrained + phis = rng.uniform(0, 2 * np.pi, N) + om.quats[0, :, 0] = torch.stack( + [ + qmult( + quat_from_axis_angle( + torch.tensor([np.cos(p), np.sin(p), 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(0.15), dtype=torch.float64), + ), + q_true[i], + ) + for i, p in enumerate(phis) + ] + ) + om.corr[0, :, 0] = 1.0 + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(progress_bar=False) + res = bloch.refine_dynamical( + pm, + thicknesses_A=np.arange(150, 800, 25.0), + tilt_stages=((0.25, 0.025), (0.04, 0.005)), + power_intensity=0.5, + sg_max=0.06, + k_max=1.0, + progress_bar=False, + ) + # positions with too few beams cannot constrain a 2x2 deformation + valid = np.array([peaks[0, i].numpy().shape[0] >= 6 for i in range(N)]) + assert valid.sum() >= 4 + err = misorientation_angle_deg(q_expect, om.quats[0, :, 0], xtl.sym_quats).numpy()[valid] + t_err = np.abs(res["thickness"][0].numpy() - t_true.numpy())[valid] + A_err = (res["deformation"][0, valid, 0] - A_true[None]).abs().max() + assert float(A_err) < 1e-3 + assert np.median(err) < 0.03 + assert (t_err <= 25).sum() >= valid.sum() - 1 + # crystal-frame strain of an unstrained position is zero, of a strained + # one has the right magnitude + sc = bloch.strain_crystal_frame(res["deformation"][0, :, 0], om.quats[0, :, 0]) + eps = sc["eps_crystal"] + assert torch.allclose(eps, eps.transpose(-1, -2)) + assert float(eps.abs().max()) < 0.02 + + +def test_image_refinement_round_trip(): + """Rendered patterns with known disk shape: fit_disk_shape recovers the + radius and edge, and the image refinement keeps a correct thickness.""" + from types import SimpleNamespace + + from quantem.core.datastructures.vector import Vector + from quantem.diffraction.orientation import OrientationMap + from quantem.diffraction.phase import PhaseMap + from quantem.diffraction.rotations import qnormalize + + energy_ev = 200e3 + xtl = _si(absorptive=True) + torch.manual_seed(2) + rng = np.random.default_rng(2) + N = 4 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + t_true = torch.tensor([300.0, 450.0, 600.0, 400.0]) + shape = (64, 64) + pixel_size, rot, ellipse = 0.04, 15.0, (0.003, -0.002) + disk_r, edge = 3.0, 0.75 + origins = np.full((1, N, 2), 32.0) + imgs = np.zeros((1, N) + shape) + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + im, _, _ = bloch.render_pattern_image( + xtl, + q_true[i], + [float(t_true[i])], + energy_ev, + shape, + origins[0, i], + pixel_size, + rot, + ellipse, + None, + disk_r, + edge, + sg_max=0.06, + k_max=1.0, + ) + imgs[0, i] = rng.poisson(im[0, 0].numpy() * 1e5 + 20) + inten, g_xy, _ = bloch._cbed_amplitudes( + xtl, + q_true[i], + torch.zeros((1, 2), dtype=torch.float64), + t_true[i : i + 1], + energy_ev, + 0.06, + 1.0, + progress_bar=False, + fast_absorption=True, + ) + inten_np = inten[0, 0, 1:].numpy() + keep = inten_np > 1e-3 * inten_np.max() + peaks[0, i] = np.column_stack([g_xy[1:].numpy()[keep], inten_np[keep]]) + peaks.metadata["rotation_ccw_deg"] = rot + dataset = SimpleNamespace(array=imgs, shape=imgs.shape) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=energy_ev) + om.build_plan(angle_step_zone_axis_deg=3.0, angle_step_in_plane_deg=5.0, verbose=False) + om.match_orientations(progress_bar=False) + om.quats[0, :, 0] = q_true + om.corr[0, :, 0] = 1.0 + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(progress_bar=False) + res = bloch.refine_dynamical( + pm, + thicknesses_A=np.arange(200, 700, 50.0), + tilt_stages=((0.05, 0.05),), + power_intensity=0.5, + sg_max=0.06, + k_max=1.0, + progress_bar=False, + ) + valid = [i for i in range(N) if peaks[0, i].numpy().shape[0] >= 6] + assert len(valid) >= 2 + shape_fit = bloch.fit_disk_shape( + dataset, + pm, + res, + origins, + pixel_size, + rot, + ellipse, + positions=[(0, i) for i in valid], + radii_px=np.array([2.0, 2.5, 3.0, 3.5, 4.0]), + edges_px=np.array([0.5, 0.75, 1.0, 1.5]), + sg_max=0.06, + k_max=1.0, + progress_bar=False, + ) + assert shape_fit["disk_radius_px"] == disk_r + assert shape_fit["edge_px"] == edge + img = bloch.refine_dynamical_image( + dataset, + pm, + res, + origins, + pixel_size, + disk_r, + edge, + rot, + ellipse, + thickness_half_range_A=50, + thickness_step_A=25, + tilt_stage=(0.02, 0.01), + sg_max=0.06, + k_max=1.0, + progress_bar=False, + ) + t_err = np.abs(img["thickness"][0].numpy() - t_true.numpy())[valid] + assert (t_err <= 25).sum() >= len(valid) - 1 + assert np.all(np.isfinite(img["cost"][0].numpy()[valid])) + assert pm.metadata["dynamical_image"]["disk_radius_px"] == disk_r + pm.apply_dynamical(res) + assert pm.metadata["dynamical_applied"]["precession_deg"] == 0.0 + + +def test_coupling_lookup_cannot_alias(): + # a difference vector outside the stored factor box must come back as a + # missing factor (zero), never as another reflection's factor + from types import SimpleNamespace + + crystal = SimpleNamespace( + hkl_dyn=torch.tensor([[1, 0, 0], [-1, 0, 0], [-1, 1, 0]]), + U_dyn=torch.tensor([1, 1, 7], dtype=torch.complex128), + ) + U, _, _ = bloch._coupling_matrix(crystal, torch.tensor([[1, 0, 0], [-1, 0, 0]]), 1.0) + assert U[0, 1] == 0 and U[1, 0] == 0 + + +def test_illumination_nodes_moments(): + # the convergence disk is integrated with the uniform-area measure: the + # second moment of a disk of radius R is R^2 / 4 per axis; the ring is + # normalized and its mean vanishes + lam = bloch.electron_wavelength_angstrom(200e3) + k0 = 1.0 / lam + t, w = bloch.illumination_nodes(200e3, semiconv_mrad=5.0, n_disk_radial=3, n_disk_azimuthal=16) + R = k0 * np.sin(5e-3) + assert np.isclose(float(w.sum()), 1.0) + assert np.isclose(float((w * t[:, 0] ** 2).sum()), R**2 / 4, rtol=1e-10) + t, w = bloch.illumination_nodes(200e3, precession_deg=0.5, n_precession=16) + assert np.isclose(float(w.sum()), 1.0) and float(t.mean(0).abs().max()) < 1e-12 + assert np.allclose(torch.linalg.norm(t, dim=1).numpy(), k0 * np.sin(np.deg2rad(0.5))) + t, w = bloch.illumination_nodes(200e3) + assert t.shape == (1, 2) and float(w[0]) == 1.0 + + +def test_mean_absorption_and_forbidden_beam(): + """Pure mean absorption damps the total intensity as exp(-2 pi u0 z/k0); + a glide-forbidden reflection (Si 200) acquires intensity through double + diffraction, which requires it to be in the beam list.""" + si = _si(absorptive=True) + q = _zone_110() + z = torch.tensor([400.0, 800.0], dtype=torch.float64) + inten, g_xy, hkl = bloch._cbed_amplitudes( + si, q, torch.zeros((1, 2), dtype=torch.float64), z, 200e3, sg_max=0.06, k_max=1.0 + ) + keys = [tuple(h) for h in hkl.tolist()] + assert (0, 0, 2) in keys or (2, 0, 0) in keys or (0, 2, 0) in keys + i200 = next(i for i, h in enumerate(keys) if sorted(abs(v) for v in h) == [0, 0, 2]) + assert float(inten[0, 1, i200]) > 1e-4 # populated by multiple scattering + # mean absorption alone: strip the off-diagonal absorptive part + U, u0, absorptive = bloch._coupling_matrix(si, hkl, bloch.relativistic_gamma(200e3)) + Uel = 0.5 * (U + U.conj().T) + lam = bloch.electron_wavelength_angstrom(200e3) + k0 = 1.0 / lam + s_t = torch.zeros((1, hkl.shape[0]), dtype=torch.float64) + gl = bloch.qrotate(q, hkl[1:].to(torch.float64) @ si.lat_recip) + s_t[0, 1:] = (2 * gl[:, 2] - lam * (gl**2).sum(1)) / (2 - 2 * lam * gl[:, 2]) + inten_np = bloch._bloch_solve(Uel, u0, True, s_t, k0, z, fast_absorption=False) + total = inten_np[0].sum(dim=1).numpy() + assert np.allclose(total, np.exp(-2 * np.pi * u0 * z.numpy() / k0), rtol=1e-8) + + +def test_fourier_ring_matches_quadrature(): + """The harmonic propagation of the centered precession ring reproduces a + converged azimuthal quadrature, with the full complex coupling.""" + si = _si(absorptive=True) + q = _zone_110() + z = torch.tensor([300.0, 600.0]) + trial = torch.zeros((1, 2), dtype=torch.float64) + beams = bloch.select_dynamical_beams(si, q, 200e3, np.deg2rad(0.4), 0.06, 1.0) + ring, w = bloch.illumination_nodes(200e3, precession_deg=0.4, n_precession=96) + inten, g_xy, _ = bloch._cbed_amplitudes( + si, q, ring, z, 200e3, 0.06, 1.0, tilt_batch=128, beams=beams + ) + ref = (inten * w[:, None, None]).sum(0) + got, g2 = bloch.average_bloch_fourier(si, q, trial, z, 200e3, 0.4, 0.06, 1.0, beams=beams) + assert np.allclose(g2.numpy(), g_xy.numpy()) + assert np.allclose(got[0].numpy(), ref.numpy(), atol=1e-11, rtol=1e-9) + # a displaced ring through the sampled coefficients + trial = torch.tensor([[0.15, -0.1]], dtype=torch.float64) + inten, _, _ = bloch._cbed_amplitudes( + si, q, ring + trial, z, 200e3, 0.06, 1.0, tilt_batch=128, beams=beams + ) + ref = (inten * w[:, None, None]).sum(0) + got, _ = bloch.average_bloch_fourier( + si, q, trial, z, 200e3, 0.4, 0.06, 1.0, beams=beams, n_geometry=128 + ) + assert np.allclose(got[0].numpy(), ref.numpy(), atol=1e-9, rtol=1e-7) + + +def test_refine_dynamical_reported_cost_reproducible(): + """The stored cost and thickness belong to the stored orientation.""" + from quantem.core.datastructures.vector import Vector + from quantem.diffraction.orientation import OrientationMap + from quantem.diffraction.phase import PhaseMap + from quantem.diffraction.rotations import qnormalize + + energy_ev = 200e3 + xtl = _si(absorptive=True) + torch.manual_seed(3) + q_true = qnormalize(torch.randn(3, 4, dtype=torch.float64)) + peaks = Vector.from_shape( + (1, 3), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(3): + inten, g_xy, _ = bloch._cbed_amplitudes( + xtl, + q_true[i], + torch.zeros((1, 2), dtype=torch.float64), + torch.tensor([450.0]), + energy_ev, + 0.06, + 1.0, + progress_bar=False, + ) + inten_np = inten[0, 0, 1:].numpy() + keep = inten_np > 1e-3 * inten_np.max() + peaks[0, i] = np.column_stack([g_xy[1:].numpy()[keep], inten_np[keep]]) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=energy_ev) + om.build_plan(angle_step_zone_axis_deg=3.0, verbose=False) + om.match_orientations(progress_bar=False) + om.quats[0, :, 0] = q_true + om.corr[0, :, 0] = 1.0 + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(progress_bar=False) + res = bloch.refine_dynamical( + pm, + thicknesses_A=np.arange(300, 600, 50.0), + tilt_stages=((0.1, 0.05),), + sg_max=0.06, + k_max=1.0, + progress_bar=False, + ) + for i in range(3): + if not torch.isfinite(res["cost"][0, i, 0]) or peaks[0, i].numpy().shape[0] < 5: + continue + q = res["quats"][0, i, 0] + d3 = torch.eye(3, dtype=torch.float64) + d3[:2, :2] = res["deformation"][0, i, 0] + beams = bloch.select_dynamical_beams( + xtl, res["quats_base"][0, i, 0], energy_ev, np.deg2rad(0.1) * np.sqrt(2), 0.06, 1.0, d3 + ) + inten, g_xy, _ = bloch._cbed_amplitudes( + xtl, + q, + torch.zeros((1, 2), dtype=torch.float64), + np.arange(300, 600, 50.0), + energy_ev, + 0.06, + 1.0, + fast_absorption=False, + deform=d3, + beams=beams, + ) + data = peaks[0, i].numpy().astype(np.float64) + qxy = torch.as_tensor(data[:, :2]) + im = torch.as_tensor(data[:, 2]).clamp_min(0) ** 0.25 + cost, _, _, _ = bloch._dynamical_cost(inten[:, :, 1:], g_xy[1:], qxy, im, 0.05, 0.25, 0.02) + t_idx = int(np.argmin(np.abs(np.arange(300, 600, 50.0) - float(res["thickness"][0, i])))) + assert np.isclose(float(cost[0, t_idx]), float(res["cost"][0, i, 0]), rtol=1e-6, atol=1e-9) + + +def test_image_cost_radial_background(): + """A quadratic radial floor is removed by the radial background model + and biases the constant one.""" + torch.manual_seed(0) + shape = (48, 48) + yy, xx = np.mgrid[0:48, 0:48] + radius = torch.as_tensor(np.hypot(yy - 24.0, xx - 24.0)) + centers = torch.tensor([[24.0, 24.0], [30.0, 35.0], [15.0, 20.0], [36.0, 12.0]]) + inten = torch.tensor([0.8, 0.05, 0.02, 0.01], dtype=torch.float64) + sim = bloch.render_disks(centers, inten, shape, 3.0, 0.7) + sim_wrong = bloch.render_disks( + centers, inten * torch.tensor([1.0, 0.5, 2.0, 1.0]), shape, 3.0, 0.7 + ) + floor = 5.0 + 0.2 * radius - 0.004 * radius**2 + meas = 1000 * sim + floor + mask = bloch._image_mask(shape, (24.0, 24.0), None, 4.5) + c_const = bloch._image_cost(meas, torch.stack([sim, sim_wrong]), mask, 0.5, "constant") + c_rad = bloch._image_cost(meas, torch.stack([sim, sim_wrong]), mask, 0.5, "radial", radius) + assert float(c_rad[0]) < 1e-12 # exact model with the right background + assert float(c_const[0]) > 1e-4 # the constant background cannot absorb it + assert float(c_rad[1]) > float(c_rad[0]) + + +def test_refine_dynamical_with_precession_and_convergence(): + """End-to-end recovery with a precession ring and a convergence disk: + the ground truth is integrated with denser illumination nodes than + the model uses.""" + from quantem.core.datastructures.vector import Vector + from quantem.diffraction.orientation import OrientationMap + from quantem.diffraction.phase import PhaseMap + from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qmult, + qnormalize, + quat_from_axis_angle, + ) + + energy_ev = 200e3 + xtl = _si(absorptive=True) + torch.manual_seed(5) + rng = np.random.default_rng(5) + N = 3 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + t_true = torch.tensor([300.0, 450.0, 600.0]) + ring, w = bloch.illumination_nodes( + energy_ev, + precession_deg=0.4, + n_precession=32, + semiconv_mrad=1.5, + n_disk_radial=3, + n_disk_azimuthal=12, + ) + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + inten, g_xy, _ = bloch._cbed_amplitudes( + xtl, + q_true[i], + ring, + t_true[i : i + 1], + energy_ev, + 0.06, + 1.0, + tilt_batch=256, + progress_bar=False, + ) + I_avg = (inten[:, 0, 1:] * w[:, None]).sum(0).numpy() + keep = I_avg > 1e-3 * I_avg.max() + peaks[0, i] = np.column_stack([g_xy[1:].numpy()[keep], I_avg[keep]]) + om = OrientationMap.from_vectors( + peaks, xtl, energy_ev=energy_ev, precession_deg=0.4, semiconv_mrad=1.5 + ) + # the matched orientations are replaced below, so a coarse plan will do + # (the precession-integrated library is the costly part) + om.build_plan(angle_step_zone_axis_deg=10.0, angle_step_in_plane_deg=10.0, verbose=False) + om.match_orientations(progress_bar=False) + phis = rng.uniform(0, 2 * np.pi, N) + om.quats[0, :, 0] = torch.stack( + [ + qmult( + quat_from_axis_angle( + torch.tensor([np.cos(p), np.sin(p), 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(0.12), dtype=torch.float64), + ), + q_true[i], + ) + for i, p in enumerate(phis) + ] + ) + om.corr[0, :, 0] = 1.0 + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(progress_bar=False) + res = bloch.refine_dynamical( + pm, + thicknesses_A=np.arange(200, 700, 25.0), + tilt_stages=((0.15, 0.05), (0.03, 0.01)), + n_precession=12, + n_precession_search=12, + n_disk_radial=2, + n_disk_azimuthal=6, + power_intensity=0.5, + sg_max=0.06, + k_max=1.0, + progress_bar=False, + ) + assert res["metadata"]["precession_deg"] == 0.4 and res["metadata"]["semiconv_mrad"] == 1.5 + valid = np.array([peaks[0, i].numpy().shape[0] >= 6 for i in range(N)]) + assert valid.sum() >= 2 + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy()[valid] + t_err = np.abs(res["thickness"][0].numpy() - t_true.numpy())[valid] + assert np.median(err) < 0.03 + assert (t_err <= 25).sum() >= valid.sum() - 1 + + +def _smooth_map_setup(n, tilt_start_deg, corrupt=None): + """A 1 x n 'map' of one grain: orientations a few hundredths of a degree + apart, thickness varying slowly, strained cell; starts tilted by + tilt_start_deg (and one position by `corrupt` degrees).""" + from quantem.core.datastructures.vector import Vector + from quantem.diffraction.orientation import OrientationMap + from quantem.diffraction.phase import PhaseMap + from quantem.diffraction.rotations import qmult, quat_from_axis_angle + + energy_ev = 200e3 + xtl = _si(absorptive=True) + torch.manual_seed(7) + rng = np.random.default_rng(7) + # a well-populated pattern: 1.5 degrees off the [110] zone axis + base = qmult( + quat_from_axis_angle( + torch.tensor([0.6, 0.8, 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(1.5), dtype=torch.float64), + ), + _zone_110(), + ) + q_true = torch.stack( + [ + qmult( + quat_from_axis_angle( + torch.tensor([1.0, 0.3, 0.0], dtype=torch.float64) / np.hypot(1, 0.3), + torch.tensor(np.deg2rad(0.03 * i), dtype=torch.float64), + ), + base, + ) + for i in range(n) + ] + ) + t_true = torch.tensor([400.0 + 25.0 * i for i in range(n)]) + A_true = torch.tensor([[1.008, 0.002], [0.002, 0.996]], dtype=torch.float64) + deform3 = torch.eye(3, dtype=torch.float64) + deform3[:2, :2] = A_true + peaks = Vector.from_shape( + (1, n), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(n): + inten, g_xy, _ = bloch._cbed_amplitudes( + xtl, + q_true[i], + torch.zeros((1, 2), dtype=torch.float64), + t_true[i : i + 1], + energy_ev, + 0.06, + 1.0, + progress_bar=False, + deform=deform3, + ) + inten_np = inten[0, 0, 1:].numpy() + keep = inten_np > 1e-3 * inten_np.max() + peaks[0, i] = np.column_stack([g_xy[1:].numpy()[keep], inten_np[keep]]) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=energy_ev) + om.build_plan(angle_step_zone_axis_deg=3.0, verbose=False) + om.match_orientations(progress_bar=False) + phis = rng.uniform(0, 2 * np.pi, n) + starts = [] + for i, p in enumerate(phis): + ang = tilt_start_deg if (corrupt is None or i != corrupt[0]) else corrupt[1] + starts.append( + qmult( + quat_from_axis_angle( + torch.tensor([np.cos(p), np.sin(p), 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(ang), dtype=torch.float64), + ), + q_true[i], + ) + ) + om.quats[0, :, 0] = torch.stack(starts) + om.corr[0, :, 0] = 1.0 + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(progress_bar=False) + return xtl, om, pm, q_true, t_true, A_true + + +def test_refine_dynamical_warm_start_matches_cold(): + from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qconj, + qmult, + quat_from_axis_angle, + ) + + n = 5 + kw = dict( + thicknesses_A=np.arange(300, 700, 25.0), + tilt_stages=((0.25, 0.05), (0.04, 0.01)), + power_intensity=0.5, + sg_max=0.06, + k_max=1.0, + neighbor_rescue=False, + progress_bar=False, + ) + out = {} + for warm in (False, True): + xtl, om, pm, q_true, t_true, A_true = _smooth_map_setup(n, 0.1) + q_match = om.quats[0, :, 0].clone() + res = bloch.refine_dynamical(pm, warm_start=warm, **kw) + out[warm] = (res, om.quats[0, :, 0].clone(), q_true, t_true, xtl) + res_c, q_c, q_true, t_true, xtl = out[False] + res_w, q_w, _, _, _ = out[True] + assert not res_c["warm_started"].any() + assert res_w["warm_started"][0, 1:, 0].all() and not res_w["warm_started"][0, 0, 0] + # same solution from both routes, both correct + d = misorientation_angle_deg(q_c, q_w, xtl.sym_quats).numpy() + assert d.max() < 0.03 + assert np.allclose(res_c["thickness"][0].numpy(), res_w["thickness"][0].numpy()) + err = misorientation_angle_deg(q_true, q_w, xtl.sym_quats).numpy() + assert err.max() < 0.03 + assert np.abs(res_w["thickness"][0].numpy() - t_true.numpy()).max() <= 25 + + # the reported tilt is measured from the position's own matched + # orientation on either route: every start was 0.1 degrees off the truth + for res in (res_c, res_w): + tilt = res["tilt_deg"][0].numpy() # (n, 2) + assert np.all(np.abs(np.hypot(tilt[:, 0], tilt[:, 1]) - 0.1) < 0.03) + base = res["quats_base"][0, :, 0] + # the base is the matched orientation turned about the beam only + # (the in-plane rotation of the deformation fit) + dq = qmult(base, qconj(q_match)) + assert torch.all(dq[:, 1:3].abs() < 1e-9) + assert misorientation_angle_deg(base, q_match).max() < 0.05 + # and tilt x base gives the solution + for i in range(n): + wx, wy = np.deg2rad(tilt[i]) + ang = np.hypot(wx, wy) + tq = quat_from_axis_angle( + torch.tensor([wx / ang, wy / ang, 0.0], dtype=torch.float64), + torch.tensor(ang, dtype=torch.float64), + ) + q_rebuilt = qmult(tq, base[i]) + assert torch.allclose( + q_rebuilt * torch.sign(q_rebuilt @ res["quats"][0, i, 0]), + res["quats"][0, i, 0], + atol=1e-9, + ) + # the zero-tilt cost is at the matched orientation, above the final + assert torch.all(res["cost_zero_tilt"][0, :, 0] > res["cost"][0, :, 0]) + assert np.allclose(res_c["tilt_deg"].numpy(), res_w["tilt_deg"].numpy(), atol=0.01) + assert np.allclose(res_c["cost_zero_tilt"].numpy(), res_w["cost_zero_tilt"].numpy(), rtol=0.05) + + +def test_refine_dynamical_neighbor_rescue(): + from quantem.diffraction.rotations import misorientation_angle_deg + + n = 5 + # position 2 starts 0.45 degrees off: outside the coarse stage's reach, + # so its cold search settles in a wrong basin; its neighbors are right + xtl, om, pm, q_true, t_true, A_true = _smooth_map_setup(n, 0.1, corrupt=(2, 0.45)) + kw = dict( + thicknesses_A=np.arange(300, 700, 25.0), + tilt_stages=((0.25, 0.05), (0.04, 0.01)), + power_intensity=0.5, + sg_max=0.06, + k_max=1.0, + warm_start=False, + progress_bar=False, + ) + res = bloch.refine_dynamical(pm, neighbor_rescue=False, **kw) + err0 = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + assert err0[2] > 0.1 # the cold start fails there + xtl, om, pm, q_true, t_true, A_true = _smooth_map_setup(n, 0.1, corrupt=(2, 0.45)) + res = bloch.refine_dynamical(pm, neighbor_rescue=True, **kw) + err1 = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + assert bool(res["rescued"][0, 2]) + assert err1[2] < 0.03 + assert abs(float(res["thickness"][0, 2]) - float(t_true[2])) <= 25 + # map-level outputs + maps = bloch.dynamical_maps(res, pm, crystal_index=0) + assert maps["mask"][0].all() + assert set(maps["strain"]) == {"aa", "bb", "cc", "ab", "ac", "bc"} + assert torch.isfinite(maps["gain"][0]).all() + + +def test_refine_dynamical_threads_match_one_worker(): + from quantem.diffraction.rotations import misorientation_angle_deg + + # 32 positions: two blocks of 16 on two threads, one warm-start chain + # broken at the block boundary. Both runs start from the same matched + # orientations (update_orientations=False leaves them in place) + n = 32 + kw = dict( + thicknesses_A=np.arange(300, 700, 25.0), + tilt_stages=((0.25, 0.05), (0.04, 0.01)), + power_intensity=0.5, + sg_max=0.06, + k_max=1.0, + neighbor_rescue=False, + update_orientations=False, + progress_bar=False, + ) + xtl, om, pm, q_true, _, _ = _smooth_map_setup(n, 0.1) + res_1 = bloch.refine_dynamical(pm, num_workers=1, **kw) + res_2 = bloch.refine_dynamical(pm, num_workers=2, **kw) + q_1, q_2 = res_1["quats"][0, :, 0], res_2["quats"][0, :, 0] + assert int(res_1["warm_started"].sum()) == n - 1 + assert int(res_2["warm_started"].sum()) == n - 2 + # the first block is the same computation on either route + assert torch.equal(q_1[:16], q_2[:16]) + assert torch.equal(res_1["thickness"][0, :16], res_2["thickness"][0, :16]) + # breaking the warm-start chain costs nothing in accuracy + e1 = misorientation_angle_deg(q_true, q_1, xtl.sym_quats).numpy() + e2 = misorientation_angle_deg(q_true, q_2, xtl.sym_quats).numpy() + assert e2.mean() <= e1.mean() + 0.01 + + +def _fake_dynamical_result(): + """A 2 x 2 refine_dynamical result with one candidate; position (1, 1) + was not refined.""" + from types import SimpleNamespace + + from quantem.diffraction.rotations import quat_from_axis_angle + + R, C, F = 2, 2, 1 + q = quat_from_axis_angle( + torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64), torch.tensor(0.3, dtype=torch.float64) + ) + deform = torch.eye(2, dtype=torch.float64).repeat(R, C, F, 1, 1) + deform[0, 0, 0] = torch.tensor([[1.01, 0.0], [0.0, 0.99]], dtype=torch.float64) + cost = torch.tensor([[0.10, 0.20], [0.15, torch.nan]], dtype=torch.float64)[..., None] + done = torch.isfinite(cost[..., 0]) + result = { + "candidate": torch.where(done, 0, -1), + "phase_index": torch.where(done, 0, -1), + "quats": q.repeat(R, C, F, 1), + "deformation": deform, + "cost": cost, + "cost_zero_tilt": cost + 0.01, + "thickness": torch.where(done, 400.0, torch.nan).to(torch.float64), + "thickness_contrast": torch.full((R, C), 0.1, dtype=torch.float64), + "tilt_deg": torch.full((R, C, 2), 0.05, dtype=torch.float64), + } + return result, SimpleNamespace(candidates=[(0, 0)]) + + +def test_dynamical_maps_mask_unrefined(): + result, pm = _fake_dynamical_result() + maps = bloch.dynamical_maps(result, pm) + assert maps["mask"].tolist() == [[True, True], [True, False]] + assert int(maps["phase_index"][1, 1]) == -1 and int(maps["phase_index"][0, 0]) == 0 + assert torch.isnan(maps["quats"][1, 1]).all() and torch.isfinite(maps["quats"][0, 0]).all() + assert torch.isnan(maps["deformation"][1, 1]).all() + assert torch.isnan(maps["thickness"][1, 1]) and float(maps["thickness"][0, 0]) == 400.0 + assert np.isclose(float(maps["gain"][0, 0]), 0.01) + assert np.isclose(float(maps["tilt_deg"][0, 0]), 0.05 * np.sqrt(2)) + # the strain of the strained position, in the crystal frame (a pure + # rotation about the beam keeps the normal strains on the diagonal) + eps = maps["strain"] + assert float(eps["aa"][0, 0]) < 0 < float(eps["bb"][0, 0]) + assert torch.isnan(eps["cc"][1, 1]) + # restricted to a crystal that won nowhere: empty mask + assert not bloch.dynamical_maps(result, pm, crystal_index=1)["mask"].any() + + +def test_plot_dynamical_and_strain_maps(): + import matplotlib.pyplot as plt + + result, pm = _fake_dynamical_result() + maps = bloch.dynamical_maps(result, pm) + fig, axs = bloch.plot_dynamical_maps(maps) + assert np.asarray(axs).size >= 4 + plt.close(fig) + strain = {k: v.numpy() for k, v in maps["strain"].items()} + fig, axs = bloch.plot_strain_crystal_frame(strain, mask=maps["mask"].numpy()) + assert np.asarray(axs).size >= 6 + plt.close(fig) diff --git a/tests/diffraction/test_calibrate.py b/tests/diffraction/test_calibrate.py new file mode 100644 index 000000000..a2d6f79ee --- /dev/null +++ b/tests/diffraction/test_calibrate.py @@ -0,0 +1,157 @@ +"""calibrate(), DiffractionCalibration and the scan rotation measurement.""" + +from types import SimpleNamespace + +import numpy as np +import pytest +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.core.io.serialize import load +from quantem.diffraction import calibration +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.rotations import qnormalize + +PIXEL_SIZE = 0.0123 +E_TRUE = np.array([0.012, -0.008]) + + +@pytest.fixture(scope="module") +def ti(): + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.9505, c=4.6855), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + return xtl + + +@pytest.fixture(scope="module") +def peaks_px(ti): + """Ti patterns in detector pixels with a known ellipse and pixel size.""" + torch.manual_seed(0) + rng = np.random.default_rng(0) + A_inv = np.linalg.inv(calibration._ellipse_matrix(E_TRUE)) + cells = [] + for _ in range(60): + q = qnormalize(torch.randn(4, dtype=torch.float64)) + pat = ti.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + qxy = np.stack([pat["qx"].numpy(), pat["qy"].numpy()], axis=1) + qxy = qxy @ A_inv.T + rng.normal(0, 0.002, qxy.shape) + cells.append(np.column_stack([qxy / PIXEL_SIZE, pat["intensity"].numpy()])) + return Vector.from_data( + [cells[:30], cells[30:]], + fields=["q_row", "q_col", "intensity"], + units=["px", "px", "counts"], + name="synthetic", + ) + + +@pytest.mark.parametrize("n_iter", [1, 2, 3]) +def test_calibrate_stable_in_n_iter(ti, peaks_px, n_iter): + # each round must refine the ellipse, not replace it with the residual + cal = calibration.calibrate(peaks_px, ti, 0.011, n_iter=n_iter) + assert abs(cal.pixel_size / PIXEL_SIZE - 1) < 2e-3 + assert np.allclose(cal.ellipse, E_TRUE, atol=1.5e-3), cal.ellipse + assert cal.metadata["reliable"] + assert cal.metadata["n_rings"] >= 3 + + +def test_calibrate_returnfig(ti, peaks_px): + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + cal, fig, axs = calibration.calibrate(peaks_px, ti, 0.011, plot=True, returnfig=True) + assert isinstance(cal, calibration.DiffractionCalibration) + assert axs.shape == (2, 2) + plt.close(fig) + + +def test_compose_ellipse(): + a, b = np.array([0.01, -0.004]), np.array([0.003, 0.002]) + c = calibration._compose_ellipse(a, b) + # to first order the components add + assert np.allclose(c, a + b, atol=1e-4) + assert np.allclose(calibration._compose_ellipse(a, None), a) + assert np.allclose(calibration._compose_ellipse(np.zeros(2), b), b) + + +def test_diffraction_calibration_apply_rebin_save(tmp_path): + cells = [[np.array([[10.0, 0.0, 1.0], [0.0, -20.0, 2.0]]), np.zeros((0, 3))]] + peaks_px = Vector.from_data( + cells, fields=["q_row", "q_col", "intensity"], units=["px", "px", "counts"], name="p" + ) + ellipse = np.array([0.01, -0.02]) + cal = calibration.DiffractionCalibration( + 0.01, ellipse, rotation_ccw_deg=90.0, metadata={"binning": 1, "n_rings": 4} + ) + + out = cal.apply(peaks_px) + assert out.fields == ["qx", "qy", "intensity"] + assert out.metadata["pixel_size"] == pytest.approx(0.01) + flat = out.numpy().astype(np.float64) + q = np.array([[0.1, 0.0], [0.0, -0.2]]) @ calibration._ellipse_matrix(ellipse).T + rot = np.array([[0.0, -1.0], [1.0, 0.0]]) + assert np.allclose(flat[:, :2], q @ rot.T, atol=1e-6) + assert np.allclose(flat[:, 2], [1.0, 2.0]) + + binned = cal.rebin(2) + assert binned.pixel_size == pytest.approx(0.02) + assert binned.metadata["binning"] == 2 + assert np.allclose(binned.ellipse, ellipse) + assert binned.rotation_ccw_deg == cal.rotation_ccw_deg + assert cal.metadata["binning"] == 1 + + path = tmp_path / "cal.zip" + cal.save(path, mode="o") + cal2 = load(path) + assert isinstance(cal2, calibration.DiffractionCalibration) + assert cal2.pixel_size == pytest.approx(cal.pixel_size) + assert np.allclose(cal2.ellipse, cal.ellipse) + assert cal2.rotation_ccw_deg == pytest.approx(90.0) + assert cal2.metadata["n_rings"] == 4 + assert np.allclose(cal2.apply(peaks_px).numpy(), out.numpy()) + + +def test_refine_calibration_empty_raises(): + with pytest.raises(ValueError, match="empty"): + calibration.refine_calibration([]) + sm = SimpleNamespace( + u_array=np.full((2, 2, 2), np.nan), + v_array=np.full((2, 2, 2), np.nan), + g1_array=np.full((2, 2, 2), np.nan), + g2_array=np.full((2, 2, 2), np.nan), + ) + with pytest.raises(ValueError, match="no positions"): + calibration.refine_calibration([sm]) + + +@pytest.mark.parametrize("theta_deg", [30.0, 210.0, 125.0]) +def test_measure_scan_rotation_sign(theta_deg): + """A gradient field seen on a detector rotated by -theta returns theta mod 180. + + The returned angle uses the convention of peaks_to_calibrated: rotating + the detector field by +theta brings it back into the scan frame. + """ + R = C = 16 + H = W = 32 + ry, rx = np.mgrid[0:R, 0:C] / R + # CoM shift in the scan frame: gradient of an asymmetric potential + phi = np.sin(2 * np.pi * ry) * np.cos(np.pi * rx) + 0.5 * rx**2 + g_r, g_c = np.gradient(phi) + g = np.stack([g_r, g_c], axis=-1) + g *= 1.5 / np.abs(g).max() + th = np.deg2rad(theta_deg) + rot = np.array([[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]]) + d = g @ rot # detector frame: rot.T applied to each scan-frame vector + rows = np.arange(H)[:, None] + cols = np.arange(W)[None, :] + arr = np.zeros((R, C, H, W)) + for i in range(R): + for j in range(C): + r0, c0 = H / 2 + d[i, j, 0], W / 2 + d[i, j, 1] + arr[i, j] = np.exp(-((rows - r0) ** 2 + (cols - c0) ** 2) / (2 * 1.5**2)) + angle = calibration.measure_scan_rotation(SimpleNamespace(array=arr)) + expected = theta_deg % 180 + diff = (angle - expected + 90) % 180 - 90 + assert abs(diff) < 1.0, (angle, expected) diff --git a/tests/diffraction/test_calibration_refine.py b/tests/diffraction/test_calibration_refine.py new file mode 100644 index 000000000..01f035d57 --- /dev/null +++ b/tests/diffraction/test_calibration_refine.py @@ -0,0 +1,68 @@ +"""Recovery of a global calibration residual from matched orientations.""" + +import numpy as np +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction import calibration +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.rotations import qnormalize + + +def test_refine_calibration_recovers_distortion(): + torch.manual_seed(1) + rng = np.random.default_rng(1) + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + + # global calibration error: 1.5% scale, 0.8% ellipticity, 0.3 deg rotation + scale_true = 1.015 + e11_true, e12_true = 0.008, -0.004 + th = np.deg2rad(0.3) + M_true = ( + scale_true + * np.array([[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]]) + @ np.array([[1 + e11_true, e12_true], [e12_true, 1 - e11_true]]) + ) + + R, C = 6, 10 + q_true = qnormalize(torch.randn(R * C, 4, dtype=torch.float64)) + cells = [] + for i in range(R * C): + p = xtl.generate_pattern(q_true[i], energy_ev=200e3, sigma_excitation=0.02) + q2 = np.stack([p["qx"].numpy(), p["qy"].numpy()], axis=1) @ M_true.T + q2 += rng.normal(0, 0.002, q2.shape) + cells.append(np.column_stack([q2, p["intensity"].numpy()])) + nested = [cells[r * C : (r + 1) * C] for r in range(R)] + peaks = Vector.from_data( + nested, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(power_intensity=0.0) + om.match_orientations(progress_bar=False) + om.refine_orientations(progress_bar=False, neighbor_rescue=False) + sm = om.calculate_strain(progress_bar=False) + + res = calibration.refine_calibration([sm]) + assert res["num_positions"] > 30 + assert abs(res["scale"] - scale_true) < 3e-3 + # a global rotation is absorbed into the refined orientations and must + # read as ~0 here (it is measured independently via the scan rotation) + assert abs(res["rotation_deg"]) < 0.1 + assert abs(res["ellipse"][0] - e11_true) < 2e-3 + assert abs(res["ellipse"][1] - e12_true) < 2e-3 + + # applying the correction must bring the residual to the identity + peaks_fixed = calibration.transform_peaks(peaks, res["correction"]) + om2 = OrientationMap.from_vectors(peaks_fixed, xtl, energy_ev=200e3) + om2.build_plan(power_intensity=0.0) + om2.match_orientations(progress_bar=False) + om2.refine_orientations(progress_bar=False, neighbor_rescue=False) + sm2 = om2.calculate_strain(progress_bar=False) + res2 = calibration.refine_calibration([sm2]) + assert abs(res2["scale"] - 1.0) < 2e-3 + assert abs(res2["ellipse"][0]) < 1.5e-3 + assert abs(res2["ellipse"][1]) < 1.5e-3 diff --git a/tests/diffraction/test_crystal.py b/tests/diffraction/test_crystal.py new file mode 100644 index 000000000..e9e8214f7 --- /dev/null +++ b/tests/diffraction/test_crystal.py @@ -0,0 +1,598 @@ +"""Tests for quantem.diffraction.crystal.""" + +import numpy as np +import pytest +import torch +from ase import Atoms +from ase.build import bulk + +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.rotations import quat_from_zone_axis + + +@pytest.fixture +def ti_beta(): + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), name="Ti beta") + xtl.calculate_structure_factors(k_max=1.5) + return xtl + + +def test_symmetry_detection(ti_beta): + assert ti_beta.pointgroup == "m-3m" + assert ti_beta.laue_group == "m-3m" + assert ti_beta.sym_quats.shape[0] == 24 # proper rotations of m-3m + + +def test_hcp_symmetry(): + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686)) + assert xtl.pointgroup == "6/mmm" + assert xtl.sym_quats.shape[0] == 12 + + +def test_bcc_absences(ti_beta): + # h + k + l odd forbidden in bcc + parity = ti_beta.hkl.sum(dim=1) % 2 + assert (parity == 0).all() + + +def test_ring_positions(ti_beta): + # (110) ring at sqrt(2)/a + g110 = np.sqrt(2) / 3.31 + assert np.isclose(float(ti_beta.g_len.min()), g110, atol=1e-6) + + +def test_pseudo_symmetry(): + ortho = Atoms("Au", positions=[[0, 0, 0]], cell=[4.000, 4.001, 4.002], pbc=True) + exact = Crystal.from_ase(ortho, pseudo_symmetry_tol=None) + pseudo = Crystal.from_ase(ortho, pseudo_symmetry_tol=0.01) # 0.04 A on a 4 A cell + assert exact.pointgroup_matching == "mmm" + assert pseudo.pointgroup_matching == "m-3m" + assert pseudo.sym_quats_matching.shape[0] == 24 + # exact group is retained for reporting/refinement + assert pseudo.pointgroup == "mmm" + + +def test_zone_axis_wedge_anchored_001(ti_beta): + wedge = ti_beta.zone_axis_wedge() + assert torch.allclose(wedge[0], torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64)) + + +def test_generate_pattern(ti_beta): + q = quat_from_zone_axis(torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64)) + p = ti_beta.generate_pattern(q, energy_ev=200e3) + # [001] zone: peaks on a square grid of 110-type spacings + assert p["qx"].shape[0] > 4 + qr = torch.hypot(p["qx"], p["qy"]) + assert float(qr.min()) > 0.4 # no direct beam + # pattern symmetric under 90 degree rotation + rot = torch.stack((-p["qy"], p["qx"]), dim=1) + orig = torch.stack((p["qx"], p["qy"]), dim=1) + d = torch.cdist(rot, orig).min(dim=1).values + assert float(d.max()) < 1e-6 + + +def _images_in_wedge(xtl, n=3000, tol=1e-9): + """Count, per random direction, its symmetry images inside the wedge.""" + from quantem.diffraction.rotations import quat_to_matrix + + rng = np.random.default_rng(0) + d = rng.normal(size=(n, 3)) + d = torch.as_tensor(d / np.linalg.norm(d, axis=1, keepdims=True)) + c = xtl.zone_axis_wedge() + Rs = quat_to_matrix(xtl.sym_quats_matching) + imgs = torch.einsum("sij,nj->nsi", Rs, d) + imgs = torch.cat([imgs, -imgs], dim=1).reshape(-1, 3) + ok = torch.ones(imgs.shape[0], dtype=torch.bool) + for i in range(3): + nrm = torch.cross(c[i], c[(i + 1) % 3], dim=0) + ok &= (imgs @ nrm) * torch.sign(nrm @ c[(i + 2) % 3]) >= -tol + return ok.reshape(n, -1).sum(dim=1) + + +@pytest.mark.parametrize( + "label,spacegroup,symbols,basis,cellpar,laue", + [ + ("Si", 227, ["Si"], [(0, 0, 0)], [5.43] * 3 + [90] * 3, "m-3m"), + ( + "pyrite", + 205, + ["Fe", "S"], + [(0, 0, 0), (0.385, 0.385, 0.385)], + [5.42] * 3 + [90] * 3, + "m-3", + ), + ("Ti", 194, ["Ti"], [(1 / 3, 2 / 3, 0.25)], [2.95, 2.95, 4.68, 90, 90, 120], "6/mmm"), + ( + "CdI2 -3m1", + 164, + ["Cd", "I"], + [(0, 0, 0), (1 / 3, 2 / 3, 0.25)], + [4.24, 4.24, 6.84, 90, 90, 120], + "-3m", + ), + ("Bi R-3m", 166, ["Bi"], [(0, 0, 0.234)], [4.55, 4.55, 11.86, 90, 90, 120], "-3m"), + ( + "P-31m", + 162, + ["Cu", "O"], + [(1 / 3, 2 / 3, 0), (0.4, 0, 0.3)], + [5.0, 5.0, 7.0, 90, 90, 120], + "-3m", + ), + ( + "ilmenite", + 148, + ["Fe", "Ti", "O"], + [(0, 0, 0.355), (0, 0, 0.146), (0.317, 0.023, 0.245)], + [5.09, 5.09, 14.09, 90, 90, 120], + "-3", + ), + ( + "rutile", + 136, + ["Ti", "O"], + [(0, 0, 0), (0.305, 0.305, 0)], + [4.59, 4.59, 2.96, 90, 90, 90], + "4/mmm", + ), + ( + "Pnma", + 62, + ["Fe", "C"], + [(0.18, 0.25, 0.33), (0.04, 0.25, 0.87)], + [5.0, 6.7, 4.5, 90, 90, 90], + "mmm", + ), + ], +) +def test_wedge_is_fundamental_domain(label, spacegroup, symbols, basis, cellpar, laue): + from ase.spacegroup import crystal as ase_crystal + + atoms = ase_crystal(symbols, basis=basis, spacegroup=spacegroup, cellpar=cellpar) + xtl = Crystal.from_ase(atoms, name=label, verbose=False) + assert xtl.laue_group_matching == laue + # exactly one symmetry image of every direction lies in the wedge: the + # wedge covers all of orientation space once (the -3m1 setting used to + # get a wedge rotated by 30 degrees, covering half the directions twice) + hits = _images_in_wedge(xtl) + assert int(hits.min()) == 1 and int(hits.max()) == 1 + labels = xtl.zone_axis_wedge_labels() + assert len(labels) == 3 and all(len(t) > 2 for t in labels) + + +def test_wedge_follows_cell_setting(): + # a rotated Cartesian setting moves the symmetry axes; the wedge follows + atoms = bulk("Si", "diamond", a=5.431, cubic=True) + atoms.rotate(37, "z", rotate_cell=True) + atoms.rotate(20, "x", rotate_cell=True) + xtl = Crystal.from_ase(atoms, verbose=False) + hits = _images_in_wedge(xtl) + assert int(hits.min()) == 1 and int(hits.max()) == 1 + assert xtl.zone_axis_wedge_labels(mathtext=False) == ["[001]", "[011]", "[111]"] + + +def test_pseudo_symmetry_default_and_warning(): + ortho = Atoms("Au", positions=[[0, 0, 0]], cell=[4.000, 4.001, 4.002], pbc=True) + xtl = Crystal.from_ase(ortho, verbose=False) # default tolerance 1% of the cell + assert xtl.pointgroup == "mmm" and xtl.pointgroup_matching == "m-3m" + assert xtl.pseudo_symmetry_report["intensity_mismatch"] < 0.05 + msg = xtl.matching_symmetry_warning() + assert msg is not None and "pseudo_symmetry_tol=None" in msg + # the pseudo group's operators are exact rotations (orthonormalized), so + # its wedge is a fundamental domain up to the cell distortion + hits = _images_in_wedge(xtl, tol=1e-3) + assert int(hits.min()) >= 1 + exact = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + assert exact.matching_symmetry_warning() is None + + +def test_pseudo_symmetry_dimensionless_and_intensity_check(): + # an almost body-centered cell: the center atom 0.002 A off (0.5, 0.5, 0.5) + # is body centered at the default tolerance, and its 100/010/001 + # patterns are identical within any measurable intensity + almost_bcc = Atoms( + "Fe2", scaled_positions=[[0, 0, 0], [0.5, 0.5, 0.5005]], cell=[4.0, 4.0, 4.0], pbc=True + ) + xtl = Crystal.from_ase(almost_bcc, verbose=False) + assert xtl.pointgroup_matching == "m-3m" + assert xtl.pseudo_symmetry_report["intensity_mismatch"] < 1e-3 + # the same cell at an unmeasurably tight distance tolerance keeps its + # own (lower) symmetry; the tolerance is a fraction of the lattice + tight = Crystal.from_ase(almost_bcc, pseudo_symmetry_tol=1e-7, verbose=False) + assert tight.pointgroup_matching == tight.pointgroup + # a candidate whose intensities do not match within the intensity + # tolerance is rejected and the cell keeps its own symmetry + strict = Crystal.from_ase(almost_bcc, pseudo_symmetry_intensity_tol=1e-9, verbose=False) + assert strict.pseudo_symmetry_report.get("candidate") == "m-3m" + assert strict.pseudo_symmetry_report.get("rejected") is True + assert strict.pointgroup_matching == strict.pointgroup + assert "rejected" in strict.symmetry_summary() + + +def _l10(other: str, a: float = 3.58) -> Atoms: + """Two species ordered in alternating (001) layers of an fcc lattice. + + The lattice stays cubic and every atom sits exactly on its site, so no + relaxation of the positions recovers the cubic parent: only the + diffracted intensities can say whether the ordering is visible. + """ + at = Atoms( + "Ni4", + scaled_positions=[[0, 0, 0], [0.5, 0.5, 0], [0.5, 0, 0.5], [0, 0.5, 0.5]], + cell=[a, a, a], + pbc=True, + ) + at.symbols = ["Ni", "Ni", other, other] + return at + + +def test_pseudo_symmetry_from_weak_ordering(): + """Ordering of species that scatter alike is found through the lattice. + + Transition metals next to each other in the periodic table (the Ni, Co, + Mn of a cathode) give superlattice reflections far too weak to index, so + the orientation library must fold the variants together. The relaxed + position search cannot find this: the atoms are already where they + belong and only the species differ. + """ + weak = Crystal.from_ase(_l10("Co"), verbose=False) + assert weak.pointgroup == "4/mmm" + assert weak.pointgroup_matching == "m-3m" + assert weak.sym_quats_matching.shape[0] == 3 * weak.sym_quats.shape[0] + assert weak.pseudo_symmetry_report["route"] == "lattice" + assert weak.pseudo_symmetry_report["intensity_mismatch"] < 0.01 + + # a light partner makes the same ordering plainly visible, and the + # candidate is rejected + strong = Crystal.from_ase(_l10("Li"), verbose=False) + assert strong.pointgroup_matching == strong.pointgroup + assert strong.pseudo_symmetry_report["rejected"] is True + assert strong.pseudo_symmetry_report["intensity_mismatch"] > 0.1 + + # the intensity tolerance is the decision, and it is the user's + borderline = Crystal.from_ase(_l10("Al"), verbose=False) + assert borderline.pointgroup_matching == borderline.pointgroup + loose = Crystal.from_ase(_l10("Al"), pseudo_symmetry_intensity_tol=0.1, verbose=False) + assert loose.pointgroup_matching == "m-3m" + + +def test_true_symmetry_cells_are_unchanged(): + """The lattice route must not disturb cells that are already at their + lattice's symmetry, nor accept a lattice symmetry the structure breaks.""" + from ase.build import bulk + + for atoms, pg in ( + (bulk("Si", "diamond", a=5.43), "m-3m"), + (bulk("Ti", "hcp", a=2.95, c=4.686), "6/mmm"), + (bulk("Ti", "bcc", a=3.26, cubic=True), "m-3m"), + ): + xtl = Crystal.from_ase(atoms, verbose=False) + assert xtl.pointgroup_matching == pg + assert xtl.sym_quats_matching.shape[0] == xtl.sym_quats.shape[0] + + # corundum sits on a hexagonal lattice but its structure is only -3m; + # the lattice route proposes 6/mmm and the intensities reject it + from ase.spacegroup import crystal as ase_crystal + + al2o3 = ase_crystal( + ("Al", "O"), + basis=[(0, 0, 0.3522), (0.3064, 0, 0.25)], + spacegroup=167, + cellpar=[4.7607, 4.7607, 12.9947, 90, 90, 120], + ) + xtl = Crystal.from_ase(al2o3, verbose=False) + assert xtl.pointgroup_matching == "-3m" + assert xtl.pseudo_symmetry_report["candidate"] == "6/mmm" + assert xtl.pseudo_symmetry_report["rejected"] is True + + +def test_projected_rotation_order(): + """Apparent zero-layer symmetry, which limits in-plane indexing.""" + bcc = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), verbose=False + ).calculate_structure_factors(k_max=1.5) + hcp = Crystal.from_ase( + bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False + ).calculate_structure_factors(k_max=1.5) + + def cartesian(xtl, uvw): + d = torch.as_tensor(np.asarray(uvw, dtype=float), dtype=torch.float64) @ xtl.lat_real + return (d / torch.linalg.norm(d)).numpy() + + # the zero-layer net of {110} along <111> is hexagonal, so the pattern + # repeats every 60 degrees while the crystal repeats every 120 + assert bcc.projected_rotation_order(cartesian(bcc, (1, 1, 1))) == 6 + assert bcc.projected_rotation_order(cartesian(bcc, (0, 0, 1))) == 4 + assert bcc.projected_rotation_order(cartesian(bcc, (0, 1, 1))) == 2 + # a general zone axis keeps the two-fold that Friedel's law provides + assert bcc.projected_rotation_order(cartesian(bcc, (1, 2, 3))) == 2 + assert hcp.projected_rotation_order(cartesian(hcp, (0, 0, 1))) == 6 + assert hcp.projected_rotation_order(cartesian(hcp, (1, 0, 0))) == 2 + + # vectorized over a stack + axes = np.stack([cartesian(bcc, u) for u in ((1, 1, 1), (0, 0, 1), (0, 1, 1))]) + assert list(bcc.projected_rotation_order(axes)) == [6, 4, 2] + + +def test_projected_order_matches_pattern_degeneracy(): + """The reported order is the rotation that leaves the pattern unchanged.""" + from quantem.diffraction.rotations import ( + qmult, + qnormalize, + quat_from_axis_angle, + quat_from_zone_axis, + ) + + bcc = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), verbose=False + ).calculate_structure_factors(k_max=1.5) + d = torch.tensor([1.0, 1.0, 1.0], dtype=torch.float64) @ bcc.lat_real + axis = d / torch.linalg.norm(d) + n = bcc.projected_rotation_order(axis.numpy()) + q = quat_from_zone_axis(axis) + beam = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + spun = qnormalize(qmult(quat_from_axis_angle(beam, torch.tensor(2 * np.pi / n)), q)) + + a = bcc.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + b = bcc.generate_pattern(spun, energy_ev=200e3, sigma_excitation=0.02) + pa = torch.stack([a["qx"], a["qy"]], dim=1) + pb = torch.stack([b["qx"], b["qy"]], dim=1) + assert pa.shape == pb.shape + # every peak of one pattern sits on a peak of the other, same intensity + dist = torch.cdist(pa, pb) + dmin, j = dist.min(dim=1) + assert float(dmin.max()) < 1e-6 + rel = (a["intensity"] - b["intensity"][j]).abs().max() / a["intensity"].max() + assert float(rel) < 1e-6 + + +_LFSO_PRISTINE_CIF = """data_ +_cell_length_a 5.1749 +_cell_length_b 8.9426 +_cell_length_c 5.1721 +_cell_angle_alpha 90 +_cell_angle_beta 109.697 +_cell_angle_gamma 90 +_symmetry_space_group_name_H-M C2/m +loop_ +_symmetry_equiv_pos_as_xyz + 'x, y, z' + '-x, y, -z' + 'x, -y, z' + '-x, -y, -z' + 'x+1/2, y+1/2, z' + '-x+1/2, y+1/2, -z' + 'x+1/2, -y+1/2, z' + '-x+1/2, -y+1/2, -z' +loop_ +_atom_site_label +_atom_site_type_symbol +_atom_site_fract_x +_atom_site_fract_y +_atom_site_fract_z +_atom_site_occupancy +Fe1 Fe 0 0.3380 0.5 0.389 +Li1 Li 0 0.3380 0.5 0.611 +Sb1 Sb 0 0 0.5 0.360 +Fe2 Fe 0 0 0.5 0.640 +Li2 Li 0 0.5 0 1 +Li3 Li 0 0.1662 0 1 +O1 O 0.7359 0.5 0.2706 1 +O2 O 0.7665 0.8415 0.2722 1 +""" + + +def test_from_cif_keeps_partial_occupancy(tmp_path): + # ASE reads a shared site as its majority species alone: this structure + # would load with no Sb at all, and Sb is its strongest scatterer + path = tmp_path / "lfso_pristine.cif" + path.write_text(_LFSO_PRISTINE_CIF) + xtl = Crystal.from_cif(path, verbose=False) + content: dict[str, float] = {} + for s, f in zip(xtl.atoms.get_chemical_symbols(), xtl.occupancy.numpy()): + content[s] = content.get(s, 0.0) + float(f) + assert content["Sb"] == pytest.approx(0.72, abs=1e-6) + assert content["Fe"] == pytest.approx(2.836, abs=1e-6) + assert content["Li"] == pytest.approx(8.444, abs=1e-6) + assert content["O"] == pytest.approx(12.0, abs=1e-6) + assert xtl.spacegroup.startswith("C2/m") + + +def test_pseudo_symmetry_names_the_breaking_reflection(tmp_path): + # the three 120 degree twin variants of the honeycomb-ordered cell differ + # only by the (020) superstructure reflection at ~0.07 of the strongest; + # at the default tolerance that rejects the layered parent, and the + # report must name the reflection responsible + path = tmp_path / "lfso_pristine.cif" + path.write_text(_LFSO_PRISTINE_CIF) + strict = Crystal.from_cif(path, verbose=False) + rep = strict.pseudo_symmetry_report + assert rep["rejected"] and rep["candidate"] == "-3m" + assert 0.05 < rep["intensity_mismatch"] < 0.1 + (h0, i0, h1, i1) = rep["broken_by"] + assert tuple(abs(v) for v in h0) == (0, 2, 0) + assert "broken by" in strict.symmetry_summary() + + +def test_pseudo_symmetry_from_parent_lattice(tmp_path): + # the honeycomb superstructure sits on a layered R-3m parent, itself on a + # rocksalt parent; neither parent's rotations map the monoclinic cell onto + # itself, so only the parent-lattice route can find them + path = tmp_path / "lfso_pristine.cif" + path.write_text(_LFSO_PRISTINE_CIF) + layered = Crystal.from_cif(path, pseudo_symmetry_intensity_tol=0.1, verbose=False) + assert layered.pointgroup_matching == "-3m" + assert layered.sym_quats_matching.shape[0] == 6 + assert "parent lattice" in layered.pseudo_symmetry_report["route"] + # the layer normal is ~[103] in the monoclinic cell + assert any("103" in lab for lab in layered.zone_axis_wedge_labels(mathtext=False)) + + rocksalt = Crystal.from_cif(path, pseudo_symmetry_intensity_tol=0.4, verbose=False) + assert rocksalt.pointgroup_matching == "m-3m" + assert rocksalt.sym_quats_matching.shape[0] == 24 + # the rocksalt parent is broken only by the (001) layer-ordering reflection + h0, i0, h1, i1 = rocksalt.pseudo_symmetry_report["broken_by"] + assert tuple(abs(v) for v in h0) == (0, 0, 1) + assert 0.3 < rocksalt.pseudo_symmetry_report["intensity_mismatch"] < 0.35 + + +_LFSO_CHARGED_CIF = """data_ +_cell_length_a 5.04848 +_cell_length_b 5.04848 +_cell_length_c 9.4279 +_cell_angle_alpha 90 +_cell_angle_beta 90 +_cell_angle_gamma 120 +_symmetry_space_group_name_H-M P-31c +loop_ +_symmetry_equiv_pos_as_xyz + 'x, y, z' + '-x, -y, -z' + '-x+y, -x, z' + '-x+y, y, -z+1/2' + '-y, -x, -z+1/2' + '-y, x-y, z' + 'y, -x+y, -z' + 'y, x, z+1/2' + 'x-y, -y, z+1/2' + 'x-y, x, -z' + '-x, -x+y, z+1/2' + 'x, x-y, -z+1/2' +loop_ +_atom_site_label +_atom_site_type_symbol +_atom_site_fract_x +_atom_site_fract_y +_atom_site_fract_z +_atom_site_occupancy +Me1 Li 0.3333333 0.6666667 0.75 0.673 +Me11 Sb 0.3333333 0.6666667 0.75 0.327 +Me2 Li 0.3333333 0.6666667 0.25 0.327 +Me22 Sb 0.3333333 0.6666667 0.25 0.673 +Me3 Fe 0 0 0.75 1 +O1 O 0.7409 0.7329 0.8660 1 +""" + + +def test_hexagonal_pseudo_symmetry_can_be_adopted(tmp_path): + # a hexagonal cell given to five decimals reproduces rotation products to + # only ~1e-6; a closure test that strict rejected every hexagonal parent + # group whatever the intensity tolerance + path = tmp_path / "lfso_charged.cif" + path.write_text(_LFSO_CHARGED_CIF) + strict = Crystal.from_cif(path, pseudo_symmetry_intensity_tol=0.4, verbose=False) + assert strict.pointgroup_matching == "-3m" + assert strict.pseudo_symmetry_report["candidate"] == "6/mmm" + loose = Crystal.from_cif(path, pseudo_symmetry_intensity_tol=1.0, verbose=False) + assert loose.pointgroup_matching == "6/mmm" + assert loose.sym_quats_matching.shape[0] == 12 + + +def test_direction_vector(): + cubic = Crystal.from_ase(bulk("Au", "fcc", a=4.08, cubic=True), verbose=False) + v = cubic.direction_vector([1, 1, 0]) + assert torch.allclose(v, torch.tensor([1.0, 1.0, 0.0], dtype=torch.float64) / np.sqrt(2)) + hcp = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False) + a1 = hcp.lat_real[0] / torch.linalg.norm(hcp.lat_real[0]) + c = hcp.lat_real[2] / torch.linalg.norm(hcp.lat_real[2]) + assert torch.allclose(hcp.direction_vector([2, -1, -1, 0]), a1) + assert torch.allclose(hcp.direction_vector([0, 0, 0, 1]), c) + # 3- and 4-index forms of the same direction agree + assert torch.allclose(hcp.direction_vector([1, 0, 0]), hcp.direction_vector([2, -1, -1, 0])) + with pytest.raises(ValueError): + hcp.direction_vector([1, 0]) + + +def test_miller_bravais_round_trip(): + from quantem.diffraction.crystal import miller_bravais_to_miller, miller_to_miller_bravais + + assert miller_to_miller_bravais([1, 0, 0]).tolist() == [2, -1, -1, 0] + assert miller_to_miller_bravais([1, 1, 0]).tolist() == [1, 1, -2, 0] + assert miller_to_miller_bravais([0, 0, 1]).tolist() == [0, 0, 0, 1] + assert miller_bravais_to_miller([2, -1, -1, 0]).tolist() == [1, 0, 0] + rng = np.random.default_rng(0) + uvw = rng.integers(-4, 5, size=(200, 3)) + uvw = uvw[np.abs(uvw).sum(axis=1) > 0] + uvtw = miller_to_miller_bravais(uvw) + assert np.all(uvtw[:, 2] == -(uvtw[:, 0] + uvtw[:, 1])) + back = miller_bravais_to_miller(uvtw) + reduced = uvw // np.gcd.reduce(np.abs(uvw), axis=1)[:, None] + assert np.array_equal(back, reduced) + + +def test_format_direction(): + from quantem.diffraction.crystal import format_direction + + bar = "̅" + assert format_direction(None) == "" + assert format_direction([1, -1, 0], mathtext=False) == "[11" + bar + "0]" + assert format_direction([1, -1, 0]) == "[1$\\bar{1}$0]" + assert format_direction([1, 0, 0], hexagonal=True, mathtext=False) == ( + "[21" + bar + "1" + bar + "0]" + ) + + +def test_spglib_no_deprecation_warnings(): + import warnings + + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False) + assert xtl.pointgroup == "6/mmm" + + +def test_generate_pattern_validates_excitation_model(ti_beta): + q = quat_from_zone_axis(torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64)) + with pytest.raises(ValueError, match="excitation_model"): + ti_beta.generate_pattern(q, excitation_model="slabb") + with pytest.raises(ValueError, match="thickness_A"): + ti_beta.generate_pattern(q, excitation_model="slab") + + +def test_generate_pattern_foil_normal(): + from quantem.diffraction.illumination import excitation_coefficients + from quantem.diffraction.rotations import qrotate + + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + c_axis = xtl.direction_vector([0, 0, 0, 1]) + + # foil normal along the beam: identical to the default geometry + q0 = quat_from_zone_axis(c_axis, in_plane_deg=10.0) + p0 = xtl.generate_pattern(q0, energy_ev=200e3) + p1 = xtl.generate_pattern(q0, energy_ev=200e3, foil_normal=(0, 0, 0, 1)) + for key in ("qx", "qy", "intensity", "s_g"): + assert torch.allclose(p0[key], p1[key], atol=1e-12) + + # a tilted flake: each spot sits where its rod meets the Ewald sphere + tilt = xtl.direction_vector([0, 1, -1, 6]) + q = quat_from_zone_axis(tilt, in_plane_deg=10.0) + p = xtl.generate_pattern(q, energy_ev=200e3, foil_normal=(0, 0, 0, 1)) + assert p["qx"].shape[0] > 5 + n_lab = qrotate(q, c_axis[None])[0] + assert float(n_lab[2]) < 0.999 # really tilted + g_lab = qrotate(q, p["hkl"].to(torch.float64) @ xtl.lat_recip) + spot = g_lab - p["s_g"][:, None] * n_lab[None] + assert torch.allclose(spot[:, :2], torch.stack([p["qx"], p["qy"]], dim=1), atol=1e-12) + s_spot, _, _ = excitation_coefficients(spot, 200e3) + # first order in s_g: the residual is far below the excitation error + s_g = p["s_g"].numpy() + big = np.abs(s_g) > 1e-3 + assert np.all(np.abs(s_spot[big]) < 0.05 * np.abs(s_g[big])) + # and moves off the projection of g, which is where it sits by default + assert float((spot[:, :2] - g_lab[:, :2]).abs().max()) > 1e-4 + + +def test_wk_factor_without_thermal_motion(): + from quantem.diffraction.wk_scattering_factors import compute_WK_factor + + g = np.linspace(0.0, 3.0, 61) + f0 = compute_WK_factor(g, 29, 200e3, thermal_sigma=None) + assert f0.dtype == np.complex128 and f0.shape == g.shape + assert np.all(np.isfinite(f0)) + # the phonon absorption vanishes continuously as the displacement -> 0 + f_small = compute_WK_factor(g, 29, 200e3, thermal_sigma=1e-3) + assert np.allclose(f_small, f0, rtol=1e-3, atol=1e-7) + # elastic part is monotonic in g, imaginary part positive + assert np.all(np.diff(f0.real) < 0) + assert np.all(compute_WK_factor(g, 29, 200e3, thermal_sigma=0.08).imag > 0) diff --git a/tests/diffraction/test_crystal_map.py b/tests/diffraction/test_crystal_map.py new file mode 100644 index 000000000..61919272f --- /dev/null +++ b/tests/diffraction/test_crystal_map.py @@ -0,0 +1,229 @@ +"""CrystalMap and PhaseMap behaviour on small synthetic alpha/beta titanium scans.""" + +import matplotlib + +matplotlib.use("Agg") + +import matplotlib.pyplot as plt +import numpy as np +import pytest +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.crystal_map import CrystalMap +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.rotations import misorientation_angle_deg, quat_from_axis_angle + +Q_ALPHA = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) +Q_BETA = quat_from_axis_angle( + torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(20.0), dtype=torch.float64), +) + + +def _crystals(): + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + ti_b = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), name="Ti beta", verbose=False + ).calculate_structure_factors(k_max=1.5) + return ti_a, ti_b + + +def _pattern(xtl, q, rng): + p = xtl.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + arr = np.column_stack([p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()]) + arr[:, :2] += rng.normal(0, 0.003, (arr.shape[0], 2)) + arr[:, 2] *= rng.lognormal(0, 0.3, arr.shape[0]) + return arr + + +def _vacuum(rng): + """Direct beam and a few detections around it: enough peaks to be + matched, but nothing diffracted.""" + arr = np.zeros((5, 3)) + arr[0, 2] = 100.0 + arr[1:, :2] = rng.uniform(-0.03, 0.03, (4, 2)) + arr[1:, 2] = 1.0 + return arr + + +def _peaks(cells): + return Vector.from_data(cells, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + + +@pytest.fixture(scope="module") +def two_phase(): + """Alpha on the left, beta on the right, one column of vacuum.""" + rng = np.random.default_rng(11) + ti_a, ti_b = _crystals() + R, C = 3, 7 + cells = [] + for _ in range(R): + row = [] + for c in range(C): + if c == C - 1: + row.append(_vacuum(rng)) + elif c < 3: + row.append(_pattern(ti_a, Q_ALPHA, rng)) + else: + row.append(_pattern(ti_b, Q_BETA, rng)) + cells.append(row) + cm = CrystalMap.from_vectors(_peaks(cells), [ti_a, ti_b], energy_ev=200e3) + cm.build_plan(angle_step_zone_axis_deg=3.0, verbose=False, progress_bar=False) + cm.match_orientations(progress_bar=False) + cm.fit(progress_bar=False) + return cm + + +def test_null_hypothesis_leaves_vacuum_unindexed(two_phase): + cm = two_phase + ph = cm.phase_index + # the vacuum column was matched (enough peaks) but diffracts nothing + assert bool(cm[0].computed[:, -1].all()) + assert (ph[:, -1] == -1).all() + assert (ph[:, :3] == 0).all() and (ph[:, 3:-1] == 1).all() + # unfit and unindexed positions carry NaN reliability, not zero + assert np.isnan(cm.phases.reliability.numpy()[:, -1]).all() + # per-crystal weights are normalized where indexed, zero elsewhere + w = cm.phases.crystal_weights.numpy() + assert np.allclose(w[ph >= 0].sum(-1), 1.0) + assert np.allclose(w[ph < 0], 0.0) + + +def test_crystal_map_mask_and_phase_fractions(two_phase): + cm = two_phase + ph = cm.phase_index + m_a = cm.mask("Ti alpha") + assert np.array_equal(m_a, cm.mask(0)) + assert (m_a[ph != 0] == 0).all() and (m_a[ph == 0] > 0).all() + m_all = cm.mask() + assert (m_all[ph == -1] == 0).all() and (m_all[ph >= 0] > 0).all() + with pytest.raises(KeyError): + cm.mask("Ti gamma") + with pytest.raises(KeyError): + cm.mask(5) + + frac = cm.phase_fractions() + assert set(frac) == {"unindexed", "Ti alpha", "Ti beta"} + assert np.isclose(sum(frac.values()), 1.0) + assert np.isclose(frac["unindexed"], (ph == -1).mean()) + assert np.isclose(frac["Ti beta"], (ph == 1).mean()) + + +def test_plot_phase_majority_filter_blacks_out_removed_positions(two_phase): + cm = two_phase + pm = cm.phases + saved = pm.phase_index.clone() + try: + # one isolated indexed position in an unindexed field: the filter + # removes it, and it must then be black, not the last phase color + lone = torch.full_like(saved, -1) + lone[1, 1] = 0 + pm.phase_index = lone + for radius, lit in ((0, True), (1, False)): + fig, ax = pm.plot_phase(majority_filter=radius, shade_by="none", scalebar=None) + rgb = np.asarray(ax.images[0].get_array()) + assert (rgb[1, 1].sum() > 0) == lit + plt.close(fig) + # and on the real decision the filter runs and draws vacuum black + pm.phase_index = saved + fig, ax = cm.plot_phase(majority_filter=1, scalebar=None) + rgb = np.asarray(ax.images[0].get_array()) + assert np.allclose(rgb[:, -1], 0.0) + plt.close(fig) + finally: + pm.phase_index = saved + + +def test_plot_phase_cycles_colors_past_the_palette(two_phase): + from quantem.diffraction.orientation_visualization import ( + DEFAULT_PHASE_COLORS, + phase_color_cycle, + ) + + n = len(DEFAULT_PHASE_COLORS) + 2 + colors = phase_color_cycle(n) + assert colors.shape == (n, 3) + assert np.allclose(colors[len(DEFAULT_PHASE_COLORS)], DEFAULT_PHASE_COLORS[0]) + # a short palette of names cycles as well + assert np.allclose(phase_color_cycle(3, ["red", "blue"])[2], (1.0, 0.0, 0.0)) + fig, ax = two_phase.plot_phase(phase_colors=np.array([[1.0, 0.0, 0.0]]), shade_by="none") + plt.close(fig) + + +def test_crystal_map_plots_smoke(two_phase): + cm = two_phase + figs = cm.plot_orientation() + assert isinstance(figs, list) and len(figs) == 4 + fig, ax = cm.plot_orientation(phase="Ti beta", direction="z") + assert len(cm.plot_orientation(phase=0)) == 2 + poles = cm.plot_pole_figure() + assert isinstance(poles, list) and len(poles) == 2 + fig, ax = cm.plot_pole_figure(pole=(0, 0, 0, 1), phase="Ti alpha", color_by="ipf") + fig, axs = cm.plot_correlation() + assert np.asarray(axs).shape == (2, 2) + fig, axs = cm.plot_correlation(mask=True, shared_scale=False) + # a single match per crystal: the default matches=(0, 1) skips the + # second, so there is one panel per crystal + fig, axs = cm.plot_matches([(0, 0), (0, 4)]) + assert axs.shape == (2, 2) + fig, axs = cm.plot_matches([(0, 0)], phase="Ti alpha") + assert axs.shape == (1, 1) + plt.close("all") + + +def test_refine_overrides_name_known_crystals(two_phase): + with pytest.raises(KeyError, match="unknown crystals"): + two_phase.refine_orientations(overrides={"Ti gamma": {}}, progress_bar=False) + + +def test_fit_without_any_paired_peak_is_unindexed(): + # every measured peak sits between the crystal's rings, so the best model + # has zero weight; argmax of zero weights must not name the first phase + ti_a, _ = _crystals() + rng = np.random.default_rng(3) + ring = np.zeros((7, 3)) + ring[0, 2] = 100.0 + phi = np.linspace(0, 2 * np.pi, 6, endpoint=False) + ring[1:, 0], ring[1:, 1], ring[1:, 2] = 0.2 * np.cos(phi), 0.2 * np.sin(phi), 5.0 + cells = [[_pattern(ti_a, Q_ALPHA, rng), ring]] + om = OrientationMap.from_vectors(_peaks(cells), ti_a, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=4.0, verbose=False, progress_bar=False) + om.match_orientations(progress_bar=False) + om.corr[0, 1, 0] = 0.5 # force the candidate into the fit + cm = CrystalMap.from_orientation_maps([om]) + cm.fit(progress_bar=False, k_max=1.0) + assert cm.phase_index[0, 0] == 0 + assert cm.phase_index[0, 1] == -1 + + +def test_single_crystal_map_end_to_end(tmp_path): + rng = np.random.default_rng(5) + ti_a, _ = _crystals() + cells = [[_pattern(ti_a, Q_ALPHA, rng) for _ in range(4)] + [_vacuum(rng)] for _ in range(2)] + cm = CrystalMap.from_vectors(_peaks(cells), ti_a, energy_ev=200e3) + assert len(cm) == 1 and cm.names == ["Ti alpha"] + cm.build_plan(angle_step_zone_axis_deg=3.0, verbose=False, progress_bar=False) + cm.match_orientations(progress_bar=False) + cm.refine_orientations(progress_bar=False) + cm.fit(progress_bar=False) + assert "refined" in repr(cm) and "phase fit" in repr(cm) + ph = cm.phase_index + assert (ph[:, :4] == 0).all() and (ph[:, 4] == -1).all() + err = misorientation_angle_deg(Q_ALPHA, cm[0].quats[:, :4, 0].reshape(-1, 4), ti_a.sym_quats) + assert float(err.max()) < 1.5 + # no runner-up crystal: reliability is NaN, and examples fall back to + # the correlation + assert np.isnan(cm.phases.reliability.numpy()).all() + picks = cm.example_positions(num=2, min_distance=1) + assert len(picks) == 2 and all(ph[p] == 0 for p in picks) + assert np.isclose(cm.phase_fractions()["Ti alpha"], 0.8) + fig, ax = cm.plot_phase() + fig, ax = cm.plot_orientation(phase=0, direction="z") + fig, ax = cm.plot_pole_figure(phase="Ti alpha") + plt.close("all") + cm.save(tmp_path / "cm.zip", mode="o") diff --git a/tests/diffraction/test_digital_dark_field.py b/tests/diffraction/test_digital_dark_field.py new file mode 100644 index 000000000..b569f82bf --- /dev/null +++ b/tests/diffraction/test_digital_dark_field.py @@ -0,0 +1,240 @@ +"""Digital dark field: apertures, polar selection and grain labels on synthetic peaks.""" + +import numpy as np +import pytest + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction import digital_dark_field as ddf + + +def _lattice_peaks(R=4, C=5, fields=("q_row", "q_col", "intensity")): + """Square lattice g1=(10,0), g2=(0,10); cells with c >= 3 also carry (5,5).""" + nested = [] + for r in range(R): + row = [] + for c in range(C): + pts = [[10 * i, 10 * j, 1.0 + r] for i in (-1, 0, 1) for j in (-1, 0, 1)] + if c >= 3: + pts.append([5.0, 5.0, 2.0]) + row.append(np.asarray(pts, dtype=float)) + nested.append(row) + return Vector.from_data(nested, fields=list(fields)) + + +def test_aperture_array_modes(): + g1, g2 = (10.0, 0.0), (0.0, 10.0) + arr = ddf.aperture_array(g1, g2, n1_range=(-1, 1), n2_range=(-1, 1)) + assert arr.shape == (9, 2) + no_center = ddf.aperture_array( + g1, g2, n1_range=(-1, 1), n2_range=(-1, 1), radius_range=(1, np.inf) + ) + assert no_center.shape == (8, 2) + line = ddf.aperture_array(g1, mode="line", n1_range=(-2, 2), center=(50, 50)) + np.testing.assert_allclose(line[:, 1], 50.0) + single = ddf.aperture_array(g1, g2, mode="single", shift=(1, 2)) + np.testing.assert_allclose(single, [[10.0, 20.0]]) + clipped = ddf.aperture_array(g1, g2, center=(15, 15), shape=(30, 30), edge=6) + np.testing.assert_allclose(clipped, [[15.0, 15.0]]) # 5 and 25 lie within the edge + with pytest.raises(ValueError): + ddf.aperture_array(g1, g2, mode="bad") + + +def test_aperture_subtract_and_image(): + fine = ddf.aperture_array((5.0, 0.0), (0.0, 5.0), n1_range=(-2, 2), n2_range=(-2, 2)) + coarse = ddf.aperture_array((10.0, 0.0), (0.0, 10.0), n1_range=(-1, 1), n2_range=(-1, 1)) + super_only = ddf.aperture_array_subtract(fine, coarse, tol=1.0) + assert super_only.shape == (25 - 9, 2) + + peaks = _lattice_peaks() + image = ddf.aperture_ddf_image(peaks, super_only, radius=1.0) + assert image.shape == (4, 5) + assert np.all(image[:, :3] == 0) + np.testing.assert_allclose(image[:, 3:], 2.0) + + # overlapping apertures count each peak once + image_full = ddf.aperture_ddf_image(peaks, np.vstack([coarse, coarse]), radius=1.0) + np.testing.assert_allclose(image_full[2, 0], 9 * 3.0) + + +def test_polar_fields_and_mask(): + peaks = _lattice_peaks(fields=("qx", "qy", "intensity")) + polar = ddf.add_polar_fields(peaks) + assert polar.fields[-2:] == ["qr", "qphi"] + flat = polar.select_fields("qx", "qy", "qr", "qphi").numpy() + np.testing.assert_allclose(flat[:, 2], np.hypot(flat[:, 0], flat[:, 1]), atol=1e-5) + # (qx, qy) = (-10, 0) is straight up on screen: +90 degrees + up = (flat[:, 0] == -10) & (flat[:, 1] == 0) + np.testing.assert_allclose(flat[up, 3], 90.0) + + ring = ddf.polar_mask(peaks, 10.0, tol=0.5) + assert ring.sum() == 4 * 20 + upper = ddf.polar_mask(peaks, 10.0, tol=0.5, phi_range=(45, 135)) + assert upper.sum() == 20 + wrapped = ddf.polar_mask(peaks, 10.0, tol=0.5, phi_range=(135, -135)) # 180 degrees + assert wrapped.sum() == 20 + image = ddf.radial_ddf_image(peaks, 10.0, tol=0.5) + np.testing.assert_allclose(image[1], 4 * 2.0) + + +def test_assign_grain_labels(): + peaks = _lattice_peaks(R=1, C=1, fields=("qx", "qy", "intensity")) + labels = np.array([0, 0, 1, 1, -1, 2, 2, 2, -1]) + labeled = peaks.copy() + labeled.add_fields("cluster", values=labels[:, None]) + out = ddf.assign_grain_labels(labeled, grain_labels=np.array([3, -1, 4])) + grains = out.select_fields("grain_label").numpy()[:, 0] + np.testing.assert_array_equal(grains, [3, 3, -1, -1, -2, 4, 4, 4, -2]) + + +def test_refine_lattice_vectors_and_group_images(): + rng = np.random.default_rng(0) + g1, g2 = np.array([20.0, 3.0]), np.array([4.0, 21.0]) + nested = [] + for r in range(3): + row = [] + for c in range(3): + n = np.array([[i, j] for i in (-2, -1, 0, 1, 2) for j in (-2, -1, 0, 1, 2)], float) + q = n @ np.stack([g1, g2]) + rng.normal(0, 0.2, (len(n), 2)) + row.append(np.concatenate([q, np.ones((len(n), 1))], axis=1)) + nested.append(row) + peaks = Vector.from_data(nested, fields=["q_row", "q_col", "intensity"]) + f1, f2 = ddf.refine_lattice_vectors( + peaks, + (19.0, 2.0), + (5.0, 20.0), + center=(0.0, 0.0), + radius=4.0, + n1_range=(-2, 2), + n2_range=(-2, 2), + ) + np.testing.assert_allclose(f1, g1, atol=0.1) + np.testing.assert_allclose(f2, g2, atol=0.1) + + a = np.zeros((4, 4)) + a[:2] = 1 + b = np.zeros((4, 4)) + b[2:] = 1 + images = np.stack([a, 2 * a, a + 0.05 * b, b, 3 * b, np.eye(4)]) + labels = ddf.group_ddf_images(images, min_correlation=0.9) + assert labels[0] == labels[1] == labels[2] >= 0 + assert labels[3] == labels[4] >= 0 and labels[3] != labels[0] + assert labels[5] == -1 + + +def test_cluster_centers_and_lattice_distance(): + peaks = _lattice_peaks(R=1, C=2) + labeled = peaks.copy() + n = labeled.total_rows + labels = np.full(n, -1) + labels[:9] = 0 # the 3x3 lattice in cell (0, 0) + labeled.add_fields("cluster", values=labels[:, None]) + centers = ddf.cluster_centers(labeled) + np.testing.assert_allclose(centers, [[0.0, 0.0]], atol=1e-6) + + d = ddf.lattice_distance([[10.0, 10.0], [5.0, 5.0], [11.0, 0.0]], (10.0, 0.0), (0.0, 10.0)) + np.testing.assert_allclose(d, [0.0, np.hypot(5, 5), 1.0]) + + +def _pixel_lattice_peaks(origin=(32.0, 32.0)): + """3 x 3 scan of a pixel lattice g1=(20, 3), g2=(4, 21) around origin.""" + g = np.array([[20.0, 3.0], [4.0, 21.0]]) + n = np.array([[i, j] for i in (-1, 0, 1) for j in (-1, 0, 1)], float) + q = np.asarray(origin) + n @ g + pts = np.concatenate([q, np.ones((len(n), 1))], axis=1) + peaks = Vector.from_data([[pts] * 3 for _ in range(3)], fields=["q_row", "q_col", "intensity"]) + return peaks, g + + +def test_refine_lattice_vectors_uses_stored_origin_for_pixel_peaks(): + peaks, g = _pixel_lattice_peaks() + with pytest.raises(ValueError, match="origin_ref"): + ddf.refine_lattice_vectors(peaks, g[0] + 1, g[1] - 1, radius=4.0) + peaks.metadata["origin_ref"] = (32.0, 32.0) + f1, f2 = ddf.refine_lattice_vectors( + peaks, g[0] + 1, g[1] - 1, radius=4.0, n1_range=(-1, 1), n2_range=(-1, 1) + ) + np.testing.assert_allclose(f1, g[0], atol=1e-4) + np.testing.assert_allclose(f2, g[1], atol=1e-4) + + # the polar functions resolve the same origin + ring = ddf.polar_mask(peaks, np.hypot(20.0, 3.0), tol=0.5) + assert ring.sum() == 9 * 2 + np.testing.assert_array_equal( + ring, ddf.polar_mask(peaks, np.hypot(20, 3), 0.5, center=(32, 32)) + ) + + +def test_aperture_array_needs_g2_for_array_mode(): + with pytest.raises(ValueError, match="g2"): + ddf.aperture_array((10.0, 0.0)) + with pytest.raises(ValueError): + ddf.aperture_array((10.0, 0.0), mode="2-beam") + half = ddf.aperture_array((10.0, 0.0), (0.0, 10.0), mode="single", shift=(0.5, 0.5)) + np.testing.assert_allclose(half, [[5.0, 5.0]]) + + +def test_ddf_images_and_cluster_coms(): + # 2 x 3 scan; cluster 0 lives in column 0, cluster 1 in row 1 + nested = [] + for r in range(2): + row = [] + for c in range(3): + pts = [[0.0, 0.0, 1.0, -1]] + if c == 0: + pts.append([5.0, 0.0, 2.0, 0]) + if r == 1: + pts.append([0.0, 5.0, 1.0 + c, 1]) + row.append(np.asarray(pts, dtype=float)) + nested.append(row) + labeled = Vector.from_data(nested, fields=["qx", "qy", "intensity", "cluster"]) + + images = ddf.ddf_images(labeled, [0, 1]) + assert images.shape == (2, 2, 3) + np.testing.assert_allclose(images[0], [[2, 0, 0], [2, 0, 0]]) + np.testing.assert_allclose(images[1], [[0, 0, 0], [1, 2, 3]]) + + coms, sizes = ddf.cluster_coms(labeled) + np.testing.assert_array_equal(sizes, [2, 3]) + np.testing.assert_allclose(coms[0], [0.5, 0.0]) + np.testing.assert_allclose(coms[1], [1.0, (0 * 1 + 1 * 2 + 2 * 3) / 6]) + coms_u, _ = ddf.cluster_coms(labeled, weighted=False) + np.testing.assert_allclose(coms_u[1], [1.0, 1.0]) + + centers = ddf.cluster_centers(labeled) + np.testing.assert_allclose(centers, [[5.0, 0.0], [0.0, 5.0]]) + + +def test_cluster_functions_accept_empty_vector(): + empty = Vector.from_data( + [[np.empty((0, 4)), np.empty((0, 4))]], fields=["qx", "qy", "intensity", "cluster"] + ) + coms, sizes = ddf.cluster_coms(empty) + assert coms.shape == (0, 2) and sizes.shape == (0,) + assert ddf.cluster_centers(empty).shape == (0, 2) + assert ddf.ddf_images(empty, [0]).shape == (1, 1, 2) + fig, ax = ddf.plot_cluster_scatter(empty) + import matplotlib.pyplot as plt + + plt.close(fig) + + +def test_plot_cluster_scatter_center(): + import matplotlib.pyplot as plt + + peaks, _ = _pixel_lattice_peaks() + labeled = peaks.copy() + labeled.add_fields("cluster", values=np.zeros((labeled.total_rows, 1))) + labeled.metadata["origin_ref"] = (32.0, 32.0) + fig, ax = ddf.plot_cluster_scatter(labeled) + assert np.isclose(np.mean(ax.get_xlim()), 32.0) + plt.close(fig) + + calibrated = Vector.from_data( + [[np.array([[0.5, 0.1, 1.0, 0], [-0.5, -0.1, 1.0, 0]])]], + fields=["qx", "qy", "intensity", "cluster"], + ) + calibrated.metadata["origin_ref"] = (128.0, 128.0) # detector pixels: ignored + fig, ax = ddf.plot_cluster_scatter(calibrated) + assert np.isclose(np.mean(ax.get_xlim()), 0.0) + assert np.isclose(np.mean(ax.get_ylim()), 0.0) + plt.close(fig) diff --git a/tests/diffraction/test_disk_detection.py b/tests/diffraction/test_disk_detection.py new file mode 100644 index 000000000..6e948de77 --- /dev/null +++ b/tests/diffraction/test_disk_detection.py @@ -0,0 +1,105 @@ +"""Correlation options of the Bragg disk detection.""" + +import numpy as np +import torch + +from quantem.diffraction.disk_detection import ( + detect_disks, + detect_disks_batch, + template_fourier, +) + +H = W = 64 +_YY, _XX = np.mgrid[0:H, 0:W] + + +def _disk(cy, cx, radius, amp): + return amp / (1 + np.exp((np.hypot(_YY - cy, _XX - cx) - radius) / 0.7)) + + +def _pattern(): + """Four disks spanning three decades of brightness on a bright halo.""" + planted = [(32, 32, 1000.0), (32, 44, 60.0), (20, 32, 12.0), (44, 20, 4.0)] + dp = sum(_disk(cy, cx, 3, a) for cy, cx, a in planted) + dp = dp + 30 * np.exp(-((_YY - 32) ** 2 + (_XX - 32) ** 2) / (2 * 25.0**2)) + dp = dp + np.random.default_rng(0).normal(0, 0.5, (H, W)) + template = torch.as_tensor(np.fft.ifftshift(_disk(32, 32, 3, 1.0)), dtype=torch.float) + return torch.as_tensor(dp, dtype=torch.float), template_fourier(template), planted + + +def _found(peaks, planted, tol=1.5): + """How many planted disks a peak list recovers.""" + if peaks.shape[0] == 0: + return 0 + return sum( + bool((np.hypot(peaks[:, 0] - cy, peaks[:, 1] - cx) < tol).any()) for cy, cx, _ in planted + ) + + +def test_corr_power_finds_weak_disks(): + """Hybrid correlation recovers disks the plain cross-correlation misses.""" + dp, tft, planted = _pattern() + common = dict(min_spacing=4.0, edge_boundary=2, max_num_peaks=50) + plain = detect_disks(dp, tft, corr_power=1.0, **common) + hybrid = detect_disks(dp, tft, corr_power=0.5, **common) + assert _found(plain, planted) < len(planted) + assert _found(hybrid, planted) == len(planted) + + +def test_batched_matches_single_with_correlation_options(): + """The batched path reproduces the per-pattern result for every option.""" + dp, tft, _ = _pattern() + common = dict(min_spacing=4.0, edge_boundary=2, max_num_peaks=50) + cases = [ + {}, + dict(corr_power=0.5), + dict(corr_power=0.0), + dict(sigma_cc=1.5), + dict(corr_power=0.7, sigma_cc=1.0, background_sigma=2.0), + ] + for kw in cases: + for subpixel in ("upsample", "parabolic"): + single = detect_disks(dp, tft, subpixel=subpixel, **common, **kw) + batch = detect_disks_batch( + torch.stack([dp, dp, dp]), tft, subpixel=subpixel, **common, **kw + )[1] + assert single.shape == batch.shape, (kw, subpixel) + # float32 round-off only: the half-spectrum and full-spectrum + # products differ in summation order once corr_power != 1 + assert np.allclose(single, batch, rtol=1e-4, atol=1e-5), (kw, subpixel) + + +def test_defaults_are_plain_cross_correlation(): + """corr_power=1 with no smoothing leaves the old behaviour untouched.""" + dp, tft, _ = _pattern() + common = dict(min_spacing=4.0, edge_boundary=2, max_num_peaks=50) + a = detect_disks(dp, tft, **common) + b = detect_disks(dp, tft, corr_power=1.0, sigma_cc=None, **common) + assert np.allclose(a, b) + + +def test_measure_origins_off_centre_beam(): + # a beam further than search_radius from the detector centre used to give + # an all-NaN measurement, whose plane fit silently returned zeros + import pytest + + from quantem.core.datastructures import Dataset4dstem + from quantem.diffraction import BraggVectors + + rows, cols = np.mgrid[0:48, 0:48] + arr = np.zeros((6, 6, 48, 48), dtype=np.float32) + for r in range(6): + for c in range(6): + cy, cx = 24.0 + 0.1 * r, 34.0 - 0.1 * c + arr[r, c] = 100 * np.exp(-((rows - cy) ** 2 + (cols - cx) ** 2) / 4.0) + arr[r, c] += 20 * np.exp(-((rows - cy - 12) ** 2 + (cols - cx) ** 2) / 4.0) + bv = BraggVectors.from_dataset(Dataset4dstem.from_array(arr)) + bv.make_template_synthetic(radius=1.5, edge=1.0) + bv.detect_disks(min_abs_intensity=1.0, min_spacing=4.0, progressbar=False) + + with pytest.raises(ValueError, match="direct beam is elsewhere"): + bv.measure_origins(search_radius=6.0) + + origins = bv.measure_origins(search_radius=6.0, center=(24.0, 34.0)) + assert np.abs(origins[..., 0] - (24.0 + 0.1 * np.arange(6)[:, None])).max() < 0.2 + assert np.abs(origins[..., 1] - (34.0 - 0.1 * np.arange(6)[None, :])).max() < 0.2 diff --git a/tests/diffraction/test_ellipse.py b/tests/diffraction/test_ellipse.py new file mode 100644 index 000000000..f0cf162b4 --- /dev/null +++ b/tests/diffraction/test_ellipse.py @@ -0,0 +1,39 @@ +"""Elliptic distortion recovery from radial histogram sharpness.""" + +import numpy as np +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction import calibration +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.rotations import qnormalize + + +def test_ellipse_recovery(): + torch.manual_seed(0) + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.9505, c=4.6855), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + + e_true = np.array([0.012, -0.008]) + A = np.array([[1 + e_true[0], e_true[1]], [e_true[1], 1 - e_true[0]]]) + A_inv = np.linalg.inv(A) + + rng = np.random.default_rng(0) + cells = [] + for i in range(60): + q = qnormalize(torch.randn(4, dtype=torch.float64)) + pat = xtl.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + qxy = np.stack([pat["qx"].numpy(), pat["qy"].numpy()], axis=1) + # distort (the inverse of the correction) plus detection noise + qxy = qxy @ A_inv.T + rng.normal(0, 0.003, qxy.shape) + cells.append(np.column_stack([qxy, pat["intensity"].numpy()])) + peaks = Vector.from_data( + [cells], + fields=["qx", "qy", "intensity"], + units=["A^-1", "A^-1", "counts"], + name="synthetic", + ) + + e_fit = calibration.calibrate_ellipse(peaks) + assert np.allclose(e_fit, e_true, atol=2e-3), (e_fit, e_true) diff --git a/tests/diffraction/test_illumination.py b/tests/diffraction/test_illumination.py new file mode 100644 index 000000000..cefb0d364 --- /dev/null +++ b/tests/diffraction/test_illumination.py @@ -0,0 +1,166 @@ +"""Illumination-averaged excitation envelopes (precession ring, convergence disk).""" + +import numpy as np +import torch +from ase.build import bulk + +from quantem.diffraction import bloch +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.illumination import ( + excitation_coefficients, + gaussian_envelope, + gaussian_envelope_ring_series, + ring_disk_quadrature, + slab_envelope, +) + + +def _quad_reference(fn, c, a, b, n_phi=256, n_r=12, n_psi=64): + """Positive angular quadrature of fn(s) over s = c + ring_x + disk_x.""" + t, w = ring_disk_quadrature(a, b, n_phi=n_phi, n_r=n_r, n_psi=n_psi) + return w @ fn(c + t[:, 0]) + + +def test_gaussian_envelope_limits_and_quadrature(): + sigma = 0.025 + c = np.linspace(-0.1, 0.1, 21) + # no illumination: the static envelope, exactly + assert np.allclose(gaussian_envelope(c, 0.0, 0.0, sigma), np.exp(-0.5 * (c / sigma) ** 2)) + # ring, disk, and both, against positive quadrature of the static envelope + for a, b in ((0.08, 0.0), (0.0, 0.02), (0.08, 0.02), (0.01, 0.005)): + for ci in (0.0, 0.03, 0.07): + ref = _quad_reference(lambda s: np.exp(-0.5 * (s / sigma) ** 2), ci, a, b) + assert abs(gaussian_envelope(ci, a, b, sigma) - ref) < 1e-9 + # the ring series agrees with the transform + assert np.allclose( + gaussian_envelope_ring_series(c, 0.009, 0.04), + gaussian_envelope(c, 0.009, 0.0, 0.04), + atol=1e-11, + ) + # torch in, torch out + out = gaussian_envelope_ring_series(torch.as_tensor(c), torch.full((21,), 0.009), 0.04) + assert isinstance(out, torch.Tensor) and out.shape == (21,) + + +def test_slab_envelope_limits_and_quadrature(): + z = 80.0 + c = np.linspace(-0.1, 0.1, 21) + assert np.allclose(slab_envelope(c, 0.0, 0.0, z), np.sinc(c * z) ** 2, atol=1e-10) + assert np.isclose(slab_envelope(np.array([0.05]), 0.0, 0.0, 0.0)[0], 1.0) + for a, b in ((0.08, 0.0), (0.0, 0.02), (0.08, 0.02)): + for ci in (0.0, 0.03): + ref = _quad_reference(lambda s: np.sinc(s * z) ** 2, ci, a, b) + assert abs(slab_envelope(np.array([ci]), a, b, z)[0] - ref) < 1e-9 + + +def test_centered_ring_coefficients_exact(): + # s_g(phi) = c_g + a_g cos(phi - delta_g) exactly for a centered ring + energy_ev, prec = 200e3, 0.5 + lam = bloch.electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + r = k0 * np.sin(np.deg2rad(prec)) + g = np.array([[0.8, 0.1, 0.02], [-0.3, 0.65, -0.015], [0.2, -0.9, 0.05]]) + c, a, b = excitation_coefficients(g, energy_ev, prec, 0.0) + assert np.all(b == 0) + phi = 2 * np.pi * (np.arange(360) + 0.3) / 360 + t = r * np.stack([np.cos(phi), np.sin(phi)], axis=1) + kz = np.sqrt(k0**2 - (t**2).sum(1))[:, None] + s = (2 * kz * g[:, 2] - 2 * (t @ g[:, :2].T) - (g**2).sum(1)) / (2 * (kz - g[:, 2])) + delta = np.arctan2(-g[:, 1], -g[:, 0]) + model = c[None] + a[None] * np.cos(phi[:, None] - delta[None]) + assert np.abs(s - model).max() < 1e-12 + # zero illumination: c is the static excitation error + c0, a0, b0 = excitation_coefficients(g, energy_ev, 0.0, 0.0) + s0 = (2 * k0 * g[:, 2] - (g**2).sum(1)) / (2 * (k0 - g[:, 2])) + assert np.allclose(c0, s0) and np.all(a0 == 0) and np.all(b0 == 0) + + +def test_slab_pattern_is_thin_bloch_limit(): + """The slab model with elastic couplings equals the Bloch calculation + for a very thin crystal (first Born), including the 000-relative + normalization, at zero and at nonzero precession.""" + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=2.0) # elastic Lobato factors only + torch.manual_seed(1) + from quantem.diffraction.rotations import qnormalize + + q = qnormalize(torch.randn(4, dtype=torch.float64)) + z = 6.0 + for prec in (0.0, 0.4): + pat = xtl.generate_pattern( + q, + 200e3, + tol_excitation_mult=4.0, + k_max=1.0, + precession_deg=prec, + excitation_model="slab", + thickness_A=z, + ) + ring, w = bloch.illumination_nodes(200e3, precession_deg=prec, n_precession=32) + inten, g_xy, hkl = bloch._cbed_amplitudes( + xtl, q, ring, torch.tensor([z]), 200e3, 0.3, 1.0, tilt_batch=64 + ) + dyn = (inten[:, 0, :] * w[:, None]).sum(0) + lut = {tuple(h): i for i, h in enumerate(hkl.tolist())} + common = [k for k, h in enumerate(pat["hkl"].tolist()) if tuple(h) in lut] + idx = [lut[tuple(pat["hkl"][k].tolist())] for k in common] + born = pat["intensity"].numpy()[common] + strong = born > 0.2 * born.max() + assert strong.sum() >= 4 + ratio = dyn[idx].numpy()[strong] / born[strong] + # first Born holds to a few percent for the strong reflections at + # 6 A (the remainder is the second-order multi-beam term); the + # illumination average is shared exactly by both sides + assert np.abs(ratio - 1).max() < 0.05 + + +def test_refine_batched_matches_loop(): + from quantem.core.datastructures.vector import Vector + from quantem.diffraction.orientation import OrientationMap + from quantem.diffraction.rotations import misorientation_angle_deg, qnormalize + + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + torch.manual_seed(2) + N = 6 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + p = xtl.generate_pattern(q_true[i], 200e3, sigma_excitation=0.02, precession_deg=0.5) + peaks[0, i] = np.stack([p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1) + out = [] + for batched in (True, False): + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3, precession_deg=0.5) + om.build_plan(angle_step_zone_axis_deg=3.0, verbose=False) + om.match_orientations(progress_bar=False) + om.refine_orientations(batched=batched, neighbor_rescue=False, progress_bar=False) + out.append(om.quats[0, :, 0].clone()) + d = misorientation_angle_deg(out[0], out[1], xtl.sym_quats).numpy() + assert d.max() < 1e-4 + + +def test_relrod_factor(): + from quantem.diffraction.illumination import relrod_factor + + rng = np.random.default_rng(0) + g = rng.normal(0, 0.5, (40, 3)) + g[:, 2] *= 0.05 + # normal along the beam: the excitation error is already along the rod + assert np.allclose(relrod_factor(g, np.array([0.0, 0.0, 1.0]), 200e3), 1.0) + # tilted normal: g + t n with t = -f s_g lies on the Ewald sphere to + # first order in s_g + n = np.array([np.sin(0.3), 0.0, np.cos(0.3)]) + f = relrod_factor(g, n, 200e3, precession_deg=0.0) + s, _, _ = excitation_coefficients(g, 200e3) + spot = g - (f * s)[:, None] * n[None] + s_spot, _, _ = excitation_coefficients(spot, 200e3) + big = np.abs(s) > 1e-3 + assert np.all(np.abs(s_spot[big]) < 0.05 * np.abs(s[big])) + # torch in, torch out, same values + f_t = relrod_factor(torch.as_tensor(g), torch.as_tensor(n), 200e3) + assert isinstance(f_t, torch.Tensor) and np.allclose(f_t.numpy(), f) + # an edge-on plate is never excited + edge = relrod_factor(np.array([[0.0, 0.0, 0.0]]), np.array([1.0, 0.0, 0.0]), 200e3) + assert edge[0] >= 1e6 diff --git a/tests/diffraction/test_orientation.py b/tests/diffraction/test_orientation.py new file mode 100644 index 000000000..d6e6c3228 --- /dev/null +++ b/tests/diffraction/test_orientation.py @@ -0,0 +1,888 @@ +"""Round-trip tests for quantem.diffraction.orientation.""" + +import numpy as np +import pytest +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.rotations import misorientation_angle_deg, qnormalize + + +def _make_peaks(xtl, q_true, sigma=0.02): + N = q_true.shape[0] + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + p = xtl.generate_pattern(q_true[i], energy_ev=200e3, sigma_excitation=sigma) + peaks[0, i] = np.stack([p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1) + return peaks + + +@pytest.mark.parametrize( + "builder,kwargs", + [ + (bulk, dict(name="Ti", crystalstructure="bcc", a=3.31, cubic=True)), + (bulk, dict(name="Ti", crystalstructure="hcp", a=2.95, c=4.686)), + ], +) +def test_roundtrip_matching(builder, kwargs): + torch.manual_seed(3) + xtl = Crystal.from_ase(builder(**kwargs)) + xtl.calculate_structure_factors(k_max=1.5) + N = 15 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0, power_intensity=0.0) + om.match_orientations(progress_bar=False) + # noiseless synthetic data: the envelope tilt is exact, so allow the + # full grid-scale correction (the default trust region is sized for + # noisy measured intensities) + om.refine_orientations(zone_max_total_deg=1.5, progress_bar=False) + + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + # majority recovered to well below the grid step; a small number of + # kinematically (near-)degenerate orientations may land elsewhere + assert np.median(err) < 0.1 + assert (err < 1.0).mean() >= 0.7 + + +def test_normalized_scores_and_reliability(): + torch.manual_seed(0) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + q_true = qnormalize(torch.randn(6, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, angle_step_in_plane_deg=3.0) + om.match_orientations(progress_bar=False) + + assert float(om.corr.max()) <= 1.0 + 1e-9 + assert float(om.corr.min()) >= 0.0 + assert om.reliability is not None + assert (om.reliability[0] > 0).all() + + +def test_mirror_channel(): + """Orientations in the opposite hemisphere are matched via the mirror.""" + torch.manual_seed(5) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + q_true = qnormalize(torch.randn(10, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0) + om.match_orientations(progress_bar=False) + om.refine_orientations(progress_bar=False) + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + used_mirror = om.mirror[0, :, 0].numpy() + # both channels appear and mirror matches are as accurate as direct ones + assert used_mirror.any() + assert (~used_mirror).any() + ok = err < 5 + assert ok.mean() >= 0.7 + assert np.median(err[ok & used_mirror]) < 0.5 + + +def test_square_detector_correction(): + """Peaks clipped by a square detector: the aperture-normalized match + recovers the orientation as well as the unclipped case.""" + torch.manual_seed(7) + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686)) + xtl.calculate_structure_factors(k_max=1.5) + q_true = qnormalize(torch.randn(10, 4, dtype=torch.float64)) + q_det = 0.9 # detector half-width < k_max: corners clipped + + peaks = Vector.from_shape( + (1, 10), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(10): + p = xtl.generate_pattern(q_true[i], energy_ev=200e3, sigma_excitation=0.02) + keep = (p["qx"].abs() < q_det) & (p["qy"].abs() < q_det) + peaks[0, i] = np.stack( + [p["qx"][keep].numpy(), p["qy"][keep].numpy(), p["intensity"][keep].numpy()], + axis=1, + ) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan( + angle_step_zone_axis_deg=2.0, + angle_step_in_plane_deg=2.0, + detector_q_max=q_det, + ) + om.match_orientations(progress_bar=False) + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + assert (err < 5).mean() >= 0.7 + # with the aperture correction, kernel leakage at the hard detector edge + # can push the normalized score a few percent above 1 + assert float(om.corr.max()) <= 1.05 + + +def _ase(spacegroup, symbols, basis, cellpar): + from ase.spacegroup import crystal as ase_crystal + + return ase_crystal(symbols, basis=basis, spacegroup=spacegroup, cellpar=cellpar) + + +@pytest.mark.parametrize( + "label,atoms,step", + [ + ( + "Bi -3m", + lambda: _ase(166, ["Bi"], [(0, 0, 0.234)], [4.55, 4.55, 11.86, 90, 90, 120]), + 2.0, + ), + # ilmenite's projections are nearly mirror symmetric, so the flipped + # orientation is a close rival and needs a finer zone grid than Bi + ( + "ilmenite -3", + lambda: _ase( + 148, + ["Fe", "Ti", "O"], + [(0, 0, 0.355), (0, 0, 0.146), (0.317, 0.023, 0.245)], + [5.09, 5.09, 14.09, 90, 90, 120], + ), + 1.5, + ), + ], +) +def test_roundtrip_low_symmetry(label, atoms, step): + # low-symmetry crystals see errors that cubic and hexagonal symmetry + # hides: a wrong wedge (trigonal) or a redundant library + torch.manual_seed(5) + xtl = Crystal.from_ase(atoms(), name=label, verbose=False) + xtl.calculate_structure_factors(k_max=1.3) + N = 20 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=step, angle_step_in_plane_deg=2.0, verbose=False) + om.match_orientations(progress_bar=False) + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + assert (err < 2.5).mean() >= 0.8 + # symmetry copies of the best zone must not count as the second best + assert float(np.median(om.reliability[0].numpy())) > 0.02 + + +def test_reliability_with_hemisphere_library(): + # Laue 2/m has no wedge: the hemisphere library holds every zone twice, + # and reliability must still see past the symmetry copy + torch.manual_seed(2) + atoms = _ase( + 14, + ["Zr", "O", "O"], + [(0.275, 0.040, 0.208), (0.070, 0.332, 0.345), (0.450, 0.758, 0.479)], + [5.15, 5.21, 5.32, 90, 99.2, 90], + ) + xtl = Crystal.from_ase(atoms, name="ZrO2", pseudo_symmetry_tol=None, verbose=False) + xtl.calculate_structure_factors(k_max=1.2) + assert xtl.zone_axis_wedge() is None + q_true = qnormalize(torch.randn(8, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, angle_step_in_plane_deg=3.0, verbose=False) + om.match_orientations(progress_bar=False) + assert float(np.median(om.reliability[0].numpy())) > 0.02 + + +def test_pseudo_symmetry_warning_on_plan(): + import warnings + + from ase import Atoms + + ortho = Atoms("Au", positions=[[0, 0, 0]], cell=[4.000, 4.001, 4.002], pbc=True) + xtl = Crystal.from_ase(ortho, verbose=False) + xtl.calculate_structure_factors(k_max=1.2) + q_true = qnormalize(torch.randn(2, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + om.build_plan(angle_step_zone_axis_deg=5.0, angle_step_in_plane_deg=5.0, verbose=False) + assert any("pseudo-symmetry" in str(x.message) for x in w) + + +def test_metadata_inheritance(): + # each stage records its hyperparameters; later stages inherit what is + # left as None, so one tuned value propagates through the whole chain + from quantem.diffraction import bloch + from quantem.diffraction.phase import PhaseMap + + torch.manual_seed(1) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + q_true = qnormalize(torch.randn(4, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3, precession_deg=0.7) + om.build_plan( + angle_step_zone_axis_deg=3.0, corr_kernel_size=0.04, power_intensity=0.3, verbose=False + ) + om.match_orientations(progress_bar=False) + om.refine_orientations(progress_bar=False) + assert om.metadata["plan"]["pair_distance"] == 0.04 + assert om.metadata["refine"]["pair_distance"] == 0.04 + assert om.metadata["match"]["min_number_peaks"] == 5 + pm = PhaseMap.from_orientation_maps([om]) + pm.fit(progress_bar=False) + assert pm.metadata["fit"]["pair_distance"] == 0.04 + assert pm.metadata["fit"]["power_intensity"] == 0.3 + assert pm.metadata["fit"]["min_number_peaks"] == 5 + xtl.calculate_dynamical_structure_factors(energy_ev=200e3, k_max=2.0) + res = bloch.refine_dynamical( + pm, + thicknesses_A=[300.0, 400.0], + tilt_stages=((0.1, 0.1),), + n_precession=4, + k_max=1.0, + mask=np.array([[True, False, False, False]]), + progress_bar=False, + ) + md = res["metadata"] + assert md["precession_deg"] == 0.7 and md["pair_distance"] == 0.04 + assert md["power_intensity"] == 0.3 and md["min_number_peaks"] == 5 + assert pm.metadata["dynamical"] is md + # explicit values still win + res2 = bloch.refine_dynamical( + pm, + thicknesses_A=[300.0], + tilt_stages=((0.1, 0.1),), + n_precession=4, + precession_deg=0.0, + pair_distance=0.06, + k_max=1.0, + mask=np.array([[True, False, False, False]]), + progress_bar=False, + ) + assert res2["metadata"]["precession_deg"] == 0.0 and res2["metadata"]["pair_distance"] == 0.06 + + +def test_precession_envelope_matches_quadrature(): + # the analytic ring-averaged envelope equals the positive quadrature of + # the static envelope over the exact excitation errors on the ring + from quantem.diffraction.illumination import ring_disk_quadrature + from quantem.diffraction.rotations import qrotate + + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + torch.manual_seed(4) + q = qnormalize(torch.randn(4, dtype=torch.float64)) + energy_ev, sigma, prec = 200e3, 0.04, 0.6 + from quantem.core.utils.utils import electron_wavelength_angstrom + + lam = electron_wavelength_angstrom(energy_ev) + k0 = 1.0 / lam + pat = xtl.generate_pattern(q, energy_ev, sigma_excitation=sigma, precession_deg=prec) + g = qrotate(q, xtl.g_vec) + hkl_map = {tuple(h): i for i, h in enumerate(xtl.hkl.tolist())} + idx = torch.tensor([hkl_map[tuple(h)] for h in pat["hkl"].tolist()]) + gs = g[idx].numpy() + r = k0 * np.sin(np.deg2rad(prec)) + t, w = ring_disk_quadrature(r, 0.0, n_phi=256) + kz = np.sqrt(k0**2 - (t**2).sum(1))[:, None] + s = (2 * kz * gs[:, 2] - 2 * (t @ gs[:, :2].T) - (gs**2).sum(1)) / (2 * (kz - gs[:, 2])) + ref = (w[:, None] * np.exp(-0.5 * (s / sigma) ** 2)).sum(0) * xtl.struct_factors_int[ + idx + ].numpy() + assert np.allclose(pat["intensity"].numpy(), ref, rtol=1e-6, atol=1e-9) + # without precession the static envelope is recovered exactly + pat0 = xtl.generate_pattern(q, energy_ev, sigma_excitation=sigma) + s0 = pat0["s_g"].numpy() + assert np.allclose( + pat0["intensity"].numpy(), + xtl.struct_factors_int[[hkl_map[tuple(h)] for h in pat0["hkl"].tolist()]].numpy() + * np.exp(-0.5 * (s0 / sigma) ** 2), + ) + + +def test_roundtrip_with_precession(): + # library, matching and refinement with the precession-averaged + # envelope: patterns simulated with precession are recovered + torch.manual_seed(6) + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + N = 12 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + p = xtl.generate_pattern( + q_true[i], energy_ev=200e3, sigma_excitation=0.02, precession_deg=0.7 + ) + peaks[0, i] = np.stack([p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3, precession_deg=0.7) + om.build_plan(angle_step_zone_axis_deg=2.0, verbose=False) + om.match_orientations(progress_bar=False) + om.refine_orientations(zone_max_total_deg=1.5, progress_bar=False) + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + assert np.median(err) < 0.3 + assert (err < 1.5).mean() >= 0.75 + assert om.metadata["precession_deg"] == 0.7 + + +def test_staged_positions_subset(): + """Matching a few positions leaves the rest untouched, and the later + stages follow the subset without repeating it.""" + torch.manual_seed(5) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + N = 6 + q_true = qnormalize(torch.randn(N, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0, power_intensity=0.0) + + test_pos = [(0, 1), (0, 4)] + om.match_orientations(positions=test_pos, progress_bar=False) + assert om.computed.sum() == len(test_pos) + assert bool(om.computed[0, 1]) and bool(om.computed[0, 4]) + assert float(om.corr[0, 0, 0]) == 0.0 # not requested, untouched + assert float(om.corr[0, 1, 0]) > 0.5 + + # refinement follows `computed` with no position list of its own + before = om.quats.clone() + om.refine_orientations(progress_bar=False, zone_max_total_deg=1.5) + untouched = torch.allclose(before[0, 0], om.quats[0, 0]) + assert untouched + err = misorientation_angle_deg(q_true[[1, 4]], om.quats[0, [1, 4], 0], xtl.sym_quats).numpy() + assert np.all(err < 1.0) + + # the full run then covers everything + om.match_orientations(progress_bar=False) + assert bool(om.computed.all()) + assert float(om.corr[0, 0, 0]) > 0.5 + + +def test_fiber_zone_axis_range(): + """A fiber plan of zero half angle samples one zone axis and still + recovers the in-plane angle exactly.""" + torch.manual_seed(7) + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.9505, c=4.6855)) + xtl.calculate_structure_factors(k_max=1.5) + N = 6 + gam = torch.rand(N, dtype=torch.float64) * 2 * np.pi + q_true = qnormalize( + torch.stack( + [torch.cos(gam / 2), torch.zeros(N), torch.zeros(N), torch.sin(gam / 2)], dim=1 + ) + ) + peaks = _make_peaks(xtl, q_true) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan( + zone_axis_range="fiber", + fiber_axis=[0, 0, 0, 1], # Miller-Bravais [0001] + fiber_angle_deg=0.0, + angle_step_in_plane_deg=2.0, + power_intensity=0.0, + verbose=False, + ) + assert om.zone_axes.shape[0] == 1 + om.match_orientations(progress_bar=False) + om.refine_orientations(progress_bar=False, zone_max_total_deg=0.5) + err = misorientation_angle_deg(q_true, om.quats[0, :, 0], xtl.sym_quats).numpy() + assert np.all(err < 0.2) + + # a cap of a few degrees covers a spread of tilts, and the hemisphere + # fallback and the symmetry wedge both stay available + om2 = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om2.build_plan( + zone_axis_range="fiber", fiber_axis=[0, 0, 1], fiber_angle_deg=5.0, verbose=False + ) + assert om2.zone_axes.shape[0] > 1 + om3 = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om3.build_plan(zone_axis_range="full", angle_step_zone_axis_deg=4.0, verbose=False) + om4 = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om4.build_plan(angle_step_zone_axis_deg=4.0, verbose=False) + assert om3.zone_axes.shape[0] > om4.zone_axes.shape[0] + + +def test_power_intensity_experiment_is_separate(): + """The measured-intensity exponent defaults to the library one and can + be set independently.""" + torch.manual_seed(11) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + q_true = qnormalize(torch.randn(3, 4, dtype=torch.float64)) + peaks = _make_peaks(xtl, q_true) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, power_intensity=0.25, verbose=False) + assert om.power_intensity_experiment == 0.25 + + om2 = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om2.build_plan( + angle_step_zone_axis_deg=3.0, + power_intensity=0.25, + power_intensity_experiment=0.0, + verbose=False, + ) + assert om2.power_intensity_experiment == 0.0 + assert om2.metadata["plan"]["power_intensity_experiment"] == 0.0 + om2.match_orientations(progress_bar=False) + assert float(om2.corr[0, 0, 0]) > 0.3 + + +def test_in_plane_angle_auto_fold(): + """The automatic fold removes the in-plane ambiguity of a <111> zone. + + Two orientations 60 degrees apart about a body-centered cubic <111> beam + give the same zero-layer pattern, so matching returns one or the other at + random. Folding by the projected order makes the reported angle the same + for both, which is what keeps an in-plane map continuous. + """ + from quantem.diffraction.rotations import ( + qmult, + quat_from_axis_angle, + quat_from_zone_axis, + ) + + torch.manual_seed(2) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.26, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + d = torch.tensor([1.0, 1.0, 1.0], dtype=torch.float64) @ xtl.lat_real + q0 = quat_from_zone_axis(d / torch.linalg.norm(d)) + beam = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + + # N in-plane angles, each also present as its 60 degree twin + N = 5 + spin = torch.linspace(0.0, 1.0, N, dtype=torch.float64) + q_a = qnormalize(qmult(quat_from_axis_angle(beam, spin), q0)) + q_b = qnormalize(qmult(quat_from_axis_angle(beam, torch.tensor(np.deg2rad(60.0))), q_a)) + q_true = torch.stack([q_a, q_b], dim=0) # (2, N, 4) + + peaks = Vector.from_shape( + (2, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(2): + for j in range(N): + p = xtl.generate_pattern(q_true[i, j], energy_ev=200e3, sigma_excitation=0.02) + peaks[i, j] = np.stack( + [p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1 + ) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0, power_intensity=0.0) + om.match_orientations(progress_bar=False) + om.refine_orientations(progress_bar=False, zone_max_total_deg=1.5) + + assert xtl.projected_rotation_order((d / torch.linalg.norm(d)).numpy()) == 6 + folded = om.in_plane_angle_deg(mod_deg="auto").numpy() + assert folded.max() <= 60.0 + 1e-6 + # the twin rows must agree once folded, to well under the library step + delta = np.abs(folded[0] - folded[1]) % 60.0 + delta = np.minimum(delta, 60.0 - delta) + assert np.all(delta < 1.0), delta + # explicit values and None still behave as before + assert om.in_plane_angle_deg(mod_deg=90.0).max() <= 90.0 + assert om.in_plane_angle_deg(mod_deg=None).max() > 60.0 + + +def test_fold_in_plane_collapses_degenerate_variants(): + """Two orientations with the same zero-layer pattern get the same color.""" + from quantem.diffraction.orientation_visualization import fold_in_plane, ipf_color + from quantem.diffraction.rotations import ( + qmult, + quat_from_axis_angle, + quat_from_zone_axis, + ) + + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.26, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + d = torch.tensor([1.0, 1.0, 1.0], dtype=torch.float64) @ xtl.lat_real + q0 = quat_from_zone_axis(d / torch.linalg.norm(d)) + beam = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + tilt_axis = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64) + + torch.manual_seed(4) + for tilt in (0.0, 1.8, 3.0): + base = qnormalize( + qmult(quat_from_axis_angle(tilt_axis, torch.tensor(np.deg2rad(tilt))), q0) + ) + spin = torch.rand(8, dtype=torch.float64) * 2 * np.pi + q_a = qnormalize(qmult(quat_from_axis_angle(beam, spin), base)) + q_b = qnormalize(qmult(quat_from_axis_angle(beam, torch.tensor(np.deg2rad(60.0))), q_a)) + folded_a = fold_in_plane(q_a, xtl) + folded_b = fold_in_plane(q_b, xtl) + assert torch.allclose(torch.abs(folded_a), torch.abs(folded_b), atol=1e-8) + c_a = ipf_color(folded_a, xtl, "r") + c_b = ipf_color(folded_b, xtl, "r") + assert np.abs(c_a - c_b).max() < 1e-6, tilt + # the out-of-plane color never depended on the in-plane angle + assert np.abs(ipf_color(q_a, xtl, "z") - ipf_color(q_b, xtl, "z")).max() < 1e-6 + + # far from the pole the ambiguity is gone and nothing is folded + far = qnormalize(qmult(quat_from_axis_angle(tilt_axis, torch.tensor(np.deg2rad(20.0))), q0)) + assert torch.allclose(fold_in_plane(far[None], xtl)[0], far, atol=1e-12) + + +def test_smooth_orientations(): + """Bilateral smoothing averages noise inside a grain, not across variants.""" + from quantem.diffraction.rotations import qmult, quat_from_axis_angle + + torch.manual_seed(6) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + + # a 6 x 6 patch of one orientation with half a degree of scatter, plus a + # second grain 30 degrees away filling the right-hand columns + R, C = 6, 6 + base = qnormalize(torch.randn(4, dtype=torch.float64)) + axis = torch.randn(R, C, 3, dtype=torch.float64) + axis = axis / axis.norm(dim=-1, keepdim=True) + noise = quat_from_axis_angle(axis.reshape(-1, 3), torch.deg2rad(0.5 * torch.randn(R * C))) + q = qnormalize(qmult(noise, base)).reshape(R, C, 4) + other = qnormalize( + qmult( + quat_from_axis_angle(torch.tensor([0.0, 0.0, 1.0]), torch.tensor(np.deg2rad(30.0))), + base, + ) + ) + q[:, 4:] = other + + peaks = Vector.from_shape( + (R, C), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(R): + for j in range(C): + p = xtl.generate_pattern(q[i, j], energy_ev=200e3, sigma_excitation=0.02) + peaks[i, j] = np.stack( + [p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1 + ) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, verbose=False) + om.quats = q.clone()[..., None, :] + om.corr = torch.ones((R, C, 1), dtype=torch.float64) + om.computed = torch.ones((R, C), dtype=torch.bool) + + before = misorientation_angle_deg(base, om.quats[:, :4, 0].reshape(-1, 4), xtl.sym_quats) + om.smooth_orientations(sigma_px=1.0, sigma_deg=1.0, max_angle_deg=5.0) + after = misorientation_angle_deg(base, om.quats[:, :4, 0].reshape(-1, 4), xtl.sym_quats) + assert float(after.mean()) < float(before.mean()), (float(before.mean()), float(after.mean())) + + # the second grain is 30 degrees away, beyond max_angle_deg, so it is + # neither pulled toward the first nor allowed to pull on it + kept = misorientation_angle_deg(other, om.quats[:, 5, 0], xtl.sym_quats) + assert float(kept.max()) < 1e-6 + + +def test_display_smoothing_needs_both_widths(): + """A bare number is refused: the angular tolerance must not default silently.""" + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + torch.manual_seed(8) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + R, C = 4, 4 + q = qnormalize(torch.randn(4, dtype=torch.float64)).expand(R, C, 4).clone() + peaks = Vector.from_shape( + (R, C), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(R): + for j in range(C): + p = xtl.generate_pattern(q[i, j], energy_ev=200e3, sigma_excitation=0.02) + peaks[i, j] = np.stack( + [p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1 + ) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=4.0, verbose=False) + om.quats = q.clone()[..., None, :] + om.corr = torch.ones((R, C, 1), dtype=torch.float64) + om.computed = torch.ones((R, C), dtype=torch.bool) + + with pytest.raises(TypeError, match="angular tolerance"): + om.plot_orientation(smooth=1.0) + + before = om.quats.clone() + om.plot_orientation(smooth={"sigma_px": 1.0, "sigma_deg": 1.0, "max_angle_deg": 5.0}) + om.plot_orientation(smooth=True) + plt.close("all") + # smoothing for display must never touch the stored orientations + assert torch.equal(before, om.quats) + + +def test_rescue_breaks_friedel_ties_by_neighbours(): + """At a zone axis the pattern rotated 180 degrees about the beam is the + same pattern, so a scattered twin can only be undone by its neighbours.""" + from quantem.diffraction.rotations import qmult, quat_from_zone_axis + + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686)) + xtl.calculate_structure_factors(k_max=1.5) + zone = torch.tensor([1.0, 0.0, 1.0], dtype=torch.float64) @ xtl.lat_real.to(torch.float64) + q = quat_from_zone_axis(zone, 20.0) + twin = qmult(torch.tensor([0.0, 0.0, 0.0, 1.0], dtype=torch.float64), q) + assert float(misorientation_angle_deg(q, twin, xtl.sym_quats_matching)) > 10.0 + + R, C = 5, 5 + peaks = Vector.from_shape( + (R, C), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + p = xtl.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + for i in range(R): + for j in range(C): + peaks[i, j] = np.stack( + [p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], axis=1 + ) + flipped = torch.zeros((R, C), dtype=torch.bool) + flipped[1, 1] = flipped[2, 3] = flipped[3, 1] = True + + def run(consensus_tol): + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, verbose=False, progress_bar=False) + om.quats = torch.where(flipped[..., None], twin, q).clone()[..., None, :] + om.corr = torch.ones((R, C, 1), dtype=torch.float64) + om.computed = torch.ones((R, C), dtype=torch.bool) + om.refine_orientations(consensus_tol=consensus_tol, progress_bar=False) + return misorientation_angle_deg(q, om.quats[..., 0, :], xtl.sym_quats_matching) + + # each pattern alone cannot tell the twin apart ... + assert float(run(0.0)[flipped].min()) > 10.0 + # ... but it is a tie, and every neighbour holds the other variant + assert float(run(0.01).max()) < 2.0 + + +def test_ipf_key_is_smooth_with_exact_corners_and_white_centre(): + from quantem.diffraction.orientation_visualization import IPF_CORNER_COLORS, _bary_to_rgb + + assert np.allclose(_bary_to_rgb(np.eye(3)), IPF_CORNER_COLORS) + assert np.allclose(_bary_to_rgb(np.ones(3) / 3), 1) + # no creases: along lines across the wedge, including across the lines + # where one corner takes over from another, the color turns gently + t = np.linspace(0, 1, 801)[:, None] + for p0, p1 in ( + ((0.55, 0.40, 0.05), (0.20, 0.10, 0.70)), + ((0.90, 0.05, 0.05), (0.05, 0.90, 0.05)), + ((0.70, 0.30, 0.00), (0.00, 0.30, 0.70)), + ): + c = _bary_to_rgb(np.array(p0) * (1 - t) + np.array(p1) * t) + assert np.abs(np.diff(c, 2, axis=0)).max() < 1e-4 + # the edges stay colored all along: midway between two corners is vivid + for i, j in ((0, 1), (1, 2), (2, 0)): + w = np.zeros(3) + w[i] = w[j] = 0.5 + rgb = _bary_to_rgb(w) + assert rgb.max() - rgb.min() > 0.5 + + +def test_wedge_labels_name_zone_axes_with_positive_leading_index(): + import re + + for xtl in ( + Crystal.from_ase(bulk("Cu", "fcc", a=3.6, cubic=True), verbose=False), + Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.68), verbose=False), + ): + for label in xtl.zone_axis_wedge_labels(mathtext=False): + # the first nonzero digit carries no overbar + m = re.search("[1-9]", label) + assert m is not None + assert label[m.end() : m.end() + 1] != "\u0305", label + + +def test_plot_matches_background_norm(): + import matplotlib + + matplotlib.use("Agg") + from types import SimpleNamespace + + from quantem.diffraction.orientation_visualization import plot_pattern_matches + + torch.manual_seed(0) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True)) + xtl.calculate_structure_factors(k_max=1.5) + peaks = _make_peaks(xtl, qnormalize(torch.randn(2, 4, dtype=torch.float64))) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, angle_step_in_plane_deg=3.0) + om.match_orientations(progress_bar=False) + img = np.random.default_rng(0).random((1, 2, 32, 32)) ** 4 + dataset = SimpleNamespace(array=img, shape=img.shape) + shown = [] + for norm in (None, {"power": 0.5, "upper_quantile": 0.9}): + fig, axs = plot_pattern_matches( + om, [(0, 0)], dataset=dataset, pixel_size=0.05, matches=(0,), norm=norm + ) + # show_2d draws the pattern in the panel, extended to q units + ax = axs[0, 0] + im = ax.images[0] + assert np.allclose(im.get_extent()[:2], (-0.5 * 0.05 - 16 * 0.05, 31.5 * 0.05 - 16 * 0.05)) + # the panel shows the recorded area and no marker outside it + assert np.allclose(ax.get_xlim(), im.get_extent()[:2]) + x0, x1, y1, y0 = im.get_extent() + for coll in ax.collections: + xy = coll.get_offsets() + assert ( + (xy[:, 0] >= x0) & (xy[:, 0] <= x1) & (xy[:, 1] >= y0) & (xy[:, 1] <= y1) + ).all() + shown.append(np.asarray(im.get_array())[..., 0]) + matplotlib.pyplot.close(fig) + # gray_r: a lower upper quantile saturates more of the pattern to black + assert (shown[1] <= shown[1].min() + 1e-6).mean() > (shown[0] <= shown[0].min() + 1e-6).mean() + + +def _quat_deg(axis, angle_deg): + from quantem.diffraction.rotations import quat_from_axis_angle + + a = torch.tensor(axis, dtype=torch.float64) + return quat_from_axis_angle(a / a.norm(), torch.tensor(np.deg2rad(angle_deg))) + + +def _two_grain_peaks(xtl, q1, q2, n=3): + """n positions, each the sum of two grains' patterns.""" + peaks = Vector.from_shape( + (1, n), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(n): + rows = [] + for q in (q1, q2): + p = xtl.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + rows.append(np.stack([p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()], 1)) + peaks[0, i] = np.concatenate(rows) + return peaks + + +def test_second_match_indexes_second_grain(): + """With deflation, the second match fits the peaks the first leaves.""" + from quantem.diffraction.rotations import quat_from_zone_axis + + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + q1 = quat_from_zone_axis(xtl.direction_vector((0, 0, 1)), 10.0) + q2 = quat_from_zone_axis(xtl.direction_vector((1, 1, 1)), 35.0) + peaks = _two_grain_peaks(xtl, q1, q2) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=2.0, angle_step_in_plane_deg=2.0, verbose=False) + om.match_orientations(num_matches=2, suppress_matched=1.0, progress_bar=False) + assert om.quats.shape[2] == 2 and om.corr_residual.shape == om.corr.shape + sym = xtl.sym_quats + for i in range(peaks.shape[1]): + e = [ + [float(misorientation_angle_deg(q, om.quats[0, i, m], sym)) for q in (q1, q2)] + for m in range(2) + ] + # the two matches are the two grains, one each + assert min(e[0][0] + e[1][1], e[0][1] + e[1][0]) < 6.0, e + + +def test_match_residual(): + """The residual of one crystal is re-matched by the other.""" + torch.manual_seed(1) + ti_a = Crystal.from_ase(bulk("Ti", "hcp", a=2.9505, c=4.6855), name="a", verbose=False) + ti_b = Crystal.from_ase(bulk("Ti", "bcc", a=3.26, cubic=True), name="b", verbose=False) + for x in (ti_a, ti_b): + x.calculate_structure_factors(k_max=1.5) + q_a = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + q_b = _quat_deg((1.0, 0.0, 0.0), 20.0) + N = 3 + peaks = Vector.from_shape( + (1, N), fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + for i in range(N): + rows = [] + for x, q, s in ((ti_a, q_a, 1.0), (ti_b, q_b, 0.5)): + p = x.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + rows.append( + np.stack([p["qx"].numpy(), p["qy"].numpy(), s * p["intensity"].numpy()], 1) + ) + peaks[0, i] = np.concatenate(rows) + oms = {} + for x in (ti_a, ti_b): + om = OrientationMap.from_vectors(peaks, x, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, verbose=False, progress_bar=False) + om.match_orientations(progress_bar=False) + oms[x.name] = om + om_b = oms["b"] + om_b.match_residual(oms["a"], progress_bar=False) + assert om_b.quats.shape[2] == 2 + for name in ("corr", "corr_residual", "mirror"): + assert getattr(om_b, name).shape == (1, N, 2), name + # with alpha's peaks removed, beta is found by one of its two matches + err = misorientation_angle_deg(q_b, om_b.quats[0], ti_b.sym_quats).amin(dim=-1) + assert float(err.max()) < 2.0, err + assert float(om_b.corr[0, :, 1].min()) > 0 + assert "match_residual" in om_b.metadata + + # nothing left once every peak is deleted: no error, an empty second match + om_a = oms["a"] + om_a.match_residual(om_a, delete_radius=10.0, progress_bar=False) + assert om_a.quats.shape[2] == 2 + assert float(om_a.corr[..., 1].abs().max()) == 0.0 + assert om_a.corr_residual.shape == om_a.corr.shape + + +def test_plot_pattern_matches_defaults_with_one_match(): + import matplotlib + + matplotlib.use("Agg") + from types import SimpleNamespace + + from quantem.diffraction.orientation_visualization import plot_pattern_matches + + torch.manual_seed(0) + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + peaks = _make_peaks(xtl, qnormalize(torch.randn(2, 4, dtype=torch.float64))) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(angle_step_zone_axis_deg=3.0, verbose=False, progress_bar=False) + om.match_orientations(progress_bar=False) + # default matches=(0, 1) with a single match: one panel, no IndexError + fig, axs = plot_pattern_matches(om, [(0, 0), (0, 1)]) + assert axs.shape == (2, 1) + with pytest.raises(ValueError, match="none of matches"): + plot_pattern_matches(om, [(0, 0)], matches=(3,)) + img = np.ones((1, 2, 16, 16)) + with pytest.raises(ValueError, match="pixel_size"): + plot_pattern_matches(om, [(0, 0)], dataset=SimpleNamespace(array=img, shape=img.shape)) + matplotlib.pyplot.close("all") + + +def test_misorientation_map_and_cluster_plots(): + import matplotlib + + matplotlib.use("Agg") + from quantem.diffraction.orientation_visualization import _pole_family + + xtl = Crystal.from_ase(bulk("Ti", "hcp", a=2.95, c=4.686), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + q0 = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + q1 = _quat_deg((1.0, 0.0, 0.0), 25.0) + R, C = 4, 6 + q = torch.where((torch.arange(C) < 3)[None, :, None], q0, q1).expand(R, C, 4).clone() + peaks = _make_peaks(xtl, q0[None]) + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.quats = q[..., None, :] + om.corr = torch.ones((R, C, 1), dtype=torch.float64) + + mis = om.misorientation_map() + assert mis.shape == (R, C) + assert float(mis[:, :3].max()) < 1e-6 and np.allclose(mis[:, 3:].numpy(), 25.0, atol=1e-6) + assert np.allclose(om.misorientation_map(reference=q1)[:, 3:].numpy(), 0.0, atol=1e-6) + + clusters = om.cluster_orientations(min_cluster_size=2) + assert clusters["sizes"].tolist() == [12, 12] + fig, ax = om.plot_cluster_map(clusters) + assert len(ax.get_legend().get_texts()) == 2 + fig, ax = om.plot_cluster_pole_figure(clusters, pole=(0, 0, 0, 1), pole_label="[0001]") + assert len(ax.collections) == 2 + # a mask selects positions above one half + half = np.zeros((R, C)) + half[:, :3] = 0.9 + half[:, 3:] = 0.4 + assert om.cluster_orientations(mask=half, min_cluster_size=2)["sizes"].tolist() == [12] + matplotlib.pyplot.close("all") + + # poles are Miller indices: hexagonal [110] is 60 degrees from [100], + # not the 45 degrees of the Cartesian (1, 1, 0) + fam = _pole_family(xtl, (1, 1, 0)) + d = xtl.direction_vector((1, 1, 0)) + assert float((fam @ d).max()) > 1 - 1e-5 + assert torch.allclose(_pole_family(xtl, (0, 0, 1)), _pole_family(xtl, (0, 0, 0, 1))) + a = xtl.direction_vector((1, 0, 0)) + assert np.isclose(float(a @ d), 0.5) diff --git a/tests/diffraction/test_reverse_monte_carlo.py b/tests/diffraction/test_reverse_monte_carlo.py new file mode 100644 index 000000000..819f8d9df --- /dev/null +++ b/tests/diffraction/test_reverse_monte_carlo.py @@ -0,0 +1,373 @@ +import numpy as np +import pytest +import torch + +from quantem.diffraction import Crystal, ReverseMonteCarlo +from quantem.diffraction.reverse_monte_carlo import cubic_rotations + +CIF = """data_VNb +_cell_length_a 3.2 +_cell_length_b 3.2 +_cell_length_c 3.2 +_cell_angle_alpha 90 +_cell_angle_beta 90 +_cell_angle_gamma 90 +_symmetry_space_group_name_H-M 'I m -3 m' +_symmetry_Int_Tables_number 229 +loop_ +_atom_site_label +_atom_site_type_symbol +_atom_site_fract_x +_atom_site_fract_y +_atom_site_fract_z +_atom_site_occupancy +Nb1 Nb 0 0 0 0.4 +V1 V 0 0 0 0.3 +Zr1 Zr 0 0 0 0.3 +""" + + +@pytest.fixture +def rmc(tmp_path): + path = tmp_path / "vnb.cif" + path.write_text(CIF) + rng = np.random.default_rng(1) + images = [rng.random((64, 64)) + 1.0 for _ in range(2)] + out = ReverseMonteCarlo.from_images( + images, zone_axes=[(0, 0, 1), (0, 1, 1)], sampling=0.02, bin_factor=2 + ) + out.set_crystal(Crystal.from_cif(path, verbose=False)) + out.geometry = dict( + centers=[np.array([32.0, 32.0])] * 2, + matrices=[np.eye(2) / 0.02, np.array([[0.0, -1.0], [1.0, 0.0]]) / 0.02], + tilts=[np.zeros(2), np.array([0.01, 0.0])], + ) + out.set_mask(bragg_radius=0.06, q_max=0.6, center_radius=0.1, edge_px=2) + out.build_supercell(cells=4, seed=0, device="cpu") + out.fit_background() + return out + + +def test_cubic_rotations(): + ops = cubic_rotations() + assert ops.shape == (24, 3, 3) + assert np.allclose([np.linalg.det(o) for o in ops], 1) + assert len({o.tobytes() for o in ops}) == 24 + + +def test_ternary_site_and_composition(rmc): + assert rmc.species == ["Nb", "V", "Zr"] + assert np.allclose(rmc.concentrations, [0.4, 0.3, 0.3]) + assert len(rmc.site_x) == 2 * 4**3 + counts = np.bincount(rmc.species_index, minlength=3) + assert counts.sum() == 128 + assert np.all(np.abs(counts - 128 * rmc.concentrations) <= 1) + + +def _score(rmc, d_fr, d_fi): + keep = rmc._keep / len(rmc.site_x) + d_int = (2 * (rmc._Fr * d_fr + rmc._Fi * d_fi) + d_fr**2 + d_fi**2) * keep + dm = rmc._scale * rmc._u * rmc._read(d_int[:, rmc._sym_index].mean(dim=1)) + return float((rmc._w * (dm**2 - 2 * rmc._r * dm)).sum()) + + +def test_swap_score_matches_recompute(rmc): + """The incremental loss change of one swap equals the loss after recomputing F from scratch.""" + loss0 = rmc._update_residual() + spec = rmc.species_index + j1 = int(np.nonzero(spec == 0)[0][0]) + j2 = int(np.nonzero(spec == 1)[0][0]) + c1, s1 = rmc._phases(rmc._positions(np.array([j1]))) + c2, s2 = rmc._phases(rmc._positions(np.array([j2]))) + df = rmc._fs[1] - rmc._fs[0] + dL = _score(rmc, df * (c1 - c2), df * (s2 - s1)) + spec[j1], spec[j2] = 1, 0 + rmc._recompute_F() + assert rmc._update_residual() - loss0 == pytest.approx(dL, rel=1e-3, abs=1e-4 * loss0) + + +def test_displacement_score_matches_recompute(rmc): + loss0 = rmc._update_residual() + j = np.array([5]) + new = np.array([[1, -1, 1]]) + co, so = rmc._phases(rmc._positions(j)) + cn, sn = rmc._phases(rmc._positions(j, new)) + f = rmc._fs[int(rmc.species_index[5])] + dL = _score(rmc, f * (cn - co), f * (so - sn)) + rmc.displacement[5] = new[0] + rmc._recompute_F() + assert rmc._update_residual() - loss0 == pytest.approx(dL, rel=1e-3, abs=1e-4 * loss0) + + +def test_run_lowers_loss_and_keeps_composition(rmc): + counts = np.bincount(rmc.species_index) + rmc.run(n_sweeps=3, batch=8, progress=False) + assert np.array_equal(np.bincount(rmc.species_index), counts) + assert rmc.loss_history[-1] <= rmc.loss_history[0] + 1e-6 + + +def test_warren_cowley_random_is_near_zero(rmc): + sro = rmc.warren_cowley(n_shells=2) + assert np.allclose(sro["radius"], [3.2 * np.sqrt(3) / 2, 3.2], atol=1e-3) + assert sro["alpha"].shape == (2, 3, 3) + assert np.all(np.abs(sro["alpha"]) < 0.35) + + +def test_coarse_diffuse_grid_matches_sections(rmc): + fine = rmc.diffuse_grid() + coarse = rmc.diffuse_grid(max_size=rmc.grid_size // 2) + assert coarse.shape[0] == fine.shape[0] // 2 + a, _, d = rmc.diffuse_section((0, 0, 1), extent=1.0, grid=fine) + b, _, _ = rmc.diffuse_section((0, 0, 1), extent=1.0, grid=coarse) + assert np.allclose(a[d > 0.2], b[d > 0.2], rtol=1e-4) # no displacements: identical + + +def test_mask_is_zero_on_bragg_peaks(rmc): + w = rmc.mask["w"][0] + b = rmc.bin_factor + for r, c in rmc.bragg_positions(0): + r, c = int(r // b), int(c // b) + if 0 <= r < w.shape[0] and 0 <= c < w.shape[1]: + assert w[r, c] < 0.3 + + +def test_sro_section_random_is_near_laue(rmc): + img, _, node_dist = rmc.diffuse_section((0, 0, 1), extent=1.0, smooth=True) + between = img[node_dist > 0.2] + assert 0.5 < between.mean() < 1.5 + + +def test_omega_embryo_is_a_collapsed_row_and_scores_exactly(rmc): + sites, new = rmc._omega_proposals(4) + assert sites.shape[1] == 3 and len(set(sites.ravel())) == sites.size + xc = rmc.site_x[sites] // rmc.refine + step = np.mod(xc[:, 1] - xc[:, 0], rmc.cells * 2) + step = np.where(step > rmc.cells, step - 2 * rmc.cells, step) + assert np.all(np.abs(step) == 1) # nearest neighbours along <111> + vec = new + assert np.all(vec[:, 0] == 0) and np.all(vec[:, 1] == -vec[:, 2]) + + loss0 = rmc._update_residual() + d_fr = torch.zeros(len(rmc._needed)) + d_fi = torch.zeros(len(rmc._needed)) + for k in range(3): + j = sites[:1, k] + co, so = rmc._phases(rmc._positions(j)) + cn, sn = rmc._phases(rmc._positions(j, new[:1, k])) + f = rmc._fs[int(rmc.species_index[j[0]])] + d_fr += (f * (cn - co))[0] + d_fi += (f * (so - sn))[0] + dL = _score(rmc, d_fr[None], d_fi[None]) + rmc.displacement[sites[0]] = new[0] + rmc._recompute_F() + assert rmc._update_residual() - loss0 == pytest.approx(dL, rel=1e-3, abs=1e-4 * loss0) + + +def test_fitted_envelope_does_not_raise_loss(rmc): + loss_measured = rmc._update_residual() + rmc.envelope = "fitted" + rmc._setup_forward() + rmc._solve_linear(rmc._model_diffuse()) + assert rmc._update_residual() <= loss_measured * (1 + 1e-6) + assert all(np.all(p >= 0) for p in rmc._env_p) + + +def test_static_b_cap_holds(rmc): + rmc.run(n_sweeps=3, batch=8, omega_fraction=1.0, max_static_b=0.2, progress=False) + assert rmc.static_b() <= 0.2 + 1e-9 + assert len(rmc._omega_vectors) == 16 # two amplitudes x eight <111> senses + + +def test_size_effect_chi_is_odd_about_bragg_nodes(rmc): + a = rmc.lattice_parameter + g = np.array([1.0, 1.0, 0.0]) / a # allowed BCC reflection + kappa = np.array([[0.02, 0.005, -0.01], [0.0, 0.03, 0.01]]) + plus = rmc._size_chi(g + kappa) + minus = rmc._size_chi(g - kappa) + assert np.all(np.abs(plus) > 0) + assert np.allclose(plus, -minus, rtol=0.2) + far = rmc._size_chi(g + 4 * kappa) + assert np.all(np.abs(far) < np.abs(plus)) # grows toward the node + + +def test_size_effect_radii_orders_species(rmc): + rmc.set_size_effect("radii") + eta = dict(zip(rmc.species, rmc.size_eta)) + assert eta["V"] < eta["Nb"] < eta["Zr"] + assert abs((rmc.concentrations * rmc.size_eta).sum()) < 1e-12 + + +def test_autoserialize_round_trip_rebuilds_model(rmc, tmp_path): + from quantem.core.io import load + + rmc.envelope = "fitted" + rmc._setup_forward() + rmc._solve_linear(rmc._model_diffuse()) + loss = rmc._update_residual() + model = rmc.model_images() + path = tmp_path / "rmc.zip" + rmc.save(path, mode="o", skip=rmc.DERIVED_ATTRIBUTES) + back = load(path) + assert back._update_residual() == pytest.approx(loss, rel=1e-5) + for a, b in zip(model, back.model_images()): + assert np.allclose(a, b, rtol=1e-4, atol=1e-6) + assert np.array_equal(back.species_index, rmc.species_index) + + +def test_displacement_correlations_see_omega(rmc): + sites, new = rmc._omega_proposals(8) + rmc.displacement[sites.reshape(-1)] = new.reshape(-1, 3) + c = rmc.displacement_correlations(n_shells=2) + assert c["longitudinal"][0] < 0 # collapsing nearest-neighbour pairs move toward each other + dist = rmc.displacement_distributions() + assert set(dist) == {"<100>", "<110>", "<111>"} + + +def test_random_displacement_scores_exactly_and_stays_bounded(rmc): + loss0 = rmc._update_residual() + sites, new = rmc._random_proposals(16) + assert np.all(np.abs(new) <= rmc._max_steps) + assert np.all(np.any(new[:, 0] != rmc.displacement[sites[:, 0]], axis=1)) + j = sites[:1, 0] + co, so = rmc._phases(rmc._positions(j)) + cn, sn = rmc._phases(rmc._positions(j, new[:1, 0])) + f = rmc._fs[int(rmc.species_index[j[0]])] + dL = _score(rmc, f * (cn - co), f * (so - sn)) + rmc.displacement[j[0]] = new[0, 0] + rmc._recompute_F() + assert rmc._update_residual() - loss0 == pytest.approx(dL, rel=1e-3, abs=1e-4 * loss0) + + +def test_shell_labels(rmc): + c = rmc.displacement_correlations(n_shells=4) + assert c["shell"] == ["1/2<111>", "<100>", "<110>", "1/2<311>"] + + +def test_set_mask_edge_px_zero_keeps_detector(rmc): + rmc.set_mask(bragg_radius=0.06, q_max=0.6, center_radius=0.1, edge_px=0) + assert all(np.any(w > 0.5) for w in rmc.mask["w"]) + + +def test_bloch_envelope_needs_fit_thickness(rmc): + with pytest.raises(RuntimeError, match="fit_thickness"): + rmc.set_envelope("bloch") + assert rmc.envelope == "measured" + with pytest.raises(ValueError, match="unknown envelope"): + rmc.set_envelope("nope") + assert rmc.envelope == "measured" + + +def test_fit_size_effect_runs_on_model_device(rmc): + out = rmc.fit_size_effect(verbose=False) + assert set(out["eta"]) == set(rmc.species) + assert out["loss"][1] <= out["loss"][0] * (1 + 1e-6) + assert abs((rmc.concentrations * rmc.size_eta).sum()) < 1e-12 + + +def _synthetic_pattern(rmc, zone, center, theta, shape=(160, 160)): + """Gaussian spots at the kinematic intensities on a broad halo about the direct beam.""" + from quantem.diffraction.reverse_monte_carlo import _rot2 + + _, g2, inten = rmc._zone_reflections(zone, 1.6) + A = _rot2(theta) / rmc.sampling + pos = np.vstack([center, np.asarray(center) + g2 @ A.T]) + inten = np.concatenate([[2 * inten.max()], inten]) # direct beam first + rr, cc = np.mgrid[0 : shape[0], 0 : shape[1]].astype(float) + im = 50.0 * np.exp(-((rr - center[0]) ** 2 + (cc - center[1]) ** 2) / (2 * 40.0**2)) + for (pr, pc), i in zip(pos, inten / inten.max()): + im += 1e3 * i * np.exp(-((rr - pr) ** 2 + (cc - pc) ** 2) / (2 * 1.2**2)) + return im + 1.0 + + +def test_fit_geometry_on_synthetic_lattice(tmp_path): + path = tmp_path / "vnb.cif" + path.write_text(CIF) + zones = [(0, 0, 1), (0, 1, 1)] + centers = [np.array([78.3, 81.6]), np.array([80.5, 79.2])] + out = ReverseMonteCarlo.from_images( + [np.zeros((160, 160))] * 2, zone_axes=zones, sampling=0.03, bin_factor=4 + ) + out.set_crystal(Crystal.from_cif(path, verbose=False)) + out.images = [ + _synthetic_pattern(out, z, c, th).astype(np.float32) + for z, c, th in zip(zones, centers, (0.3, -0.7)) + ] + out.fit_geometry(scale_range=(0.97, 1.03), verbose=False) + geo = out.geometry + for c, c_fit in zip(centers, geo["centers"]): + assert np.allclose(c_fit, c, atol=0.3) + # the halo biases the peak centroids slightly on so small a detector + assert out.lattice_parameter == pytest.approx(3.2, rel=1e-2) + assert all(np.rad2deg(np.linalg.norm(t)) < 0.5 for t in geo["tilts"]) + assert all(rms < 0.3 for rms in geo["rms_px"]) + assert len(geo["_inten_kin"]) == len(geo["excitation_width"]) == 2 + + out.fit_thickness( + thickness=(100.0, 200.0), step=50.0, tilt_range_deg=0.0, k_max=1.0, depth_samples=4, + verbose=False, + ) # fmt: skip + assert out.geometry["thickness_k_max"] == 1.0 + assert all(g is not None for g in out.geometry["bloch_g"]) + + out.set_mask(bragg_radius=0.08, q_max=1.0, center_radius=0.15, edge_px=4) + out.build_supercell(cells=3, seed=0, device="cpu") + out.fit_background() + for envelope in ("kinematic", "bloch"): + out.set_envelope(envelope) + assert out.envelope == envelope + assert np.isfinite(out._update_residual()) + + +def test_shell_correlations_match_warren_cowley(rmc): + sc = rmc.shell_correlations(n_shells=4) + sro = rmc.warren_cowley(n_shells=4) + ratio = sc["ratio"] + np.testing.assert_allclose(sc["radius"], sro["radius"]) + # pair counts are symmetric, and an unlike pair has ratio = 1 - alpha + np.testing.assert_allclose(ratio, np.swapaxes(ratio, 1, 2), atol=1e-9) + np.testing.assert_allclose(ratio[:, 0, 1], 1 - sro["alpha"][:, 0, 1], atol=0.02) + # a random arrangement sits near 1 in every shell + assert np.abs(ratio - 1).max() < 0.2 + assert sc["shell"][:3] == ["1/2<111>", "<100>", "<110>"] + # the default radius reaches five lattice parameters, or half the supercell + far = rmc.shell_correlations() + assert far["radius"][-1] <= min(5, rmc.cells / 2) * rmc.lattice_parameter + 1e-6 + assert len(far["shell"]) == len(far["radius"]) + + +def test_plot_shell_correlations(rmc): + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + fig, axs = rmc.plot_shell_correlations() + K = len(rmc.species) + assert len(axs) == K * (K + 1) // 2 + plt.close(fig) + + +def test_to_atoms_and_cif_round_trip(rmc, tmp_path): + from ase.io import read + + atoms = rmc.to_atoms() + n = len(rmc.site_x) + assert len(atoms) == n + assert np.allclose(atoms.cell.lengths(), rmc._a_crystal * rmc.cells) + counts = {s: int((np.array(atoms.get_chemical_symbols()) == s).sum()) for s in rmc.species} + expected = np.bincount(rmc.species_index, minlength=len(rmc.species)) + assert [counts[s] for s in rmc.species] == expected.tolist() + # displacements move the atoms off the ideal sites by the fitted amount + first = int( + np.argsort(np.array([rmc.species[k] for k in rmc.species_index]), kind="stable")[0] + ) + rmc.displacement[first] = [1, 0, 0] + step = rmc._a_crystal * rmc.cells / rmc.grid_size + shifted = rmc.to_atoms().get_positions()[0] - atoms.get_positions()[0] + assert np.allclose(shifted, [step, 0, 0]) + rmc.displacement[first] = 0 + path = rmc.to_cif(tmp_path / "rmc.cif") + back = read(path) + assert len(back) == n + assert np.allclose(back.cell.lengths(), atoms.cell.lengths()) + assert np.allclose(back.get_scaled_positions(), atoms.get_scaled_positions(), atol=1e-5) diff --git a/tests/diffraction/test_rotation_convention.py b/tests/diffraction/test_rotation_convention.py new file mode 100644 index 000000000..8ff8e6ddd --- /dev/null +++ b/tests/diffraction/test_rotation_convention.py @@ -0,0 +1,61 @@ +"""Self-consistency of the detector-to-scan rotation convention. + +Simulates a cubic crystal rotated in-plane by a known angle, records its +peaks in a detector frame that is rotated relative to the scan frame, and +checks that peaks_to_calibrated(rotation_ccw_deg=...) plus orientation +matching recovers the ground-truth in-plane angle in the scan frame. +""" + +import numpy as np +import pytest +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction import calibration +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.rotations import quat_from_axis_angle + + +@pytest.mark.parametrize("rot_scan_deg", [0.0, 34.0, -34.0]) +def test_rotation_roundtrip(rot_scan_deg): + xtl = Crystal.from_ase(bulk("Al", "fcc", a=4.05, cubic=True), verbose=False) + xtl.calculate_structure_factors(k_max=1.5) + + # ground truth: [001] zone, 30 deg in-plane rotation, in the SCAN frame + angle_true = 30.0 + q_true = quat_from_axis_angle( + torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64), + torch.tensor(np.deg2rad(angle_true), dtype=torch.float64), + ) + pat = xtl.generate_pattern(q_true, energy_ev=200e3, sigma_excitation=0.02) + q_scan = np.stack([pat["qx"].numpy(), pat["qy"].numpy()], axis=1) + + # what the detector records: scan-frame vectors rotated by -rot_scan + # (peaks_to_calibrated undoes this with rotation_ccw_deg=+rot_scan) + th = np.deg2rad(-rot_scan_deg) + rot_back = np.array([[np.cos(th), -np.sin(th)], [np.sin(th), np.cos(th)]]) + q_det = q_scan @ rot_back.T + + pixel_size = 0.01 + data = np.column_stack([q_det / pixel_size, pat["intensity"].numpy()]) + peaks_px = Vector.from_data( + [[data]], + fields=["q_row", "q_col", "intensity"], + units=["px", "px", "counts"], + name="synthetic", + ) + peaks = calibration.peaks_to_calibrated(peaks_px, pixel_size, rotation_ccw_deg=rot_scan_deg) + + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan(power_intensity=0.0) + om.match_orientations(progress_bar=False) + om.refine_orientations(progress_bar=False) + + # the recovered orientation must equal the scan-frame ground truth + # (modulo crystal symmetry) -- independent of the detector rotation + from quantem.diffraction.rotations import misorientation_angle_deg + + err = float(misorientation_angle_deg(q_true, om.quats[0, 0, 0], xtl.sym_quats)) + assert err < 1.0, f"misorientation {err:.2f} deg at rot {rot_scan_deg}" diff --git a/tests/diffraction/test_rotations.py b/tests/diffraction/test_rotations.py new file mode 100644 index 000000000..2dbe27243 --- /dev/null +++ b/tests/diffraction/test_rotations.py @@ -0,0 +1,226 @@ +"""Tests for quantem.diffraction.rotations.""" + +import numpy as np +import pytest +import torch + +from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qconj, + qmult, + qnormalize, + qrotate, + quat_from_axis_angle, + quat_from_euler_zxz, + quat_from_matrix, + quat_from_zone_axis, + quat_to_euler_zxz, + quat_to_matrix, + sample_zone_axes, + zone_axis_from_quat, +) + + +@pytest.fixture +def random_quats(): + torch.manual_seed(0) + return qnormalize(torch.randn(100, 4, dtype=torch.float64)) + + +def test_matrix_roundtrip(random_quats): + R = quat_to_matrix(random_quats) + assert torch.allclose(quat_from_matrix(R), random_quats, atol=1e-10) + + +def test_euler_roundtrip(random_quats): + e = quat_to_euler_zxz(random_quats) + assert torch.allclose(qnormalize(quat_from_euler_zxz(e)), random_quats, atol=1e-8) + + +def test_rotate_matches_matrix(random_quats): + v = torch.randn(100, 3, dtype=torch.float64) + R = quat_to_matrix(random_quats) + assert torch.allclose(qrotate(random_quats, v), (R @ v[..., None]).squeeze(-1), atol=1e-10) + + +def test_mult_conj_identity(random_quats): + q = random_quats + ident = qmult(q, qconj(q)) + expect = torch.zeros_like(q) + expect[:, 0] = 1.0 + assert torch.allclose(ident, expect, atol=1e-10) + + +def test_zone_axis_roundtrip(): + torch.manual_seed(1) + v = torch.randn(50, 3, dtype=torch.float64) + v = v / torch.linalg.norm(v, dim=-1, keepdim=True) + q = quat_from_zone_axis(v, in_plane_deg=25.0) + assert torch.allclose(zone_axis_from_quat(q), v, atol=1e-10) + + +def test_axis_angle(): + axis = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + q = quat_from_axis_angle(axis, torch.tensor(np.pi / 2, dtype=torch.float64)) + v = torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64) + assert torch.allclose( + qrotate(q, v), torch.tensor([0.0, 1.0, 0.0], dtype=torch.float64), atol=1e-9 + ) + + +def test_misorientation_symmetry(): + # 90 degree rotation about z is a cubic symmetry: misorientation 0 + from ase.build import bulk + + from quantem.diffraction.crystal import Crystal + + xtl = Crystal.from_ase(bulk("Au", "fcc", a=4.08, cubic=True)) + qa = torch.tensor([1.0, 0, 0, 0], dtype=torch.float64) + qb = quat_from_axis_angle( + torch.tensor([0.0, 0, 1.0], dtype=torch.float64), + torch.tensor(np.pi / 2, dtype=torch.float64), + ) + ang = misorientation_angle_deg(qa, qb, xtl.sym_quats) + assert float(ang) < 1e-4 + ang_nosym = misorientation_angle_deg(qa, qb) + assert abs(float(ang_nosym) - 90.0) < 1e-6 + + +def test_sample_zone_axes_wedge(): + corners = torch.tensor([[0, 0, 1], [0, 1, 1], [1, 1, 1]], dtype=torch.float64) + corners = corners / torch.linalg.norm(corners, dim=-1, keepdim=True) + v, inds = sample_zone_axes(corners, 2.0) + assert torch.allclose(torch.linalg.norm(v, dim=-1), torch.ones(v.shape[0], dtype=v.dtype)) + # corners present + for c in corners: + assert torch.linalg.norm(v - c, dim=-1).min() < 1e-8 + + +def test_sample_zone_axes_isotropic(): + # the requested step holds in every direction, whatever the apex angle + for apex_deg in (30.0, 120.0): + a = np.deg2rad(apex_deg) + corners = torch.tensor( + [[0, 0, 1], [1, 0, 0], [np.cos(a), np.sin(a), 0]], dtype=torch.float64 + ) + v, _ = sample_zone_axes(corners, 2.0) + dots = (v @ v.T).clamp(-1, 1) + ang = torch.rad2deg(torch.acos(dots)) + ang.fill_diagonal_(1e9) + nn = ang.min(dim=1).values + # points on the equator row: spacing within 25% of the step + eq = v[:, 2].abs() < 1e-9 + assert nn[eq].min() > 1.5 and nn[eq].max() < 2.5 + + +def test_symmetry_reduced_zone_angles(): + from ase.build import bulk + + from quantem.diffraction.crystal import Crystal + from quantem.diffraction.rotations import symmetry_reduced_zone_angles + + xtl = Crystal.from_ase(bulk("Ti", "bcc", a=3.31, cubic=True), verbose=False) + za = torch.tensor( + [[0, 0, 1.0], [1.0, 0, 0], [0, 1.0, 0], [1.0, 1.0, 1.0], [-1.0, 1.0, 1.0]], + dtype=torch.float64, + ) + za = za / torch.linalg.norm(za, dim=1, keepdim=True) + ang = symmetry_reduced_zone_angles(za, xtl.sym_quats) + # cubic axes are symmetry equivalent, as are the two <111> directions + assert float(ang[0, 1]) < 1e-6 and float(ang[0, 2]) < 1e-6 + assert float(ang[3, 4]) < 1e-6 + assert abs(float(ang[0, 3]) - 54.7356) < 1e-3 + + +def test_misorientation_axis_angle(): + from quantem.diffraction.rotations import misorientation_axis_angle + + axis = torch.tensor([1.0, 2.0, 2.0], dtype=torch.float64) / 3 + qa = qnormalize(torch.tensor([0.9, 0.1, -0.3, 0.2], dtype=torch.float64)) + dq = quat_from_axis_angle(axis, torch.tensor(np.deg2rad(35.0), dtype=torch.float64)) + qb = qmult(qa, dq) # R(qb) = R(qa) R(dq): dq is in the crystal frame of qa + ax, ang = misorientation_axis_angle(qa, qb) + assert abs(float(ang) - 35.0) < 1e-8 + assert torch.allclose(ax, axis, atol=1e-8) + + # with cubic symmetry, a 90 degree turn about [001] plus 10 degrees + # about the same axis reduces to 10 degrees, axis unchanged up to sign + sym = _cubic_sym_quats() + z = torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64) + qb = qmult(qa, quat_from_axis_angle(z, torch.tensor(np.deg2rad(100.0), dtype=torch.float64))) + ax, ang = misorientation_axis_angle(qa, qb, sym) + assert abs(float(ang) - 10.0) < 1e-6 + assert abs(abs(float(ax @ z)) - 1.0) < 1e-6 + assert abs(float(ang) - float(misorientation_angle_deg(qa, qb, sym))) < 1e-6 + + # broadcasting over a batch + qs = qnormalize(torch.randn(7, 4, dtype=torch.float64)) + ax, ang = misorientation_axis_angle(qs, qs[:1], sym) + assert ax.shape == (7, 3) and ang.shape == (7,) + assert torch.allclose(ang, misorientation_angle_deg(qs, qs[:1], sym), atol=1e-4) + + +def _cubic_sym_quats(): + import itertools + + from quantem.diffraction.rotations import symmetry_quaternions + + # the 24 proper rotations of m-3m as signed permutation matrices + mats = [] + for perm in itertools.permutations(range(3)): + for signs in itertools.product((1, -1), repeat=3): + M = np.zeros((3, 3), dtype=int) + for i, (j, s) in enumerate(zip(perm, signs)): + M[i, j] = s + if round(np.linalg.det(M)) == 1: + mats.append(M) + return symmetry_quaternions(np.array(mats), np.eye(3)) + + +def test_symmetry_aligned(): + from quantem.diffraction.rotations import symmetry_aligned + + torch.manual_seed(3) + sym = _cubic_sym_quats() + assert sym.shape == (24, 4) + ref = qnormalize(torch.randn(4, dtype=torch.float64)) + # small perturbations of the reference, each moved to a random symmetry + # branch: alignment must bring every one back next to the reference + eps = quat_from_axis_angle( + torch.randn(10, 3, dtype=torch.float64), torch.full((10,), 0.05, dtype=torch.float64) + ) + near = qmult(ref.expand(10, 4), eps) + far = qmult(near, sym[torch.randint(1, 24, (10,))]) + out = symmetry_aligned(ref, far, sym) + assert out.shape == (10, 4) + assert torch.allclose(out, qnormalize(near), atol=1e-10) + # the same orientations as the inputs + assert torch.allclose( + misorientation_angle_deg(out, far, sym), torch.zeros(10, dtype=torch.float64), atol=1e-4 + ) + + +def test_sample_zone_axis_cap(): + from quantem.diffraction.rotations import sample_zone_axis_cap + + axis = torch.tensor([1.0, -1.0, 2.0], dtype=torch.float64) + unit = axis / torch.linalg.norm(axis) + pts = sample_zone_axis_cap(axis, 10.0, 2.0) + assert pts.ndim == 2 and pts.shape[1] == 3 + assert torch.allclose( + torch.linalg.norm(pts, dim=1), torch.ones(pts.shape[0], dtype=torch.float64), atol=1e-12 + ) + ang = torch.rad2deg(torch.acos((pts @ unit).clamp(-1, 1))) + assert float(ang.max()) <= 10.0 + 1e-9 + # equal-area count: cap area / step^2 + n_expect = 2 * np.pi * (1 - np.cos(np.deg2rad(10))) / np.deg2rad(2) ** 2 + assert abs(pts.shape[0] - np.ceil(n_expect)) <= 1 + # covers the cap: every direction in it has a sample within ~step + rng = np.random.default_rng(0) + probe = sample_zone_axis_cap(axis, 9.0, 0.5)[rng.choice(300, 50)] + d = torch.rad2deg(torch.acos((probe @ pts.T).clamp(-1, 1))).min(dim=1).values + assert float(d.max()) < 2.0 + # zero half angle: the axis alone; the poles take the short path + assert torch.allclose(sample_zone_axis_cap(axis, 0.0, 1.0), unit[None]) + down = sample_zone_axis_cap(torch.tensor([0.0, 0.0, -1.0]), 5.0, 1.0) + assert float(down[:, 2].max()) < -np.cos(np.deg2rad(5.0)) + 1e-9 diff --git a/tests/diffraction/test_strain_bragg_vectors.py b/tests/diffraction/test_strain_bragg_vectors.py new file mode 100644 index 000000000..c3d3cefdb --- /dev/null +++ b/tests/diffraction/test_strain_bragg_vectors.py @@ -0,0 +1,174 @@ +"""BraggVectors lattice fit and StrainMap on synthetic peaks (no disk detection).""" + +import numpy as np +import pytest + +from quantem.core.datastructures import Dataset4dstem +from quantem.core.datastructures.vector import Vector +from quantem.diffraction import BraggVectors +from quantem.diffraction.strain import StrainMap + +ORIGIN = np.array([32.0, 32.0]) +G1 = np.array([9.0, 1.0]) +G2 = np.array([-1.5, 11.0]) + + +def _synthetic_bragg_vectors(deform, R=8, C=6, metadata=None): + """BraggVectors whose peaks sit on origin + a M g1 + b M g2, M = deform(r, c). + + The peaks are written straight into ``bv.peaks``; the dataset is only a + zero cube that sets the scan and detector shapes. Peaks are stored as + float32, so fitted vectors match to about 1e-6 pixels. + """ + nested = [] + for r in range(R): + row = [] + for c in range(C): + M = np.asarray(deform(r, c), dtype=float) + g1, g2 = M @ G1, M @ G2 + pts = [] + for a in range(-2, 3): + for b in range(-2, 3): + q = ORIGIN + a * g1 + b * g2 + pts.append([q[0], q[1], 10.0 if a == b == 0 else 1.0]) + row.append(np.asarray(pts)) + nested.append(row) + + ds = Dataset4dstem.from_array(np.zeros((R, C, 64, 64), dtype=np.float32)) + if metadata: + ds.metadata.update(metadata) + bv = BraggVectors.from_dataset(ds) + bv.peaks = Vector.from_data(nested, fields=["q_row", "q_col", "intensity"]) + bv.compute_bvm() + bv.choose_basis_vectors(origin=ORIGIN, g1=G1, g2=G2, plot=False) + bv.index_peaks(plot=False) + bv.fit_lattice(min_num_peaks=5, progressbar=False, plot=False) + return bv + + +def test_reciprocal_scaling_ramp_gives_compressive_strain(): + # reciprocal vectors grow with row, so the real-space lattice shrinks: + # e = 1 / s - 1, negative and decreasing down the scan + s = 1.0 + 0.002 * np.arange(8) + bv = _synthetic_bragg_vectors(lambda r, c: s[r] * np.eye(2)) + np.testing.assert_allclose(bv.g1_array[:, 0], s[:, None] * G1[None, :], atol=2e-5) + assert np.all(bv.mask_weight > 0.99) + + with pytest.warns(UserWarning, match="no detector rotation"): + sm = bv.calculate_strain_map(g1_ref=G1, g2_ref=G2) + expected = (1.0 / s - 1.0)[:, None] * np.ones((1, 6)) + np.testing.assert_allclose(sm.e_rr.array, expected, atol=2e-5) + np.testing.assert_allclose(sm.e_cc.array, expected, atol=2e-5) + np.testing.assert_allclose(sm.e_rc.array, 0.0, atol=2e-5) + np.testing.assert_allclose(sm.phi.array, 0.0, atol=2e-5) + assert sm.e_rr.array[-1, 0] < 0 + assert np.all(np.diff(sm.e_rr.array[:, 0]) < 0) + + # automatic reference (weighted median over the scan): same ramp, offset + with pytest.warns(UserWarning): + sm_auto = bv.calculate_strain_map() + ramp = sm_auto.e_rr.array[:, 0] + assert np.all(np.diff(ramp) < 0) + assert ramp.min() < 0 < ramp.max() + + +def test_real_and_reciprocal_vectors_give_same_strain(): + rng = np.random.default_rng(0) + R, C = 4, 5 + F = np.eye(2)[None, None] + 0.01 * rng.normal(size=(R, C, 2, 2)) + A0 = np.array([[3.0, 0.5], [-0.4, 2.5]]) # real-space basis, columns a1, a2 + G0 = np.linalg.inv(A0).T # reciprocal basis, columns g1, g2 (G0.T @ A0 = I) + A = F @ A0 + G = np.linalg.inv(F).transpose(0, 1, 3, 2) @ G0 + + common = dict(ds_shape=(R, C), q_to_r_rotation_ccw_deg=0.0, q_transpose=False) + sm_real = StrainMap( + g1_array=A[..., :, 0], + g2_array=A[..., :, 1], + real_space=True, + g1_ref=A0[:, 0], + g2_ref=A0[:, 1], + **common, + ) + sm_recip = StrainMap( + g1_array=G[..., :, 0], + g2_array=G[..., :, 1], + real_space=False, + g1_ref=G0[:, 0], + g2_ref=G0[:, 1], + **common, + ) + expected = { + "e_rr": F[..., 0, 0] - 1, + "e_cc": F[..., 1, 1] - 1, + "e_rc": 0.5 * (F[..., 0, 1] + F[..., 1, 0]), + "phi": 0.5 * (F[..., 1, 0] - F[..., 0, 1]), + } + for name, value in expected.items(): + np.testing.assert_allclose(getattr(sm_real, name).array, value, atol=1e-12) + np.testing.assert_allclose(getattr(sm_recip, name).array, value, atol=1e-12) + + +def test_counterclockwise_lattice_rotation_gives_positive_phi(): + theta = np.deg2rad(0.5) + rot = np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]) + g = rot @ np.stack([G1, G2], axis=1) + sm = StrainMap( + g1_array=np.broadcast_to(g[:, 0], (2, 2, 2)).copy(), + g2_array=np.broadcast_to(g[:, 1], (2, 2, 2)).copy(), + ds_shape=(2, 2), + real_space=False, + g1_ref=G1, + g2_ref=G2, + ) + np.testing.assert_allclose(sm.phi.array, np.sin(theta), atol=1e-12) + np.testing.assert_allclose(sm.e_rc.array, 0.0, atol=1e-12) + + +def test_q_to_r_rotation_is_applied(): + # uniaxial real-space stretch of 1 % along the detector row axis + stretch = np.diag([1.0 / 1.01, 1.0]) # reciprocal vectors shrink along rows + + bv = _synthetic_bragg_vectors(lambda r, c: stretch, R=3, C=3) + with pytest.warns(UserWarning): + sm0 = bv.calculate_strain_map(g1_ref=G1, g2_ref=G2) + np.testing.assert_allclose(sm0.e_rr.array, 0.01, atol=2e-5) + np.testing.assert_allclose(sm0.e_cc.array, 0.0, atol=2e-5) + + # a 90 degree detector-to-scan rotation moves the stretch onto the scan columns + sm90 = bv.calculate_strain_map( + g1_ref=G1, g2_ref=G2, q_to_r_rotation_ccw_deg=90.0, q_transpose=False + ) + np.testing.assert_allclose(sm90.e_rr.array, 0.0, atol=2e-5) + np.testing.assert_allclose(sm90.e_cc.array, 0.01, atol=2e-5) + assert bv.metadata["q_to_r_rotation_ccw_deg"] == 90.0 + + # the same rotation read from the dataset metadata + bv_md = _synthetic_bragg_vectors( + lambda r, c: stretch, + R=3, + C=3, + metadata={"q_to_r_rotation_ccw_deg": 90.0, "q_transpose": False}, + ) + with pytest.warns(UserWarning, match="using Dataset4dstem metadata"): + sm_md = bv_md.calculate_strain_map(g1_ref=G1, g2_ref=G2) + np.testing.assert_allclose(sm_md.e_cc.array, 0.01, atol=2e-5) + np.testing.assert_allclose(sm_md.e_rr.array, 0.0, atol=2e-5) + + # a transpose alone also swaps rows and columns + sm_t = bv.calculate_strain_map( + g1_ref=G1, g2_ref=G2, q_to_r_rotation_ccw_deg=0.0, q_transpose=True + ) + np.testing.assert_allclose(sm_t.e_cc.array, 0.01, atol=2e-5) + + +def test_estimate_strain_precision_is_quiet_by_default(capsys): + rng = np.random.default_rng(1) + g1 = G1[None, None] + 0.01 * rng.normal(size=(12, 12, 2)) + g2 = G2[None, None] + 0.01 * rng.normal(size=(12, 12, 2)) + sm = StrainMap(g1_array=g1, g2_array=g2, ds_shape=(12, 12), real_space=False) + out = sm.estimate_strain_precision(plot=False) + assert capsys.readouterr().out == "" + assert np.isfinite(out["precision"]["combined"]) + sm.estimate_strain_precision(plot=False, verbose=True) + assert "Strain precision" in capsys.readouterr().out diff --git a/tests/diffraction/test_two_phase_map.py b/tests/diffraction/test_two_phase_map.py new file mode 100644 index 000000000..efa4b90f9 --- /dev/null +++ b/tests/diffraction/test_two_phase_map.py @@ -0,0 +1,325 @@ +"""End-to-end two-phase mapping on a synthetic alpha/beta titanium scan. + +Left region: bcc beta in a fixed orientation. Right region: hcp alpha in a +Burgers-related orientation ((110)beta || (0001)alpha). A two-column band in +the middle contains both patterns superimposed, as at a lath boundary. +""" + +import numpy as np +import torch +from ase.build import bulk + +from quantem.core.datastructures.vector import Vector +from quantem.diffraction.crystal import Crystal +from quantem.diffraction.orientation import OrientationMap +from quantem.diffraction.phase import PhaseMap +from quantem.diffraction.rotations import ( + misorientation_angle_deg, + qmult, + quat_from_axis_angle, +) + + +def _pattern(xtl, q, rng): + p = xtl.generate_pattern(q, energy_ev=200e3, sigma_excitation=0.02) + arr = np.column_stack([p["qx"].numpy(), p["qy"].numpy(), p["intensity"].numpy()]) + arr[:, :2] += rng.normal(0, 0.003, (arr.shape[0], 2)) + arr[:, 2] *= rng.lognormal(0, 0.3, arr.shape[0]) + return arr + + +def test_two_phase_map(): + torch.manual_seed(2) + rng = np.random.default_rng(2) + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + ti_b = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), name="Ti beta", verbose=False + ).calculate_structure_factors(k_max=1.5) + + # beta along [111] zone; alpha along [0001]: the Burgers-related pair + # shares the hexagonal net, the hard case for phase mapping + q_beta = quat_from_axis_angle( + torch.tensor([1.0, -1.0, 0.0], dtype=torch.float64) / np.sqrt(2), + torch.tensor(np.arccos(1 / np.sqrt(3)), dtype=torch.float64), + ) + q_alpha = qmult( + quat_from_axis_angle( + torch.tensor([0.0, 0.0, 1.0], dtype=torch.float64), + torch.tensor(np.deg2rad(14.0), dtype=torch.float64), + ), + torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64), + ) + + R, C = 6, 11 + band = (5, 6) # columns with both phases + cells = [] + truth = np.zeros((R, C), dtype=int) # 0 alpha, 1 beta + for r in range(R): + row = [] + for c in range(C): + if c < band[0]: + arr = _pattern(ti_b, q_beta, rng) + truth[r, c] = 1 + elif c > band[1]: + arr = _pattern(ti_a, q_alpha, rng) + truth[r, c] = 0 + else: + a = _pattern(ti_a, q_alpha, rng) + b = _pattern(ti_b, q_beta, rng) + a[:, 2] *= 0.5 + b[:, 2] *= 0.5 + arr = np.concatenate([a, b]) + truth[r, c] = 2 + row.append(arr) + cells.append(row) + peaks = Vector.from_data(cells, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + + oms = [] + for xtl in (ti_a, ti_b): + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan() + om.match_orientations(num_matches=2, progress_bar=False) + om.refine_orientations(progress_bar=False) + oms.append(om) + + # orientation recovery in the pure regions + err_a = misorientation_angle_deg( + q_alpha, oms[0].quats[:, band[1] + 1 :, 0].reshape(-1, 4), ti_a.sym_quats + ).numpy() + # along [111], beta and its 60 degree twin about [111] give identical + # kinematical patterns, so either is a correct match; which one wins is + # decided by round-off and differs between platforms + q_beta_twin = qmult( + q_beta, + quat_from_axis_angle( + torch.tensor([1.0, 1.0, 1.0], dtype=torch.float64) / np.sqrt(3), + torch.tensor(np.pi / 3, dtype=torch.float64), + ), + ) + q_found = oms[1].quats[:, : band[0], 0].reshape(-1, 4) + err_b = np.minimum( + misorientation_angle_deg(q_beta, q_found, ti_b.sym_quats).numpy(), + misorientation_angle_deg(q_beta_twin, q_found, ti_b.sym_quats).numpy(), + ) + assert np.median(err_a) < 1.0 + assert np.median(err_b) < 1.0 + + pm = PhaseMap.from_orientation_maps(oms) + pm.fit(max_patterns=2, progress_bar=False) + pi = pm.phase_index.numpy() + + pure_a = pi[:, band[1] + 1 :] + pure_b = pi[:, : band[0]] + assert (pure_a == 0).mean() > 0.85 + assert (pure_b == 1).mean() > 0.85 + + # overlap band: every position must be assigned one of the two true + # phases with a valid orientation (either is acceptable) + band_pi = pi[:, band[0] : band[1] + 1] + assert np.isin(band_pi, [0, 1]).all() + + +def test_crystal_map_sets_k_max_once(): + # two phases simulated to different ranges are not compared fairly, and + # setting k_max per crystal invites exactly that mistake + import numpy as np + import pytest + from ase.build import bulk + + from quantem.core.datastructures import Vector + from quantem.diffraction import Crystal, CrystalMap + + peaks = Vector.from_data( + [[np.array([[0.0, 0.0, 1.0], [0.3, 0.1, 0.5], [-0.2, 0.4, 0.3]])]], + fields=["qx", "qy", "intensity"], + name="p", + ) + au = Crystal.from_ase(bulk("Au", "fcc", a=4.08, cubic=True), verbose=False) + fe = Crystal.from_ase(bulk("Fe", "bcc", a=2.87, cubic=True), verbose=False) + + with pytest.raises(ValueError, match="pass k_max"): + CrystalMap.from_vectors(peaks, [au, fe]) + + au.calculate_structure_factors(k_max=1.4) + fe.calculate_structure_factors(k_max=2.0) + with pytest.raises(ValueError, match="different k_max"): + CrystalMap.from_vectors(peaks, [au, fe]) + + cm = CrystalMap.from_vectors(peaks, [au, fe], k_max=1.8) + assert cm.k_max == 1.8 + assert au.k_max == fe.k_max == 1.8 + assert float(au.g_len.max()) <= 1.8 and float(fe.g_len.max()) <= 1.8 + + +def test_dynamical_update_is_local_and_examples_are_spread(): + from quantem.diffraction.crystal_map import CrystalMap + + rng = np.random.default_rng(4) + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + ti_b = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), name="Ti beta", verbose=False + ).calculate_structure_factors(k_max=1.5) + q_a = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + q_b = quat_from_axis_angle( + torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(20.0), dtype=torch.float64), + ) + R, C = 3, 8 + cells = [ + [_pattern(ti_a, q_a, rng) if c < 4 else _pattern(ti_b, q_b, rng) for c in range(C)] + for _ in range(R) + ] + peaks = Vector.from_data(cells, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + oms = [] + for xtl in (ti_a, ti_b): + om = OrientationMap.from_vectors(peaks, xtl, energy_ev=200e3) + om.build_plan() + om.match_orientations(progress_bar=False) + oms.append(om) + cm = CrystalMap.from_orientation_maps(oms) + cm.fit(progress_bar=False) + before = cm.phase_index.copy() + assert (before[:, :4] == 0).all() and (before[:, 4:] == 1).all() + + # a dynamical result that reached one position and flipped it + F = len(cm.phases.candidates) + cost = torch.full((R, C, F), torch.nan, dtype=torch.float64) + cost[0, 0] = torch.tensor([0.9, 0.1]) + cm.phases.apply_dynamical({"cost": cost, "phase_index": cost.nan_to_num(9).argmin(-1)}) + after = cm.phase_index + assert after[0, 0] == 1 + changed = after != before + changed[0, 0] = False + assert not changed.any() + assert (cm.phases.metadata["kinematical"]["phase_index"].numpy() == before).all() + + picks = cm.example_positions(phase="Ti beta", num=3, min_distance=3) + assert all(cm.phase_index[p] == 1 for p in picks) + assert all( + (a[0] - b[0]) ** 2 + (a[1] - b[1]) ** 2 >= 9 + for i, a in enumerate(picks) + for b in picks[:i] + ) + + +def test_loaded_crystal_map_shares_orientation_maps(tmp_path): + from quantem.core.io.serialize import load + from quantem.diffraction.crystal_map import CrystalMap + + rng = np.random.default_rng(5) + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + q = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + cells = [[_pattern(ti_a, q, rng) for _ in range(3)] for _ in range(2)] + peaks = Vector.from_data(cells, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + om = OrientationMap.from_vectors(peaks, ti_a, energy_ev=200e3) + om.build_plan() + om.match_orientations(progress_bar=False) + cm = CrystalMap.from_orientation_maps([om]) + cm.fit(progress_bar=False) + cm.save(tmp_path / "cm.zip", mode="o") + cm2 = load(tmp_path / "cm.zip") + # one set of maps: what a refinement writes through the phase map is + # what the crystal map shows + assert cm2.phases.orientation_maps[0] is cm2.orientation_maps[0] + assert cm2.phases.orientation_maps[0].crystal is cm2[0].crystal + + +def test_plot_calibration_shows_every_crystal(): + import matplotlib + + matplotlib.use("Agg") + from quantem.diffraction.crystal_map import CrystalMap + + rng = np.random.default_rng(6) + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + ti_b = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), name="Ti beta", verbose=False + ).calculate_structure_factors(k_max=1.5) + q = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + cells = [[_pattern(ti_a, q, rng) for _ in range(3)] for _ in range(2)] + peaks = Vector.from_data(cells, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + cm = CrystalMap.from_vectors(peaks, [ti_a, ti_b], energy_ev=200e3) + fig, axs = cm.plot_calibration() + # one column per crystal, the histogram above the azimuth panel + assert axs.shape == (2, 2) + assert [ax.get_title() for ax in axs[0]] == ["Ti alpha", "Ti beta"] + assert all(len(ax.collections) > 0 for ax in axs[1]) + matplotlib.pyplot.close(fig) + fig, ax = cm.plot_bragg_rings() + labels = [t.get_text() for t in ax.get_legend().get_texts()] + assert labels == ["Ti alpha rings", "Ti beta rings"] + matplotlib.pyplot.close(fig) + + +def test_refine_skips_crystals_out_of_the_running(): + from quantem.diffraction.crystal_map import CrystalMap + + rng = np.random.default_rng(7) + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + ti_b = Crystal.from_ase( + bulk("Ti", "bcc", a=3.26, cubic=True), name="Ti beta", verbose=False + ).calculate_structure_factors(k_max=1.5) + q_a = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + q_b = quat_from_axis_angle( + torch.tensor([1.0, 0.0, 0.0], dtype=torch.float64), + torch.tensor(np.deg2rad(20.0), dtype=torch.float64), + ) + cells = [ + [_pattern(ti_a, q_a, rng) if c < 3 else _pattern(ti_b, q_b, rng) for c in range(6)] + for _ in range(2) + ] + peaks = Vector.from_data(cells, fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t") + fractions = [] + for margin in (None, 0.1): + cm = CrystalMap.from_vectors(peaks, [ti_a, ti_b], energy_ev=200e3) + cm.build_plan() + cm.match_orientations(progress_bar=False) + corr = np.stack([om.corr[..., 0].numpy() for om in cm]) + q_lib = [om.quats.clone() for om in cm] + cm.refine_orientations(competitive_margin=margin, progress_bar=False) + cm.fit(progress_bar=False) + fractions.append(cm.phase_index.copy()) + if margin is not None: + out = corr < corr.max(axis=0) - margin + for i, om in enumerate(cm): + # where a crystal is out of the running its library match stays + assert torch.equal(om.quats[out[i]], q_lib[i][out[i]]) + assert out.any() + assert (fractions[0] == fractions[1]).all() + + +def test_zone_reflections_keep_one_zone(): + from quantem.diffraction.calibration import plot_calibration, zone_reflections + + ti_a = Crystal.from_ase( + bulk("Ti", "hcp", a=2.9505, c=4.6855), name="Ti alpha", verbose=False + ).calculate_structure_factors(k_max=1.5) + basal = zone_reflections(ti_a, (0, 0, 0, 1)) + # [0001] keeps exactly the hk0 reflections, and 3- and 4-index agree + assert (basal.hkl[:, 2] == 0).all() + assert len(basal.hkl) == int((ti_a.hkl[:, 2] == 0).sum()) + assert torch.equal(zone_reflections(ti_a, (0, 0, 1)).hkl, basal.hkl) + assert basal.name == "Ti alpha [0001]" and ti_a.name == "Ti alpha" + assert len(ti_a.hkl) > len(basal.hkl) + # the plots accept it and title the column by the zone + import matplotlib + + matplotlib.use("Agg") + rng = np.random.default_rng(8) + q = torch.tensor([1.0, 0.0, 0.0, 0.0], dtype=torch.float64) + peaks = Vector.from_data( + [[_pattern(ti_a, q, rng)]], fields=["qx", "qy", "intensity"], units=["A^-1"] * 3, name="t" + ) + fig, axs = plot_calibration(peaks, ti_a, zone_axis=(0, 0, 0, 1)) + assert axs[0, 0].get_title() == "Ti alpha [0001]" + matplotlib.pyplot.close(fig) diff --git a/uv.lock b/uv.lock index f9c7b5c21..d6dc44c36 100644 --- a/uv.lock +++ b/uv.lock @@ -125,6 +125,23 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ed/c9/d7977eaacb9df673210491da99e6a247e93df98c715fc43fd136ce1d3d33/arrow-1.4.0-py3-none-any.whl", hash = "sha256:749f0769958ebdc79c173ff0b0670d59051a535fa26e8eba02953dc19eb43205", size = 68797 }, ] +[[package]] +name = "ase" +version = "3.29.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "matplotlib" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "numpy", version = "2.5.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "scipy", version = "1.18.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/aa/36/cfe17324ea701d2caa479f25b134c54aeb36998443fd88c13e4f1af06f52/ase-3.29.0.tar.gz", hash = "sha256:ef4e2caa38169e3fbbc4164764a060d1877a6692519a4bed82521328eeb0d9aa", size = 2466051 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/66/a6/a3d29378618d3d9ba0962473f8bbf09a5ad15eb8b210c5153fb88d0b8aad/ase-3.29.0-py3-none-any.whl", hash = "sha256:7b9dd103f007810339c24acfee2f6b677c0c48443b21d3c98e52959246cf4ebf", size = 3014724 }, +] + [[package]] name = "asttokens" version = "3.0.2" @@ -3622,6 +3639,7 @@ name = "quantem" version = "0.1.9" source = { editable = "." } dependencies = [ + { name = "ase" }, { name = "cmasher" }, { name = "colorspacious" }, { name = "dill" }, @@ -3636,6 +3654,7 @@ dependencies = [ { name = "scikit-image" }, { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, { name = "scipy", version = "1.18.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "spglib" }, { name = "tensorboard" }, { name = "torch" }, { name = "torchinfo" }, @@ -3668,6 +3687,7 @@ test = [ [package.metadata] requires-dist = [ + { name = "ase", specifier = ">=3.23" }, { name = "cmasher", specifier = ">=1.9.2" }, { name = "colorspacious" }, { name = "dill" }, @@ -3681,6 +3701,7 @@ requires-dist = [ { name = "rosettasciio", specifier = ">=0.8.0" }, { name = "scikit-image", specifier = ">=0.25.2" }, { name = "scipy" }, + { name = "spglib", specifier = ">=2.5" }, { name = "tensorboard", specifier = ">=2.19.0" }, { name = "torch", specifier = ">=2.7.0" }, { name = "torchinfo", specifier = ">=1.8.0" }, @@ -4227,6 +4248,43 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/66/87/5ed59e1d0290564e2027ed066c52b059892c4637741337fc0967183c8d4d/soupsieve-2.10-py3-none-any.whl", hash = "sha256:8596eb8967d744174820280fa62b4542a2e955bfaccca73ed8a13c6eb8e9b502", size = 38019 }, ] +[[package]] +name = "spglib" +version = "2.7.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "numpy", version = "2.5.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/a9/06/7964acb4c444191376bd87f91579475fbe7623ca943cce40cee8fb7f2c36/spglib-2.7.0.tar.gz", hash = "sha256:c40907a42c9dc45572f46740bf95412f84fb0eda30267e31665d104a4bde6627", size = 2366134 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/44/a8/d841ae7743c58227af277f7f16aa844376fa11c426090d6ae35e7e93af76/spglib-2.7.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:3cf0ff80c01d8631ef4b9f1b78da79ff2044834e6e2d870f7f20c8579c921136", size = 910793 }, + { url = "https://files.pythonhosted.org/packages/e8/02/11baf94cf682cdaafa046b72d4b2adcf944e19e2b2741454e329dedb2fc2/spglib-2.7.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a7b29d2cfca6ac53e927686ca0b91257126e47f6abfa26451723a5cd40070352", size = 944977 }, + { url = "https://files.pythonhosted.org/packages/ce/fa/6d1bc8f8cb08945ca8c37c95b42bf336b6b9a8a737eced1ce64f0cebe9ce/spglib-2.7.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f892ecce2dd1bc636b14a4e5bc13aabb73b008bd37a4d23636882c8971c432a0", size = 960531 }, + { url = "https://files.pythonhosted.org/packages/b0/79/2fd5e33b431cd0afcdd441bd10704c11cdf74c09b721249297284e5bf0b2/spglib-2.7.0-cp311-cp311-win_amd64.whl", hash = "sha256:468879702577124dcde0607a75396576e256f1cfa2d8fe48da4a928fbb27abc6", size = 669827 }, + { url = "https://files.pythonhosted.org/packages/f7/e9/4e07c9c1bda40df54e09bd686eae0dc13d46e76a5ef4d43582971a86eb32/spglib-2.7.0-cp311-cp311-win_arm64.whl", hash = "sha256:ceb6730a2324d0c83579c803f3782e28bd41e79bbfe0c3dfdbf30e3d3a6320e5", size = 649076 }, + { url = "https://files.pythonhosted.org/packages/68/0e/36720beeca8452530e50ab8a16b91e8721e34c0f97fd25e9c4ddd8b9324b/spglib-2.7.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:ef70132e23dfcc7ab6813742e0edab3f9906e61cd11c857f014bd5610a8bc88c", size = 911009 }, + { url = "https://files.pythonhosted.org/packages/47/a0/24df91cbde6a3237d54cfb21602cc8ebb4102cd4e3ec9497c66135c2b190/spglib-2.7.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:59f134e74f7f488de4bf5579ee6a35af25cb2c478c138de664fea1e14f3efbaf", size = 946821 }, + { url = "https://files.pythonhosted.org/packages/67/4c/75ac6f7ade28019b216c7333322f2886e1c0105202cd74506f530664bf26/spglib-2.7.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:6913906fd9108e7bb2ce06a810513a95a82d801530f10230979bf3427bb7e771", size = 962531 }, + { url = "https://files.pythonhosted.org/packages/a7/5f/4e283139af178bb445eedff281a90e66ceff1b814ace70a9d90a2197acc3/spglib-2.7.0-cp312-cp312-win_amd64.whl", hash = "sha256:d5729ff0040baae764c17249302cd99f0eb4e73449612a8c69d3e60a215f062e", size = 671111 }, + { url = "https://files.pythonhosted.org/packages/a5/b9/20e52d46e33bf69ceba4fc86602f006c06ce4ab10e3c930f4722fb270b02/spglib-2.7.0-cp312-cp312-win_arm64.whl", hash = "sha256:54f4b6e789475384c62e759c618172707f261c0eae8017949fe4994b6b8cc779", size = 646679 }, + { url = "https://files.pythonhosted.org/packages/2c/1c/a0fe8c0523a0e7d608f49f09895e5c599329265c9bfacd269a21458b7564/spglib-2.7.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:ab061ea6a3c3c25a1d0018b09c333c0458792036d3f45d892bd52793ed1f1bda", size = 911085 }, + { url = "https://files.pythonhosted.org/packages/2a/34/cb3c522c4aaf6ce319b37bbec71d373b9e2cf0bcfe7d42c365cd6c113b4b/spglib-2.7.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:be28673e90f7a6c7770f73c57e529d2bdbb373d06d26ee5e90991b548e9238aa", size = 946857 }, + { url = "https://files.pythonhosted.org/packages/9e/64/3b1213f2f655ff143ed142292b47ec3f1f9bda8641e659a7e33c4cf0e8a9/spglib-2.7.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f627a4ed6f2396ed6e3e8eaf33a53ad143c8ffb8756a84a640f4569ac5ffa2a7", size = 962470 }, + { url = "https://files.pythonhosted.org/packages/5c/3a/c51883ce739a00f9f60196f3dcb4ed91b690299a4ec64defd8ec5b2c5899/spglib-2.7.0-cp313-cp313-win_amd64.whl", hash = "sha256:c76411bc1b96cd87c8733994747c7692512b583bb4ef89a65463ff4255221c11", size = 671073 }, + { url = "https://files.pythonhosted.org/packages/35/78/3f9ec6ae93a48527dce0eceb6eeab74e6ad1fb2977adb5cbdfc03d43193c/spglib-2.7.0-cp313-cp313-win_arm64.whl", hash = "sha256:0d8ecf030d13d67c4cc272423e5652b74eda57f86a0b118e007f6d12974cc256", size = 646711 }, + { url = "https://files.pythonhosted.org/packages/1b/47/86e3c15c3e1c252bde40a794eea4742c142f23fc5f9c3d7551f083c1fa20/spglib-2.7.0-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:95e3dd7ef992ff8a88f6ef2e5909aaa60ecb479004cc1f73c1e6285d54227960", size = 911712 }, + { url = "https://files.pythonhosted.org/packages/05/61/ab2447bb47fa69934adc2fc2d13f771dedd3b2fd3171c95307446c948f01/spglib-2.7.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:97e0fcea2db3915bd973fdd2cc0a757b1f99bda71ce815da333d75ad1ffc3eb1", size = 947528 }, + { url = "https://files.pythonhosted.org/packages/9d/69/898d9e005131b0b1c7e5dce2b79f36aeb20ec4d3a88cca596b522a0fa4df/spglib-2.7.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:39b978c08ef2ebc0eaba833c488fc4c0f9b1fc0f50d4a8584f176741eea69376", size = 962474 }, + { url = "https://files.pythonhosted.org/packages/c8/56/7b25ee5348722dc93ca245ed950f1a89f8a944906140629055f394c072a4/spglib-2.7.0-cp314-cp314-win_amd64.whl", hash = "sha256:5f334b4b66c8aafd583fafab5b15a56e27efdd2dc6cb1064dfcd0fe59ae130f4", size = 679679 }, + { url = "https://files.pythonhosted.org/packages/20/37/eda9a34f25b13e47298fa1b94cc4dfd8b0fcfc46c7d63ea046aa1bf91fe7/spglib-2.7.0-cp314-cp314-win_arm64.whl", hash = "sha256:b032842fc223de46d2ef7d220459e1a61ed90329ac2e72818c605f1fc87451b8", size = 656403 }, + { url = "https://files.pythonhosted.org/packages/39/af/1c8d0f98d07969b7fa7323d522732124d88caf4ee3b680ef59120bd7b229/spglib-2.7.0-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:b7e29c796cfdadcc3857aef330acc19b9bc50c83e9911fb23b28390e7c80bae5", size = 920791 }, + { url = "https://files.pythonhosted.org/packages/b5/c6/89a3f31f831efc4108a19f110873559990b72186745cd3e151de28b256cc/spglib-2.7.0-cp314-cp314t-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9b6ca88bb6e604bc8f63efe87b3b2470c2e25f56988b775bd332cefa8866f5c5", size = 946881 }, + { url = "https://files.pythonhosted.org/packages/7e/e9/1ca63db2cebd381bd6b27ae309f25d270e70928359a6f0360db09b77894e/spglib-2.7.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:50629939a9cd6fa3df5a12f6f025ceb3c78534284f875371574c360e4ccaf5e1", size = 963803 }, + { url = "https://files.pythonhosted.org/packages/28/97/459b37c3802633f77c883883c75f5d4429b601ae8d930410b999c4e1dafb/spglib-2.7.0-cp314-cp314t-win_amd64.whl", hash = "sha256:cb77daaf9dd5d48d523a888f37cebd47fa63ff28dfcf1aac2b031b914f9ed55a", size = 696536 }, +] + [[package]] name = "sqlalchemy" version = "2.1.1"