From b152cc6f70d04b6eaf915a5296fa14e707831439 Mon Sep 17 00:00:00 2001 From: cophus Date: Sun, 14 Jun 2026 14:59:37 -0700 Subject: [PATCH 1/6] initial atomic tracing --- src/quantem/__init__.py | 1 + src/quantem/core/datastructures/dataset3d.py | 34 +- src/quantem/tomography/__init__.py | 1 + src/quantem/tomography/atom_analysis.py | 0 src/quantem/tomography/atom_trace.py | 920 +++++++++++++++++++ widget/src/quantem/widget/__init__.py | 3 +- widget/src/quantem/widget/show3d_atoms.py | 122 +++ 7 files changed, 1067 insertions(+), 14 deletions(-) create mode 100644 src/quantem/tomography/atom_analysis.py create mode 100644 src/quantem/tomography/atom_trace.py create mode 100644 widget/src/quantem/widget/show3d_atoms.py diff --git a/src/quantem/__init__.py b/src/quantem/__init__.py index ba70f629f..d369aa568 100644 --- a/src/quantem/__init__.py +++ b/src/quantem/__init__.py @@ -10,5 +10,6 @@ from quantem import imaging as imaging from quantem import diffractive_imaging as diffractive_imaging +from quantem import tomography as tomography __version__ = version("quantem") diff --git a/src/quantem/core/datastructures/dataset3d.py b/src/quantem/core/datastructures/dataset3d.py index 1af7b4e62..ba2583d55 100644 --- a/src/quantem/core/datastructures/dataset3d.py +++ b/src/quantem/core/datastructures/dataset3d.py @@ -187,12 +187,13 @@ def show( start: int = 0, end: int | None = None, step: int = 1, - max: int | None = 20, + max_indices: int | None = 20, ncols: int = 4, scalebar: ScalebarConfig | bool = False, title_prefix: str | None = None, suptitle: str | None = None, returnfig: bool = False, + same_scale: bool = True, **kwargs, ) -> tuple[Figure, Axes] | None: """ @@ -203,10 +204,10 @@ def show( start : int, default 0 First frame index. Supports negative indexing. end : int or None, optional - End frame index (exclusive). If None, determined by max. + End frame index (exclusive). If None, determined by max_indices. step : int, default 1 Step between frames. Negative step shows frames in reverse order. - max : int or None, default 20 + max_indices : int or None, default 20 Maximum number of frames to show. Prevents memory issues. Set to None to show all frames. ncols : int, default 4 @@ -220,6 +221,10 @@ def show( Figure super title displayed above all subplots. returnfig : bool, default False If True, returns (fig, axes). + same_scale : bool, default True + If True, all slices share one intensity range (the volume's global + min/max) so they are directly comparable. Ignored when you pass + ``vmin``/``vmax``/``norm``. Set False for per-slice auto-contrast. **kwargs : dict Keyword arguments for show_2d (cmap, cbar, vmin, vmax, norm, etc.). @@ -232,7 +237,7 @@ def show( Raises ------ ValueError - If step is zero, ncols < 1, max < 1, start is out of bounds, + If step is zero, ncols < 1, max_indices < 1, start is out of bounds, or the specified range has no frames to display. Examples @@ -240,12 +245,12 @@ def show( Basic usage: >>> data.show() # first 20 frames - >>> data.show(max=None) # all frames (use with caution) + >>> data.show(max_indices=None) # all frames (use with caution) Single frame: - >>> data.show(start=5, max=1) # frame 5 - >>> data.show(start=-1, max=1) # last frame + >>> data.show(start=5, max_indices=1) # frame 5 + >>> data.show(start=-1, max_indices=1) # last frame Frame range: @@ -257,7 +262,7 @@ def show( Grid layout: >>> data.show(ncols=2) # 2 columns - >>> data.show(ncols=5, max=10) # 5x2 grid + >>> data.show(ncols=5, max_indices=10) # 5x2 grid Titles: @@ -278,8 +283,8 @@ def show( raise ValueError("Step cannot be zero.") if ncols < 1: raise ValueError(f"ncols must be >= 1, got {ncols}.") - if max is not None and max < 1: - raise ValueError(f"max must be >= 1 or None, got {max}.") + if max_indices is not None and max_indices < 1: + raise ValueError(f"max_indices must be >= 1 or None, got {max_indices}.") if start < 0: start = total_frames + start if start < 0 or start >= total_frames: @@ -299,9 +304,9 @@ def show( if step > 0: end_idx = min(end_idx, total_frames) - # Apply max limit to avoid creating huge list - if max is not None: - max_end = start + max * step + # Apply max_indices limit to avoid creating huge list + if max_indices is not None: + max_end = start + max_indices * step if step > 0: end_idx = min(end_idx, max_end) elif max_end > end_idx: @@ -328,6 +333,9 @@ def show( labels.extend([""] * pad_count) image_grid = [images[i : i + ncols] for i in range(0, len(images), ncols)] label_grid = [labels[i : i + ncols] for i in range(0, len(labels), ncols)] + if same_scale and not any(key in kwargs for key in ("vmin", "vmax", "norm")): + kwargs["vmin"] = float(np.nanmin(self.array)) + kwargs["vmax"] = float(np.nanmax(self.array)) fig, axes = show_2d(image_grid, scalebar=scalebar, title=label_grid, **kwargs) if pad_count > 0: for ax in np.array(axes).flat[-pad_count:]: diff --git a/src/quantem/tomography/__init__.py b/src/quantem/tomography/__init__.py index e69de29bb..2ec8d60c0 100644 --- a/src/quantem/tomography/__init__.py +++ b/src/quantem/tomography/__init__.py @@ -0,0 +1 @@ +from quantem.tomography.atom_trace import Atoms as Atoms diff --git a/src/quantem/tomography/atom_analysis.py b/src/quantem/tomography/atom_analysis.py new file mode 100644 index 000000000..e69de29bb diff --git a/src/quantem/tomography/atom_trace.py b/src/quantem/tomography/atom_trace.py new file mode 100644 index 000000000..e0975ec3c --- /dev/null +++ b/src/quantem/tomography/atom_trace.py @@ -0,0 +1,920 @@ +"""Atom tracing from 3D tomographic volumes via differentiable Gaussian splatting. + +This module decomposes a reconstructed 3D volume (``Dataset3d``) into a set of +atomic sites, each modeled as a 3D Gaussian. Rather than fitting one atom at a +time (the classic Levenberg-Marquardt approach), the whole volume is treated as a +single differentiable model and *all* parameters are optimized jointly with Adam. +Overlap between neighboring atoms is handled automatically by the joint fit. + +Forward model +------------- +The volume is decomposed into a sharp atomic part and a smooth background, both +non-negative:: + + volume = volume_atoms + volume_background, volume_atoms >= 0, volume_background >= 0 + +``volume_atoms`` + A sum of 3D Gaussians. Intensities are kept >= 0, so the atomic part is >= 0. +``volume_background`` + A low-degree tensor-product Bezier (Bernstein) field over a small control + lattice. With non-negative control points it is non-negative *everywhere* + (Bernstein partition-of-unity) and slowly varying by construction, so it + cannot absorb sharp atomic features -- the model class itself separates the + two. This is the hook for future background regularization. + +Atom models +----------- +isotropic + 5 degrees of freedom per atom: ``x, y, z`` position, intensity ``I``, and a + single width ``sigma``. +anisotropic (planned) + 10 degrees of freedom: ``x, y, z``, ``I`` and a Cholesky-parameterized + precision matrix ``Lambda = L @ L.T`` (positive-definite by construction, so + the Gaussian can never diverge). + +Efficiency +---------- +The renderer never evaluates every Gaussian over every voxel. Each Gaussian only +contributes to a small ``(2*window_radius + 1)**3`` window around its rounded +center, accumulated with ``scatter_add``. Cost is ``O(n_atoms * window_volume)`` +rather than ``O(n_atoms * n_voxels)``. A window of 3-4 sigma is the right +accuracy/speed trade-off for tracing (relative truncation error ~1e-3 at 3 sigma, +~6e-6 at 4 sigma); larger windows are only needed to *measure* truncation. + +Conventions +----------- +All internal site coordinates are in **voxel / array-index units** matching the +volume's axes (axis 0, 1, 2). Conversion to physical units happens only at the +public ``sites`` (Vector) boundary, using the volume's ``sampling``/``origin``. + +.. note:: + Under construction. The differentiable forward model below (atom renderer + + Bezier background) is complete and verified; the ``Atoms`` class, seeding, + schedule, and regularizers follow. +""" + +from __future__ import annotations + +import math + +import numpy as np +import torch +import torch.nn.functional as F +from torch import Tensor +from tqdm.auto import tqdm + +from quantem.core import config +from quantem.core.datastructures import Dataset3d, Vector +from quantem.core.io.serialize import AutoSerialize + +__all__ = [ + "Atoms", + "render_isotropic", + "bernstein_basis", + "render_background", + "gaussian_blur3d", + "seed_peaks", +] + + +def render_isotropic( + positions: Tensor, + intensities: Tensor, + sigmas: Tensor | float, + volume_shape: tuple[int, int, int], + window_radius: int, +) -> Tensor: + """Render a sum of isotropic 3D Gaussians into a dense volume by local splatting. + + Each Gaussian contributes only to a ``(2*window_radius + 1)**3`` window around + its rounded center. The window *indices* come from ``round(positions)`` (held + constant w.r.t. gradients), while the Gaussian is evaluated at the continuous + ``positions``, so gradients flow to ``positions``, ``intensities`` and + ``sigmas``. Choose ``window_radius >~ 3-4 * sigma_max`` so that truncation at + the window edge is negligible for tracing. + + Parameters + ---------- + positions : Tensor + ``(N, 3)`` float tensor of site centers in voxel/array-index coordinates. + intensities : Tensor + ``(N,)`` float tensor of Gaussian amplitudes (kept >= 0 by the caller). + sigmas : Tensor or float + Gaussian width(s) in voxels. Scalar / shape ``(1,)`` broadcasts to all + atoms; otherwise shape ``(N,)``. + volume_shape : tuple[int, int, int] + Output volume shape ``(D0, D1, D2)``. + window_radius : int + Half-width (in voxels) of the cubic splat window per atom. + + Returns + ------- + Tensor + ``(D0, D1, D2)`` rendered atomic volume, differentiable w.r.t. + ``positions``, ``intensities`` and ``sigmas``. + """ + device = positions.device + dtype = positions.dtype + n_atoms = positions.shape[0] + d0, d1, d2 = volume_shape + + if n_atoms == 0: + return torch.zeros(volume_shape, device=device, dtype=dtype) + + sigmas = torch.as_tensor(sigmas, device=device, dtype=dtype) + if sigmas.ndim == 0: + sigmas = sigmas.reshape(1) + if sigmas.shape[0] == 1: + sigmas = sigmas.expand(n_atoms) + + # Integer window centers -- detached so they carry no gradient. + centers = torch.round(positions.detach()).long() # (N, 3) + + # Cubic window offsets, (W**3, 3). + rng = torch.arange(-window_radius, window_radius + 1, device=device) + o0, o1, o2 = torch.meshgrid(rng, rng, rng, indexing="ij") + offsets = torch.stack((o0.reshape(-1), o1.reshape(-1), o2.reshape(-1)), dim=1) + + # Integer voxel coordinates for every (atom, window-voxel): (N, W**3, 3). + vox = centers[:, None, :] + offsets[None, :, :] + + # Continuous displacement from the (sub-voxel) atom center to each voxel. + diff = vox.to(dtype) - positions[:, None, :] # (N, W**3, 3) + r2 = (diff * diff).sum(dim=-1) # (N, W**3) + val = intensities[:, None] * torch.exp(-0.5 * r2 / (sigmas[:, None] ** 2)) + + # Zero out-of-bounds contributions; clamp their indices to a valid slot so the + # scatter is safe (those entries add 0.0). + shape_t = torch.tensor(volume_shape, device=device) + valid = ((vox >= 0) & (vox < shape_t)).all(dim=-1) # (N, W**3) + val = val * valid.to(dtype) + vox = torch.minimum(torch.maximum(vox, torch.zeros_like(shape_t)), shape_t - 1) + + flat_idx = (vox[..., 0] * (d1 * d2) + vox[..., 1] * d2 + vox[..., 2]).reshape(-1) + out = torch.zeros(d0 * d1 * d2, device=device, dtype=dtype) + out = out.scatter_add(0, flat_idx, val.reshape(-1)) + return out.reshape(volume_shape) + + +def bernstein_basis( + num_samples: int, + degree: int, + device: torch.device | str | None = None, + dtype: torch.dtype = torch.float32, +) -> Tensor: + """Bernstein basis matrix sampled on a uniform grid over ``[0, 1]``. + + ``B[s, i] = C(degree, i) * t**i * (1 - t)**(degree - i)`` with + ``t = s / (num_samples - 1)``. Each row sums to 1 (partition of unity), which + makes a Bezier field a convex combination of its control points. + + Parameters + ---------- + num_samples : int + Number of evenly spaced sample points (the axis length of the volume). + degree : int + Polynomial degree; the control lattice has ``degree + 1`` points on this + axis. Low degree -> smoother, more slowly varying field. + device, dtype + Torch device/dtype for the returned matrix. + + Returns + ------- + Tensor + ``(num_samples, degree + 1)`` basis matrix. + """ + t = torch.linspace(0.0, 1.0, num_samples, device=device, dtype=dtype)[:, None] + i = torch.arange(degree + 1, device=device, dtype=dtype)[None, :] + log_binom = ( + torch.lgamma(torch.tensor(degree + 1.0, device=device, dtype=dtype)) + - torch.lgamma(i + 1.0) + - torch.lgamma(degree - i + 1.0) + ) + # torch.pow gives 0**0 == 1, so the endpoints interpolate the corner controls. + return torch.exp(log_binom) * t.pow(i) * (1.0 - t).pow(degree - i) + + +def render_background( + control_points: Tensor, + bases: tuple[Tensor, Tensor, Tensor], +) -> Tensor: + """Evaluate a tensor-product Bezier (Bernstein) field over a dense volume. + + With non-negative ``control_points`` the field is non-negative everywhere and + bounded by ``[control_points.min(), control_points.max()]`` (partition of + unity), and is smooth/slowly varying for low control-lattice degree. + + Parameters + ---------- + control_points : Tensor + ``(n0 + 1, n1 + 1, n2 + 1)`` control lattice (kept >= 0 by the caller). + bases : tuple of Tensor + Per-axis Bernstein bases ``(B0, B1, B2)`` from :func:`bernstein_basis`, + with shapes ``(D0, n0+1)``, ``(D1, n1+1)``, ``(D2, n2+1)``. + + Returns + ------- + Tensor + ``(D0, D1, D2)`` background field, differentiable w.r.t. ``control_points``. + """ + b0, b1, b2 = bases + f = torch.einsum("ijk,ai->ajk", control_points, b0) + f = torch.einsum("ajk,bj->abk", f, b1) + f = torch.einsum("abk,ck->abc", f, b2) + return f + + +def gaussian_blur3d(volume: Tensor, sigma: float) -> Tensor: + """Separable 3D Gaussian blur with replicate padding. + + Parameters + ---------- + volume : Tensor + ``(D0, D1, D2)`` volume. + sigma : float + Standard deviation in voxels. ``sigma <= 0`` returns the input unchanged. + + Returns + ------- + Tensor + Blurred ``(D0, D1, D2)`` volume. + """ + if sigma <= 0: + return volume + radius = max(1, int(math.ceil(3.0 * sigma))) + x = torch.arange(-radius, radius + 1, device=volume.device, dtype=volume.dtype) + kernel = torch.exp(-0.5 * (x / sigma) ** 2) + kernel = kernel / kernel.sum() + v = volume[None, None] # (1, 1, D0, D1, D2) + for axis in range(3): + shape = [1, 1, 1, 1, 1] + shape[2 + axis] = kernel.numel() + pad = [0, 0, 0, 0, 0, 0] # F.pad order is (W_lo, W_hi, H_lo, H_hi, D_lo, D_hi) + pad[(2 - axis) * 2] = radius + pad[(2 - axis) * 2 + 1] = radius + v = F.pad(v, pad, mode="replicate") + v = F.conv3d(v, kernel.reshape(shape)) + return v[0, 0] + + +def seed_peaks( + volume: Tensor, + blur_sigmas: tuple[float, float] | float = (1.0, 2.0), + threshold: float | None = None, + threshold_pct: float = 99.0, + min_distance: float = 0.0, + max_peaks: int | None = None, + progress: bool = False, +) -> tuple[Tensor, Tensor]: + """Seed candidate atom sites by difference-of-Gaussians peak detection. + + Bandpass-filters the volume (single blur, or a difference of two blurs to also + suppress smooth background), finds 3x3x3 local maxima above a threshold, + refines each to sub-voxel accuracy with a per-axis parabolic fit, and + optionally enforces a minimum spacing (greedy, brightest-first). + + Parameters + ---------- + volume : Tensor + ``(D0, D1, D2)`` volume. + blur_sigmas : tuple[float, float] or float + Two sigmas -> difference-of-Gaussians (bandpass). One sigma -> single + matched-filter blur. In voxels. + threshold : float or None + Absolute threshold on the filtered response. Overrides ``threshold_pct``. + threshold_pct : float + Percentile (0-100) of the filtered response used as the threshold when + ``threshold`` is None. Robust to outliers and to data scale; default 99.0. + min_distance : float + Minimum spacing between seeds in voxels; closer (dimmer) peaks are + dropped. 0 disables. + max_peaks : int or None + Keep at most this many seeds (brightest first). + progress : bool + Show a tqdm progress bar over the detection stages. + + Returns + ------- + positions : Tensor + ``(M, 3)`` sub-voxel seed coordinates in voxel/array-index units. + intensities : Tensor + ``(M,)`` volume value sampled at each seed (integer voxel). + """ + # Peak detection runs on CPU for cross-device determinism: the local-maximum + # test relies on exact float equality (resp == max_pool(resp)), which is not + # reliable on MPS. Seeding is fast and one-time, so this is cheap. + bar = tqdm(total=4, desc="find_initial", disable=not progress) + volume = volume.detach().to("cpu") + dtype = volume.dtype + if isinstance(blur_sigmas, (tuple, list)): + s_lo, s_hi = blur_sigmas + resp = gaussian_blur3d(volume, float(s_lo)) - gaussian_blur3d(volume, float(s_hi)) + else: + resp = gaussian_blur3d(volume, float(blur_sigmas)) + if threshold is None: + threshold = float(np.percentile(resp.numpy(), threshold_pct)) + bar.set_postfix_str("DoG bandpass") + bar.update(1) + + # 3x3x3 local maxima above threshold (max_pool uses -inf padding, so the + # equality test is exact for the window argmax). + pooled = F.max_pool3d(resp[None, None], kernel_size=3, stride=1, padding=1)[0, 0] + is_peak = (resp == pooled) & (resp > threshold) + # Drop the outer shell so the parabolic fit always has neighbors. + is_peak[[0, -1], :, :] = False + is_peak[:, [0, -1], :] = False + is_peak[:, :, [0, -1]] = False + idx = is_peak.nonzero(as_tuple=False) # (M, 3) long + bar.set_postfix_str(f"{idx.shape[0]} maxima") + bar.update(1) + if idx.shape[0] == 0: + bar.close() + empty = torch.zeros((0, 3), dtype=dtype) + return empty, empty[:, 0] + + i, j, k = idx[:, 0], idx[:, 1], idx[:, 2] + f0 = resp[i, j, k] + offsets = torch.zeros_like(idx, dtype=dtype) + eps = torch.finfo(dtype).eps + neighbor_pairs = ( + (resp[i - 1, j, k], resp[i + 1, j, k]), + (resp[i, j - 1, k], resp[i, j + 1, k]), + (resp[i, j, k - 1], resp[i, j, k + 1]), + ) + for axis, (f_lo, f_hi) in enumerate(neighbor_pairs): + denom = f_lo - 2.0 * f0 + f_hi + shift = torch.where( + denom.abs() > eps, 0.5 * (f_lo - f_hi) / denom, torch.zeros_like(denom) + ) + offsets[:, axis] = shift.clamp(-0.5, 0.5) + + positions = idx.to(dtype) + offsets + intensities = volume[i, j, k] + # Sort brightest-first. + order = torch.argsort(intensities, descending=True) + positions, intensities = positions[order], intensities[order] + bar.set_postfix_str("subpixel refine") + bar.update(1) + + # Greedy minimum-distance suppression (brightest wins). + if min_distance > 0 and positions.shape[0] > 1: + from scipy.spatial import cKDTree + + pts = positions.numpy() + neighbors = cKDTree(pts).query_ball_point(pts, r=float(min_distance)) + suppressed = np.zeros(pts.shape[0], dtype=bool) + keep: list[int] = [] + for a in tqdm(range(pts.shape[0]), desc="dedup", disable=not progress, leave=False): + if suppressed[a]: + continue + keep.append(a) + for b in neighbors[a]: + if b > a: + suppressed[b] = True + keep_t = torch.as_tensor(keep) + positions, intensities = positions[keep_t], intensities[keep_t] + bar.set_postfix_str(f"{positions.shape[0]} sites") + bar.update(1) + bar.close() + + if max_peaks is not None and positions.shape[0] > max_peaks: + positions, intensities = positions[:max_peaks], intensities[:max_peaks] + + return positions, intensities + + +def _inverse_softplus(y: Tensor) -> Tensor: + """Numerically stable inverse of ``softplus`` (no overflow for large y). + + ``log(exp(y) - 1) = y + log(-expm1(-y))``; for large y the second term -> 0, + so the result stays finite (the naive ``log(expm1(y))`` overflows in float32 + around y ~ 88). + """ + y = y.clamp(min=1e-6) + return y + torch.log(-torch.expm1(-y)) + + +class Atoms(AutoSerialize): + """Atomic sites traced from a 3D volume by differentiable Gaussian splatting. + + The volume is modeled as ``volume_atoms + volume_background`` (both >= 0) and + all parameters are optimized jointly. Atom parameters are stored internally in + **voxel/array-index** coordinates; the public ``sites`` Vector reports them in + the source volume's physical units. + + Construct with :meth:`from_dataset`, then call :meth:`find_initial` (and, + soon, ``trace``/``optimize``) to populate and refine the sites. + + Parameters + ---------- + volume : Dataset3d + Input 3D volume (e.g. a tomographic reconstruction). + model : {"isotropic"} + Atom model. Only isotropic (x, y, z, intensity, sigma) is implemented; + anisotropic is planned. + sigma_init : float + Initial Gaussian width in voxels. + sigma_cutoff : float + Splat window half-width in units of sigma (3-4 is the speed/accuracy + sweet spot). + background : bool + If True, model a smooth non-negative Bezier background. + background_degree : int or tuple[int, int, int] + Per-axis degree of the background control lattice (lower = smoother). + device : str, int, or None + Compute device; None selects the quantem default (cuda > mps > cpu). + """ + + _token = object() + + def __init__( + self, + volume: Dataset3d, + *, + model: str = "isotropic", + sigma_init: float = 2.0, + sigma_cutoff: float = 4.0, + background: bool = True, + background_degree: int | tuple[int, int, int] = 4, + device: str | int | None = None, + _token: object | None = None, + ): + if _token is not self._token: + raise RuntimeError("Use Atoms.from_dataset() to instantiate this class.") + super().__init__() + if model != "isotropic": + raise NotImplementedError( + "Only model='isotropic' is implemented; anisotropic is planned." + ) + dev, _ = config.validate_device(device) + self.device = dev + self._model = model + self._sigma_init = float(sigma_init) + if self._sigma_init <= 0: + raise ValueError(f"sigma_init must be > 0, got {self._sigma_init}.") + self._sigma_cutoff = float(sigma_cutoff) + self._dtype = torch.float32 + + # Calibration from the source volume (for voxel <-> physical conversion). + self._source = volume + self._sampling = np.asarray(volume.sampling, dtype=float) + self._origin = np.asarray(volume.origin, dtype=float) + self._units = list(volume.units) + self._signal_units = volume.signal_units + + # Raw volume as a torch tensor on the chosen device. + arr = volume.array + vol = torch.as_tensor(np.asarray(arr)) if arr is not None else volume.tensor + self._volume = vol.to(device=self.device, dtype=self._dtype) + self._shape = tuple(int(s) for s in self._volume.shape) + + # Atom parameters (empty until seeded). Stored raw/unconstrained. + self._positions = torch.zeros((0, 3), dtype=self._dtype, device=self.device) + self._raw_intensity = torch.zeros((0,), dtype=self._dtype, device=self.device) + self._raw_sigma = torch.zeros((0,), dtype=self._dtype, device=self.device) + + # Bezier background control lattice + per-axis Bernstein bases. + self._has_background = bool(background) + if self._has_background: + degs = ( + (background_degree,) * 3 + if isinstance(background_degree, int) + else tuple(int(d) for d in background_degree) + ) + self._bg_degrees = degs + self._bases = tuple( + bernstein_basis(self._shape[a], degs[a], device=self.device, dtype=self._dtype) + for a in range(3) + ) + init_bg = max(float(self._volume.min()), 1e-4) + self._raw_background = _inverse_softplus( + torch.full( + tuple(d + 1 for d in degs), init_bg, dtype=self._dtype, device=self.device + ) + ) + else: + self._bg_degrees = None + self._bases = None + self._raw_background = None + + # ------------------------------------------------------------------ # + # Construction + # ------------------------------------------------------------------ # + @classmethod + def from_dataset( + cls, + volume: Dataset3d, + *, + model: str = "isotropic", + sigma_init: float = 2.0, + sigma_cutoff: float = 4.0, + background: bool = True, + background_degree: int | tuple[int, int, int] = 4, + device: str | int | None = None, + ) -> "Atoms": + """Create an :class:`Atoms` tracer from a 3D :class:`Dataset3d` volume.""" + if not isinstance(volume, Dataset3d): + raise TypeError(f"volume must be a Dataset3d, got {type(volume).__name__}.") + return cls( + volume, + model=model, + sigma_init=sigma_init, + sigma_cutoff=sigma_cutoff, + background=background, + background_degree=background_degree, + device=device, + _token=cls._token, + ) + + # ------------------------------------------------------------------ # + # Constrained parameter views (raw -> physical-meaning) + # ------------------------------------------------------------------ # + @property + def _intensity(self) -> Tensor: + return F.softplus(self._raw_intensity) + + @property + def _sigma(self) -> Tensor: + return F.softplus(self._raw_sigma) + + @property + def _background_control(self) -> Tensor | None: + return None if not self._has_background else F.softplus(self._raw_background) + + def _window_radius(self) -> int: + smax = self._sigma_init if self.num_sites == 0 else float(self._sigma.detach().max()) + return max(1, int(math.ceil(self._sigma_cutoff * smax))) + + # ------------------------------------------------------------------ # + # State + # ------------------------------------------------------------------ # + @property + def num_sites(self) -> int: + """Number of atomic sites.""" + return int(self._positions.shape[0]) + + @property + def model(self) -> str: + """Atom model name.""" + return self._model + + # ------------------------------------------------------------------ # + # Seeding + # ------------------------------------------------------------------ # + def find_initial( + self, + blur_sigmas: tuple[float, float] | float | None = None, + threshold: float | None = None, + threshold_pct: float = 99.0, + min_distance: float | None = None, + max_peaks: int | None = None, + progress: bool = True, + ) -> "Atoms": + """Find an initial set of sites by difference-of-Gaussians peak detection. + + This is the first step of tracing (the seed for later refinement and + densification). Defaults derive from ``sigma_init``: the DoG bandpass + brackets it, the threshold is the ``threshold_pct`` percentile of the + response, and ``min_distance`` is ``2 * sigma_init``. Pass an absolute + ``threshold`` to override. Set ``progress=False`` to silence the bar. + See :func:`seed_peaks` for the detector. + """ + if blur_sigmas is None: + blur_sigmas = (0.75 * self._sigma_init, 1.5 * self._sigma_init) + if min_distance is None: + min_distance = 2.0 * self._sigma_init + pos, inten = seed_peaks( + self._volume, + blur_sigmas=blur_sigmas, + threshold=threshold, + threshold_pct=threshold_pct, + min_distance=min_distance, + max_peaks=max_peaks, + progress=progress, + ) + self._positions = pos.to(device=self.device, dtype=self._dtype) + self._raw_intensity = _inverse_softplus(inten.clamp(min=1e-6)).to( + device=self.device, dtype=self._dtype + ) + self._raw_sigma = _inverse_softplus( + torch.full((pos.shape[0],), self._sigma_init, dtype=self._dtype, device=self.device) + ) + return self + + # ------------------------------------------------------------------ # + # Refinement (joint gradient optimization + site add / remove / merge) + # ------------------------------------------------------------------ # + def refine( + self, + num_iterations: int = 100, + learning_rate: float = 0.05, + loss: str = "huber", + sigma_bounds: tuple[float, float] | None = None, + intensity_min: float | None = None, + add: bool = True, + remove: bool = True, + merge: bool = True, + update_every: int = 25, + min_distance: float | None = None, + progress: bool = True, + ) -> "Atoms": + """Refine the atomic model by joint gradient descent. + + All site parameters (positions, intensities, widths) and the background + are optimized together against the measured volume. By default the site + set is also maintained during refinement: new sites are *added* at peaks + of the residual, weak sites are *removed*, and duplicates are *merged*. + + Parameters + ---------- + num_iterations : int + Number of Adam steps. + learning_rate : float + Base step size (positions, in voxels). Intensity/background steps are + scaled internally by the data magnitude; widths use half this rate. + loss : {"huber", "mse"} + Data-fidelity term. Huber is robust to reconstruction artifacts. + sigma_bounds : tuple[float, float] or None + Hard (min, max) bounds on widths in voxels. Default + ``(0.5, 2.0) * sigma_init``. + intensity_min : float or None + Sites dimmer than this are removed. Default 10% of the median site + intensity. + add, remove, merge : bool + Enable adding sites at residual peaks / removing weak sites / merging + close sites during refinement. + update_every : int + Apply the add/remove/merge maintenance every this many iterations. + min_distance : float or None + Minimum site spacing for merge/add, in voxels. Default + ``2 * sigma_init``. + progress : bool + Show a tqdm progress bar (loss and site count). + """ + if self.num_sites == 0: + raise RuntimeError("No sites to refine; call find_initial() first.") + if loss not in ("huber", "mse"): + raise ValueError("loss must be 'huber' or 'mse'.") + sig_lo, sig_hi = sigma_bounds or (0.5 * self._sigma_init, 2.0 * self._sigma_init) + if min_distance is None: + min_distance = 2.0 * self._sigma_init + self._intensity_scale = float(self._intensity.detach().median().clamp(min=1e-3)) + if intensity_min is None: + intensity_min = 0.1 * self._intensity_scale + delta = self._intensity_scale + + self._enable_grad() + opt = self._build_optimizer(learning_rate) + bar = tqdm(range(num_iterations), desc="refine", disable=not progress) + for it in bar: + opt.zero_grad(set_to_none=True) + model = self.render() + if loss == "huber": + data_term = F.huber_loss(model, self._volume, delta=delta) + else: + data_term = F.mse_loss(model, self._volume) + data_term.backward() + opt.step() + # Hard constraints: widths in band, positions inside the volume. + with torch.no_grad(): + clamped = F.softplus(self._raw_sigma).clamp(sig_lo, sig_hi) + self._raw_sigma.copy_(_inverse_softplus(clamped)) + self._positions.clamp_(min=0.0) + for a in range(3): + self._positions[:, a].clamp_(max=float(self._shape[a] - 1)) + # Periodic site-set maintenance. + if update_every and (it + 1) % update_every == 0 and it + 1 < num_iterations: + changed = False + with torch.no_grad(): + if remove: + changed |= self.remove_sites(intensity_min, sigma_bounds=(sig_lo, sig_hi)) + if merge: + changed |= self.merge_sites(min_distance) + if add: + changed |= self.add_sites(min_distance=min_distance) + if changed: + self._enable_grad() + opt = self._build_optimizer(learning_rate) + bar.set_postfix(loss=f"{float(data_term.detach()):.4g}", n=self.num_sites) + self._disable_grad() + return self + + def _enable_grad(self) -> None: + for p in (self._positions, self._raw_intensity, self._raw_sigma): + p.requires_grad_(True) + if self._has_background: + self._raw_background.requires_grad_(True) + + def _disable_grad(self) -> None: + for p in (self._positions, self._raw_intensity, self._raw_sigma): + p.requires_grad_(False) + if self._has_background: + self._raw_background.requires_grad_(False) + + def _build_optimizer(self, learning_rate: float) -> torch.optim.Optimizer: + # Per-group rates: positions in voxels, intensity/background scaled by the + # data magnitude, widths slower for stability. + scale = getattr(self, "_intensity_scale", 1.0) + groups = [ + {"params": [self._positions], "lr": learning_rate}, + {"params": [self._raw_sigma], "lr": 0.5 * learning_rate}, + {"params": [self._raw_intensity], "lr": learning_rate * scale}, + ] + if self._has_background: + groups.append({"params": [self._raw_background], "lr": learning_rate * scale}) + return torch.optim.Adam(groups) + + def remove_sites(self, intensity_min: float, sigma_bounds=None) -> bool: + """Remove sites dimmer than ``intensity_min`` (and outside ``sigma_bounds``). + + Returns True if any site was removed. + """ + keep = self._intensity.detach() >= intensity_min + if sigma_bounds is not None: + sig = self._sigma.detach() + keep = keep & (sig >= sigma_bounds[0]) & (sig <= sigma_bounds[1]) + if bool(keep.all()): + return False + self._positions = self._positions.detach()[keep] + self._raw_intensity = self._raw_intensity.detach()[keep] + self._raw_sigma = self._raw_sigma.detach()[keep] + return True + + def merge_sites(self, min_distance: float) -> bool: + """Merge sites closer than ``min_distance`` voxels, keeping the brightest. + + Returns True if any site was merged away. + """ + if self.num_sites < 2: + return False + from scipy.spatial import cKDTree + + pos = self._positions.detach().cpu().numpy() + inten = self._intensity.detach().cpu().numpy() + neighbors = cKDTree(pos).query_ball_point(pos, r=float(min_distance)) + suppressed = np.zeros(pos.shape[0], dtype=bool) + for a in np.argsort(inten)[::-1]: + if suppressed[a]: + continue + for b in neighbors[a]: + if b != a and inten[b] <= inten[a]: + suppressed[b] = True + if not suppressed.any(): + return False + keep = torch.as_tensor(np.nonzero(~suppressed)[0], device=self.device) + self._positions = self._positions.detach()[keep] + self._raw_intensity = self._raw_intensity.detach()[keep] + self._raw_sigma = self._raw_sigma.detach()[keep] + return True + + def add_sites(self, min_distance: float, threshold_pct: float = 99.9, max_add=None) -> bool: + """Add sites at peaks of the positive residual not already covered. + + Returns True if any site was added. + """ + residual = (self._volume - self.render()).detach().clamp(min=0.0) + blur = (0.75 * self._sigma_init, 1.5 * self._sigma_init) + new_pos, new_int = seed_peaks( + residual, + blur_sigmas=blur, + threshold_pct=threshold_pct, + min_distance=min_distance, + progress=False, + ) + if new_pos.shape[0] == 0: + return False + if self.num_sites > 0: + from scipy.spatial import cKDTree + + existing = self._positions.detach().cpu().numpy() + dist, _ = cKDTree(existing).query(new_pos.cpu().numpy(), k=1) + far = torch.as_tensor(dist >= float(min_distance)) + new_pos, new_int = new_pos[far], new_int[far] + if new_pos.shape[0] == 0: + return False + if max_add is not None: + new_pos, new_int = new_pos[:max_add], new_int[:max_add] + new_pos = new_pos.to(device=self.device, dtype=self._dtype) + new_raw_int = _inverse_softplus( + new_int.to(device=self.device, dtype=self._dtype).clamp(min=1e-6) + ) + new_raw_sig = _inverse_softplus( + torch.full((new_pos.shape[0],), self._sigma_init, dtype=self._dtype, device=self.device) + ) + self._positions = torch.cat([self._positions.detach(), new_pos], dim=0) + self._raw_intensity = torch.cat([self._raw_intensity.detach(), new_raw_int], dim=0) + self._raw_sigma = torch.cat([self._raw_sigma.detach(), new_raw_sig], dim=0) + return True + + # ------------------------------------------------------------------ # + # Rendering / decomposition + # ------------------------------------------------------------------ # + def _render_atoms(self) -> Tensor: + return render_isotropic( + self._positions, self._intensity, self._sigma, self._shape, self._window_radius() + ) + + def _render_background(self) -> Tensor: + if not self._has_background: + return torch.zeros(self._shape, dtype=self._dtype, device=self.device) + return render_background(self._background_control, self._bases) + + def render(self) -> Tensor: + """Render the full model ``volume_atoms + volume_background`` as a tensor.""" + return self._render_atoms() + self._render_background() + + def _as_dataset(self, tensor: Tensor, name: str) -> Dataset3d: + arr = tensor.detach().to("cpu").numpy().astype(np.float32) + return Dataset3d.from_array( + arr, + name=name, + origin=self._origin, + sampling=self._sampling, + units=self._units, + signal_units=self._signal_units, + ) + + @property + def volume(self) -> Dataset3d: + """The source volume.""" + return self._source + + @property + def volume_atoms(self) -> Dataset3d: + """Rendered atomic part (>= 0).""" + return self._as_dataset(self._render_atoms(), "volume_atoms") + + @property + def volume_background(self) -> Dataset3d: + """Rendered smooth background (>= 0).""" + return self._as_dataset(self._render_background(), "volume_background") + + @property + def volume_model(self) -> Dataset3d: + """Full model ``volume_atoms + volume_background``.""" + return self._as_dataset(self.render(), "volume_model") + + @property + def residual(self) -> Dataset3d: + """Source volume minus the model.""" + return self._as_dataset(self._volume - self.render(), "residual") + + # ------------------------------------------------------------------ # + # Sites as a Vector (physical units) + # ------------------------------------------------------------------ # + @property + def sites(self) -> Vector: + """Atomic sites as a :class:`Vector` with fields x, y, z, intensity, sigma. + + Positions and sigma are in the source volume's physical units (sigma uses + the mean sampling). + """ + fields = ["x", "y", "z", "intensity", "sigma"] + units = [self._units[0], self._units[1], self._units[2], self._signal_units, self._units[0]] + v = Vector.from_shape(shape=(), fields=fields, units=units, name="atoms") + n = self.num_sites + if n == 0: + return v + pos_vox = self._positions.detach().to("cpu").numpy() + pos_phys = self._origin[None, :] + pos_vox * self._sampling[None, :] + inten = self._intensity.detach().to("cpu").numpy() + sigma_phys = self._sigma.detach().to("cpu").numpy() * float(self._sampling.mean()) + data = np.column_stack([pos_phys, inten, sigma_phys]).astype(np.float32) + v[...] = data + return v + + # ------------------------------------------------------------------ # + # Interactive 3D widget + # ------------------------------------------------------------------ # + def show_3d_atoms(self, **kwargs): + """Open the interactive 3D slice + atom-overlay widget. + + Renders an orthogonal slice (xy / xz / yz) through the volume with the + atomic sites overlaid (marker size scaled by intensity, opacity fading + with distance from the slice). Sites are passed in voxel coordinates so + they register with the volume. Requires the ``quantem.widget`` package; + keyword arguments are forwarded to :class:`quantem.widget.Show3DAtoms`. + """ + from quantem.widget import Show3DAtoms + + if self.num_sites: + sites = np.column_stack( + [ + self._positions.detach().to("cpu").numpy(), + self._intensity.detach().to("cpu").numpy(), + self._sigma.detach().to("cpu").numpy(), + ] + ).astype(np.float32) + else: + sites = np.zeros((0, 5), dtype=np.float32) + volume = self._volume.detach().to("cpu").numpy().astype(np.float32) + kwargs.setdefault("title", str(self._source.name)) + kwargs.setdefault("sampling", tuple(float(s) for s in self._sampling)) + return Show3DAtoms(volume, sites=sites, **kwargs) + + def __repr__(self) -> str: + return ( + f"quantem.Atoms(model={self._model}, num_sites={self.num_sites}, " + f"volume_shape={self._shape}, device={self.device}, " + f"background={self._has_background})" + ) diff --git a/widget/src/quantem/widget/__init__.py b/widget/src/quantem/widget/__init__.py index 96d8aebc3..f28ada793 100644 --- a/widget/src/quantem/widget/__init__.py +++ b/widget/src/quantem/widget/__init__.py @@ -2,6 +2,7 @@ from quantem.widget.show2d import Show2D from quantem.widget.show4dstem import Show4DSTEM +from quantem.widget.show3d_atoms import Show3DAtoms try: __version__ = version("quantem.widget") @@ -9,4 +10,4 @@ # Source-tree imports (e.g. `PYTHONPATH=src pytest`) skip pip install. __version__ = "0.0.0+local" -__all__ = ["Show2D", "Show4DSTEM"] +__all__ = ["Show2D", "Show4DSTEM", "Show3DAtoms"] diff --git a/widget/src/quantem/widget/show3d_atoms.py b/widget/src/quantem/widget/show3d_atoms.py new file mode 100644 index 000000000..0c4071556 --- /dev/null +++ b/widget/src/quantem/widget/show3d_atoms.py @@ -0,0 +1,122 @@ +"""show3d_atoms: interactive 3D volume slice viewer with an atomic-site overlay. + +Shows a single orthogonal slice (xy / xz / yz) through a 3D volume in one panel, +with a movable slice position, an intensity histogram with an adjustable color +range, and the traced atomic sites overlaid. Marker size scales with site +intensity and marker opacity falls off with distance from the slice, so you can +judge tracing quality and pick thresholds slice by slice. + +All coordinates are in voxel / array-index space so the sites register exactly +with the volume. +""" + +import pathlib + +import anywidget +import numpy as np +import traitlets + + +class Show3DAtoms(anywidget.AnyWidget): + """Interactive orthogonal-slice viewer with atomic-site overlay. + + Parameters + ---------- + volume : ndarray or Dataset3d + 3D scalar volume ``(n0, n1, n2)``. + sites : ndarray, optional + ``(N, >=3)`` array of sites in voxel coordinates; columns are + ``[a0, a1, a2, intensity, sigma]`` (intensity/sigma optional, padded + with zeros). + sampling : sequence of float, optional + Voxel size per axis (for axis labels). Default ``(1, 1, 1)``. + title : str, optional + Title shown above the panel. + cmap : str, default "gray" + Colormap name. + """ + + _esm = pathlib.Path(__file__).parent / "static" / "show3d_atoms.js" + + # Data (voxel/array-index space). + volume_bytes = traitlets.Bytes(b"").tag(sync=True) + sites_bytes = traitlets.Bytes(b"").tag(sync=True) + n0 = traitlets.Int(0).tag(sync=True) + n1 = traitlets.Int(0).tag(sync=True) + n2 = traitlets.Int(0).tag(sync=True) + num_sites = traitlets.Int(0).tag(sync=True) + sampling = traitlets.List(traitlets.Float(), default_value=[1.0, 1.0, 1.0]).tag(sync=True) + + # Display / interaction state. + title = traitlets.Unicode("").tag(sync=True) + cmap = traitlets.Unicode("gray").tag(sync=True) + plane = traitlets.Unicode("xy").tag(sync=True) # "xy" | "xz" | "yz" + slice_index = traitlets.Int(0).tag(sync=True) + slice_thickness = traitlets.Float(3.0).tag(sync=True) # voxels at full opacity + opacity_falloff = traitlets.Float(3.0).tag(sync=True) # extra voxels fading to 0 + vmin_pct = traitlets.Float(0.0).tag(sync=True) + vmax_pct = traitlets.Float(100.0).tag(sync=True) + marker_scale = traitlets.Float(1.0).tag(sync=True) + show_sites = traitlets.Bool(True).tag(sync=True) + show_slice = traitlets.Bool(True).tag(sync=True) + canvas_size = traitlets.Int(520).tag(sync=True) + + def __init__( + self, + volume, + sites=None, + *, + sampling=None, + title="", + cmap="gray", + **kwargs, + ): + super().__init__(**kwargs) + vol = self._to_numpy3d(volume) + self._volume = vol + self.n0, self.n1, self.n2 = (int(s) for s in vol.shape) + self.volume_bytes = np.ascontiguousarray(vol, dtype=np.float32).tobytes() + + if sites is None: + sites = np.zeros((0, 5), dtype=np.float32) + sites = np.asarray(sites, dtype=np.float32) + if sites.ndim != 2 or sites.shape[1] < 3: + raise ValueError( + f"sites must be (N, >=3) [a0,a1,a2,(intensity),(sigma)], got {sites.shape}" + ) + if sites.shape[1] < 5: + sites = np.pad(sites, ((0, 0), (0, 5 - sites.shape[1])), constant_values=0.0) + self._sites = sites[:, :5] + self.num_sites = int(sites.shape[0]) + self.sites_bytes = np.ascontiguousarray(self._sites, dtype=np.float32).tobytes() + + if sampling is not None: + self.sampling = [float(s) for s in sampling] + self.title = title + self.cmap = cmap + # xy plane fixes axis 2; start in the middle. + self.slice_index = int(vol.shape[2] // 2) + + @staticmethod + def _to_numpy3d(volume): + if isinstance(volume, np.ndarray): + arr = volume + else: + arr = getattr(volume, "array", None) + if arr is None: + if hasattr(volume, "detach"): # torch tensor + arr = volume.detach().cpu().numpy() + elif hasattr(volume, "numpy"): # Dataset + arr = volume.numpy() + else: + arr = volume + arr = np.asarray(arr) + if arr.ndim != 3: + raise ValueError(f"volume must be 3D, got shape {arr.shape}") + return arr + + def __repr__(self) -> str: + return ( + f"Show3DAtoms(shape=({self.n0}, {self.n1}, {self.n2}), " + f"{self.num_sites} sites, plane={self.plane})" + ) From 6aaa32bcde86f38f448f48f0cfd59b887e23a6c1 Mon Sep 17 00:00:00 2001 From: cophus Date: Sun, 14 Jun 2026 21:17:47 -0700 Subject: [PATCH 2/6] adding widget for 3d plotting --- src/quantem/tomography/atom_trace.py | 309 ++++++++++++++-- widget/js/show3d_atoms/index.tsx | 415 ++++++++++++++++++++++ widget/scripts/build.mjs | 1 + widget/src/quantem/widget/show3d_atoms.py | 2 + 4 files changed, 699 insertions(+), 28 deletions(-) create mode 100644 widget/js/show3d_atoms/index.tsx diff --git a/src/quantem/tomography/atom_trace.py b/src/quantem/tomography/atom_trace.py index e0975ec3c..4c4be7f07 100644 --- a/src/quantem/tomography/atom_trace.py +++ b/src/quantem/tomography/atom_trace.py @@ -74,6 +74,7 @@ "render_background", "gaussian_blur3d", "seed_peaks", + "estimate_nn_spacing", ] @@ -261,7 +262,7 @@ def seed_peaks( volume: Tensor, blur_sigmas: tuple[float, float] | float = (1.0, 2.0), threshold: float | None = None, - threshold_pct: float = 99.0, + threshold_fraction: float = 0.99, min_distance: float = 0.0, max_peaks: int | None = None, progress: bool = False, @@ -281,10 +282,15 @@ def seed_peaks( Two sigmas -> difference-of-Gaussians (bandpass). One sigma -> single matched-filter blur. In voxels. threshold : float or None - Absolute threshold on the filtered response. Overrides ``threshold_pct``. - threshold_pct : float - Percentile (0-100) of the filtered response used as the threshold when - ``threshold`` is None. Robust to outliers and to data scale; default 99.0. + Absolute cutoff on the filtered (difference-of-Gaussians) response: a voxel + is a candidate peak only if its response exceeds this. Overrides + ``threshold_fraction``. + threshold_fraction : float + Detection threshold as a quantile (0-1) of the filtered response, used when + ``threshold`` is None. Only voxels above this quantile are kept, so + ``0.99`` keeps the brightest ~1% of the response (more, weaker sites) and + ``0.999`` the brightest ~0.1% (fewer, stronger sites). Robust to outliers + and to the absolute data scale. min_distance : float Minimum spacing between seeds in voxels; closer (dimmer) peaks are dropped. 0 disables. @@ -312,7 +318,7 @@ def seed_peaks( else: resp = gaussian_blur3d(volume, float(blur_sigmas)) if threshold is None: - threshold = float(np.percentile(resp.numpy(), threshold_pct)) + threshold = float(np.quantile(resp.numpy(), threshold_fraction)) bar.set_postfix_str("DoG bandpass") bar.update(1) @@ -383,6 +389,72 @@ def seed_peaks( return positions, intensities +def estimate_nn_spacing( + volume: Tensor, + r_min: float = 2.0, + r_max: float | None = None, + return_profile: bool = False, +) -> float | tuple[float, np.ndarray, np.ndarray]: + """Estimate the nearest-neighbor spacing (in voxels) from the volume autocorrelation. + + The autocorrelation ``IFFT(|FFT(v)|^2)`` is radially averaged; the radius of + its first peak beyond the central self-peak is the first coordination shell, + i.e. the typical nearest-neighbor spacing. Using ``~0.75 *`` this value as the + seeding ``min_distance`` strongly suppresses duplicate/false-positive sites. + + Parameters + ---------- + volume : Tensor or ndarray + 3D volume. + r_min : float + Ignore peaks closer than this radius (excludes the central self-peak). + r_max : float or None + Largest radius to consider. Default ``min(shape) // 2``. + return_profile : bool + If True, also return the radial-autocorrelation profile (normalized so the + zero-lag value is 1) for plotting/inspection. + + Returns + ------- + float or tuple + The nearest-neighbor spacing in voxels; or, if ``return_profile``, + ``(spacing, radii, radial)`` with ``radii``/``radial`` as 1D arrays. + + Raises + ------ + RuntimeError + If no autocorrelation shell peak is found. + """ + from scipy.signal import find_peaks + + vt = torch.as_tensor(volume).detach().to(device="cpu", dtype=torch.float32) + vt = vt - vt.mean() + power = torch.fft.fftn(vt).abs() ** 2 + ac = torch.fft.fftshift(torch.fft.ifftn(power).real).numpy() + shape = ac.shape + center = [s // 2 for s in shape] + sq = [((np.arange(s) - c) ** 2).astype(np.float32) for s, c in zip(shape, center)] + radius = np.sqrt(sq[0][:, None, None] + sq[1][None, :, None] + sq[2][None, None, :]) + r_int = np.rint(radius).astype(np.int64).ravel() + counts = np.bincount(r_int) + radial = np.bincount(r_int, weights=ac.ravel().astype(np.float64)) / np.maximum(counts, 1) + cutoff = int(r_max) if r_max is not None else min(shape) // 2 + radial = radial[:cutoff] + if radial[0] > 0: + radial = radial / radial[0] # normalize so the zero-lag value is 1 + peaks, _ = find_peaks(radial) + shell = [int(p) for p in peaks if p >= max(1, int(round(r_min)))] + if not shell: + raise RuntimeError( + "Could not estimate NN spacing from the autocorrelation; " + "pass min_distance to find_initial() explicitly." + ) + spacing = float(shell[0]) + if return_profile: + return spacing, np.arange(len(radial), dtype=float), radial + return spacing + + def _inverse_softplus(y: Tensor) -> Tensor: """Numerically stable inverse of ``softplus`` (no overflow for large y). @@ -561,33 +633,129 @@ def model(self) -> str: # ------------------------------------------------------------------ # # Seeding # ------------------------------------------------------------------ # + def estimate_spacing( + self, + r_min: float = 2.0, + recompute: bool = False, + plot: bool = False, + returnfig: bool = False, + ) -> float | tuple: + """Estimate (and cache) the nearest-neighbor atomic spacing, in voxels. + + Uses the volume autocorrelation (:func:`estimate_nn_spacing`). The result + seeds the default ``min_distance`` for :meth:`find_initial`, which strongly + suppresses duplicate / false-positive sites (especially around weak ones). + + Parameters + ---------- + r_min : float, default 2.0 + Ignore autocorrelation peaks closer than this radius (excludes the + central self-peak). + recompute : bool, default False + Recompute even if a cached value exists. + plot : bool, default False + Plot the radial autocorrelation profile with the detected first-shell + peak (= the spacing) and the resulting ``min_distance`` marked, to help + you judge whether the estimate is sensible. + returnfig : bool, default False + If True, return ``(spacing, fig, ax)`` instead of just the spacing. + + Returns + ------- + float or tuple + The spacing in voxels, or ``(spacing, fig, ax)`` if ``returnfig``. + """ + if recompute or getattr(self, "_nn_spacing", None) is None: + self._nn_spacing, radii, radial = estimate_nn_spacing( + self._volume, r_min=r_min, return_profile=True + ) + self._nn_profile = (radii, radial) + + if not plot: + return self._nn_spacing + + import matplotlib.pyplot as plt + + radii, radial = self._nn_profile + spacing = self._nn_spacing + # Show the shell structure beyond the central self-peak. + r_hi = min(len(radial) - 1, max(int(round(4 * spacing)), 12)) + fig, ax = plt.subplots(figsize=(6, 3.2)) + ax.plot(radii[1 : r_hi + 1], radial[1 : r_hi + 1], color="0.2", lw=1.2) + ax.axvline(spacing, color="tab:red", ls="--", label=f"NN spacing = {spacing:.1f} voxels") + ax.axvline( + 0.75 * spacing, color="tab:blue", ls=":", + label=f"min_distance = {0.75 * spacing:.2f} voxels", + ) + ax.set_xlabel("radius (voxels)") + ax.set_ylabel("radial autocorrelation") + ax.set_title("volume autocorrelation") + ax.legend(fontsize=9) + fig.tight_layout() + return (self._nn_spacing, fig, ax) if returnfig else self._nn_spacing + def find_initial( self, blur_sigmas: tuple[float, float] | float | None = None, threshold: float | None = None, - threshold_pct: float = 99.0, + threshold_fraction: float = 0.99, min_distance: float | None = None, + spacing_factor: float = 0.75, max_peaks: int | None = None, progress: bool = True, ) -> "Atoms": - """Find an initial set of sites by difference-of-Gaussians peak detection. - - This is the first step of tracing (the seed for later refinement and - densification). Defaults derive from ``sigma_init``: the DoG bandpass - brackets it, the threshold is the ``threshold_pct`` percentile of the - response, and ``min_distance`` is ``2 * sigma_init``. Pass an absolute - ``threshold`` to override. Set ``progress=False`` to silence the bar. - See :func:`seed_peaks` for the detector. + """Find an initial set of atomic sites by difference-of-Gaussians detection. + + Band-pass filters the volume to enhance atom-sized blobs, keeps the local + maxima above a threshold as candidate sites, refines each to sub-voxel + accuracy, and drops duplicates closer than ``min_distance``. This is the + first step of tracing and seeds :meth:`refine`. + + Parameters + ---------- + blur_sigmas : tuple[float, float] or float or None + Difference-of-Gaussians widths in voxels, ``(small, large)``; the + band-pass highlights features at the atom scale. Default brackets + ``sigma_init`` as ``(0.75, 1.5) * sigma_init``. + threshold : float or None + Absolute cutoff on the filtered response (a voxel is a candidate only + if its response exceeds this). Overrides ``threshold_fraction``. + threshold_fraction : float, default 0.99 + Detection threshold as a quantile (0-1) of the filtered response: only + voxels whose response is above this quantile are kept. ``0.99`` keeps + the brightest ~1% of the response (more sites, including weaker ones); + ``0.999`` keeps the brightest ~0.1% (fewer, stronger sites). Being a + quantile, it is robust to outliers and to the absolute data scale. + min_distance : float or None + Minimum spacing between sites in voxels; of two sites closer than this, + the dimmer is removed. Defaults to ``spacing_factor`` times the + nearest-neighbor spacing from :meth:`estimate_spacing` -- the main lever + against duplicate / false-positive sites. + spacing_factor : float, default 0.75 + Fraction of the estimated nearest-neighbor spacing used for the default + ``min_distance`` (ignored when ``min_distance`` is given). + max_peaks : int or None + If set, keep only the brightest ``max_peaks`` sites. + progress : bool, default True + Show a tqdm progress bar. + + Returns + ------- + Atoms + ``self``, with ``sites`` populated. """ if blur_sigmas is None: blur_sigmas = (0.75 * self._sigma_init, 1.5 * self._sigma_init) if min_distance is None: - min_distance = 2.0 * self._sigma_init + try: + min_distance = spacing_factor * self.estimate_spacing() + except RuntimeError: + min_distance = 2.0 * self._sigma_init pos, inten = seed_peaks( self._volume, blur_sigmas=blur_sigmas, threshold=threshold, - threshold_pct=threshold_pct, + threshold_fraction=threshold_fraction, min_distance=min_distance, max_peaks=max_peaks, progress=progress, @@ -614,6 +782,9 @@ def refine( add: bool = True, remove: bool = True, merge: bool = True, + add_threshold_fraction: float = 0.98, + min_neighbors: int = 2, + isolation_radius: float | None = None, update_every: int = 25, min_distance: float | None = None, progress: bool = True, @@ -641,8 +812,22 @@ def refine( Sites dimmer than this are removed. Default 10% of the median site intensity. add, remove, merge : bool - Enable adding sites at residual peaks / removing weak sites / merging - close sites during refinement. + Enable adding sites at residual peaks / removing weak (low-intensity) + sites / merging close sites during refinement. + add_threshold_fraction : float, default 0.98 + Detection quantile (0-1) for ``add``: lower values recover weaker atoms + from the residual (raise toward 1 to add only obvious ones). Pair with + ``min_neighbors`` to reject the extra noise this admits. + min_neighbors : int, default 2 + Remove sites with fewer than this many neighbors within + ``isolation_radius`` -- atoms do not float alone, so isolated detections + are almost always false positives. This is what makes a low + ``add_threshold_fraction`` usable for finding weak atoms (real weak + atoms sit on the lattice and are kept; isolated noise is dropped). Set + 0 to disable. + isolation_radius : float or None + Neighbor-search radius (voxels) for ``min_neighbors``. Default + ``1.5 *`` the estimated nearest-neighbor spacing. update_every : int Apply the add/remove/merge maintenance every this many iterations. min_distance : float or None @@ -657,7 +842,10 @@ def refine( raise ValueError("loss must be 'huber' or 'mse'.") sig_lo, sig_hi = sigma_bounds or (0.5 * self._sigma_init, 2.0 * self._sigma_init) if min_distance is None: - min_distance = 2.0 * self._sigma_init + try: + min_distance = 0.75 * self.estimate_spacing() + except RuntimeError: + min_distance = 2.0 * self._sigma_init self._intensity_scale = float(self._intensity.detach().median().clamp(min=1e-3)) if intensity_min is None: intensity_min = 0.1 * self._intensity_scale @@ -686,12 +874,16 @@ def refine( if update_every and (it + 1) % update_every == 0 and it + 1 < num_iterations: changed = False with torch.no_grad(): - if remove: - changed |= self.remove_sites(intensity_min, sigma_bounds=(sig_lo, sig_hi)) + if add: + changed |= self.add_sites( + min_distance=min_distance, threshold_fraction=add_threshold_fraction + ) if merge: changed |= self.merge_sites(min_distance) - if add: - changed |= self.add_sites(min_distance=min_distance) + if min_neighbors > 0: + changed |= self.remove_isolated(isolation_radius, min_neighbors) + if remove: + changed |= self.remove_sites(intensity_min, sigma_bounds=(sig_lo, sig_hi)) if changed: self._enable_grad() opt = self._build_optimizer(learning_rate) @@ -767,17 +959,78 @@ def merge_sites(self, min_distance: float) -> bool: self._raw_sigma = self._raw_sigma.detach()[keep] return True - def add_sites(self, min_distance: float, threshold_pct: float = 99.9, max_add=None) -> bool: - """Add sites at peaks of the positive residual not already covered. + def remove_isolated(self, radius: float | None = None, min_neighbors: int = 3) -> bool: + """Remove isolated sites (fewer than ``min_neighbors`` neighbors within ``radius``). + + Counts how many other sites lie within ``radius`` voxels of each site and + drops those below ``min_neighbors``. Physically, atoms do not float alone + in vacuum, so isolated detections are almost always false positives. This + is also what makes detecting *weak* atoms practical: lower the + ``find_initial`` / ``add`` threshold to admit weak sites (and noise), then + remove the noise here -- real weak atoms sit on the lattice with many + neighbors and are kept, while spurious peaks are isolated and removed. + + Parameters + ---------- + radius : float or None + Neighbor-search radius in voxels. Default ``1.5 *`` the estimated + nearest-neighbor spacing (:meth:`estimate_spacing`). + min_neighbors : int, default 3 + Minimum neighbors within ``radius`` required to keep a site. + + Returns + ------- + bool + True if any site was removed. + """ + if self.num_sites < 2: + return False + if radius is None: + radius = 1.5 * self.estimate_spacing() + from scipy.spatial import cKDTree + + pts = self._positions.detach().cpu().numpy() + counts = cKDTree(pts).query_ball_point(pts, r=float(radius), return_length=True) + keep = (counts - 1) >= min_neighbors # subtract the site's own match + if bool(keep.all()): + return False + keep_t = torch.as_tensor(np.nonzero(keep)[0], device=self.device) + self._positions = self._positions.detach()[keep_t] + self._raw_intensity = self._raw_intensity.detach()[keep_t] + self._raw_sigma = self._raw_sigma.detach()[keep_t] + return True + + def add_sites( + self, min_distance: float, threshold_fraction: float = 0.999, max_add: int | None = None + ) -> bool: + """Add new sites at peaks of the positive residual (densification). - Returns True if any site was added. + Detects peaks in ``volume - model`` (clamped to >= 0) that lie at least + ``min_distance`` voxels from existing sites, and appends them. Used during + :meth:`refine` to recover atoms the current model is missing. + + Parameters + ---------- + min_distance : float + Minimum spacing in voxels, both among new sites and from existing ones. + threshold_fraction : float, default 0.999 + Detection quantile (0-1) on the residual response (see + :func:`seed_peaks`); high by default so only clear, missed atoms are + added. + max_add : int or None + If set, cap the number of sites added in this call. + + Returns + ------- + bool + True if any site was added. """ residual = (self._volume - self.render()).detach().clamp(min=0.0) blur = (0.75 * self._sigma_init, 1.5 * self._sigma_init) new_pos, new_int = seed_peaks( residual, blur_sigmas=blur, - threshold_pct=threshold_pct, + threshold_fraction=threshold_fraction, min_distance=min_distance, progress=False, ) diff --git a/widget/js/show3d_atoms/index.tsx b/widget/js/show3d_atoms/index.tsx new file mode 100644 index 000000000..0d5ff760d --- /dev/null +++ b/widget/js/show3d_atoms/index.tsx @@ -0,0 +1,415 @@ +/** + * Show3DAtoms - orthogonal-slice viewer of a 3D volume with an atomic-site overlay. + * + * One panel shows an xy / xz / yz slice (movable along its normal), an intensity + * histogram (percentile-clipped so outliers don't dominate) with an adjustable + * color range, and the traced sites overlaid: marker size scales with intensity, + * marker opacity fades with distance from the slice. + * + * Interaction: left-drag = box zoom, middle-drag = pan, double-click = reset view + * and all settings. All site coordinates are in voxel/array-index space. + */ + +import * as React from "react"; +import { createRender, useModelState } from "@anywidget/react"; +import Box from "@mui/material/Box"; +import Stack from "@mui/material/Stack"; +import Typography from "@mui/material/Typography"; +import Slider from "@mui/material/Slider"; +import Select from "@mui/material/Select"; +import MenuItem from "@mui/material/MenuItem"; +import Switch from "@mui/material/Switch"; +import { useTheme } from "../theme"; +import { extractFloat32 } from "../format"; +import { COLORMAPS, COLORMAP_NAMES, applyColormap } from "../colormaps"; +import { findDataRange, sliderRange, percentileClip } from "../stats"; + +const sliderStyles = { + py: 0, + "& .MuiSlider-thumb": { width: 11, height: 11 }, + "& .MuiSlider-rail": { height: 2 }, + "& .MuiSlider-track": { height: 2 }, +}; + +const MARKER_BASE = 1.5; // half the previous default size +const MIN_ZOOM_VOX = 4; // smallest viewport extent (voxels) +const AXIS_LABELS = ["x", "y", "z"]; + +type PlaneInfo = { + normalAxis: number; rowAxis: number; colAxis: number; + rows: number; cols: number; depth: number; normalLabel: string; +}; +type Viewport = { row0: number; row1: number; col0: number; col1: number }; + +function planeInfo(plane: string, n0: number, n1: number, n2: number): PlaneInfo { + let normalAxis: number, rowAxis: number, colAxis: number; + if (plane === "yz") { normalAxis = 0; rowAxis = 1; colAxis = 2; } + else if (plane === "xz") { normalAxis = 1; rowAxis = 0; colAxis = 2; } + else { normalAxis = 2; rowAxis = 0; colAxis = 1; } // xy + const dims = [n0, n1, n2]; + return { + normalAxis, rowAxis, colAxis, + rows: dims[rowAxis], cols: dims[colAxis], depth: dims[normalAxis], + normalLabel: AXIS_LABELS[normalAxis], + }; +} + +function extractSlice(vol: Float32Array, n1: number, n2: number, info: PlaneInfo, k: number): Float32Array { + const { rows, cols, normalAxis } = info; + const out = new Float32Array(rows * cols); + const stride0 = n1 * n2; + if (normalAxis === 2) { + for (let r = 0; r < rows; r++) for (let c = 0; c < cols; c++) out[r * cols + c] = vol[r * stride0 + c * n2 + k]; + } else if (normalAxis === 1) { + for (let r = 0; r < rows; r++) for (let c = 0; c < cols; c++) out[r * cols + c] = vol[r * stride0 + k * n2 + c]; + } else { + const base = k * stride0; + for (let r = 0; r < rows; r++) for (let c = 0; c < cols; c++) out[r * cols + c] = vol[base + r * n2 + c]; + } + return out; +} + +/** Histogram binned over a fixed [lo, hi] range, clipping outliers into the end bins. */ +function histogramInRange(data: Float32Array, lo: number, hi: number, nbins = 96): number[] { + const bins = new Array(nbins).fill(0); + const range = hi > lo ? hi - lo : 1; + const scale = nbins / range; + for (let i = 0; i < data.length; i++) { + const v = data[i]; + if (!isFinite(v)) continue; + let b = Math.floor((v - lo) * scale); + if (b < 0) b = 0; else if (b >= nbins) b = nbins - 1; + bins[b]++; + } + const mx = Math.max(...bins, 1e-9); + for (let i = 0; i < nbins; i++) bins[i] /= mx; + return bins; +} + +function Show3DAtoms() { + const { colors: tc } = useTheme(); + + const [volumeBytes] = useModelState("volume_bytes"); + const [sitesBytes] = useModelState("sites_bytes"); + const [n0] = useModelState("n0"); + const [n1] = useModelState("n1"); + const [n2] = useModelState("n2"); + const [numSites] = useModelState("num_sites"); + const [title] = useModelState("title"); + const [cmap, setCmap] = useModelState("cmap"); + const [plane, setPlane] = useModelState("plane"); + const [sliceIndex, setSliceIndex] = useModelState("slice_index"); + const [thickness, setThickness] = useModelState("slice_thickness"); + const [falloff, setFalloff] = useModelState("opacity_falloff"); + const [vminPct, setVminPct] = useModelState("vmin_pct"); + const [vmaxPct, setVmaxPct] = useModelState("vmax_pct"); + const [markerScale, setMarkerScale] = useModelState("marker_scale"); + const [markerLinewidth, setMarkerLinewidth] = useModelState("marker_linewidth"); + const [markerFilled, setMarkerFilled] = useModelState("marker_filled"); + const [showSites, setShowSites] = useModelState("show_sites"); + const [showSlice, setShowSlice] = useModelState("show_slice"); + const [canvasSize] = useModelState("canvas_size"); + + const volume = React.useMemo(() => extractFloat32(volumeBytes), [volumeBytes]); + const sites = React.useMemo(() => extractFloat32(sitesBytes), [sitesBytes]); + const info = React.useMemo(() => planeInfo(plane, n0, n1, n2), [plane, n0, n1, n2]); + const k = Math.max(0, Math.min(info.depth - 1, sliceIndex)); + + // Robust intensity range (percentile-clipped) for histogram + color mapping. + const baseRange = React.useMemo(() => { + if (!volume) return { lo: 0, hi: 1 }; + const { vmin, vmax, min, max } = percentileClip(volume, 0.5, 99.5); + return vmax > vmin ? { lo: vmin, hi: vmax } : findDataRange(volume).max > findDataRange(volume).min + ? { lo: min, hi: max } : { lo: min, hi: min + 1 }; + }, [volume]); + const histBins = React.useMemo( + () => (volume ? histogramInRange(volume, baseRange.lo, baseRange.hi) : null), + [volume, baseRange], + ); + const maxIntensity = React.useMemo(() => { + if (!sites || numSites === 0) return 1; + let m = 0; + for (let i = 0; i < numSites; i++) m = Math.max(m, sites[i * 5 + 3]); + return m > 0 ? m : 1; + }, [sites, numSites]); + + const scale = canvasSize / Math.max(info.rows, info.cols, 1); + const canvasW = Math.round(info.cols * scale); + const canvasH = Math.round(info.rows * scale); + + const [viewport, setViewport] = React.useState(null); + const [drag, setDrag] = React.useState<{ mode: "zoom" | "pan"; sx: number; sy: number; startVp: Viewport } | null>(null); + const [dragBox, setDragBox] = React.useState<{ x0: number; y0: number; x1: number; y1: number } | null>(null); + + const fullVp = React.useCallback((): Viewport => ({ row0: 0, row1: info.rows, col0: 0, col1: info.cols }), [info]); + const vp = viewport ?? fullVp(); + + // Reset viewport + clamp slice when the plane changes. + const prevPlane = React.useRef(plane); + React.useEffect(() => { + if (prevPlane.current !== plane) { + prevPlane.current = plane; + setViewport(null); + if (sliceIndex > info.depth - 1) setSliceIndex(Math.floor(info.depth / 2)); + } + }, [plane, info.depth, sliceIndex, setSliceIndex]); + + // ---- Effect A: colormap the full slice into an offscreen canvas (heavy) ---- + const offRef = React.useRef(null); + const [sliceVersion, setSliceVersion] = React.useState(0); + React.useEffect(() => { + if (!volume) return; + let off = offRef.current; + if (!off) { off = document.createElement("canvas"); offRef.current = off; } + off.width = info.cols; off.height = info.rows; + const octx = off.getContext("2d")!; + if (showSlice) { + const sliceData = extractSlice(volume, n1, n2, info, k); + const { vmin, vmax } = sliderRange(baseRange.lo, baseRange.hi, vminPct, vmaxPct); + const rgba = new Uint8ClampedArray(info.rows * info.cols * 4); + applyColormap(sliceData, rgba, COLORMAPS[cmap] || COLORMAPS.gray, vmin, vmax); + const img = octx.createImageData(info.cols, info.rows); + img.data.set(rgba); + octx.putImageData(img, 0, 0); + } else { + octx.fillStyle = "#000"; + octx.fillRect(0, 0, info.cols, info.rows); + } + setSliceVersion((v) => v + 1); + }, [volume, n1, n2, info, k, cmap, baseRange, vminPct, vmaxPct, showSlice]); + + // ---- Effect B: draw offscreen (viewport crop) + atom overlay + rubber band (light) ---- + const canvasRef = React.useRef(null); + React.useEffect(() => { + const canvas = canvasRef.current; + const off = offRef.current; + if (!canvas || !off) return; + const ctx = canvas.getContext("2d"); + if (!ctx) return; + const dpr = window.devicePixelRatio || 1; + canvas.width = canvasW * dpr; canvas.height = canvasH * dpr; + ctx.setTransform(dpr, 0, 0, dpr, 0, 0); + ctx.clearRect(0, 0, canvasW, canvasH); + + ctx.imageSmoothingEnabled = false; + ctx.drawImage(off, vp.col0, vp.row0, vp.col1 - vp.col0, vp.row1 - vp.row0, 0, 0, canvasW, canvasH); + + if (showSites && sites && numSites > 0) { + const half = thickness / 2; + const fade = Math.max(falloff, 1e-6); + const sxv = canvasW / (vp.col1 - vp.col0); + const syv = canvasH / (vp.row1 - vp.row0); + for (let i = 0; i < numSites; i++) { + const o = i * 5; + const d = Math.abs(sites[o + info.normalAxis] - k); + if (d > half + fade) continue; + const opacity = d <= half ? 1 : Math.max(0, 1 - (d - half) / fade); + if (opacity <= 0.01) continue; + // +0.5 voxel so markers sit on pixel centers (drawImage maps voxel i to [i, i+1)). + const x = (sites[o + info.colAxis] + 0.5 - vp.col0) * sxv; + const y = (sites[o + info.rowAxis] + 0.5 - vp.row0) * syv; + const r = MARKER_BASE * markerScale * (0.4 + 1.6 * Math.sqrt(Math.min(1, sites[o + 3] / maxIntensity))); + if (x < -r || x > canvasW + r || y < -r || y > canvasH + r) continue; + ctx.beginPath(); + ctx.arc(x, y, r, 0, 2 * Math.PI); + ctx.lineWidth = markerLinewidth; + if (markerFilled) { + ctx.fillStyle = `rgba(255, 70, 70, ${0.85 * opacity})`; + ctx.fill(); + if (markerLinewidth > 0) { + ctx.strokeStyle = `rgba(255, 255, 255, ${0.5 * opacity})`; + ctx.stroke(); + } + } else { + ctx.strokeStyle = `rgba(255, 70, 70, ${opacity})`; + ctx.stroke(); + } + } + } + + if (dragBox) { + ctx.lineWidth = 1; + ctx.strokeStyle = tc.accent; + ctx.setLineDash([4, 3]); + ctx.strokeRect( + Math.min(dragBox.x0, dragBox.x1), Math.min(dragBox.y0, dragBox.y1), + Math.abs(dragBox.x1 - dragBox.x0), Math.abs(dragBox.y1 - dragBox.y0), + ); + ctx.setLineDash([]); + } + }, [sliceVersion, viewport, vp, sites, numSites, info, k, showSites, thickness, falloff, + markerScale, markerLinewidth, markerFilled, maxIntensity, canvasW, canvasH, dragBox, tc]); + + // ---- Mouse interaction: left=box zoom, middle=pan, dblclick=reset ---- + const relPx = (e: { clientX: number; clientY: number }) => { + const rect = canvasRef.current!.getBoundingClientRect(); + return { x: e.clientX - rect.left, y: e.clientY - rect.top }; + }; + const onMouseDown = (e: React.MouseEvent) => { + const p = relPx(e); + if (e.button === 0) { + setDrag({ mode: "zoom", sx: p.x, sy: p.y, startVp: vp }); + setDragBox({ x0: p.x, y0: p.y, x1: p.x, y1: p.y }); + } else if (e.button === 1) { + e.preventDefault(); + setDrag({ mode: "pan", sx: p.x, sy: p.y, startVp: vp }); + } + }; + React.useEffect(() => { + if (!drag) return; + const onMove = (e: MouseEvent) => { + const p = relPx(e); + if (drag.mode === "zoom") { + setDragBox((b) => (b ? { ...b, x1: p.x, y1: p.y } : b)); + } else { + const ext = drag.startVp; + const dCol = ((p.x - drag.sx) / canvasW) * (ext.col1 - ext.col0); + const dRow = ((p.y - drag.sy) / canvasH) * (ext.row1 - ext.row0); + let c0 = ext.col0 - dCol, r0 = ext.row0 - dRow; + const cw = ext.col1 - ext.col0, ch = ext.row1 - ext.row0; + c0 = Math.max(0, Math.min(info.cols - cw, c0)); + r0 = Math.max(0, Math.min(info.rows - ch, r0)); + setViewport({ col0: c0, col1: c0 + cw, row0: r0, row1: r0 + ch }); + } + }; + const onUp = () => { + if (drag.mode === "zoom") { + setDragBox((b) => { + if (b && Math.abs(b.x1 - b.x0) > 4 && Math.abs(b.y1 - b.y0) > 4) { + const toData = (x: number, y: number) => ({ + col: vp.col0 + (x / canvasW) * (vp.col1 - vp.col0), + row: vp.row0 + (y / canvasH) * (vp.row1 - vp.row0), + }); + const a = toData(b.x0, b.y0), c = toData(b.x1, b.y1); + let col0 = Math.min(a.col, c.col), col1 = Math.max(a.col, c.col); + let row0 = Math.min(a.row, c.row), row1 = Math.max(a.row, c.row); + // Enforce min size + match canvas aspect (no pixel distortion). + if (col1 - col0 < MIN_ZOOM_VOX) { const m = (col0 + col1) / 2; col0 = m - MIN_ZOOM_VOX / 2; col1 = m + MIN_ZOOM_VOX / 2; } + if (row1 - row0 < MIN_ZOOM_VOX) { const m = (row0 + row1) / 2; row0 = m - MIN_ZOOM_VOX / 2; row1 = m + MIN_ZOOM_VOX / 2; } + const aspect = canvasW / canvasH; + let w = col1 - col0, h = row1 - row0; + if (w / h > aspect) { const nh = w / aspect, m = (row0 + row1) / 2; row0 = m - nh / 2; row1 = m + nh / 2; } + else { const nw = h * aspect, m = (col0 + col1) / 2; col0 = m - nw / 2; col1 = m + nw / 2; } + col0 = Math.max(0, col0); col1 = Math.min(info.cols, col1); + row0 = Math.max(0, row0); row1 = Math.min(info.rows, row1); + setViewport({ row0, row1, col0, col1 }); + } + return null; + }); + } + setDrag(null); + }; + window.addEventListener("mousemove", onMove); + window.addEventListener("mouseup", onUp); + return () => { window.removeEventListener("mousemove", onMove); window.removeEventListener("mouseup", onUp); }; + }, [drag, canvasW, canvasH, info, vp]); + + const resetAll = () => { + setViewport(null); + setPlane("xy"); + setSliceIndex(Math.floor(n2 / 2)); + setVminPct(0); setVmaxPct(100); + setThickness(3); setFalloff(3); + setMarkerScale(1); setMarkerLinewidth(1); setMarkerFilled(true); + setShowSites(true); setShowSlice(true); + }; + + const fmt = (v: number) => (Math.abs(v) >= 1000 ? v.toExponential(1) : v.toFixed(1)); + const labelSx = { fontSize: 11, color: tc.text, whiteSpace: "nowrap" } as const; + const valSx = { fontSize: 10, fontFamily: "monospace", color: tc.textMuted } as const; + const selSx = { fontSize: 11, color: tc.text, bgcolor: tc.controlBg, "& .MuiSelect-select": { py: 0.4 } }; + const cr = sliderRange(baseRange.lo, baseRange.hi, vminPct, vmaxPct); + + return ( + + {title && {title}} + + + e.preventDefault()} + style={{ width: canvasW, height: canvasH, display: "block", cursor: drag?.mode === "pan" ? "grabbing" : "crosshair", + border: `1px solid ${tc.border}`, background: "#000" }} + /> + + { + if (!el || !histBins) return; + const ctx = el.getContext("2d"); if (!ctx) return; + const W = canvasW, H = 46, dpr = window.devicePixelRatio || 1; + el.width = W * dpr; el.height = H * dpr; ctx.setTransform(dpr, 0, 0, dpr, 0, 0); + ctx.clearRect(0, 0, W, H); + const nb = histBins.length, bw = W / nb; + const lo = (vminPct / 100) * nb, hi = (vmaxPct / 100) * nb; + for (let i = 0; i < nb; i++) { + const bh = histBins[i] * (H - 2); + ctx.fillStyle = i >= lo && i <= hi ? tc.accent : tc.textMuted; + ctx.fillRect(i * bw + 0.5, H - bh, Math.max(1, bw - 1), bh); + } + }} + style={{ width: canvasW, height: 46, display: "block", border: `1px solid ${tc.border}` }} + /> + { const [a, b] = v as number[]; setVminPct(Math.min(a, b - 1)); setVmaxPct(Math.max(b, a + 1)); }} + min={0} max={100} size="small" sx={{ ...sliderStyles, width: canvasW }} + /> + + {fmt(cr.vmin)} + color range + {fmt(cr.vmax)} + + + drag: zoom · middle-drag: pan · double-click: reset + + + + + plane + + cmap + + + + {info.normalLabel} slice: {k} / {info.depth - 1} + setSliceIndex(v as number)} min={0} max={Math.max(0, info.depth - 1)} size="small" sx={sliderStyles} /> + + + slice thickness: {thickness.toFixed(1)} voxels + setThickness(v as number)} min={0.5} max={20} step={0.5} size="small" sx={sliderStyles} /> + + + opacity falloff: {falloff.toFixed(1)} voxels + setFalloff(v as number)} min={0} max={20} step={0.5} size="small" sx={sliderStyles} /> + + + marker scale: {markerScale.toFixed(2)}× + setMarkerScale(v as number)} min={0.1} max={5} step={0.1} size="small" sx={sliderStyles} /> + + + line width: {markerLinewidth.toFixed(2)} + setMarkerLinewidth(v as number)} min={0} max={4} step={0.25} size="small" sx={sliderStyles} /> + + + setShowSites(e.target.checked)} size="small" /> + sites ({numSites}) + setShowSlice(e.target.checked)} size="small" /> + slice + + + setMarkerFilled(e.target.checked)} size="small" /> + filled markers ({markerFilled ? "filled" : "hollow"}) + + + + + ); +} + +export const render = createRender(Show3DAtoms); diff --git a/widget/scripts/build.mjs b/widget/scripts/build.mjs index 7c8daea95..dabd2ea5e 100644 --- a/widget/scripts/build.mjs +++ b/widget/scripts/build.mjs @@ -9,6 +9,7 @@ const watch = process.argv.includes("--watch"); const widgets = [ { name: "show2d" }, { name: "show4dstem" }, + { name: "show3d_atoms" }, ]; rmSync("src/quantem/widget/static", { recursive: true, force: true }); diff --git a/widget/src/quantem/widget/show3d_atoms.py b/widget/src/quantem/widget/show3d_atoms.py index 0c4071556..9f0a0d638 100644 --- a/widget/src/quantem/widget/show3d_atoms.py +++ b/widget/src/quantem/widget/show3d_atoms.py @@ -57,6 +57,8 @@ class Show3DAtoms(anywidget.AnyWidget): vmin_pct = traitlets.Float(0.0).tag(sync=True) vmax_pct = traitlets.Float(100.0).tag(sync=True) marker_scale = traitlets.Float(1.0).tag(sync=True) + marker_linewidth = traitlets.Float(1.0).tag(sync=True) + marker_filled = traitlets.Bool(True).tag(sync=True) show_sites = traitlets.Bool(True).tag(sync=True) show_slice = traitlets.Bool(True).tag(sync=True) canvas_size = traitlets.Int(520).tag(sync=True) From 57ee2eb12059ec378d38719a97e2c7bf8a3da6d4 Mon Sep 17 00:00:00 2001 From: cophus Date: Tue, 15 Sep 2026 14:14:43 -0700 Subject: [PATCH 3/6] atoms module --- src/quantem/__init__.py | 1 + src/quantem/atoms/__init__.py | 9 + src/quantem/atoms/atomic_model.py | 954 +++++++++++++++++++++++++++++ src/quantem/atoms/matching.py | 287 +++++++++ src/quantem/atoms/measurements.py | 266 ++++++++ src/quantem/atoms/pdf.py | 216 +++++++ src/quantem/atoms/show_atoms.py | 619 +++++++++++++++++++ src/quantem/atoms/templates.py | 350 +++++++++++ src/quantem/atoms/visualization.py | 416 +++++++++++++ tests/atoms/test_atoms.py | 238 +++++++ 10 files changed, 3356 insertions(+) create mode 100644 src/quantem/atoms/__init__.py create mode 100644 src/quantem/atoms/atomic_model.py create mode 100644 src/quantem/atoms/matching.py create mode 100644 src/quantem/atoms/measurements.py create mode 100644 src/quantem/atoms/pdf.py create mode 100644 src/quantem/atoms/show_atoms.py create mode 100644 src/quantem/atoms/templates.py create mode 100644 src/quantem/atoms/visualization.py create mode 100644 tests/atoms/test_atoms.py diff --git a/src/quantem/__init__.py b/src/quantem/__init__.py index 7f361fbee..4e89d0820 100644 --- a/src/quantem/__init__.py +++ b/src/quantem/__init__.py @@ -12,5 +12,6 @@ from quantem import spectroscopy as spectroscopy from quantem import diffractive_imaging as diffractive_imaging from quantem import tomography as tomography +from quantem import atoms as atoms __version__ = version("quantem") diff --git a/src/quantem/atoms/__init__.py b/src/quantem/atoms/__init__.py new file mode 100644 index 000000000..4449ee87e --- /dev/null +++ b/src/quantem/atoms/__init__.py @@ -0,0 +1,9 @@ +from quantem.atoms.atomic_model import AtomicModel as AtomicModel +from quantem.atoms.templates import ( + PolyhedralTemplate as PolyhedralTemplate, + TEMPLATE_NAMES as TEMPLATE_NAMES, + get_template as get_template, + template_from_crystal as template_from_crystal, +) +from quantem.atoms.matching import match_template as match_template +from quantem.atoms.visualization import PLOT_REGISTRY as PLOT_REGISTRY diff --git a/src/quantem/atoms/atomic_model.py b/src/quantem/atoms/atomic_model.py new file mode 100644 index 000000000..f284ca82a --- /dev/null +++ b/src/quantem/atoms/atomic_model.py @@ -0,0 +1,954 @@ +"""3D atomic model analysis: calibration, neighbor finding, structure classification. + +:class:`AtomicModel` holds the sites of a 3D atomic model (for example traced +from an atomic electron tomography reconstruction) in a +:class:`~quantem.core.datastructures.Vector` and provides the analysis +pipeline: + +1. :meth:`AtomicModel.compute_pdf` - radial distribution function and first + nearest-neighbor (NN) peak fit. +2. :meth:`AtomicModel.calibrate` - set the physical size of one voxel from + the measured NN distance of a reference crystal. +3. :meth:`AtomicModel.find_neighbors` - neighbor lists, coordination and bond + lengths. +4. :meth:`AtomicModel.match_templates` - fast polyhedral template matching + against ``fcc``, ``hcp``, ``bcc``, ``diamond`` ... environments. +5. :meth:`AtomicModel.classify` / :meth:`AtomicModel.segment_grains` / + :meth:`AtomicModel.compute_strain` - per-site structure, grain (sector) + labels via orientation clustering, and local strain. + +Every per-site result is stored as a named *channel* (a field of the sites +``Vector``) so it can be plotted with :meth:`AtomicModel.plot` or explored +interactively with :meth:`AtomicModel.show`. + +Coordinate convention +--------------------- +Fields ``x, y, z`` are the positions along array axes 0, 1, 2 of the source +volume, stored in their native (typically voxel) units. ``sampling`` converts +them to physical units; :attr:`AtomicModel.positions` returns calibrated +coordinates. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any, Sequence + +import numpy as np +from numpy.typing import NDArray + +from quantem.atoms import measurements as meas +from quantem.atoms.matching import TemplateMatch, match_template +from quantem.atoms.pdf import ( + find_neighbors, + fit_first_peak, + nn_distance_from_lattice, + radial_distribution, +) +from quantem.atoms.templates import PolyhedralTemplate, get_template +from quantem.atoms.visualization import PLOT_REGISTRY +from quantem.core.datastructures import Vector +from quantem.core.io.serialize import AutoSerialize + +__all__ = ["AtomicModel"] + +_POSITION_FIELDS = ("x", "y", "z") + + +class AtomicModel(AutoSerialize): + """A 3D atomic model with per-site measurement channels. + + Use the ``from_*`` constructors rather than ``__init__``. + + Parameters + ---------- + sites : Vector + 0-D ``Vector`` whose cell holds one row per site with at least the + fields ``x, y, z``. + sampling : ndarray + ``(3,)`` physical size of one native coordinate unit along each axis. + units : str + Physical length unit after calibration (e.g. ``"A"``). + name : str + Model name. + metadata : dict + Free-form metadata. + """ + + _token = object() + + def __init__( + self, + sites: Vector, + sampling: NDArray, + units: str, + name: str, + metadata: dict[str, Any] | None = None, + _token: object | None = None, + ) -> None: + if _token is not self._token: + raise RuntimeError("Use AtomicModel.from_array() or another from_* constructor.") + self._sites = sites + self._sampling = np.asarray(sampling, dtype=float).reshape(3) + self._units = str(units) + self._name = str(name) + self._metadata: dict[str, Any] = dict(metadata or {}) + self._pdf: dict[str, Any] | None = None + self._nn_fit: dict[str, Any] | None = None + self._neighbor_distances: NDArray | None = None + self._neighbor_indices: NDArray | None = None + self._matches: dict[str, TemplateMatch] = {} + self._template_specs: dict[str, dict[str, Any]] = {} + self._structure_names: list[str] = [] + self._categories: dict[str, list[str]] = {} + + # ------------------------------------------------------------------ # + # Constructors + # ------------------------------------------------------------------ # + @classmethod + def from_array( + cls, + xyz: NDArray, + sampling: float | Sequence[float] = 1.0, + units: str = "voxels", + name: str | None = None, + channels: dict[str, NDArray] | None = None, + metadata: dict[str, Any] | None = None, + ) -> "AtomicModel": + """Create a model from an ``(N, 3)`` coordinate array. + + Parameters + ---------- + xyz : ndarray + ``(N, 3)`` positions along axes 0, 1, 2. A ``(3, N)`` array is + transposed automatically. + sampling : float or sequence of float + Physical size of one coordinate unit (scalar or per axis). + units : str + Physical length unit (``"voxels"`` if uncalibrated). + name : str, optional + Model name. + channels : dict, optional + Extra per-site arrays stored as channels, e.g. ``{"species": ...}``. + metadata : dict, optional + Free-form metadata. + """ + xyz = np.asarray(xyz, dtype=float) + if xyz.ndim != 2 or 3 not in xyz.shape: + raise ValueError(f"xyz must be (N, 3) or (3, N), got {xyz.shape}") + if xyz.shape[1] != 3: + xyz = xyz.T + sites = Vector.from_shape( + shape=(), + fields=list(_POSITION_FIELDS), + units=[units] * 3, + name="sites", + ) + sites[...] = np.ascontiguousarray(xyz) + sampling_arr = np.broadcast_to(np.asarray(sampling, dtype=float), (3,)).copy() + model = cls( + sites=sites, + sampling=sampling_arr, + units=units, + name=name or "atomic model", + metadata=metadata, + _token=cls._token, + ) + for key, values in (channels or {}).items(): + model.set_channel(key, values) + return model + + @classmethod + def from_mat( + cls, + path: str | Path, + key: str | None = None, + sampling: float | Sequence[float] = 1.0, + units: str = "voxels", + name: str | None = None, + one_based: bool = False, + ) -> "AtomicModel": + """Load coordinates from a MATLAB ``.mat`` file. + + Parameters + ---------- + path : str or Path + File path (v5/v7 or v7.3 HDF5). + key : str, optional + Variable name. If omitted, the first numeric ``(N, 3)`` / ``(3, N)`` + array is used. + sampling, units, name + See :meth:`from_array`. + one_based : bool + Subtract 1 from the coordinates (MATLAB 1-based voxel indices). + """ + path = Path(path) + arrays: dict[str, NDArray] = {} + try: + import scipy.io as sio + + raw = sio.loadmat(path) + arrays = { + k: np.asarray(v) + for k, v in raw.items() + if not k.startswith("__") and isinstance(v, np.ndarray) and v.dtype.kind in "fiu" + } + except NotImplementedError: # v7.3 + import h5py + + with h5py.File(path, "r") as f: + + def _collect(g, prefix=""): + for k, v in g.items(): + if isinstance(v, h5py.Dataset) and v.dtype.kind in "fiu": + arrays[prefix + k] = np.asarray(v[()]).T + elif isinstance(v, h5py.Group): + _collect(v, prefix + k + "/") + + _collect(f) + if key is None: + candidates = [k for k, v in arrays.items() if v.ndim == 2 and 3 in v.shape] + if not candidates: + raise ValueError( + f"No (N, 3) coordinate array found in {path.name}: {list(arrays)}" + ) + key = candidates[0] + xyz = arrays[key].astype(float) + if one_based: + xyz = xyz - 1.0 + return cls.from_array( + xyz, + sampling=sampling, + units=units, + name=name or path.stem, + metadata={"source": str(path), "key": key}, + ) + + @classmethod + def from_xyz( + cls, + path: str | Path, + sampling: float | Sequence[float] = 1.0, + units: str = "A", + name: str | None = None, + ) -> "AtomicModel": + """Load an ``.xyz`` text file (``element x y z`` rows after a 2-line header).""" + path = Path(path) + lines = path.read_text().strip().splitlines() + try: + count = int(lines[0].split()[0]) + body = lines[2 : 2 + count] + except (ValueError, IndexError): + body = lines + symbols, coords = [], [] + for line in body: + parts = line.split() + if len(parts) < 4: + continue + symbols.append(parts[0]) + coords.append([float(parts[1]), float(parts[2]), float(parts[3])]) + xyz = np.asarray(coords) + names = sorted(set(symbols)) + species = np.array([names.index(s) for s in symbols], dtype=float) + model = cls.from_array(xyz, sampling=sampling, units=units, name=name or path.stem) + model.set_channel("species", species, categories=names) + return model + + @classmethod + def from_atoms(cls, atoms: Any, name: str | None = None) -> "AtomicModel": + """Create a model from a :class:`quantem.tomography.Atoms` tracing result. + + Positions are taken in voxel units together with the traced intensity + and Gaussian width, and ``sampling`` is copied from the source volume. + """ + sites = atoms.sites.array + sampling = np.asarray(atoms._sampling, dtype=float) + xyz = (sites[:, :3] - np.asarray(atoms._origin)[None, :]) / sampling[None, :] + model = cls.from_array( + xyz, + sampling=sampling, + units=str(atoms._units[0]), + name=name or f"{atoms._source.name} atoms", + ) + model.set_channel("intensity", sites[:, 3]) + model.set_channel("sigma", sites[:, 4] / float(sampling.mean())) + return model + + # ------------------------------------------------------------------ # + # Basic properties + # ------------------------------------------------------------------ # + @property + def name(self) -> str: + """Model name.""" + return self._name + + @name.setter + def name(self, value: str) -> None: + self._name = str(value) + + @property + def metadata(self) -> dict[str, Any]: + """Free-form metadata dictionary.""" + return self._metadata + + @property + def sites(self) -> Vector: + """Per-site table (0-D ``Vector``); fields ``x, y, z`` plus channels.""" + return self._sites + + @property + def num_sites(self) -> int: + """Number of atomic sites.""" + return int(self._sites.array.shape[0]) + + @property + def sampling(self) -> NDArray: + """``(3,)`` physical size of one native coordinate unit per axis.""" + return self._sampling + + @property + def units(self) -> str: + """Physical length unit of :attr:`positions`.""" + return self._units + + @property + def positions_native(self) -> NDArray: + """``(N, 3)`` positions in native (uncalibrated) units.""" + return np.array(self._sites.select_fields(*_POSITION_FIELDS).array, dtype=float) + + @positions_native.setter + def positions_native(self, value: NDArray) -> None: + value = np.asarray(value, dtype=float) + if value.shape != (self.num_sites, 3): + raise ValueError(f"positions must have shape {(self.num_sites, 3)}") + self._sites.select_fields(*_POSITION_FIELDS)[...] = value + self._invalidate() + + @property + def positions(self) -> NDArray: + """``(N, 3)`` calibrated positions (native * sampling).""" + return self.positions_native * self._sampling[None, :] + + @property + def center(self) -> NDArray: + """``(3,)`` mean calibrated position.""" + return self.positions.mean(0) + + @property + def channels(self) -> list[str]: + """Names of all per-site channels (fields other than ``x, y, z``).""" + return [f for f in self._sites.fields if f not in _POSITION_FIELDS] + + @property + def categories(self) -> dict[str, list[str]]: + """Label names for categorical channels, e.g. ``{"structure": ["fcc", "hcp"]}``.""" + return self._categories + + @property + def structure_names(self) -> list[str]: + """Template names indexed by the ``structure`` channel value.""" + return list(self._structure_names) + + @property + def templates(self) -> dict[str, PolyhedralTemplate]: + """Templates used in the last :meth:`match_templates` call.""" + return { + name: PolyhedralTemplate( + name=name, + vectors=np.asarray(spec["vectors"]), + shells=tuple(float(x) for x in spec["shells"]), + shell_counts=tuple(int(x) for x in spec["shell_counts"]), + symmetry=np.asarray(spec["symmetry"]), + ) + for name, spec in self._template_specs.items() + } + + @property + def matches(self) -> dict[str, TemplateMatch]: + """Raw per-template matching results (see :class:`TemplateMatch`).""" + return self._matches + + @property + def pdf(self) -> dict[str, Any] | None: + """Result of :meth:`compute_pdf` (native units), or ``None``.""" + return self._pdf + + @property + def nn_fit(self) -> dict[str, Any] | None: + """First-peak fit from :meth:`compute_pdf` (native units), or ``None``.""" + return self._nn_fit + + @property + def nn_distance(self) -> float: + """Mean nearest-neighbor distance in native units (requires :meth:`compute_pdf`).""" + if self._nn_fit is None: + self.compute_pdf() + assert self._nn_fit is not None + return float(self._nn_fit["r_nn"]) + + @property + def bond_length(self) -> float: + """Mean nearest-neighbor distance in calibrated units.""" + return self.nn_distance * float(self._sampling.mean()) + + @property + def neighbor_indices(self) -> NDArray: + """``(N, K)`` neighbor indices sorted by distance (``-1`` = missing).""" + if self._neighbor_indices is None: + self.find_neighbors() + assert self._neighbor_indices is not None + return self._neighbor_indices + + @property + def neighbor_distances(self) -> NDArray: + """``(N, K)`` neighbor distances in native units.""" + if self._neighbor_distances is None: + self.find_neighbors() + assert self._neighbor_distances is not None + return self._neighbor_distances + + def _invalidate(self) -> None: + self._pdf = None + self._nn_fit = None + self._neighbor_distances = None + self._neighbor_indices = None + self._matches = {} + + # ------------------------------------------------------------------ # + # Channels + # ------------------------------------------------------------------ # + def get_channel(self, name: str) -> NDArray: + """Return a per-site channel as a ``(N,)`` array.""" + if name in _POSITION_FIELDS: + return self.positions[:, _POSITION_FIELDS.index(name)] + if name not in self._sites.fields: + raise KeyError(f"Unknown channel {name!r}; available: {self.channels}") + return np.array(self._sites.select_fields(name).array[:, 0], dtype=float) + + def set_channel( + self, + name: str, + values: NDArray, + units: str = "none", + categories: Sequence[str] | None = None, + ) -> None: + """Add or overwrite a per-site channel. + + Parameters + ---------- + name : str + Channel name. + values : ndarray + ``(N,)`` values (cast to float; use integer codes for categories). + units : str + Units label. + categories : sequence of str, optional + Names for integer codes ``0, 1, ...``; marks the channel categorical. + """ + if name in _POSITION_FIELDS: + raise ValueError("Use positions_native to modify coordinates.") + values = np.asarray(values, dtype=float).reshape(-1) + if values.shape[0] != self.num_sites: + raise ValueError(f"values must have length {self.num_sites}, got {values.shape[0]}") + if name in self._sites.fields: + self._sites.select_fields(name)[...] = values[:, None] + else: + self._sites.add_fields(name, values[:, None], units) + if categories is not None: + self._categories[name] = [str(c) for c in categories] + elif name in self._categories: + del self._categories[name] + + def remove_channel(self, name: str) -> None: + """Delete a channel.""" + self._sites.remove_fields(name) + self._categories.pop(name, None) + + def __getitem__(self, name: str) -> NDArray: + return self.get_channel(name) + + # ------------------------------------------------------------------ # + # Pair distribution function and calibration + # ------------------------------------------------------------------ # + def compute_pdf( + self, + r_max: float | None = None, + dr: float | None = None, + sigma: float | None = None, + fit_radius: float = 1.25, + cutoff_sigma: float = 2.0, + r_min: float | None = None, + ) -> dict[str, Any]: + """Compute the radial distribution function and fit the first NN peak. + + All radii are in native units. + + Parameters + ---------- + r_max : float, optional + Maximum radius; default 4x an initial NN estimate. + dr : float, optional + Bin width; default ``r_max / 600``. + sigma : float, optional + Smoothing width; default ``2 * dr``. + fit_radius, cutoff_sigma, r_min + See :func:`quantem.atoms.pdf.fit_first_peak`. + + Returns + ------- + dict + RDF arrays (``r``, ``g``, ``g_smooth``, ``counts``) plus the fit. + """ + xyz = self.positions_native + if r_max is None: + d1, _ = find_neighbors(xyz, 1) + r_max = 4.0 * float(np.median(d1)) + if dr is None: + dr = r_max / 600.0 + pdf = radial_distribution(xyz, r_max=r_max, dr=dr, sigma=sigma) + fit = fit_first_peak( + pdf["r"], + pdf["g_smooth"], + fit_radius=fit_radius, + cutoff_sigma=cutoff_sigma, + r_min=r_min, + ) + self._pdf = pdf + self._nn_fit = fit + return {**pdf, **{k: v for k, v in fit.items()}} + + def calibrate( + self, + structure: str | None = None, + lattice_constant: float | None = None, + nn_distance: float | None = None, + units: str = "A", + ) -> float: + """Set ``sampling`` so the measured NN distance matches a reference. + + Provide either ``structure`` + ``lattice_constant`` or ``nn_distance``. + + Parameters + ---------- + structure : str, optional + Reference crystal (``"fcc"``, ``"bcc"``, ``"hcp"`` ...). + lattice_constant : float, optional + Lattice constant of the reference crystal in ``units``. + nn_distance : float, optional + Target NN distance in ``units``. + units : str + Physical unit of the reference. + + Returns + ------- + float + The isotropic scale (physical units per native unit). + """ + if nn_distance is None: + if structure is None or lattice_constant is None: + raise ValueError("Give structure and lattice_constant, or nn_distance.") + nn_distance = nn_distance_from_lattice(structure, lattice_constant) + scale = float(nn_distance) / self.nn_distance + self._sampling = np.full(3, scale) + self._units = units + return scale + + # ------------------------------------------------------------------ # + # Neighbors and bonds + # ------------------------------------------------------------------ # + def find_neighbors(self, num_neighbors: int = 24, cutoff: float | None = None) -> None: + """Build neighbor lists and per-site bond statistics. + + Adds channels ``num_neighbors`` (first-shell coordination), + ``bond_mean`` and ``bond_std`` (calibrated units). + + Parameters + ---------- + num_neighbors : int + Neighbors stored per site; must exceed the largest template. + cutoff : float, optional + First-shell radial cutoff in native units. Default: upper cutoff + from the RDF first-peak fit. + """ + dist, idx = find_neighbors(self.positions_native, num_neighbors) + self._neighbor_distances = dist + self._neighbor_indices = idx + if cutoff is None: + cutoff = self.first_shell_cutoff + first = dist <= cutoff + scale = float(self._sampling.mean()) + d_first = np.where(first, dist, np.nan) * scale + with np.errstate(invalid="ignore"): + self.set_channel("num_neighbors", first.sum(1)) + self.set_channel("bond_mean", np.nanmean(d_first, axis=1), self._units) + self.set_channel("bond_std", np.nanstd(d_first, axis=1), self._units) + + @property + def first_shell_cutoff(self) -> float: + """Upper radial cutoff of the first shell (native units) from the RDF fit.""" + if self._nn_fit is None: + self.compute_pdf() + assert self._nn_fit is not None + return float(self._nn_fit["cutoff"][1]) + + def neighbor_vectors(self, normalize: bool = True) -> tuple[NDArray, NDArray]: + """Neighbor displacement vectors. + + Parameters + ---------- + normalize : bool + Divide by the NN distance so bonds have length ~1. + + Returns + ------- + dxyz, dist : ndarray + ``(N, K, 3)`` vectors and ``(N, K)`` lengths (native units, or NN + units when normalized). Missing neighbors are ``inf``. + """ + idx = self.neighbor_indices + dist = self.neighbor_distances.copy() + xyz = self.positions_native + safe = np.where(idx >= 0, idx, 0) + dxyz = xyz[safe] - xyz[:, None, :] + dxyz[idx < 0] = np.inf + if normalize: + r_nn = self.nn_distance + return dxyz / r_nn, dist / r_nn + return dxyz, dist + + def bond_angles(self) -> NDArray: + """All first-shell bond angles per site, ``(N, K(K-1)/2)`` degrees with ``nan`` padding.""" + dxyz, dist = self.neighbor_vectors(normalize=False) + valid = dist <= self.first_shell_cutoff + dxyz = np.where(np.isfinite(dxyz), dxyz, 0.0) + return meas.bond_angles(dxyz, valid) + + # ------------------------------------------------------------------ # + # Template matching and classification + # ------------------------------------------------------------------ # + def match_templates( + self, + templates: Sequence[str | PolyhedralTemplate] = ("fcc", "hcp"), + score_radius: float = 0.5, + cutoff_factor: float = 1.15, + angle_tolerance: float = 30.0, + num_refine: int = 2, + chunk_size: int = 512, + device: str | None = None, + progress: bool = True, + ) -> dict[str, TemplateMatch]: + """Match polyhedral templates to every site. + + Adds channels ``score_``, ``rmsd_`` and ``matched_`` + for each template, then calls :meth:`classify` with default settings. + + Parameters + ---------- + templates : sequence of str or PolyhedralTemplate + Template names (see :data:`quantem.atoms.TEMPLATE_NAMES`) or objects. + score_radius : float + Matching radius in NN units; see :func:`quantem.atoms.matching.match_template`. + cutoff_factor : float + Neighbors farther than ``cutoff_factor * template.max_radius`` (NN + units) are ignored for that template. + angle_tolerance, num_refine, chunk_size, device, progress + Forwarded to :func:`quantem.atoms.matching.match_template`. + + Returns + ------- + dict + ``{name: TemplateMatch}``. + """ + dxyz, dist = self.neighbor_vectors(normalize=True) + dxyz = np.where(np.isfinite(dxyz), dxyz, 1e3) + self._matches = {} + self._template_specs = {} + for item in templates: + template = get_template(item) if isinstance(item, str) else item + name = template.name + valid = dist <= cutoff_factor * template.max_radius + if valid.shape[1] < template.num_neighbors: + raise ValueError( + f"Template {name!r} has {template.num_neighbors} neighbors but only " + f"{valid.shape[1]} are stored; call find_neighbors(num_neighbors=...)." + ) + result = match_template( + dxyz, + valid, + template, + score_radius=score_radius, + angle_tolerance=angle_tolerance, + num_refine=num_refine, + chunk_size=chunk_size, + device=device, + progress=progress, + ) + self._matches[name] = result + self._template_specs[name] = { + "vectors": np.asarray(template.vectors), + "shells": list(template.shells), + "shell_counts": list(template.shell_counts), + "symmetry": np.asarray(template.symmetry), + } + self.set_channel(f"score_{name}", result["score"]) + self.set_channel(f"rmsd_{name}", result["rmsd"]) + self.set_channel(f"matched_{name}", result["num_matched"]) + self.classify() + return self._matches + + def classify( + self, threshold: float = 0.5, smooth: bool = False, use_strained: bool = False + ) -> NDArray: + """Assign each site to its best-scoring template. + + Adds channels ``structure`` (categorical: template index, ``-1`` for + unclassified), ``score_max`` and, when exactly two templates were + matched, ``score_diff`` (first minus second). + + Parameters + ---------- + threshold : float + Minimum score for a site to be classified. + smooth : bool + Average each site's scores with its first-shell neighbors before + deciding (more robust for noisy models). + use_strained : bool + Use the affine-fit scores (``score_strained``) instead. + + Returns + ------- + ndarray + ``(N,)`` structure codes. + """ + if not self._matches: + raise RuntimeError("Call match_templates() first.") + names = list(self._matches) + key = "score_strained" if use_strained else "score" + scores = np.stack([self._matches[n][key] for n in names], axis=1) + if smooth: + idx = self.neighbor_indices + first = self.neighbor_distances <= self.first_shell_cutoff + safe = np.where(idx >= 0, idx, 0) + nb = scores[safe] * first[..., None] + scores = (scores + nb.sum(1)) / (1.0 + first.sum(1))[:, None] + best = scores.argmax(1) + score_max = scores.max(1) + structure = np.where(score_max >= threshold, best, -1) + self._structure_names = names + self.set_channel("structure", structure, categories=names) + self.set_channel("score_max", score_max) + if len(names) == 2: + self.set_channel("score_diff", scores[:, 0] - scores[:, 1]) + return structure + + def rotations(self, template: str | None = None) -> NDArray: + """``(N, 3, 3)`` fitted orientations (lab <- crystal) for a template. + + With ``template=None`` the rotation from each site's classified + structure is returned (identity for unclassified sites). + """ + if template is not None: + return self._matches[template]["rotation"] + structure = self.get_channel("structure").astype(int) + out = np.tile(np.eye(3), (self.num_sites, 1, 1)) + for i, name in enumerate(self._structure_names): + sel = structure == i + out[sel] = self._matches[name]["rotation"][sel] + return out + + def segment_grains( + self, + structure: str = "fcc", + angle_threshold: float = 5.0, + min_size: int = 20, + min_score: float | None = None, + ) -> NDArray: + """Cluster sites of one structure into grains by local orientation. + + Neighboring sites of the given structure whose disorientation is below + ``angle_threshold`` are connected; connected components become grains. + Adds channels ``grain`` (``-1`` = none) and ``misorientation`` (largest + disorientation to any first-shell neighbor of the same structure). + + Parameters + ---------- + structure : str + Template name to segment (e.g. ``"fcc"``). + angle_threshold : float + Maximum disorientation (degrees) inside a grain. + min_size : int + Grains with fewer sites are discarded. + min_score : float, optional + Only sites with ``score_`` above this take part; + default: sites classified as ``structure``. + + Returns + ------- + ndarray + ``(N,)`` grain labels sorted by decreasing size. + """ + if structure not in self._matches: + raise KeyError(f"No match for {structure!r}; run match_templates first.") + template = self.templates[structure] + rot = self._matches[structure]["rotation"] + idx = self.neighbor_indices + first = self.neighbor_distances <= self.first_shell_cutoff + if min_score is None: + member = self.get_channel("structure").astype(int) == self._structure_names.index( + structure + ) + else: + member = self._matches[structure]["score"] >= min_score + ang = meas.misorientation(rot, idx, template.symmetry) + safe = np.where(idx >= 0, idx, 0) + same = member[:, None] & member[safe] & first & (idx >= 0) + edge = same & (ang < angle_threshold) + labels = meas.segment_grains(idx, edge, min_size=min_size) + labels[~member] = -1 + worst = np.where(same, ang, -1.0).max(axis=1).clip(0.0) + num_grains = int(labels.max()) + 1 if labels.size else 0 + self.set_channel("grain", labels, categories=[str(i) for i in range(num_grains)]) + self.set_channel("misorientation", worst, "deg") + return labels + + def compute_strain( + self, template: str | None = None, frame: str = "lab" + ) -> dict[str, NDArray]: + """Local strain from the affine template fit. + + Adds channels ``strain_xx, strain_yy, strain_zz, strain_xy, strain_xz, + strain_yz, strain_dilation, strain_equivalent``. Strain is relative to + the mean NN distance of the model. + + Parameters + ---------- + template : str, optional + Template whose fit to use; default: each site's classified structure. + frame : {"lab", "crystal"} + Frame of the strain tensor. + """ + if template is not None: + f = self._matches[template]["deformation"] + r = self._matches[template]["rotation"] + else: + structure = self.get_channel("structure").astype(int) + f = np.tile(np.eye(3), (self.num_sites, 1, 1)) + r = f.copy() + for i, name in enumerate(self._structure_names): + sel = structure == i + f[sel] = self._matches[name]["deformation"][sel] + r[sel] = self._matches[name]["rotation"][sel] + strain = meas.strain_from_deformation(f, r, frame=frame) + for key in ("e_xx", "e_yy", "e_zz", "e_xy", "e_xz", "e_yz"): + self.set_channel("strain_" + key[2:], strain[key]) + self.set_channel("strain_dilation", strain["dilation"]) + self.set_channel("strain_equivalent", strain["equivalent"]) + return strain + + # ------------------------------------------------------------------ # + # Other measurements + # ------------------------------------------------------------------ # + def surface_distance(self) -> NDArray: + """Distance of each site to the convex hull (calibrated units); channel ``surface_distance``.""" + d = meas.convex_hull_distance(self.positions) + self.set_channel("surface_distance", d, self._units) + return d + + def sample_volume(self, volume: Any, radius: float = 1.5, name: str = "intensity") -> NDArray: + """Mean reconstruction intensity around each site; stored as a channel. + + Parameters + ---------- + volume : ndarray or Dataset3d + Source volume indexed like the native coordinates. + radius : float + Sphere radius in voxels. + name : str + Channel name. + """ + arr = getattr(volume, "array", volume) + if hasattr(arr, "detach"): + arr = arr.detach().cpu().numpy() + values = meas.sample_volume(np.asarray(arr), self.positions_native, radius=radius) + self.set_channel(name, values) + return values + + def classify_species( + self, + channel: str = "intensity", + num_species: int = 2, + names: Sequence[str] | None = None, + ) -> NDArray: + """Split a channel (e.g. intensity) into species with 1D k-means; channel ``species``.""" + labels, centers = meas.kmeans_1d(self.get_channel(channel), num_species) + if names is None: + names = [f"species_{i}" for i in range(num_species)] + self.set_channel("species", labels, categories=names) + self._metadata["species_centers"] = centers.tolist() + return labels + + # ------------------------------------------------------------------ # + # Geometry helpers + # ------------------------------------------------------------------ # + def rotate(self, rotation: NDArray, about_center: bool = True) -> None: + """Rotate all positions in place with a ``(3, 3)`` matrix (``x' = R x``).""" + rotation = np.asarray(rotation, dtype=float) + xyz = self.positions_native + c = xyz.mean(0) if about_center else np.zeros(3) + self.positions_native = (xyz - c) @ rotation.T + c + + def select(self, mask: NDArray) -> "AtomicModel": + """Return a new model containing only the sites where ``mask`` is True.""" + mask = np.asarray(mask, dtype=bool) + table = self._sites.array[mask] + sites = Vector.from_shape( + shape=(), fields=list(self._sites.fields), units=list(self._sites.units), name="sites" + ) + sites[...] = np.ascontiguousarray(table) + model = AtomicModel( + sites=sites, + sampling=self._sampling.copy(), + units=self._units, + name=self._name, + metadata=dict(self._metadata), + _token=self._token, + ) + model._categories = dict(self._categories) + model._structure_names = list(self._structure_names) + return model + + # ------------------------------------------------------------------ # + # Visualization + # ------------------------------------------------------------------ # + def plot(self, kind: str = "slab", show_docstring: bool = False, **kwargs): + """Static matplotlib plots; see :mod:`quantem.atoms.visualization`. + + Parameters + ---------- + kind : str + One of ``"pdf"``, ``"histogram"``, ``"slab"``, ``"slices"``, + ``"template"``. + show_docstring : bool + Print the plot function's docstring instead of plotting. + **kwargs + Forwarded to the plot function. + """ + if kind not in PLOT_REGISTRY: + raise ValueError(f"Unknown plot kind {kind!r}; choose from {list(PLOT_REGISTRY)}") + fn = PLOT_REGISTRY[kind] + if show_docstring: + print(fn.__doc__) + return None + return fn(self, **kwargs) + + def show(self, **kwargs): + """Open the interactive 3D viewer (:class:`quantem.atoms.ShowAtoms3D`).""" + from quantem.atoms.show_atoms import ShowAtoms3D + + return ShowAtoms3D(self, **kwargs) + + def __repr__(self) -> str: + return ( + f"AtomicModel(name={self._name!r}, num_sites={self.num_sites}, " + f"units={self._units!r}, channels={self.channels})" + ) diff --git a/src/quantem/atoms/matching.py b/src/quantem/atoms/matching.py new file mode 100644 index 000000000..1b5f4e56b --- /dev/null +++ b/src/quantem/atoms/matching.py @@ -0,0 +1,287 @@ +"""Fast polyhedral template matching for 3D atomic models. + +Every site is compared against one or more :class:`PolyhedralTemplate` objects +to find the rotation that best aligns the template's neighbor vectors with the +site's measured neighbor vectors. The search is fully vectorized with PyTorch +(CPU or GPU), so all sites are processed in chunks rather than one at a time. + +Algorithm (per site, for one template) +-------------------------------------- +1. **Trial rotations.** A reference pair of adjacent template vectors + ``(t_a, t_b)`` defines an orthonormal frame. Every ordered pair of measured + neighbor vectors ``(p_i, p_j)`` defines another frame; the rotation mapping + one frame onto the other is a trial orientation. Pairs whose angle differs + from the template pair angle by more than ``angle_tolerance`` are skipped. +2. **Scoring.** Each rotated template vector is matched to its nearest + measured neighbor; the score is the sum of ``clip(1 - d / score_radius, 0, 1)`` + over template vectors, so it ranges from 0 to the number of template + vectors. The best-scoring trial is kept. +3. **Refinement.** Template vectors are assigned one-to-one to their nearest + neighbors within ``score_radius`` and the rotation is re-fit with a batched + Kabsch (SVD) solve. The final score uses only these one-to-one matches, so + a template vector cannot be "explained" by a neighbor already used. A second least-squares solve yields the affine + deformation gradient ``F`` (template -> measured), from which local strain + is derived. + +All neighbor vectors must be given in **nearest-neighbor bond-length units** +(i.e. divided by the mean NN distance), matching the template convention. +""" + +from __future__ import annotations + +import math + +import numpy as np +import torch +from numpy.typing import NDArray +from torch import Tensor +from tqdm.auto import tqdm + +from quantem.atoms.templates import PolyhedralTemplate + +__all__ = ["match_template", "TemplateMatch"] + + +class TemplateMatch(dict): + """Result of :func:`match_template` for one template (a dict with fixed keys). + + Keys + ---- + ``score`` : ``(N,)`` soft score in ``[0, 1]`` (normalized by template size). + ``score_strained`` : ``(N,)`` score after the affine (strain-allowing) fit. + ``num_matched`` : ``(N,)`` number of template vectors matched one-to-one. + ``rmsd`` : ``(N,)`` RMS distance of matched vectors (NN units). + ``rotation`` : ``(N, 3, 3)`` rotation ``R`` with ``p ~ R t`` (lab <- template). + ``deformation`` : ``(N, 3, 3)`` deformation gradient ``F`` with ``p ~ F t``. + ``matched`` : ``(N, M)`` neighbor-slot index matched to each template vector, or -1. + """ + + +def _frames(u: Tensor, v: Tensor) -> Tensor: + """Right-handed orthonormal frames (columns) from batched vector pairs.""" + e1 = u / u.norm(dim=-1, keepdim=True).clamp_min(1e-12) + v2 = v - e1 * (v * e1).sum(-1, keepdim=True) + e2 = v2 / v2.norm(dim=-1, keepdim=True).clamp_min(1e-12) + e3 = torch.cross(e1, e2, dim=-1) + return torch.stack([e1, e2, e3], dim=-1) + + +def _soft_score(dmin: Tensor, score_radius: float) -> Tensor: + return (1.0 - dmin / score_radius).clamp(0.0, 1.0).sum(-1) + + +def _kabsch(t: Tensor, p: Tensor, w: Tensor, fallback: Tensor) -> Tensor: + """Batched Kabsch rotation with ``p ~ R t``; weights ``w`` select matches.""" + h = (t * w[..., None]).transpose(1, 2) @ p # (N,3,3) = sum w t p^T + h_cpu = h.detach().to("cpu").to(torch.float64) + u, _, vt = torch.linalg.svd(h_cpu) + v = vt.transpose(1, 2) + d = torch.det(v @ u.transpose(1, 2)) + sign = torch.ones_like(d) + sign[d < 0] = -1.0 + v = v.clone() + v[:, :, 2] *= sign[:, None] + r = (v @ u.transpose(1, 2)).to(fallback.device, fallback.dtype) + enough = (w.sum(-1) >= 3).to(fallback.device) + return torch.where(enough[:, None, None], r, fallback) + + +def _deformation(t: Tensor, p: Tensor, w: Tensor, fallback: Tensor, min_matched: int) -> Tensor: + """Batched least-squares deformation gradient ``F`` with ``p ~ F t``.""" + tw = t * w[..., None] + a = tw.transpose(1, 2) @ t # (N,3,3) sum w t t^T + b = tw.transpose(1, 2) @ p # (N,3,3) sum w t p^T + eye = torch.eye(3, device=t.device, dtype=t.dtype) + a = a + 1e-6 * eye + # F^T = A^{-1} B -> F = B^T A^{-1} + ft = torch.linalg.solve(a.to("cpu").to(torch.float64), b.to("cpu").to(torch.float64)) + f = ft.transpose(1, 2).to(t.device, t.dtype) + enough = w.sum(-1) >= min_matched + return torch.where(enough[:, None, None], f, fallback) + + +def _assign_one_to_one( + dist: Tensor, valid: Tensor, score_radius: float +) -> tuple[Tensor, Tensor, Tensor]: + """Match each template vector to its nearest valid neighbor, one-to-one. + + Parameters + ---------- + dist : Tensor + ``(N, M, K)`` distances between rotated template vectors and neighbors. + valid : Tensor + ``(N, K)`` neighbor validity mask. + + Returns + ------- + dmin, j, keep : Tensor + ``(N, M)`` nearest distance, neighbor slot index and match mask. + """ + dist = dist.masked_fill(~valid[:, None, :], float("inf")) + dmin, j = dist.min(-1) + keep = dmin < score_radius + n, _, k = dist.shape + best = torch.full((n, k), float("inf"), device=dist.device, dtype=dist.dtype) + best = best.scatter_reduce(1, j, dmin.masked_fill(~keep, float("inf")), reduce="amin") + keep = keep & (dmin <= best.gather(1, j) + 1e-6) + return dmin, j, keep + + +def match_template( + dxyz: NDArray | Tensor, + valid: NDArray | Tensor, + template: PolyhedralTemplate, + score_radius: float = 0.5, + angle_tolerance: float = 30.0, + num_pair_neighbors: int | None = None, + num_refine: int = 2, + chunk_size: int = 512, + device: str | torch.device | None = None, + progress: bool = True, +) -> TemplateMatch: + """Match one polyhedral template to every site. + + Parameters + ---------- + dxyz : array or Tensor + ``(N, K, 3)`` neighbor vectors in NN bond-length units, sorted by + distance. + valid : array or Tensor + ``(N, K)`` boolean mask of usable neighbors (e.g. inside a radial cutoff). + template : PolyhedralTemplate + Template to match. + score_radius : float + Matching radius (NN units). A template vector contributes + ``1 - d / score_radius`` to the score when its nearest neighbor is at + distance ``d``. + angle_tolerance : float + Trial neighbor pairs are only tested if their angle is within this many + degrees of the reference template pair angle. + num_pair_neighbors : int, optional + Trial rotations are built from ordered pairs among the first this many + neighbors. Default ``min(K, M + 2)`` with ``M`` the template size. + num_refine : int + Number of assign-then-Kabsch refinement iterations. + chunk_size : int + Sites processed per batch (memory ~ ``chunk_size * P * M * K`` floats). + device : str or torch.device, optional + Torch device. Default: ``quantem.core.config`` device, or CPU. + progress : bool + Show a progress bar. + + Returns + ------- + TemplateMatch + See :class:`TemplateMatch` for keys. All values are NumPy arrays. + """ + if device is None: + from quantem.core import config + + device = config.get_device() + device = torch.device(device) + dtype = torch.float32 + + p_all = torch.as_tensor(np.asarray(dxyz), dtype=dtype) + valid_all = torch.as_tensor(np.asarray(valid), dtype=torch.bool) + n, k, _ = p_all.shape + t = torch.as_tensor(template.vectors, dtype=dtype, device=device) + m = t.shape[0] + k_pair = min(k, m + 2) if num_pair_neighbors is None else min(k, int(num_pair_neighbors)) + + # reference template pair: vector 0 and its closest first-shell partner + t_r = t.norm(dim=-1) + first = t_r <= template.shells[0] * (1 + 1e-3) + cos_t = (t @ t[0]) / (t_r * t_r[0]) + cos_t[0] = -2.0 + cos_t[~first] = -2.0 + j_ref = int(torch.argmax(cos_t)) + ref_angle = float(torch.arccos(cos_t[j_ref].clamp(-1, 1))) + ref_frame = _frames(t[0][None], t[j_ref][None])[0] # (3,3) + tol = math.radians(angle_tolerance) + ref_len_a, ref_len_b = float(t_r[0]), float(t_r[j_ref]) + + ii, jj = torch.meshgrid(torch.arange(k_pair), torch.arange(k_pair), indexing="ij") + pair_mask = ii != jj + ii, jj = ii[pair_mask].to(device), jj[pair_mask].to(device) + num_pairs = ii.numel() + + score = torch.zeros(n, dtype=dtype) + score_strained = torch.zeros(n, dtype=dtype) + num_matched = torch.zeros(n, dtype=torch.int64) + rmsd = torch.full((n,), float("nan"), dtype=dtype) + rotation = torch.eye(3, dtype=dtype).repeat(n, 1, 1) + deformation = torch.eye(3, dtype=dtype).repeat(n, 1, 1) + matched = torch.full((n, m), -1, dtype=torch.int64) + min_matched_affine = max(4, m // 3) + + ranges = range(0, n, chunk_size) + for start in tqdm(ranges, desc=f"match {template.name}", disable=not progress): + sl = slice(start, min(start + chunk_size, n)) + p = p_all[sl].to(device) # (Nc,K,3) + pv = valid_all[sl].to(device) # (Nc,K) + nc = p.shape[0] + p_safe = torch.where(pv[..., None], p, torch.full_like(p, 1e3)) + + # ---- trial rotations from neighbor pairs ------------------------- + pi, pj = p[:, ii], p[:, jj] # (Nc,P,3) + li, lj = pi.norm(dim=-1), pj.norm(dim=-1) + cos_ij = (pi * pj).sum(-1) / (li * lj).clamp_min(1e-12) + ang = torch.arccos(cos_ij.clamp(-1, 1)) + ok = (ang - ref_angle).abs() < tol + ok &= pv[:, ii] & pv[:, jj] + ok &= ((li - ref_len_a).abs() < 0.35 * ref_len_a) & ( + (lj - ref_len_b).abs() < 0.35 * ref_len_b + ) + frames = _frames(pi, pj) # (Nc,P,3,3) + rot = frames @ ref_frame.T # (Nc,P,3,3): p ~ R t + t_rot = torch.einsum("npab,mb->npma", rot, t) # (Nc,P,M,3) + dist = torch.cdist( + t_rot.reshape(nc * num_pairs, m, 3), p_safe.repeat_interleave(num_pairs, 0) + ) + dmin = dist.view(nc, num_pairs, m, k).min(-1).values + trial_score = _soft_score(dmin, score_radius).masked_fill(~ok, -1.0) + best = trial_score.argmax(dim=1) + r_best = rot[torch.arange(nc, device=device), best] # (Nc,3,3) + del dist, dmin, t_rot, rot, frames + + # ---- refinement: one-to-one assignment + Kabsch ------------------ + t_b = t[None].expand(nc, m, 3) + for _ in range(max(1, num_refine)): + t_rot = t_b @ r_best.transpose(1, 2) + dmin, j, keep = _assign_one_to_one(torch.cdist(t_rot, p_safe), pv, score_radius) + w = keep.to(dtype) + p_m = p_safe.gather(1, j[..., None].expand(-1, -1, 3)) + r_best = _kabsch(t_b, p_m, w, r_best) + + t_rot = t_b @ r_best.transpose(1, 2) + dmin, j, keep = _assign_one_to_one(torch.cdist(t_rot, p_safe), pv, score_radius) + w = keep.to(dtype) + p_m = p_safe.gather(1, j[..., None].expand(-1, -1, 3)) + f_best = _deformation(t_b, p_m, w, r_best, min_matched_affine) + t_def = t_b @ f_best.transpose(1, 2) + dmin_def = ( + torch.cdist(t_def, p_safe).masked_fill(~pv[:, None, :], float("inf")).min(-1).values + ) + + n_match = keep.sum(-1) + dmin_1to1 = dmin.masked_fill(~keep, float("inf")) + _, _, keep_def = _assign_one_to_one(torch.cdist(t_def, p_safe), pv, score_radius) + dmin_def = dmin_def.masked_fill(~keep_def, float("inf")) + score[sl] = (_soft_score(dmin_1to1, score_radius) / m).cpu() + score_strained[sl] = (_soft_score(dmin_def, score_radius) / m).cpu() + num_matched[sl] = n_match.cpu() + sq = (dmin**2 * w).sum(-1) / n_match.clamp_min(1).to(dtype) + rmsd[sl] = torch.where(n_match > 0, sq.sqrt(), torch.full_like(sq, float("nan"))).cpu() + rotation[sl] = r_best.cpu() + deformation[sl] = f_best.cpu() + matched[sl] = torch.where(keep, j, torch.full_like(j, -1)).cpu() + + return TemplateMatch( + score=score.numpy(), + score_strained=score_strained.numpy(), + num_matched=num_matched.numpy(), + rmsd=rmsd.numpy(), + rotation=rotation.numpy(), + deformation=deformation.numpy(), + matched=matched.numpy(), + ) diff --git a/src/quantem/atoms/measurements.py b/src/quantem/atoms/measurements.py new file mode 100644 index 000000000..704c66c9c --- /dev/null +++ b/src/quantem/atoms/measurements.py @@ -0,0 +1,266 @@ +"""Per-site measurements derived from template matching and neighbor lists. + +Functions here operate on plain arrays so they can be tested independently of +:class:`~quantem.atoms.AtomicModel`, which wraps them. Everything is +vectorized with NumPy / SciPy; nothing loops over atoms in Python. +""" + +from __future__ import annotations + +import numpy as np +from numpy.typing import NDArray +from scipy.sparse import coo_matrix +from scipy.sparse.csgraph import connected_components +from scipy.spatial import ConvexHull + +__all__ = [ + "misorientation", + "segment_grains", + "strain_from_deformation", + "bond_angles", + "convex_hull_distance", + "sample_volume", + "kmeans_1d", + "rotation_to_quaternion", +] + + +def misorientation( + rotation: NDArray, + neighbor_index: NDArray, + symmetry: NDArray | None = None, + chunk_size: int = 2048, +) -> NDArray: + """Disorientation angle between every site and each of its neighbors. + + Parameters + ---------- + rotation : ndarray + ``(N, 3, 3)`` site orientations (lab <- crystal). + neighbor_index : ndarray + ``(N, K)`` neighbor indices; ``-1`` marks missing neighbors. + symmetry : ndarray, optional + ``(S, 3, 3)`` proper symmetry rotations of the crystal; the minimum + angle over all symmetry-equivalent descriptions is returned. + chunk_size : int + Sites per batch. + + Returns + ------- + ndarray + ``(N, K)`` angles in degrees, ``nan`` for missing neighbors. + """ + rotation = np.asarray(rotation, dtype=np.float64) + n, k = neighbor_index.shape + if symmetry is None: + symmetry = np.eye(3)[None] + out = np.full((n, k), np.nan) + idx_safe = np.where(neighbor_index >= 0, neighbor_index, 0) + for start in range(0, n, chunk_size): + sl = slice(start, min(start + chunk_size, n)) + ra = rotation[sl] # (n,3,3) + rb = rotation[idx_safe[sl]] # (n,k,3,3) + delta = np.einsum("nji,nkjl->nkil", ra, rb) # R_a^T R_b + trace = np.einsum("nkij,sji->nks", delta, symmetry).max(-1) + ang = np.degrees(np.arccos(np.clip((trace - 1.0) / 2.0, -1.0, 1.0))) + out[sl] = ang + out[neighbor_index < 0] = np.nan + return out + + +def segment_grains( + neighbor_index: NDArray, + edge_mask: NDArray, + min_size: int = 1, +) -> NDArray: + """Label connected components of a neighbor graph. + + Parameters + ---------- + neighbor_index : ndarray + ``(N, K)`` neighbor indices (``-1`` = missing). + edge_mask : ndarray + ``(N, K)`` boolean; ``True`` where site ``n`` and neighbor ``k`` belong + to the same grain (e.g. same structure and small misorientation). + min_size : int + Components smaller than this are labelled ``-1``. + + Returns + ------- + ndarray + ``(N,)`` integer labels ordered by decreasing size (0 = largest), + ``-1`` for sites in components smaller than ``min_size``. + """ + n, k = neighbor_index.shape + rows = np.repeat(np.arange(n), k) + cols = neighbor_index.ravel() + keep = edge_mask.ravel() & (cols >= 0) + graph = coo_matrix((np.ones(keep.sum()), (rows[keep], cols[keep])), shape=(n, n)) + _, labels = connected_components(graph, directed=False) + sizes = np.bincount(labels) + order = np.argsort(-sizes, kind="stable") + rank = np.empty_like(order) + rank[order] = np.arange(order.size) + out = rank[labels] + out[sizes[labels] < min_size] = -1 + return out + + +def strain_from_deformation(deformation: NDArray, rotation: NDArray, frame: str = "lab") -> dict: + """Small-strain tensor components from a deformation gradient. + + With ``p ~ F t`` and polar decomposition ``F = V R = R U`` the stretch in + the lab frame is ``V`` and in the crystal (template) frame ``U``. The + strain is ``sym(V) - I`` or ``sym(U) - I``. + + Parameters + ---------- + deformation : ndarray + ``(N, 3, 3)`` deformation gradients ``F``. + rotation : ndarray + ``(N, 3, 3)`` rotations ``R`` from the rigid fit. + frame : {"lab", "crystal"} + Frame in which to express the strain. + + Returns + ------- + dict of ndarray + ``e_xx, e_yy, e_zz, e_xy, e_xz, e_yz`` components, ``dilation`` + (mean normal strain) and ``equivalent`` (von Mises deviatoric strain), + plus the full ``(N, 3, 3)`` tensor under ``"tensor"``. + """ + f = np.asarray(deformation, dtype=np.float64) + r = np.asarray(rotation, dtype=np.float64) + if frame == "lab": + stretch = f @ np.transpose(r, (0, 2, 1)) + elif frame == "crystal": + stretch = np.transpose(r, (0, 2, 1)) @ f + else: + raise ValueError("frame must be 'lab' or 'crystal'") + e = 0.5 * (stretch + np.transpose(stretch, (0, 2, 1))) - np.eye(3)[None] + dil = np.trace(e, axis1=1, axis2=2) / 3.0 + dev = e - dil[:, None, None] * np.eye(3)[None] + equivalent = np.sqrt((2.0 / 3.0) * np.einsum("nij,nij->n", dev, dev)) + return { + "e_xx": e[:, 0, 0], + "e_yy": e[:, 1, 1], + "e_zz": e[:, 2, 2], + "e_xy": e[:, 0, 1], + "e_xz": e[:, 0, 2], + "e_yz": e[:, 1, 2], + "dilation": dil, + "equivalent": equivalent, + "tensor": e, + } + + +def bond_angles(dxyz: NDArray, valid: NDArray) -> NDArray: + """All bond angles at each site between pairs of valid neighbor vectors. + + Parameters + ---------- + dxyz : ndarray + ``(N, K, 3)`` neighbor vectors. + valid : ndarray + ``(N, K)`` mask of first-shell neighbors. + + Returns + ------- + ndarray + ``(N, K*(K-1)/2)`` angles in degrees, ``nan`` where either neighbor is + invalid. + """ + n, k, _ = dxyz.shape + unit = dxyz / np.linalg.norm(dxyz, axis=-1, keepdims=True).clip(1e-12) + iu, ju = np.triu_indices(k, 1) + cos = np.einsum("nkd,nkd->nk", unit[:, iu], unit[:, ju]) + ang = np.degrees(np.arccos(np.clip(cos, -1, 1))) + ang[~(valid[:, iu] & valid[:, ju])] = np.nan + return ang + + +def convex_hull_distance(xyz: NDArray) -> NDArray: + """Signed distance from each point to the convex hull surface (positive inside).""" + xyz = np.asarray(xyz, dtype=float) + hull = ConvexHull(xyz) + eq = hull.equations # (F, 4): n . x + d = 0, outward normals + d = -(xyz @ eq[:, :3].T + eq[:, 3][None, :]) + return d.min(axis=1) + + +def sample_volume(volume: NDArray, xyz: NDArray, radius: float = 1.5) -> NDArray: + """Mean volume intensity inside a sphere around each site (voxel units). + + Parameters + ---------- + volume : ndarray + 3D array indexed ``[x, y, z]`` matching the coordinate order of ``xyz``. + xyz : ndarray + ``(N, 3)`` site positions in voxel coordinates. + radius : float + Integration sphere radius in voxels. + + Returns + ------- + ndarray + ``(N,)`` mean intensity; sites whose sphere leaves the volume use only + the in-bounds voxels. + """ + volume = np.asarray(volume) + xyz = np.asarray(xyz, dtype=float) + r = int(np.ceil(radius)) + rng = np.arange(-r, r + 1) + off = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + off = off[np.linalg.norm(off, axis=1) <= radius] + center = np.rint(xyz).astype(int) + idx = center[:, None, :] + off[None, :, :] # (N, V, 3) + shape = np.array(volume.shape) + inside = np.all((idx >= 0) & (idx < shape), axis=-1) + idx = np.clip(idx, 0, shape - 1) + vals = volume[idx[..., 0], idx[..., 1], idx[..., 2]].astype(float) + vals[~inside] = 0.0 + counts = inside.sum(1).clip(1) + return vals.sum(1) / counts + + +def kmeans_1d( + values: NDArray, num_clusters: int = 2, num_iter: int = 50 +) -> tuple[NDArray, NDArray]: + """Simple 1D k-means with quantile initialization. + + Returns + ------- + labels, centers : ndarray + ``(N,)`` labels sorted so that cluster 0 has the smallest center, and + ``(num_clusters,)`` sorted centers. + """ + v = np.asarray(values, dtype=float) + finite = np.isfinite(v) + q = np.linspace(0, 1, num_clusters + 2)[1:-1] + centers = np.quantile(v[finite], q) + labels = np.zeros(v.shape, dtype=int) + for _ in range(num_iter): + labels = np.argmin(np.abs(v[:, None] - centers[None, :]), axis=1) + new = np.array( + [ + v[finite & (labels == c)].mean() if np.any(finite & (labels == c)) else centers[c] + for c in range(num_clusters) + ] + ) + if np.allclose(new, centers): + break + centers = new + order = np.argsort(centers) + remap = np.empty_like(order) + remap[order] = np.arange(num_clusters) + return remap[labels], centers[order] + + +def rotation_to_quaternion(rotation: NDArray) -> NDArray: + """Convert ``(N, 3, 3)`` rotation matrices to ``(N, 4)`` unit quaternions (w, x, y, z).""" + from scipy.spatial.transform import Rotation + + q = Rotation.from_matrix(np.asarray(rotation)).as_quat() # x, y, z, w + q = np.roll(q, 1, axis=1) + q[q[:, 0] < 0] *= -1 + return q diff --git a/src/quantem/atoms/pdf.py b/src/quantem/atoms/pdf.py new file mode 100644 index 000000000..7b2375af2 --- /dev/null +++ b/src/quantem/atoms/pdf.py @@ -0,0 +1,216 @@ +"""Pair distribution functions, neighbor lists and bond-length calibration. + +The radial distribution function (RDF) of a 3D atomic model is computed with a +KD-tree cumulative pair count, which is fast (tens of milliseconds for tens of +thousands of atoms) and needs no pair list. The first RDF peak is fit with an +asymmetric generalized Gaussian so the mean nearest-neighbor (NN) distance and a +first-shell cutoff can be estimated robustly, following the approach of the +original MATLAB cluster analysis scripts. + +Calibration maps the measured NN distance to the NN distance of a reference +crystal, giving the physical size of one voxel. +""" + +from __future__ import annotations + +import numpy as np +from numpy.typing import NDArray +from scipy.ndimage import gaussian_filter1d +from scipy.optimize import curve_fit +from scipy.spatial import cKDTree + +__all__ = [ + "radial_distribution", + "fit_first_peak", + "find_neighbors", + "nn_distance_from_lattice", +] + + +def radial_distribution( + xyz: NDArray, + r_max: float, + dr: float, + sigma: float | None = None, +) -> dict[str, NDArray]: + """Compute the radial distribution function of a set of 3D points. + + Parameters + ---------- + xyz : ndarray + ``(N, 3)`` coordinates. + r_max : float + Maximum radius. + dr : float + Radial bin width. + sigma : float, optional + Gaussian smoothing width (same units as ``r``). Default ``2 * dr``. + + Returns + ------- + dict + ``r`` bin centers, ``counts`` pair counts per bin (each unordered pair + counted once), ``g`` the RDF normalized by the shell volume and mean + number density, and ``g_smooth`` a Gaussian-smoothed copy of ``g``. + """ + xyz = np.asarray(xyz, dtype=float) + tree = cKDTree(xyz) + edges = np.arange(0.0, r_max + dr, dr) + pairs = tree.query_pairs(r_max, output_type="ndarray") + if pairs.shape[0]: + dist = np.linalg.norm(xyz[pairs[:, 0]] - xyz[pairs[:, 1]], axis=1) + counts = np.histogram(dist, bins=edges)[0].astype(float) + else: + counts = np.zeros(edges.size - 1) + r = 0.5 * (edges[1:] + edges[:-1]) + shell_volume = 4.0 * np.pi * r**2 * dr + extent = xyz.max(0) - xyz.min(0) + density = xyz.shape[0] / max(float(np.prod(extent)), 1e-12) + g = counts / (shell_volume * density * xyz.shape[0] / 2.0) + if sigma is None: + sigma = 2.0 * dr + g_smooth = gaussian_filter1d(g, sigma / dr, mode="nearest") if sigma > 0 else g.copy() + return {"r": r, "counts": counts, "g": g, "g_smooth": g_smooth} + + +def _asymmetric_peak(r, amp, r0, p_lo, p_hi, w_lo, w_hi): + lo = amp * np.exp(-((np.abs(r - r0) / w_lo) ** p_lo)) + hi = amp * np.exp(-((np.abs(r - r0) / w_hi) ** p_hi)) + return np.where(r < r0, lo, hi) + + +def fit_first_peak( + r: NDArray, + g: NDArray, + fit_radius: float = 1.25, + cutoff_sigma: float = 2.0, + r_min: float | None = None, +) -> dict[str, float | NDArray]: + """Fit the first RDF peak with an asymmetric generalized Gaussian. + + Parameters + ---------- + r, g : ndarray + Radial bins and (smoothed) RDF. + fit_radius : float + Fit window extends from ``r_min`` to ``fit_radius * r_peak``. + cutoff_sigma : float + First-shell cutoffs are placed ``cutoff_sigma`` widths below / above + the peak position. + r_min : float, optional + Ignore the RDF below this radius (e.g. tracing artifacts). Defaults to + one third of the peak radius. + + Returns + ------- + dict + ``r_nn`` peak position, ``width_lo``/``width_hi`` asymmetric widths, + ``cutoff`` = ``(lo, hi)`` first-shell radial cutoffs, ``fit`` the + fitted curve on ``r`` and ``coefs`` the raw fit coefficients. + """ + r = np.asarray(r, dtype=float) + g = np.asarray(g, dtype=float) + p = np.where(r > 0, g / np.maximum(r, 1e-12) ** 2, 0.0) + # first significant peak: first local maximum of g above 50% of global max + i_max = int(np.argmax(g)) + peaks = np.where((g[1:-1] > g[:-2]) & (g[1:-1] >= g[2:]) & (g[1:-1] > 0.5 * g[i_max]))[0] + 1 + i_peak = int(peaks[0]) if peaks.size else i_max + del p + r0 = float(r[i_peak]) + amp0 = float(g[i_peak]) + if r_min is None: + r_min = r0 / 3.0 + sub = (r >= r_min) & (r <= fit_radius * r0) + w0 = 0.15 * r0 + coefs0 = [amp0, r0, 2.0, 2.0, w0, w0] + lower = [0, r_min, 1.0, 1.0, 1e-3 * r0, 1e-3 * r0] + upper = [np.inf, fit_radius * r0, 8.0, 8.0, r0, r0] + try: + coefs, _ = curve_fit( + _asymmetric_peak, r[sub], g[sub], p0=coefs0, bounds=(lower, upper), maxfev=20000 + ) + except (RuntimeError, ValueError): + coefs = np.array(coefs0) + amp, r_nn, _, _, w_lo, w_hi = coefs + fit = _asymmetric_peak(r, *coefs) + return { + "r_nn": float(r_nn), + "amplitude": float(amp), + "width_lo": float(w_lo), + "width_hi": float(w_hi), + "cutoff": (float(r_nn - cutoff_sigma * w_lo), float(r_nn + cutoff_sigma * w_hi)), + "fit": fit, + "coefs": np.asarray(coefs), + } + + +def find_neighbors(xyz: NDArray, num_neighbors: int) -> tuple[NDArray, NDArray]: + """Return the ``num_neighbors`` nearest neighbors of every point. + + Parameters + ---------- + xyz : ndarray + ``(N, 3)`` coordinates. + num_neighbors : int + Neighbors per point (self excluded). + + Returns + ------- + distances, indices : ndarray + ``(N, num_neighbors)`` arrays sorted by distance. If fewer than + ``num_neighbors`` points exist, missing entries have index ``-1`` and + infinite distance. + """ + xyz = np.asarray(xyz, dtype=float) + n = xyz.shape[0] + k = min(num_neighbors + 1, n) + tree = cKDTree(xyz) + dist, idx = tree.query(xyz, k=k) + # drop the query point itself (not necessarily in column 0 when points coincide) + is_self = idx == np.arange(n)[:, None] + dist = np.where(is_self, -np.inf, dist) + order = np.argsort(dist, axis=1, kind="stable") + dist = np.take_along_axis(dist, order, axis=1) + idx = np.take_along_axis(idx, order, axis=1) + has_self = is_self.any(axis=1) + dist = np.where(has_self[:, None], dist[:, 1:], dist[:, : k - 1]) + idx = np.where(has_self[:, None], idx[:, 1:], idx[:, : k - 1]) + if k - 1 < num_neighbors: + pad = num_neighbors - (k - 1) + dist = np.pad(dist, ((0, 0), (0, pad)), constant_values=np.inf) + idx = np.pad(idx, ((0, 0), (0, pad)), constant_values=-1) + return dist, idx + + +_NN_FACTORS = { + "fcc": 1.0 / np.sqrt(2.0), + "bcc": np.sqrt(3.0) / 2.0, + "sc": 1.0, + "hcp": 1.0, + "diamond": np.sqrt(3.0) / 4.0, + "zincblende": np.sqrt(3.0) / 4.0, + "wurtzite": np.sqrt(3.0 / 8.0) * np.sqrt(8.0 / 3.0) * 3.0 / 8.0, # u*c with c = 1.633a +} + + +def nn_distance_from_lattice(structure: str, lattice_constant: float) -> float: + """Nearest-neighbor distance of a reference crystal. + + Parameters + ---------- + structure : str + ``"fcc"``, ``"bcc"``, ``"sc"``, ``"hcp"``, ``"diamond"``, ``"zincblende"`` + or ``"wurtzite"``. For the hexagonal structures ``lattice_constant`` + is ``a`` and ideal ``c/a`` is assumed. + lattice_constant : float + Cubic lattice constant ``a`` (or hexagonal ``a``). + + Returns + ------- + float + Nearest-neighbor distance in the units of ``lattice_constant``. + """ + key = structure.lower() + if key not in _NN_FACTORS: + raise ValueError(f"Unknown structure {structure!r}; choose from {list(_NN_FACTORS)}") + return float(_NN_FACTORS[key] * lattice_constant) diff --git a/src/quantem/atoms/show_atoms.py b/src/quantem/atoms/show_atoms.py new file mode 100644 index 000000000..0abf9d41b --- /dev/null +++ b/src/quantem/atoms/show_atoms.py @@ -0,0 +1,619 @@ +"""Interactive 3D viewer for :class:`~quantem.atoms.AtomicModel`. + +:class:`ShowAtoms3D` is a self-contained `anywidget `_ +(no JavaScript build step) that renders all sites of a model in a notebook: + +* drag to rotate, scroll to zoom, shift-drag (or right-drag) to pan, + double-click to reset; +* a **channel** dropdown colors sites by any per-site channel, with a live + histogram whose two handles set the color range; +* **clip** controls cut a slab along the view direction or a model axis so + flat atomic planes can be isolated, and "hide outside range" keeps only sites + whose channel value lies inside the histogram range (e.g. only twin sites); +* the current view is synced back to Python (:attr:`ShowAtoms3D.view_matrix`) + so the same projection can be reproduced with ``model.plot("slab", normal=...)``. + +The widget is optional: it requires the ``anywidget`` package. +""" + +from __future__ import annotations + +import json +from typing import Any + +import numpy as np +from numpy.typing import NDArray + +try: + import anywidget + import traitlets +except ImportError as exc: # pragma: no cover - optional dependency + raise ImportError( + "ShowAtoms3D requires the 'anywidget' package: pip install anywidget" + ) from exc + +from quantem.atoms.visualization import _CATEGORICAL_COLORS + +__all__ = ["ShowAtoms3D"] + +_DEFAULT_CMAPS = [ + "viridis", + "turbo", + "inferno", + "magma", + "plasma", + "cividis", + "coolwarm", + "RdBu_r", + "gray", +] + +_ESM = r""" +function mat3mul(a, b) { + const o = new Float64Array(9); + for (let i = 0; i < 3; i++) for (let j = 0; j < 3; j++) { + o[3*i+j] = a[3*i]*b[j] + a[3*i+1]*b[3+j] + a[3*i+2]*b[6+j]; + } + return o; +} +function rotX(t){const c=Math.cos(t),s=Math.sin(t);return [1,0,0,0,c,-s,0,s,c];} +function rotY(t){const c=Math.cos(t),s=Math.sin(t);return [c,0,s,0,1,0,-s,0,c];} +function orthonormalize(m){ + let r=[m[0],m[1],m[2]], u=[m[3],m[4],m[5]]; + const nr=Math.hypot(...r); r=r.map(v=>v/nr); + const d=u[0]*r[0]+u[1]*r[1]+u[2]*r[2]; u=[u[0]-d*r[0],u[1]-d*r[1],u[2]-d*r[2]]; + const nu=Math.hypot(...u); u=u.map(v=>v/nu); + const n=[r[1]*u[2]-r[2]*u[1], r[2]*u[0]-r[0]*u[2], r[0]*u[1]-r[1]*u[0]]; + return [...r,...u,...n]; +} + +function render({ model, el }) { + // ---------------- data ---------------- + const N = model.get("num_sites"); + const pos = new Float32Array(model.get("positions").buffer.slice(0)); + const channelNames = model.get("channel_names"); + const chanData = new Float32Array(model.get("channel_data").buffer.slice(0)); + const categories = model.get("channel_categories"); + const palettes = model.get("channel_palettes"); + const cmapNames = model.get("cmap_names"); + const luts = new Uint8Array(model.get("cmap_luts").buffer.slice(0)); + const radius0 = model.get("model_radius"); + const units = model.get("units"); + + const proj = new Float32Array(N * 3); + const order = new Uint32Array(N); + const depthKey = new Float32Array(N); + const visible = new Uint8Array(N); + let spriteCache = new Map(); + let spriteRadiusPx = -1; + let dragging = false, dragMode = 0, lastX = 0, lastY = 0, moved = false; + let hist = null, histDrag = 0; + + // ---------------- layout ---------------- + el.innerHTML = ""; + const root = document.createElement("div"); root.className = "qa-root"; el.appendChild(root); + const left = document.createElement("div"); left.className = "qa-left"; root.appendChild(left); + const title = document.createElement("div"); title.className = "qa-title"; left.appendChild(title); + const canvas = document.createElement("canvas"); canvas.className = "qa-canvas"; left.appendChild(canvas); + const status = document.createElement("div"); status.className = "qa-status"; left.appendChild(status); + const panel = document.createElement("div"); panel.className = "qa-panel"; root.appendChild(panel); + const ctx = canvas.getContext("2d"); + + function row(label) { + const r = document.createElement("div"); r.className = "qa-row"; + const l = document.createElement("label"); l.textContent = label; r.appendChild(l); + panel.appendChild(r); return r; + } + function select(label, options, key, fmt) { + const r = row(label); const s = document.createElement("select"); + options.forEach(o => { const op = document.createElement("option"); op.value = o; op.textContent = fmt ? fmt(o) : o; s.appendChild(op); }); + s.value = model.get(key); + s.onchange = () => { model.set(key, s.value); model.save_changes(); }; + r.appendChild(s); return s; + } + function slider(label, key, min, max, step, live) { + const r = row(label); const s = document.createElement("input"); s.type = "range"; + s.min = min; s.max = max; s.step = step; s.value = model.get(key); + const v = document.createElement("span"); v.className = "qa-val"; v.textContent = Number(model.get(key)).toFixed(2); + s.oninput = () => { v.textContent = Number(s.value).toFixed(2); model.set(key, Number(s.value)); if (live) draw(); }; + s.onchange = () => { model.set(key, Number(s.value)); model.save_changes(); }; + r.appendChild(s); r.appendChild(v); return { s, v }; + } + function checkbox(label, key) { + const r = row(label); const c = document.createElement("input"); c.type = "checkbox"; c.checked = model.get(key); + c.onchange = () => { model.set(key, c.checked); model.save_changes(); }; + r.appendChild(c); return c; + } + function button(parent, text, fn) { + const b = document.createElement("button"); b.className = "qa-btn"; b.textContent = text; b.onclick = fn; parent.appendChild(b); return b; + } + + const channelSel = select("channel", channelNames, "channel"); + const cmapSel = select("colormap", cmapNames, "cmap"); + const histCanvas = document.createElement("canvas"); histCanvas.className = "qa-hist"; histCanvas.width = 240; histCanvas.height = 96; panel.appendChild(histCanvas); + const hctx = histCanvas.getContext("2d"); + const rangeRow = row("range"); + const vminIn = document.createElement("input"); vminIn.type = "number"; vminIn.className = "qa-num"; vminIn.step = "any"; + const vmaxIn = document.createElement("input"); vmaxIn.type = "number"; vmaxIn.className = "qa-num"; vmaxIn.step = "any"; + rangeRow.appendChild(vminIn); rangeRow.appendChild(vmaxIn); + vminIn.onchange = () => { model.set("vmin", Number(vminIn.value)); model.set("range_auto", false); model.save_changes(); }; + vmaxIn.onchange = () => { model.set("vmax", Number(vmaxIn.value)); model.set("range_auto", false); model.save_changes(); }; + const rr = row(""); button(rr, "auto range", () => { model.set("range_auto", true); model.save_changes(); autoRange(); }); + const hideChk = checkbox("hide outside range", "hide_outside_range"); + const legend = document.createElement("div"); legend.className = "qa-legend"; panel.appendChild(legend); + + const sep1 = document.createElement("div"); sep1.className = "qa-sep"; sep1.textContent = "clip"; panel.appendChild(sep1); + const clipChk = checkbox("enable clip", "clip_enabled"); + const clipAxis = select("axis", ["view", "x", "y", "z"], "clip_axis"); + const clipCenter = slider("center", "clip_center", -radius0, radius0, radius0 / 400, true); + const clipThick = slider("thickness", "clip_thickness", 0.1, 2 * radius0, radius0 / 400, true); + + const sep2 = document.createElement("div"); sep2.className = "qa-sep"; sep2.textContent = "display"; panel.appendChild(sep2); + const sizeSl = slider("marker", "marker_size", 0.05, 2.0, 0.01, true); + const cueSl = slider("depth cue", "depth_cue", 0, 1, 0.01, true); + const edgeChk = checkbox("edges", "show_edges"); + const bgChk = checkbox("dark background", "dark_background"); + const viewRow = row("view"); + button(viewRow, "x", () => setView([0,1,0, 0,0,1, 1,0,0])); + button(viewRow, "y", () => setView([0,0,1, 1,0,0, 0,1,0])); + button(viewRow, "z", () => setView([1,0,0, 0,1,0, 0,0,1])); + button(viewRow, "reset", () => { setView([1,0,0, 0,1,0, 0,0,1]); model.set("zoom", 1.0); model.set("pan", [0, 0]); model.save_changes(); }); + const info = document.createElement("div"); info.className = "qa-info"; panel.appendChild(info); + + function setView(m) { model.set("rotation", Array.from(m)); model.save_changes(); draw(); } + + // ---------------- channel helpers ---------------- + function channelIndex() { const i = channelNames.indexOf(model.get("channel")); return i < 0 ? 0 : i; } + function channelValues() { const i = channelIndex(); return chanData.subarray(i * N, (i + 1) * N); } + function isCategorical() { return categories[model.get("channel")] !== undefined; } + function autoRange() { + const v = channelValues(); const name = model.get("channel"); + if (isCategorical()) { model.set("vmin", -0.5); model.set("vmax", categories[name].length - 0.5); model.save_changes(); return; } + const arr = Array.from(v).filter(Number.isFinite).sort((a, b) => a - b); + if (!arr.length) return; + const lo = arr[Math.floor(0.01 * (arr.length - 1))], hi = arr[Math.floor(0.99 * (arr.length - 1))]; + model.set("vmin", lo); model.set("vmax", hi === lo ? lo + 1e-9 : hi); model.save_changes(); + } + function computeHist() { + const v = channelValues(); const nb = 64; + let lo, hi; + if (isCategorical()) { lo = -1.5; hi = categories[model.get("channel")].length - 0.5; } + else { + let mn = Infinity, mx = -Infinity; + for (let i = 0; i < N; i++) { const x = v[i]; if (Number.isFinite(x)) { if (x < mn) mn = x; if (x > mx) mx = x; } } + lo = mn; hi = mx === mn ? mn + 1e-9 : mx; + } + const counts = new Float64Array(nb); + for (let i = 0; i < N; i++) { const x = v[i]; if (!Number.isFinite(x)) continue; let b = Math.floor((x - lo) / (hi - lo) * nb); if (b >= nb) b = nb - 1; if (b < 0) b = 0; counts[b]++; } + hist = { lo, hi, counts, nb }; + } + function lutColor(t) { + const ci = Math.max(0, cmapNames.indexOf(model.get("cmap"))); + const k = Math.max(0, Math.min(255, Math.round(t * 255))); + const o = (ci * 256 + k) * 3; return [luts[o], luts[o + 1], luts[o + 2]]; + } + function drawHist() { + if (!hist) computeHist(); + const W = histCanvas.width, H = histCanvas.height, pad = 4, hh = H - 22; + const dark = model.get("dark_background"); + hctx.fillStyle = dark ? "#1e1e1e" : "#f4f4f4"; hctx.fillRect(0, 0, W, H); + const vmin = model.get("vmin"), vmax = model.get("vmax"); + const cmax = Math.max(1, ...hist.counts); + const cat = isCategorical(); const name = model.get("channel"); + for (let b = 0; b < hist.nb; b++) { + const x0 = pad + (W - 2 * pad) * b / hist.nb, w = (W - 2 * pad) / hist.nb; + const val = hist.lo + (b + 0.5) / hist.nb * (hist.hi - hist.lo); + let t = (val - vmin) / (vmax - vmin); + let col; + if (cat) { const code = Math.round(val); const p = palettes[name]; col = code < 0 || !p ? [140,140,140] : p.slice(3*(code % (p.length/3)), 3*(code % (p.length/3))+3); } + else col = lutColor(Math.max(0, Math.min(1, t))); + const h = Math.log1p(hist.counts[b]) / Math.log1p(cmax) * hh; + hctx.fillStyle = `rgb(${col[0]},${col[1]},${col[2]})`; + hctx.globalAlpha = (val < vmin || val > vmax) ? 0.3 : 1.0; + hctx.fillRect(x0, pad + hh - h, Math.max(w - 1, 1), h); + } + hctx.globalAlpha = 1; + // colorbar + for (let i = 0; i < W - 2 * pad; i++) { + const val = hist.lo + i / (W - 2 * pad) * (hist.hi - hist.lo); + const t = Math.max(0, Math.min(1, (val - vmin) / (vmax - vmin))); + let col; + if (cat) { const code = Math.round(val); const p = palettes[name]; col = code < 0 || !p ? [140,140,140] : p.slice(3*(code % (p.length/3)), 3*(code % (p.length/3))+3); } + else col = lutColor(t); + hctx.fillStyle = `rgb(${col[0]},${col[1]},${col[2]})`; hctx.fillRect(pad + i, pad + hh + 4, 1, 8); + } + // handles + const xs = [vmin, vmax].map(v => pad + (W - 2 * pad) * (v - hist.lo) / (hist.hi - hist.lo)); + hctx.strokeStyle = dark ? "#fff" : "#000"; hctx.lineWidth = 1.5; + xs.forEach(x => { hctx.beginPath(); hctx.moveTo(x, pad); hctx.lineTo(x, H - 2); hctx.stroke(); }); + vminIn.value = Number(vmin.toPrecision(5)); vmaxIn.value = Number(vmax.toPrecision(5)); + // legend + legend.innerHTML = ""; + if (cat) { + const p = palettes[name] || []; + categories[name].forEach((lab, i) => { + const d = document.createElement("div"); d.className = "qa-legend-item"; + const sw = document.createElement("span"); sw.className = "qa-swatch"; sw.style.background = `rgb(${p[3*i]},${p[3*i+1]},${p[3*i+2]})`; + d.appendChild(sw); d.appendChild(document.createTextNode(`${i}: ${lab}`)); legend.appendChild(d); + }); + const d = document.createElement("div"); d.className = "qa-legend-item"; + const sw = document.createElement("span"); sw.className = "qa-swatch"; sw.style.background = "rgb(140,140,140)"; + d.appendChild(sw); d.appendChild(document.createTextNode("-1: none")); legend.appendChild(d); + } + } + function histPointer(ev) { + const rect = histCanvas.getBoundingClientRect(); const pad = 4, W = histCanvas.width; + const x = (ev.clientX - rect.left) * (W / rect.width); + return hist.lo + (x - pad) / (W - 2 * pad) * (hist.hi - hist.lo); + } + histCanvas.onpointerdown = (ev) => { + if (!hist) computeHist(); + const v = histPointer(ev); const vmin = model.get("vmin"), vmax = model.get("vmax"); + histDrag = Math.abs(v - vmin) < Math.abs(v - vmax) ? 1 : 2; histCanvas.setPointerCapture(ev.pointerId); + }; + histCanvas.onpointermove = (ev) => { + if (!histDrag) return; const v = histPointer(ev); + if (histDrag === 1) model.set("vmin", Math.min(v, model.get("vmax") - 1e-9)); else model.set("vmax", Math.max(v, model.get("vmin") + 1e-9)); + model.set("range_auto", false); drawHist(); draw(); + }; + histCanvas.onpointerup = (ev) => { histDrag = 0; histCanvas.releasePointerCapture(ev.pointerId); model.save_changes(); }; + + // ---------------- sprites ---------------- + function sprite(r, g, b, shade, radiusPx, edges) { + const key = ((r << 16) | (g << 8) | b) * 32 + shade; + let s = spriteCache.get(key); if (s) return s; + const R = Math.max(1, Math.ceil(radiusPx)); const size = 2 * R + 2; + s = document.createElement("canvas"); s.width = size; s.height = size; + const c = s.getContext("2d"); + const f = 1 - 0.85 * shade / 31; + const cx = R + 1, cy = R + 1; + const grad = c.createRadialGradient(cx - 0.35 * R, cy - 0.35 * R, 0.1 * R, cx, cy, R); + grad.addColorStop(0, `rgb(${Math.min(255, r*f+90*f)|0},${Math.min(255, g*f+90*f)|0},${Math.min(255, b*f+90*f)|0})`); + grad.addColorStop(1, `rgb(${(r*f*0.75)|0},${(g*f*0.75)|0},${(b*f*0.75)|0})`); + c.fillStyle = grad; c.beginPath(); c.arc(cx, cy, R, 0, 2 * Math.PI); c.fill(); + if (edges && R > 2) { c.strokeStyle = "rgba(0,0,0,0.6)"; c.lineWidth = Math.max(0.5, R * 0.12); c.stroke(); } + spriteCache.set(key, s); return s; + } + + // ---------------- main draw ---------------- + function draw() { + const W = model.get("canvas_size"), H = W; + if (canvas.width !== W) { canvas.width = W; canvas.height = H; } + const dark = model.get("dark_background"); + ctx.fillStyle = dark ? "#111" : "#fff"; ctx.fillRect(0, 0, W, H); + const m = model.get("rotation"); const zoom = model.get("zoom"); const pan = model.get("pan"); + const scale = zoom * (0.5 * W * 0.92) / radius0; + const cx = W / 2 + pan[0], cy = H / 2 + pan[1]; + const vals = channelValues(); const vmin = model.get("vmin"), vmax = model.get("vmax"); + const cat = isCategorical(); const name = model.get("channel"); const pal = palettes[name]; + const hide = model.get("hide_outside_range"); + const clipOn = model.get("clip_enabled"), clipAx = model.get("clip_axis"); + const cc = model.get("clip_center"), ct = model.get("clip_thickness"); + let count = 0; + let dmin = Infinity, dmax = -Infinity; + for (let i = 0; i < N; i++) { + const x = pos[3*i], y = pos[3*i+1], z = pos[3*i+2]; + const u = m[0]*x + m[1]*y + m[2]*z, v = m[3]*x + m[4]*y + m[5]*z, d = m[6]*x + m[7]*y + m[8]*z; + proj[3*i] = cx + u * scale; proj[3*i+1] = cy - v * scale; proj[3*i+2] = d; + let ok = 1; + if (clipOn) { const c = clipAx === "view" ? d : (clipAx === "x" ? x : (clipAx === "y" ? y : z)); if (Math.abs(c - cc) > ct / 2) ok = 0; } + if (ok && hide) { const val = vals[i]; if (!(val >= vmin && val <= vmax)) ok = 0; } + visible[i] = ok; + if (ok) { if (d < dmin) dmin = d; if (d > dmax) dmax = d; order[count++] = i; } + } + const sub = order.subarray(0, count); + for (let k = 0; k < count; k++) depthKey[sub[k]] = proj[3*sub[k]+2]; + sub.sort((a, b) => depthKey[a] - depthKey[b]); + const radiusPx = Math.max(0.6, model.get("marker_size") * model.get("bond_length") * 0.5 * scale); + if (radiusPx !== spriteRadiusPx) { spriteCache = new Map(); spriteRadiusPx = radiusPx; } + const edges = model.get("show_edges"); const cue = model.get("depth_cue"); + const dr = Math.max(dmax - dmin, 1e-9); + const inv = 1 / Math.max(vmax - vmin, 1e-12); + const ci = Math.max(0, cmapNames.indexOf(model.get("cmap"))); + for (let k = 0; k < count; k++) { + const i = sub[k]; const val = vals[i]; + let r, g, b; + if (cat) { const code = Math.round(val); if (code < 0 || !pal || !Number.isFinite(val)) { r = 140; g = 140; b = 140; } else { const j = 3 * (code % (pal.length / 3)); r = pal[j]; g = pal[j+1]; b = pal[j+2]; } } + else if (!Number.isFinite(val)) { r = 140; g = 140; b = 140; } + else { let t = (val - vmin) * inv; t = t < 0 ? 0 : (t > 1 ? 1 : t); const o = (ci * 256 + Math.round(t * 255)) * 3; r = luts[o]; g = luts[o+1]; b = luts[o+2]; } + const t = (proj[3*i+2] - dmin) / dr; + const shade = Math.round(cue * (1 - t) * 31); + const s = sprite(r, g, b, shade, radiusPx, edges); + ctx.drawImage(s, proj[3*i] - s.width / 2, proj[3*i+1] - s.height / 2); + } + status.textContent = `${count} / ${N} sites shown`; + } + + // ---------------- interaction ---------------- + canvas.oncontextmenu = (e) => e.preventDefault(); + canvas.onpointerdown = (ev) => { + dragging = true; moved = false; lastX = ev.clientX; lastY = ev.clientY; + dragMode = (ev.button === 2 || ev.shiftKey) ? 2 : 1; canvas.setPointerCapture(ev.pointerId); + }; + canvas.onpointermove = (ev) => { + if (!dragging) return; + const dx = ev.clientX - lastX, dy = ev.clientY - lastY; lastX = ev.clientX; lastY = ev.clientY; + if (Math.abs(dx) + Math.abs(dy) > 0) moved = true; + if (dragMode === 2) { const p = model.get("pan"); model.set("pan", [p[0] + dx, p[1] + dy]); } + else { + const m = model.get("rotation"); const k = 0.008; + const rs = mat3mul(rotX(-dy * k), rotY(dx * k)); + model.set("rotation", Array.from(orthonormalize(mat3mul(rs, m)))); + } + draw(); + }; + canvas.onpointerup = (ev) => { + dragging = false; canvas.releasePointerCapture(ev.pointerId); + if (!moved) pick(ev); model.save_changes(); + }; + canvas.onwheel = (ev) => { ev.preventDefault(); const z = model.get("zoom") * Math.exp(-ev.deltaY * 0.0015); model.set("zoom", Math.max(0.05, Math.min(50, z))); draw(); model.save_changes(); }; + canvas.ondblclick = () => { model.set("zoom", 1.0); model.set("pan", [0, 0]); model.save_changes(); draw(); }; + function pick(ev) { + const rect = canvas.getBoundingClientRect(); + const x = (ev.clientX - rect.left) * canvas.width / rect.width, y = (ev.clientY - rect.top) * canvas.height / rect.height; + let best = -1, bd = 1e9; + for (let i = 0; i < N; i++) { if (!visible[i]) continue; const dx = proj[3*i] - x, dy = proj[3*i+1] - y; const d2 = dx*dx + dy*dy - 0.02 * proj[3*i+2]; if (d2 < bd) { bd = d2; best = i; } } + if (best >= 0 && Math.sqrt(bd + 1) < 20) { + model.set("selected", best); model.save_changes(); + const v = channelValues()[best]; const name = model.get("channel"); + const lab = isCategorical() ? (v >= 0 ? categories[name][Math.round(v)] : "none") : v.toPrecision(4); + info.textContent = `site ${best}: ${name} = ${lab} (${pos[3*best].toFixed(2)}, ${pos[3*best+1].toFixed(2)}, ${pos[3*best+2].toFixed(2)}) ${units}`; + } + } + + // ---------------- model listeners ---------------- + model.on("change:channel", () => { channelSel.value = model.get("channel"); hist = null; if (model.get("range_auto")) autoRange(); drawHist(); draw(); }); + model.on("change:cmap", () => { cmapSel.value = model.get("cmap"); spriteCache = new Map(); drawHist(); draw(); }); + ["vmin", "vmax"].forEach(k => model.on("change:" + k, () => { drawHist(); draw(); })); + ["rotation", "zoom", "pan", "marker_size", "depth_cue", "show_edges", "clip_enabled", "clip_axis", "clip_center", "clip_thickness", "hide_outside_range", "canvas_size"].forEach(k => model.on("change:" + k, () => { + if (k === "clip_axis") clipAxis.value = model.get(k); + if (k === "clip_enabled") clipChk.checked = model.get(k); + if (k === "hide_outside_range") hideChk.checked = model.get(k); + if (k === "show_edges") edgeChk.checked = model.get(k); + if (k === "clip_center") { clipCenter.s.value = model.get(k); clipCenter.v.textContent = Number(model.get(k)).toFixed(2); } + if (k === "clip_thickness") { clipThick.s.value = model.get(k); clipThick.v.textContent = Number(model.get(k)).toFixed(2); } + if (k === "marker_size") { sizeSl.s.value = model.get(k); sizeSl.v.textContent = Number(model.get(k)).toFixed(2); } + if (k === "depth_cue") { cueSl.s.value = model.get(k); cueSl.v.textContent = Number(model.get(k)).toFixed(2); } + draw(); + })); + model.on("change:dark_background", () => { bgChk.checked = model.get("dark_background"); root.classList.toggle("qa-dark", model.get("dark_background")); drawHist(); draw(); }); + model.on("change:title", () => { title.textContent = model.get("title"); }); + + title.textContent = model.get("title"); + root.classList.toggle("qa-dark", model.get("dark_background")); + if (model.get("range_auto")) autoRange(); + drawHist(); draw(); +} +export default { render }; +""" + +_CSS = r""" +.qa-root { display: flex; gap: 10px; font-family: system-ui, sans-serif; font-size: 12px; color: #222; background: #fafafa; padding: 6px; border-radius: 6px; } +.qa-root.qa-dark { color: #ddd; background: #222; } +.qa-left { display: flex; flex-direction: column; gap: 4px; } +.qa-title { font-weight: 600; font-size: 13px; } +.qa-canvas { border: 1px solid #888; cursor: grab; touch-action: none; } +.qa-status, .qa-info { font-size: 11px; opacity: 0.8; min-height: 14px; } +.qa-panel { display: flex; flex-direction: column; gap: 4px; width: 250px; } +.qa-row { display: flex; align-items: center; gap: 6px; } +.qa-row label { width: 74px; flex: none; } +.qa-row select { flex: 1; min-width: 0; } +.qa-row input[type=range] { flex: 1; min-width: 0; } +.qa-val { width: 44px; text-align: right; font-variant-numeric: tabular-nums; } +.qa-num { width: 80px; } +.qa-hist { border: 1px solid #888; touch-action: none; cursor: col-resize; } +.qa-sep { margin-top: 6px; font-weight: 600; border-bottom: 1px solid #888; } +.qa-btn { padding: 2px 8px; font-size: 11px; } +.qa-legend { display: flex; flex-wrap: wrap; gap: 4px 10px; } +.qa-legend-item { display: flex; align-items: center; gap: 4px; } +.qa-swatch { width: 10px; height: 10px; border-radius: 50%; border: 1px solid #555; display: inline-block; } +""" + + +def _colormap_luts(names: list[str]) -> bytes: + import matplotlib.pyplot as plt + + out = np.zeros((len(names), 256, 3), dtype=np.uint8) + for i, n in enumerate(names): + rgba = plt.get_cmap(n)(np.linspace(0, 1, 256)) + out[i] = (rgba[:, :3] * 255).round().astype(np.uint8) + return out.tobytes() + + +class ShowAtoms3D(anywidget.AnyWidget): + """Interactive 3D site viewer for an :class:`~quantem.atoms.AtomicModel`. + + Parameters + ---------- + model : AtomicModel + Model to display; all channels are sent to the browser as float32. + channel : str, optional + Initial color channel (default: ``"structure"`` if present). + cmap : str + Initial colormap for continuous channels. + cmaps : sequence of str, optional + Colormaps offered in the dropdown (matplotlib names). + canvas_size : int + Canvas width and height in pixels. + marker_size : float + Marker diameter as a fraction of the bond length. + dark_background : bool + Dark theme. + title : str, optional + Title shown above the canvas (default: model name). + + Attributes + ---------- + rotation : list of float + Row-major ``(3, 3)`` view matrix (rows: right, up, toward-viewer). + view_matrix : ndarray + Same as ``rotation`` as an array; ``view_matrix[2]`` is the viewing + direction usable as ``model.plot("slab", normal=...)``. + selected : int + Index of the last clicked site (``-1`` if none). + """ + + _esm = _ESM + _css = _CSS + + positions = traitlets.Bytes(b"").tag(sync=True) + num_sites = traitlets.Int(0).tag(sync=True) + channel_names = traitlets.List(traitlets.Unicode()).tag(sync=True) + channel_data = traitlets.Bytes(b"").tag(sync=True) + channel_categories = traitlets.Dict().tag(sync=True) + channel_palettes = traitlets.Dict().tag(sync=True) + cmap_names = traitlets.List(traitlets.Unicode()).tag(sync=True) + cmap_luts = traitlets.Bytes(b"").tag(sync=True) + model_radius = traitlets.Float(1.0).tag(sync=True) + bond_length = traitlets.Float(1.0).tag(sync=True) + units = traitlets.Unicode("").tag(sync=True) + title = traitlets.Unicode("").tag(sync=True) + + channel = traitlets.Unicode("").tag(sync=True) + cmap = traitlets.Unicode("viridis").tag(sync=True) + vmin = traitlets.Float(0.0).tag(sync=True) + vmax = traitlets.Float(1.0).tag(sync=True) + range_auto = traitlets.Bool(True).tag(sync=True) + hide_outside_range = traitlets.Bool(False).tag(sync=True) + rotation = traitlets.List(traitlets.Float(), default_value=[1, 0, 0, 0, 1, 0, 0, 0, 1]).tag( + sync=True + ) + zoom = traitlets.Float(1.0).tag(sync=True) + pan = traitlets.List(traitlets.Float(), default_value=[0.0, 0.0]).tag(sync=True) + clip_enabled = traitlets.Bool(False).tag(sync=True) + clip_axis = traitlets.Unicode("view").tag(sync=True) + clip_center = traitlets.Float(0.0).tag(sync=True) + clip_thickness = traitlets.Float(1.0).tag(sync=True) + marker_size = traitlets.Float(0.7).tag(sync=True) + depth_cue = traitlets.Float(0.5).tag(sync=True) + show_edges = traitlets.Bool(True).tag(sync=True) + dark_background = traitlets.Bool(True).tag(sync=True) + canvas_size = traitlets.Int(640).tag(sync=True) + selected = traitlets.Int(-1).tag(sync=True) + + def __init__( + self, + model: Any, + channel: str | None = None, + cmap: str = "viridis", + cmaps: list[str] | None = None, + canvas_size: int = 640, + marker_size: float = 0.7, + dark_background: bool = True, + title: str | None = None, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self._model = model + xyz = np.asarray(model.positions, dtype=np.float32) + center = xyz.mean(0) + xyz = xyz - center + self._center = center + self.positions = np.ascontiguousarray(xyz, dtype=np.float32).tobytes() + self.num_sites = int(xyz.shape[0]) + self.model_radius = float(np.linalg.norm(xyz, axis=1).max()) if xyz.shape[0] else 1.0 + try: + self.bond_length = float(model.bond_length) + except Exception: + self.bond_length = float(self.model_radius / 20) + self.units = str(model.units) + self.title = title if title is not None else str(model.name) + self.cmap_names = list(cmaps or _DEFAULT_CMAPS) + self.cmap_luts = _colormap_luts(self.cmap_names) + self.cmap = cmap if cmap in self.cmap_names else self.cmap_names[0] + self.canvas_size = int(canvas_size) + self.marker_size = float(marker_size) + self.dark_background = bool(dark_background) + self.clip_thickness = float(3 * self.bond_length) + self.update_channels() + if channel is None: + channel = "structure" if "structure" in self.channel_names else self.channel_names[0] + self.channel = channel + + # ------------------------------------------------------------------ # + def update_channels(self) -> None: + """Re-send all channels of the model (call after adding new channels).""" + model = self._model + names = ["x", "y", "z"] + list(model.channels) + data = np.stack([model.get_channel(n) for n in names], axis=0).astype(np.float32) + cats: dict[str, list[str]] = {} + palettes: dict[str, list[int]] = {} + for name, labels in model.categories.items(): + cats[name] = list(labels) + n_cat = max( + len(labels), int(np.nanmax(model.get_channel(name))) + 1 if model.num_sites else 1 + ) + if name == "grain" or n_cat > len(_CATEGORICAL_COLORS): + rng = np.random.default_rng(0) + colors = rng.uniform(0.15, 0.95, (n_cat, 3)) + else: + colors = _CATEGORICAL_COLORS[np.arange(n_cat) % len(_CATEGORICAL_COLORS)] + palettes[name] = (colors * 255).round().astype(int).ravel().tolist() + for name in ("grain",): + if name in model.channels and name not in cats: + codes = model.get_channel(name) + n_cat = int(np.nanmax(codes)) + 1 if codes.size else 1 + cats[name] = [str(i) for i in range(n_cat)] + rng = np.random.default_rng(0) + palettes[name] = ( + (rng.uniform(0.15, 0.95, (max(n_cat, 1), 3)) * 255) + .round() + .astype(int) + .ravel() + .tolist() + ) + self.channel_names = names + self.channel_data = np.ascontiguousarray(data).tobytes() + self.channel_categories = cats + self.channel_palettes = palettes + + @property + def view_matrix(self) -> NDArray: + """``(3, 3)`` current view matrix (rows: right, up, toward-viewer).""" + return np.asarray(self.rotation, dtype=float).reshape(3, 3) + + @property + def view_direction(self) -> NDArray: + """``(3,)`` unit vector pointing from the model toward the viewer.""" + return self.view_matrix[2] + + def set_view(self, normal: str | NDArray, up: str | NDArray | None = None) -> None: + """Look along ``normal`` (as in :func:`quantem.atoms.visualization.view_matrix`).""" + from quantem.atoms.visualization import view_matrix + + self.rotation = view_matrix(normal, up).ravel().tolist() + + def set_range(self, vmin: float, vmax: float) -> None: + """Set the color range and disable auto-ranging.""" + self.range_auto = False + self.vmin, self.vmax = float(vmin), float(vmax) + + def clip( + self, axis: str = "view", center: float = 0.0, thickness: float | None = None + ) -> None: + """Enable slab clipping along ``axis`` (``"view"``, ``"x"``, ``"y"`` or ``"z"``).""" + self.clip_axis = axis + self.clip_center = float(center) + if thickness is not None: + self.clip_thickness = float(thickness) + self.clip_enabled = True + + def slab_kwargs(self) -> dict[str, Any]: + """Keyword arguments reproducing the current view with ``model.plot("slab", ...)``.""" + vm = self.view_matrix + out: dict[str, Any] = {"normal": vm[2], "up": vm[1], "channel": self.channel} + if self.clip_enabled and self.clip_axis == "view": + out.update(offset=self.clip_center, thickness=self.clip_thickness) + if not self.range_auto: + out.update(vmin=self.vmin, vmax=self.vmax) + if self.hide_outside_range: + out["hide"] = None + return out + + def __repr__(self) -> str: + return f"ShowAtoms3D({self.num_sites} sites, channel={self.channel!r}, clip={self.clip_enabled})" + + def _repr_json_state(self) -> str: # pragma: no cover - debugging aid + return json.dumps( + {k: getattr(self, k) for k in ("channel", "cmap", "vmin", "vmax", "rotation", "zoom")} + ) diff --git a/src/quantem/atoms/templates.py b/src/quantem/atoms/templates.py new file mode 100644 index 000000000..d8107a6d9 --- /dev/null +++ b/src/quantem/atoms/templates.py @@ -0,0 +1,350 @@ +"""Polyhedral templates for local structure classification of 3D atomic models. + +A template is the set of neighbor vectors around a reference site in an ideal +crystal, expressed in units of the nearest-neighbor (NN) bond length so that the +first shell has radius 1. Templates are generated from crystal definitions +(lattice vectors + basis) rather than typed by hand, so any number of neighbor +shells can be requested. + +Built-in structures +------------------- +``fcc`` + Face-centered cubic; 12 NN (cuboctahedron), 6 second neighbors at sqrt(2). +``hcp`` + Hexagonal close-packed (ideal c/a); 12 NN (anticuboctahedron). Atoms on a + coherent {111} twin plane in an FCC crystal have this local environment, + so the ``hcp`` score is also the twin-boundary detector for FCC particles. +``bcc`` + Body-centered cubic; 8 NN plus 6 second neighbors at 2/sqrt(3) = 1.155. + Two shells are used by default because the second shell is so close. +``sc`` + Simple cubic; 6 NN. +``diamond`` (alias ``zincblende``) + Cubic diamond / zincblende site; 4 NN plus 12 second neighbors at 1.633. + Two shells are used by default since 4 NN cannot distinguish it from + ``wurtzite``. +``wurtzite`` + Hexagonal diamond (ideal u = 3/8, c/a = 1.633); 4 NN plus 12 second + neighbors. Differs from ``diamond`` only in the second shell (ABAB vs ABC). +``ico`` + Icosahedral (13-atom Mackay) site; 12 NN at the vertices of an icosahedron. + +Examples +-------- +>>> from quantem.atoms.templates import get_template +>>> fcc = get_template("fcc") +>>> fcc.num_neighbors, fcc.shells +(12, (1.0,)) +>>> get_template("fcc", num_shells=2).num_neighbors +18 +""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +import numpy as np +from numpy.typing import NDArray + +__all__ = [ + "PolyhedralTemplate", + "TEMPLATE_NAMES", + "get_template", + "template_from_crystal", + "template_symmetry_rotations", +] + +_SHELL_TOL = 1e-3 + + +@dataclass(frozen=True) +class PolyhedralTemplate: + """Neighbor vectors of an ideal local environment, in NN bond-length units. + + Parameters + ---------- + name : str + Structure name, e.g. ``"fcc"``. + vectors : ndarray + ``(M, 3)`` neighbor vectors sorted by radius; first shell radius is 1. + shells : tuple of float + Radius of each neighbor shell included in ``vectors``. + shell_counts : tuple of int + Number of neighbors in each included shell. + symmetry : ndarray + ``(S, 3, 3)`` proper rotations that map the template onto itself. + Used to reduce orientations to a fundamental zone when computing + misorientations between neighboring sites. + """ + + name: str + vectors: NDArray[np.floating] + shells: tuple[float, ...] + shell_counts: tuple[int, ...] + symmetry: NDArray[np.floating] = field(repr=False, default=None) # type: ignore[assignment] + + @property + def num_neighbors(self) -> int: + """Total number of neighbor vectors in the template.""" + return int(self.vectors.shape[0]) + + @property + def max_radius(self) -> float: + """Radius of the outermost included shell (NN units).""" + return float(self.shells[-1]) + + @property + def num_symmetry(self) -> int: + """Number of proper symmetry rotations of the template.""" + return 0 if self.symmetry is None else int(self.symmetry.shape[0]) + + def __repr__(self) -> str: + shells = ", ".join(f"{r:.3f}x{n}" for r, n in zip(self.shells, self.shell_counts)) + return ( + f"PolyhedralTemplate(name={self.name!r}, num_neighbors={self.num_neighbors}, " + f"shells=[{shells}], num_symmetry={self.num_symmetry})" + ) + + +# --------------------------------------------------------------------------- # +# Crystal definitions (lattice vectors as rows; basis in fractional coordinates) +# --------------------------------------------------------------------------- # +def _hexagonal_cell(a: float, c: float) -> NDArray: + return np.array([[a, 0.0, 0.0], [-a / 2, a * np.sqrt(3) / 2, 0.0], [0.0, 0.0, c]]) + + +def _crystal_definition(name: str) -> tuple[NDArray, NDArray, int, int]: + """Return (cell, basis, site_index, default_shells) with NN distance = 1.""" + if name == "fcc": + a = np.sqrt(2.0) + cell = np.eye(3) * a + basis = np.array([[0, 0, 0], [0, 0.5, 0.5], [0.5, 0, 0.5], [0.5, 0.5, 0]]) + return cell, basis, 0, 1 + if name == "bcc": + a = 2.0 / np.sqrt(3.0) + cell = np.eye(3) * a + basis = np.array([[0, 0, 0], [0.5, 0.5, 0.5]]) + return cell, basis, 0, 2 + if name == "sc": + return np.eye(3), np.zeros((1, 3)), 0, 1 + if name == "hcp": + a = 1.0 + c = a * np.sqrt(8.0 / 3.0) + cell = _hexagonal_cell(a, c) + basis = np.array([[0, 0, 0], [1 / 3, 2 / 3, 0.5]]) + return cell, basis, 0, 1 + if name == "diamond": + a = 4.0 / np.sqrt(3.0) + cell = np.eye(3) * a + fcc = np.array([[0, 0, 0], [0, 0.5, 0.5], [0.5, 0, 0.5], [0.5, 0.5, 0]]) + basis = np.vstack([fcc, fcc + 0.25]) + return cell, basis, 0, 2 + if name == "wurtzite": + u = 3.0 / 8.0 + c_over_a = np.sqrt(8.0 / 3.0) + a = 1.0 / (u * c_over_a) # bond length u*c = 1 + c = a * c_over_a + cell = _hexagonal_cell(a, c) + basis = np.array([[0, 0, 0], [1 / 3, 2 / 3, 0.5], [0, 0, u], [1 / 3, 2 / 3, 0.5 + u]]) + return cell, basis, 0, 2 + raise ValueError(f"Unknown crystal structure {name!r}. Choose from {TEMPLATE_NAMES}.") + + +def template_from_crystal( + cell: NDArray, + basis: NDArray, + site_index: int = 0, + num_shells: int = 1, + name: str = "custom", + max_search: int = 3, +) -> PolyhedralTemplate: + """Build a template from a crystal definition by collecting neighbor shells. + + Parameters + ---------- + cell : ndarray + ``(3, 3)`` lattice vectors as rows. + basis : ndarray + ``(B, 3)`` fractional coordinates of the basis atoms. + site_index : int + Which basis atom is the reference site. + num_shells : int + Number of neighbor shells to include. + name : str + Name stored on the template. + max_search : int + Half-width (in cells) of the supercell searched for neighbors. + + Returns + ------- + PolyhedralTemplate + Neighbor vectors scaled so the first shell has radius 1. + """ + cell = np.asarray(cell, dtype=float) + basis = np.asarray(basis, dtype=float) + rng = np.arange(-max_search, max_search + 1) + ijk = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + frac = (ijk[:, None, :] + basis[None, :, :]).reshape(-1, 3) + cart = frac @ cell + origin = basis[site_index] @ cell + d = cart - origin[None, :] + r = np.linalg.norm(d, axis=1) + keep = r > 1e-8 + d, r = d[keep], r[keep] + order = np.argsort(r) + d, r = d[order], r[order] + + # group into shells + shells: list[float] = [] + counts: list[int] = [] + for radius in r: + if shells and abs(radius - shells[-1]) < _SHELL_TOL * max(1.0, radius): + counts[-1] += 1 + else: + shells.append(float(radius)) + counts.append(1) + if num_shells > len(shells): + raise ValueError(f"Only {len(shells)} shells found; increase max_search.") + n_keep = int(sum(counts[:num_shells])) + r_nn = shells[0] + vectors = d[:n_keep] / r_nn + shells_out = tuple(s / r_nn for s in shells[:num_shells]) + template = PolyhedralTemplate( + name=name, + vectors=np.ascontiguousarray(vectors), + shells=shells_out, + shell_counts=tuple(counts[:num_shells]), + ) + return _with_symmetry(template) + + +def _icosahedron_template(num_shells: int) -> PolyhedralTemplate: + if num_shells != 1: + raise ValueError("The 'ico' template only defines a single neighbor shell.") + phi = (1.0 + np.sqrt(5.0)) / 2.0 + pts = [] + for s1 in (-1, 1): + for s2 in (-1, 1): + pts.append([0.0, s1 * 1.0, s2 * phi]) + pts.append([s1 * 1.0, s2 * phi, 0.0]) + pts.append([s2 * phi, 0.0, s1 * 1.0]) + vectors = np.array(pts) + vectors /= np.linalg.norm(vectors, axis=1, keepdims=True) + template = PolyhedralTemplate(name="ico", vectors=vectors, shells=(1.0,), shell_counts=(12,)) + return _with_symmetry(template) + + +TEMPLATE_NAMES: tuple[str, ...] = ( + "fcc", + "hcp", + "bcc", + "sc", + "diamond", + "zincblende", + "wurtzite", + "ico", +) + + +def get_template(name: str, num_shells: int | None = None) -> PolyhedralTemplate: + """Return a built-in polyhedral template. + + Parameters + ---------- + name : str + One of ``TEMPLATE_NAMES``. ``"zincblende"`` is an alias of ``"diamond"``. + num_shells : int, optional + Number of neighbor shells. Defaults to 1 for close-packed structures + and 2 for ``bcc``, ``diamond`` and ``wurtzite``. + + Returns + ------- + PolyhedralTemplate + """ + key = name.lower() + if key == "zincblende": + key = "diamond" + if key == "ico": + return _icosahedron_template(1 if num_shells is None else num_shells) + cell, basis, site_index, default_shells = _crystal_definition(key) + return template_from_crystal( + cell, + basis, + site_index=site_index, + num_shells=default_shells if num_shells is None else int(num_shells), + name=key, + ) + + +# --------------------------------------------------------------------------- # +# Symmetry rotations of a template (numerical, template-agnostic) +# --------------------------------------------------------------------------- # +def _frame(u: NDArray, v: NDArray) -> NDArray: + """Right-handed orthonormal frame with e1 = u/|u| and e2 in the (u, v) plane.""" + e1 = u / np.linalg.norm(u) + v2 = v - e1 * (v @ e1) + e2 = v2 / np.linalg.norm(v2) + e3 = np.cross(e1, e2) + return np.stack([e1, e2, e3], axis=1) # columns + + +def template_symmetry_rotations(vectors: NDArray, tol: float = 1e-3) -> NDArray: + """Find all proper rotations mapping a set of vectors onto itself. + + Candidate rotations are built by mapping the frame spanned by two fixed + template vectors onto the frame spanned by every ordered pair of template + vectors with the same lengths and angle, then verified against the full set. + + Parameters + ---------- + vectors : ndarray + ``(M, 3)`` template vectors. + tol : float + Matching tolerance in the same units as ``vectors``. + + Returns + ------- + ndarray + ``(S, 3, 3)`` unique proper rotation matrices, identity first. + """ + v = np.asarray(vectors, dtype=float) + m = v.shape[0] + r = np.linalg.norm(v, axis=1) + # reference pair: vector 0 and its nearest non-collinear partner + i0 = 0 + cosines = (v @ v[i0]) / (r * r[i0]) + non_collinear = np.where(np.abs(cosines) < 1 - 1e-6)[0] + j0 = non_collinear[np.argmax(cosines[non_collinear])] + ref_frame = _frame(v[i0], v[j0]) + ref_angle = np.arccos(np.clip(cosines[j0], -1, 1)) + + rotations: list[NDArray] = [np.eye(3)] + for i in range(m): + if abs(r[i] - r[i0]) > tol: + continue + for j in range(m): + if i == j or abs(r[j] - r[j0]) > tol: + continue + ang = np.arccos(np.clip((v[i] @ v[j]) / (r[i] * r[j]), -1, 1)) + if abs(ang - ref_angle) > 1e-4: + continue + rot = _frame(v[i], v[j]) @ ref_frame.T + if np.linalg.det(rot) < 0: + continue + mapped = v @ rot.T + dist = np.linalg.norm(mapped[:, None, :] - v[None, :, :], axis=2) + if np.all(dist.min(axis=1) < tol): + if not any(np.allclose(rot, q, atol=1e-6) for q in rotations): + rotations.append(rot) + return np.stack(rotations, axis=0) + + +def _with_symmetry(template: PolyhedralTemplate) -> PolyhedralTemplate: + sym = template_symmetry_rotations(template.vectors) + return PolyhedralTemplate( + name=template.name, + vectors=template.vectors, + shells=template.shells, + shell_counts=template.shell_counts, + symmetry=sym, + ) diff --git a/src/quantem/atoms/visualization.py b/src/quantem/atoms/visualization.py new file mode 100644 index 000000000..ab31cd75a --- /dev/null +++ b/src/quantem/atoms/visualization.py @@ -0,0 +1,416 @@ +"""Static matplotlib plots for :class:`~quantem.atoms.AtomicModel`. + +Nothing here modifies model state. Functions are registered by name so they +can be called as ``model.plot(kind=...)``: + +``"pdf"`` + Radial distribution function with the first-peak fit and template shells. +``"histogram"`` + Histogram of a per-site channel. +``"slab"`` + Projection of the sites inside a slab, colored by a channel. +``"slices"`` + Grid of slab projections stepping through the model. +``"template"`` + Neighbor vectors of one site with the fitted template overlaid. +""" + +from __future__ import annotations + +from typing import Any, Callable + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.colors import ListedColormap +from numpy.typing import NDArray + +PLOT_REGISTRY: dict[str, Callable[..., Any]] = {} + +_CATEGORICAL_COLORS = np.array( + [ + [0.20, 0.45, 0.95], + [0.95, 0.60, 0.10], + [0.20, 0.70, 0.30], + [0.85, 0.20, 0.25], + [0.60, 0.35, 0.80], + [0.55, 0.35, 0.20], + [0.90, 0.45, 0.75], + [0.50, 0.50, 0.50], + [0.70, 0.75, 0.10], + [0.10, 0.75, 0.80], + ] +) + + +def _register(name: str): + def decorator(fn): + PLOT_REGISTRY[name] = fn + return fn + + return decorator + + +# --------------------------------------------------------------------------- # +# Geometry helpers +# --------------------------------------------------------------------------- # +_AXES = {"x": np.array([1.0, 0, 0]), "y": np.array([0, 1.0, 0]), "z": np.array([0, 0, 1.0])} + + +def view_matrix(normal: str | NDArray, up: str | NDArray | None = None) -> NDArray: + """Orthonormal ``(3, 3)`` matrix whose rows are (right, up, normal). + + Projecting ``(xyz - center) @ view_matrix.T`` gives image columns + ``(u, v, depth)`` with ``depth`` measured along ``normal`` (toward the + viewer). + + Parameters + ---------- + normal : str or array + Viewing direction as ``"x"``, ``"y"``, ``"z"`` or a 3-vector. + up : str or array, optional + Approximate screen-up direction; default picks the axis least aligned + with ``normal``. + """ + n = _AXES[normal] if isinstance(normal, str) else np.asarray(normal, dtype=float) + n = n / np.linalg.norm(n) + if up is None: + up_v = _AXES["y"] if abs(n[1]) < 0.9 else _AXES["x"] + if isinstance(normal, str) and normal == "y": + up_v = _AXES["x"] + else: + up_v = _AXES[up] if isinstance(up, str) else np.asarray(up, dtype=float) + u = up_v - n * (up_v @ n) + u = u / np.linalg.norm(u) + r = np.cross(u, n) + return np.stack([r, u, n], axis=0) + + +def channel_colors( + model, + channel: str, + cmap: str | None = None, + vmin: float | None = None, + vmax: float | None = None, +) -> tuple[NDArray, Any, tuple[float, float], list[str] | None]: + """Map a channel to RGB colors. + + Returns + ------- + rgb, colormap, (vmin, vmax), categories + ``(N, 3)`` colors; categories is a list of labels for categorical + channels (``-1`` codes are drawn grey), otherwise ``None``. + """ + values = model.get_channel(channel) + categories = model.categories.get(channel) + if categories is not None: + codes = values.astype(int) + n_cat = max(len(categories), int(codes.max()) + 1 if codes.size else 1) + colors = _CATEGORICAL_COLORS[np.arange(n_cat) % len(_CATEGORICAL_COLORS)] + if channel == "grain" or n_cat > len(_CATEGORICAL_COLORS): + rng = np.random.default_rng(0) + colors = rng.uniform(0.15, 0.95, (n_cat, 3)) + rgb = np.full((values.size, 3), 0.6) + ok = codes >= 0 + rgb[ok] = colors[codes[ok] % n_cat] + return rgb, ListedColormap(colors), (-0.5, n_cat - 0.5), list(categories) + finite = np.isfinite(values) + if vmin is None: + vmin = float(np.percentile(values[finite], 1)) if finite.any() else 0.0 + if vmax is None: + vmax = float(np.percentile(values[finite], 99)) if finite.any() else 1.0 + if vmax <= vmin: + vmax = vmin + 1e-9 + colormap = plt.get_cmap(cmap or ("turbo" if channel.startswith("score") else "viridis")) + t = np.clip((values - vmin) / (vmax - vmin), 0, 1) + rgb = colormap(np.nan_to_num(t))[:, :3] + rgb[~finite] = 0.6 + return rgb, colormap, (vmin, vmax), None + + +# --------------------------------------------------------------------------- # +# Plots +# --------------------------------------------------------------------------- # +@_register("pdf") +def plot_pdf( + model, show_fit: bool = True, show_templates: bool = True, ax=None, returnfig: bool = False +): + """Plot the radial distribution function (native units). + + Parameters + ---------- + show_fit : bool + Overlay the first-peak fit and first-shell cutoffs. + show_templates : bool + Mark the shell radii of matched templates (scaled by the NN distance). + """ + if model.pdf is None: + model.compute_pdf() + pdf, fit = model.pdf, model.nn_fit + if ax is None: + fig, ax = plt.subplots(figsize=(8, 4)) + else: + fig = ax.figure + ax.plot(pdf["r"], pdf["g_smooth"], color="k", lw=1.5, label="RDF") + if show_fit and fit is not None: + ax.plot(pdf["r"], fit["fit"], color="r", lw=1.5, label=f"fit r_nn = {fit['r_nn']:.3f}") + for c in fit["cutoff"]: + ax.axvline(c, color="r", ls="--", lw=1) + if show_templates and model.templates: + r_nn = model.nn_distance + for i, (name, template) in enumerate(model.templates.items()): + col = _CATEGORICAL_COLORS[i % len(_CATEGORICAL_COLORS)] + for j, s in enumerate(template.shells): + ax.axvline(s * r_nn, color=col, lw=1, alpha=0.7, label=name if j == 0 else None) + ax.set_xlabel(f"radius [{model.sites.units[0]}]") + ax.set_ylabel("g(r)") + ax.set_xlim(0, pdf["r"].max()) + ax.set_ylim(0, None) + ax.legend(loc="upper right") + ax.set_title(model.name) + return (fig, ax) if returnfig else None + + +@_register("histogram") +def plot_histogram( + model, channel: str = "score_max", bins: int = 100, ax=None, returnfig: bool = False, **kwargs +): + """Histogram of a per-site channel.""" + values = model.get_channel(channel) + values = values[np.isfinite(values)] + if ax is None: + fig, ax = plt.subplots(figsize=(6, 3.5)) + else: + fig = ax.figure + categories = model.categories.get(channel) + if categories is not None: + codes = values.astype(int) + counts = np.bincount(codes[codes >= 0], minlength=len(categories)) + ax.bar(np.arange(len(categories)), counts, color=_CATEGORICAL_COLORS[: len(categories)]) + ax.set_xticks(np.arange(len(categories)), categories) + n_none = int((codes < 0).sum()) + ax.set_title(f"{channel} ({n_none} unclassified)") + else: + ax.hist(values, bins=bins, color="0.3", **kwargs) + ax.set_xlabel(channel) + ax.set_ylabel("count") + return (fig, ax) if returnfig else None + + +def _project(model, normal, up, offset, thickness, hide=None, channel=None): + v = view_matrix(normal, up) + xyz = model.positions - model.center[None, :] + uvd = xyz @ v.T + keep = np.ones(xyz.shape[0], dtype=bool) + if thickness is not None: + keep &= np.abs(uvd[:, 2] - offset) <= thickness / 2.0 + if hide is not None and channel is not None: + vals = model.get_channel(channel) + keep &= ~((vals >= hide[0]) & (vals <= hide[1])) + return uvd, keep, v + + +@_register("slab") +def plot_slab( + model, + channel: str = "structure", + normal: str | NDArray = "z", + up: str | NDArray | None = None, + offset: float = 0.0, + thickness: float | None = None, + cmap: str | None = None, + vmin: float | None = None, + vmax: float | None = None, + marker_size: float | None = None, + depth_cue: float = 0.35, + hide: tuple[float, float] | None = None, + edgecolor: str | None = "k", + ax=None, + figsize: tuple[float, float] = (7, 7), + colorbar: bool = True, + title: str | None = None, + returnfig: bool = False, +): + """Project the sites inside a slab onto the viewing plane. + + Parameters + ---------- + channel : str + Channel used for coloring. + normal : str or array + Viewing direction (slab normal): ``"x"``, ``"y"``, ``"z"`` or a vector. + up : str or array, optional + Screen-up direction. + offset : float + Slab center along ``normal`` relative to the model center (calibrated units). + thickness : float, optional + Slab thickness (calibrated units); ``None`` shows all sites. + cmap, vmin, vmax + Color mapping for continuous channels. + marker_size : float, optional + Scatter marker area; default scales with the bond length. + depth_cue : float + Darkening of sites far from the viewer (0 = none). + hide : (lo, hi), optional + Hide sites whose channel value lies inside this range. + edgecolor : str or None + Marker edge color. + """ + uvd, keep, _ = _project(model, normal, up, offset, thickness, hide, channel) + rgb, colormap, (lo, hi), categories = channel_colors(model, channel, cmap, vmin, vmax) + uvd, rgb = uvd[keep], rgb[keep] + order = np.argsort(uvd[:, 2]) + uvd, rgb = uvd[order], rgb[order] + if depth_cue > 0 and uvd.shape[0] > 1: + d = uvd[:, 2] + t = (d - d.min()) / max(d.max() - d.min(), 1e-9) + rgb = rgb * (1 - depth_cue * (1 - t))[:, None] + if ax is None: + fig, ax = plt.subplots(figsize=figsize) + else: + fig = ax.figure + if marker_size is None: + # marker diameter ~ 0.9 bond lengths, in points, from the axes width + all_uv = (model.positions - model.center[None, :]) @ view_matrix(normal, up).T + extent = max(np.ptp(all_uv[:, 0]), np.ptp(all_uv[:, 1]), 1e-9) * 1.05 + axes_width_pt = ax.get_position().width * fig.get_size_inches()[0] * 72.0 + diameter_pt = 0.9 * model.bond_length * axes_width_pt / extent + marker_size = max(diameter_pt**2, 1.0) + sc = ax.scatter( + uvd[:, 1], + uvd[:, 0], + s=marker_size, + c=rgb, + edgecolors=edgecolor, + linewidths=0.3 if edgecolor else 0, + ) + ax.set_aspect("equal") + ax.invert_yaxis() + ax.set_xlabel(f"v [{model.units}]") + ax.set_ylabel(f"u [{model.units}]") + if title is None: + n_str = normal if isinstance(normal, str) else np.round(normal, 2) + title = f"{channel} | normal {n_str}" + if thickness is not None: + title += f" | offset {offset:.1f}, thickness {thickness:.1f}" + ax.set_title(title) + if colorbar: + if categories is not None: + from matplotlib.lines import Line2D + + handles = [ + Line2D([], [], marker="o", ls="", color=colormap(i), markeredgecolor="k", label=c) + for i, c in enumerate(categories) + ] + ax.legend(handles=handles, loc="upper right", fontsize=8) + else: + sm = plt.cm.ScalarMappable(cmap=colormap, norm=plt.Normalize(lo, hi)) + fig.colorbar(sm, ax=ax, fraction=0.04, pad=0.02, label=channel) + del sc + return (fig, ax) if returnfig else None + + +@_register("slices") +def plot_slices( + model, + channel: str = "structure", + normal: str | NDArray = "z", + num_slices: int = 6, + thickness: float | None = None, + start: float | None = None, + end: float | None = None, + ncols: int = 3, + figsize_per: float = 4.0, + returnfig: bool = False, + **kwargs, +): + """Grid of slab projections stepping along ``normal``. + + Parameters + ---------- + num_slices : int + Number of slabs. + thickness : float, optional + Slab thickness; default is the step between slabs. + start, end : float, optional + Range of slab centers along ``normal`` (relative to the model center); + default spans the model. + **kwargs + Forwarded to :func:`plot_slab`. + """ + v = view_matrix(normal, kwargs.get("up")) + depth = (model.positions - model.center[None, :]) @ v[2] + if start is None: + start = float(depth.min()) + 0.05 * np.ptp(depth) + if end is None: + end = float(depth.max()) - 0.05 * np.ptp(depth) + centers = np.linspace(start, end, num_slices) + if thickness is None: + thickness = float(centers[1] - centers[0]) if num_slices > 1 else np.ptp(depth) + nrows = int(np.ceil(num_slices / ncols)) + fig, axes = plt.subplots( + nrows, ncols, figsize=(figsize_per * ncols, figsize_per * nrows), squeeze=False + ) + kwargs.setdefault("colorbar", False) + kwargs.setdefault("depth_cue", 0.0) + for i, ax in enumerate(axes.ravel()): + if i >= num_slices: + ax.axis("off") + continue + plot_slab( + model, + channel=channel, + normal=normal, + offset=float(centers[i]), + thickness=thickness, + ax=ax, + title=f"offset {centers[i]:.1f} {model.units}", + **kwargs, + ) + fig.suptitle(f"{channel} slices along {normal}") + fig.tight_layout() + return (fig, axes) if returnfig else None + + +@_register("template") +def plot_template( + model, index: int, template: str | None = None, ax=None, returnfig: bool = False +): + """3D plot of one site's neighbor vectors with the fitted template overlaid. + + Parameters + ---------- + index : int + Site index. + template : str, optional + Template name; default is the site's classified structure (or the first + matched template). + """ + if not model.matches: + raise RuntimeError("Run match_templates() first.") + if template is None: + structure = int(model.get_channel("structure")[index]) + template = model.structure_names[structure] if structure >= 0 else list(model.matches)[0] + match = model.matches[template] + tmpl = model.templates[template] + dxyz, dist = model.neighbor_vectors(normalize=True) + p = dxyz[index] + p = p[np.isfinite(p).all(1) & (dist[index] <= 1.15 * tmpl.max_radius * 1.3)] + t = tmpl.vectors @ match["rotation"][index].T + if ax is None: + fig = plt.figure(figsize=(6, 6)) + ax = fig.add_subplot(111, projection="3d") + else: + fig = ax.figure + ax.scatter(p[:, 0], p[:, 1], p[:, 2], s=80, c="0.3", edgecolors="k", label="neighbors") + ax.scatter(t[:, 0], t[:, 1], t[:, 2], s=160, marker="+", c="r", linewidths=2, label=template) + ax.scatter([0], [0], [0], s=120, c="b", marker="x") + b = max(tmpl.max_radius, 1.0) * 1.2 + ax.set_xlim(-b, b) + ax.set_ylim(-b, b) + ax.set_zlim(-b, b) + ax.set_box_aspect((1, 1, 1)) + ax.set_title( + f"site {index}: {template} score {match['score'][index]:.2f}, matched {match['num_matched'][index]}" + ) + ax.legend() + return (fig, ax) if returnfig else None diff --git a/tests/atoms/test_atoms.py b/tests/atoms/test_atoms.py new file mode 100644 index 000000000..aed14828a --- /dev/null +++ b/tests/atoms/test_atoms.py @@ -0,0 +1,238 @@ +"""Tests for quantem.atoms: templates, RDF, template matching and the model pipeline.""" + +import numpy as np +import pytest +from scipy.spatial.transform import Rotation + +from quantem.atoms import AtomicModel, get_template, match_template +from quantem.atoms.measurements import ( + convex_hull_distance, + kmeans_1d, + misorientation, + sample_volume, + segment_grains, + strain_from_deformation, +) +from quantem.atoms.pdf import ( + find_neighbors, + fit_first_peak, + nn_distance_from_lattice, + radial_distribution, +) +from quantem.atoms.templates import _crystal_definition + + +def make_lattice(name: str, n: int = 8, noise: float = 0.03, seed: int = 0): + """Spherical particle of structure ``name`` with NN distance 1, randomly rotated.""" + cell, basis, _, _ = _crystal_definition(name) + rng = np.arange(-n, n + 1) + ijk = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + xyz = ((ijk[:, None, :] + basis[None, :, :]).reshape(-1, 3)) @ cell + radius = 0.8 * n + xyz = xyz[np.linalg.norm(xyz, axis=1) < radius] + interior = np.linalg.norm(xyz, axis=1) < radius - 2.5 + xyz = xyz @ Rotation.random(random_state=seed).as_matrix().T + xyz = xyz + np.random.default_rng(seed).normal(0, noise, xyz.shape) + return xyz, interior + + +# --------------------------------------------------------------------------- # +# templates +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize( + "name, num_neighbors, shells, num_symmetry", + [ + ("fcc", 12, (1.0,), 24), + ("hcp", 12, (1.0,), 6), + ("bcc", 14, (1.0, 2 / np.sqrt(3)), 24), + ("sc", 6, (1.0,), 24), + ("diamond", 16, (1.0, np.sqrt(8 / 3)), 12), + ("wurtzite", 16, (1.0, np.sqrt(8 / 3)), 3), + ("ico", 12, (1.0,), 60), + ], +) +def test_template_shells_and_symmetry(name, num_neighbors, shells, num_symmetry): + t = get_template(name) + assert t.num_neighbors == num_neighbors + assert np.allclose(t.shells, shells, atol=1e-6) + assert t.num_symmetry == num_symmetry + assert np.allclose(np.linalg.norm(t.vectors[: t.shell_counts[0]], axis=1), 1.0) + + +def test_template_extra_shells(): + assert get_template("fcc", num_shells=2).num_neighbors == 18 + assert get_template("zincblende").name == "diamond" + + +# --------------------------------------------------------------------------- # +# RDF / neighbors / calibration +# --------------------------------------------------------------------------- # +def test_rdf_first_peak(): + xyz, _ = make_lattice("fcc", n=8, noise=0.02) + rdf = radial_distribution(xyz * 7.0, r_max=25.0, dr=0.05) + fit = fit_first_peak(rdf["r"], rdf["g_smooth"]) + assert abs(fit["r_nn"] - 7.0) < 0.1 + assert fit["cutoff"][0] < 7.0 < fit["cutoff"][1] + + +def test_find_neighbors_padding(): + xyz = np.random.default_rng(0).random((5, 3)) + dist, idx = find_neighbors(xyz, 8) + assert dist.shape == (5, 8) and idx.shape == (5, 8) + assert np.all(idx[:, 4:] == -1) and np.all(np.isinf(dist[:, 4:])) + assert not np.any(idx[:, :4] == np.arange(5)[:, None]) + + +def test_nn_distance_from_lattice(): + assert np.isclose(nn_distance_from_lattice("fcc", 3.89), 3.89 / np.sqrt(2)) + assert np.isclose(nn_distance_from_lattice("bcc", 2.0), np.sqrt(3)) + + +# --------------------------------------------------------------------------- # +# matching +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("name", ["fcc", "hcp", "bcc", "diamond", "wurtzite"]) +def test_match_classifies_synthetic_lattice(name): + xyz, interior = make_lattice(name, n=7, noise=0.03) + dist, idx = find_neighbors(xyz, 24) + dxyz = xyz[idx] - xyz[:, None, :] + scores = {} + for tn in ["fcc", "hcp", "bcc", "diamond", "wurtzite"]: + t = get_template(tn) + valid = dist <= 1.15 * t.max_radius + scores[tn] = match_template(dxyz, valid, t, device="cpu", progress=False)["score"] + names = list(scores) + best = np.array(names)[np.stack([scores[n] for n in names], 1).argmax(1)] + assert (best[interior] == name).mean() > 0.97 + assert scores[name][interior].mean() > 0.8 + + +def test_match_recovers_rotation(): + cell, basis, _, _ = _crystal_definition("fcc") + rng = np.arange(-4, 5) + ijk = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + xyz = ((ijk[:, None, :] + basis[None, :, :]).reshape(-1, 3)) @ cell + r0 = Rotation.random(random_state=3).as_matrix() + xyz = xyz @ r0.T + center = int(np.argmin(np.linalg.norm(xyz, axis=1))) + dist, idx = find_neighbors(xyz, 14) + dxyz = xyz[idx] - xyz[:, None, :] + t = get_template("fcc") + res = match_template( + dxyz[center : center + 1], + dist[center : center + 1] < 1.15, + t, + device="cpu", + progress=False, + ) + assert res["score"][0] > 0.999 + assert res["num_matched"][0] == 12 + ang = misorientation(res["rotation"], np.array([[0]]), t.symmetry) + # rotation equals r0 up to a cubic symmetry op + delta = res["rotation"][0].T @ r0 + traces = np.einsum("ij,sji->s", delta, t.symmetry) + assert np.degrees(np.arccos(np.clip((traces.max() - 1) / 2, -1, 1))) < 0.5 + assert ang.shape == (1, 1) + + +# --------------------------------------------------------------------------- # +# measurements +# --------------------------------------------------------------------------- # +def test_misorientation_symmetry_invariance(): + t = get_template("fcc") + r = Rotation.random(4, random_state=1).as_matrix() + r_sym = np.einsum("nij,jk->nik", r, t.symmetry[7]) + rot = np.concatenate([r, r_sym]) + idx = np.array([[4], [5], [6], [7], [0], [1], [2], [3]]) + assert np.allclose(misorientation(rot, idx, t.symmetry), 0.0, atol=1e-4) + + +def test_segment_grains_and_strain(): + labels = segment_grains( + np.array([[1], [0], [3], [2], [-1]]), np.array([[1], [1], [1], [1], [0]], bool), min_size=2 + ) + assert labels.tolist() == [0, 0, 1, 1, -1] + f = np.array([[[1.04, 0, 0], [0, 1.0, 0], [0, 0, 0.98]]]) + s = strain_from_deformation(f, np.eye(3)[None]) + assert np.isclose(s["e_xx"][0], 0.04) and np.isclose(s["e_zz"][0], -0.02) + assert np.isclose(s["dilation"][0], 0.02 / 3) + + +def test_hull_sample_kmeans(): + xyz = np.random.default_rng(0).random((300, 3)) + d = convex_hull_distance(xyz) + assert d.min() >= -1e-9 and d.max() < 0.5 + vol = np.zeros((8, 8, 8)) + vol[4, 4, 4] = 27.0 + assert sample_volume(vol, np.array([[4, 4, 4]]), radius=1.0)[0] == pytest.approx(27.0 / 7) + labels, centers = kmeans_1d(np.r_[np.zeros(20), np.ones(20) * 5]) + assert labels[:20].sum() == 0 and labels[20:].sum() == 20 and np.allclose(centers, [0, 5]) + + +# --------------------------------------------------------------------------- # +# AtomicModel pipeline +# --------------------------------------------------------------------------- # +def make_twinned_fcc(n: int = 9, noise: float = 0.03): + """FCC sphere with a (111) twin: sites above the plane are mirrored.""" + cell, basis, _, _ = _crystal_definition("fcc") + rng = np.arange(-n, n + 1) + ijk = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + xyz = ((ijk[:, None, :] + basis[None, :, :]).reshape(-1, 3)) @ cell + xyz = xyz[np.linalg.norm(xyz, axis=1) < 0.8 * n] + normal = np.array([1.0, 1.0, 1.0]) / np.sqrt(3) + h = xyz @ normal + plane = h[np.argmin(np.abs(h - 0.3))] # a lattice plane slightly off-center + below = h < plane - 1e-6 + # the twin is the lower crystal reflected through the plane + mirrored = xyz[below] - 2.0 * (h[below] - plane)[:, None] * normal[None, :] + xyz = np.concatenate([xyz[~(h > plane + 1e-6)], mirrored]) + above = xyz @ normal > plane + 1e-6 + return xyz + np.random.default_rng(1).normal(0, noise, xyz.shape), above + + +def test_atomic_model_pipeline(): + xyz, above = make_twinned_fcc() + model = AtomicModel.from_array(xyz * 7.2, units="voxels", name="twin") + assert model.num_sites == xyz.shape[0] + pdf = model.compute_pdf() + assert abs(pdf["r_nn"] - 7.2) < 0.15 + scale = model.calibrate("fcc", lattice_constant=3.89) + assert np.isclose(scale * 7.2, 3.89 / np.sqrt(2), rtol=0.03) + assert model.units == "A" + model.find_neighbors(20) + assert np.median(model["num_neighbors"]) == 12 + model.match_templates(["fcc", "hcp"], progress=False, device="cpu") + assert set(model.channels) >= {"score_fcc", "score_hcp", "structure", "score_diff"} + structure = model["structure"] + hcp = structure == 1 + # twin-plane sites have an hcp environment; there is exactly one such plane + hcp_heights = (model.positions_native / 7.2) @ (np.ones(3) / np.sqrt(3)) + assert 0.05 < hcp.mean() < 0.25 + full = model["num_neighbors"] >= 12 + assert (np.abs(hcp_heights[hcp & full]) < 0.4).mean() > 0.9 + grains = model.segment_grains("fcc", angle_threshold=5.0, min_size=20) + assert grains.max() + 1 == 2 + # the two grains are the two sides of the twin plane + for g in (0, 1): + side = above[grains == g] + assert side.mean() > 0.95 or side.mean() < 0.05 + strain = model.compute_strain() + assert abs(np.nanmedian(strain["dilation"])) < 0.02 + model.surface_distance() + assert model["surface_distance"].min() >= -1e-6 + # selection keeps channels and categories + sub = model.select(structure == 0) + assert sub.num_sites == int((structure == 0).sum()) and "structure" in sub.categories + + +def test_atomic_model_channels_and_shapes(): + model = AtomicModel.from_array(np.random.default_rng(0).random((3, 50)) * 10) + assert model.positions.shape == (50, 3) + model.set_channel("foo", np.arange(50)) + assert model["foo"][-1] == 49 + model.set_channel("foo", np.zeros(50)) + assert model["foo"].sum() == 0 + with pytest.raises(ValueError): + model.set_channel("bar", np.zeros(3)) + with pytest.raises(KeyError): + model.get_channel("missing") From 1d1cb622ed7771e30b6aaceef388da171a41ac1a Mon Sep 17 00:00:00 2001 From: cophus Date: Tue, 15 Sep 2026 16:53:30 -0700 Subject: [PATCH 4/6] fixes to atomic analysis and vis --- src/quantem/atoms/atomic_model.py | 140 ++++++- src/quantem/atoms/measurements.py | 39 ++ src/quantem/atoms/show_atoms.py | 622 +---------------------------- src/quantem/atoms/visualization.py | 66 ++- tests/atoms/test_atoms.py | 17 + 5 files changed, 252 insertions(+), 632 deletions(-) diff --git a/src/quantem/atoms/atomic_model.py b/src/quantem/atoms/atomic_model.py index f284ca82a..7e3bdc8e6 100644 --- a/src/quantem/atoms/atomic_model.py +++ b/src/quantem/atoms/atomic_model.py @@ -762,25 +762,42 @@ def segment_grains( angle_threshold: float = 5.0, min_size: int = 20, min_score: float | None = None, + fill_iterations: int = 3, + fill_angle: float | None = None, ) -> NDArray: """Cluster sites of one structure into grains by local orientation. Neighboring sites of the given structure whose disorientation is below - ``angle_threshold`` are connected; connected components become grains. - Adds channels ``grain`` (``-1`` = none) and ``misorientation`` (largest - disorientation to any first-shell neighbor of the same structure). + ``angle_threshold`` are connected, and connected components become + grains. Sites of the structure left without a grain (typically sites + between two grains, or sites in components smaller than ``min_size``) + are then assigned by majority vote of their first-shell neighbors, + provided their disorientation to that grain is below ``fill_angle``; + the vote is repeated ``fill_iterations`` times so gaps close inward. + + Adds channels ``grain`` (``-1`` = none), ``misorientation`` (largest + disorientation to any first-shell neighbor of the same structure) and + ``boundary`` (categorical: ``interior``, ``grain_boundary`` for sites + with a first-shell neighbor in another grain, ``twin`` for sites of + the second matched template such as ``hcp``, and ``surface`` for sites + with fewer than 9 first-shell neighbors). Parameters ---------- structure : str Template name to segment (e.g. ``"fcc"``). angle_threshold : float - Maximum disorientation (degrees) inside a grain. + Maximum disorientation (degrees) between connected sites. min_size : int - Grains with fewer sites are discarded. + Components with fewer sites are dissolved and re-filled. min_score : float, optional Only sites with ``score_`` above this take part; default: sites classified as ``structure``. + fill_iterations : int + Number of majority-vote passes over unassigned sites (0 disables). + fill_angle : float, optional + Maximum disorientation for a vote to count; default + ``3 * angle_threshold``. Returns ------- @@ -805,10 +822,37 @@ def segment_grains( edge = same & (ang < angle_threshold) labels = meas.segment_grains(idx, edge, min_size=min_size) labels[~member] = -1 + if fill_angle is None: + fill_angle = 3.0 * angle_threshold + for _ in range(int(fill_iterations)): + labels = meas.fill_labels(labels, idx, same & (ang < fill_angle), member) + # relabel by decreasing size + if labels.max() >= 0: + sizes = np.bincount(labels[labels >= 0]) + order = np.argsort(-sizes, kind="stable") + rank = np.empty_like(order) + rank[order] = np.arange(order.size) + labels = np.where(labels >= 0, rank[np.clip(labels, 0, None)], -1) worst = np.where(same, ang, -1.0).max(axis=1).clip(0.0) num_grains = int(labels.max()) + 1 if labels.size else 0 self.set_channel("grain", labels, categories=[str(i) for i in range(num_grains)]) self.set_channel("misorientation", worst, "deg") + # boundary classification + structure_id = self.get_channel("structure").astype(int) + nb_labels = np.where(idx >= 0, labels[np.where(idx >= 0, idx, 0)], -2) + other_grain = ( + first & (nb_labels >= 0) & (nb_labels != labels[:, None]) & (labels[:, None] >= 0) + ) + boundary = np.zeros(self.num_sites, dtype=int) + boundary[other_grain.any(axis=1)] = 1 + names = self._structure_names + twin_ids = [i for i, n in enumerate(names) if n != structure] + if twin_ids: + boundary[np.isin(structure_id, twin_ids)] = 2 + boundary[first.sum(axis=1) < 9] = 3 + self.set_channel( + "boundary", boundary, categories=["interior", "grain_boundary", "twin", "surface"] + ) return labels def compute_strain( @@ -897,6 +941,78 @@ def rotate(self, rotation: NDArray, about_center: bool = True) -> None: c = xyz.mean(0) if about_center else np.zeros(3) self.positions_native = (xyz - c) @ rotation.T + c + def merge_close_sites( + self, + min_distance: float | None = None, + mode: str = "merge", + weight_channel: str | None = None, + ) -> int: + """Merge or remove sites closer than ``min_distance`` (in place). + + Sites are grouped into clusters by connecting every pair closer than + ``min_distance``. With ``mode="merge"`` each cluster is replaced by + one site at its (weighted) mean position, with all channels averaged; + with ``mode="remove"`` only the site with the largest weight (or the + first site) of each cluster is kept. Neighbor lists, template matches + and the RDF are cleared afterwards. + + Parameters + ---------- + min_distance : float, optional + Distance threshold in native units. Default: half the NN distance + from the RDF fit. + mode : {"merge", "remove"} + How to resolve each cluster. + weight_channel : str, optional + Channel used as weights (e.g. ``"intensity"``); equal weights if + omitted. + + Returns + ------- + int + Number of sites removed. + """ + from scipy.sparse import coo_matrix + from scipy.sparse.csgraph import connected_components + from scipy.spatial import cKDTree + + if min_distance is None: + min_distance = 0.5 * self.nn_distance + xyz = self.positions_native + n = xyz.shape[0] + pairs = cKDTree(xyz).query_pairs(float(min_distance), output_type="ndarray") + if pairs.shape[0] == 0: + return 0 + graph = coo_matrix((np.ones(pairs.shape[0]), (pairs[:, 0], pairs[:, 1])), shape=(n, n)) + _, labels = connected_components(graph, directed=False) + table = np.array(self._sites.array, dtype=float) + weights = ( + np.ones(n) + if weight_channel is None + else self.get_channel(weight_channel).astype(float) + ) + weights = np.where(np.isfinite(weights) & (weights > 0), weights, 1e-12) + num_clusters = int(labels.max()) + 1 + if mode == "merge": + sums = np.zeros((num_clusters, table.shape[1])) + np.add.at(sums, labels, table * weights[:, None]) + wsum = np.bincount(labels, weights=weights, minlength=num_clusters) + new_table = sums / wsum[:, None] + elif mode == "remove": + order = np.lexsort((-weights, labels)) + first = np.ones(n, dtype=bool) + first[1:] = labels[order][1:] != labels[order][:-1] + new_table = table[order][first] + else: + raise ValueError("mode must be 'merge' or 'remove'") + sites = Vector.from_shape( + shape=(), fields=list(self._sites.fields), units=list(self._sites.units), name="sites" + ) + sites[...] = np.ascontiguousarray(new_table) + self._sites = sites + self._invalidate() + return int(n - new_table.shape[0]) + def select(self, mask: NDArray) -> "AtomicModel": """Return a new model containing only the sites where ``mask`` is True.""" mask = np.asarray(mask, dtype=bool) @@ -942,9 +1058,19 @@ def plot(self, kind: str = "slab", show_docstring: bool = False, **kwargs): return fn(self, **kwargs) def show(self, **kwargs): - """Open the interactive 3D viewer (:class:`quantem.atoms.ShowAtoms3D`).""" - from quantem.atoms.show_atoms import ShowAtoms3D + """Open the interactive 3D viewer ``quantem.widget.ShowAtoms3D``. + Requires the ``quantem.widget`` package. Keyword arguments are + forwarded to the viewer (``channel``, ``cmap``, ``canvas_size``, + ``marker_size``, ``dark_background``, ``title``). + """ + try: + from quantem.widget import ShowAtoms3D + except ImportError as exc: + raise ImportError( + "AtomicModel.show() requires the quantem.widget package: " + "pip install quantem.widget" + ) from exc return ShowAtoms3D(self, **kwargs) def __repr__(self) -> str: diff --git a/src/quantem/atoms/measurements.py b/src/quantem/atoms/measurements.py index 704c66c9c..2d714c598 100644 --- a/src/quantem/atoms/measurements.py +++ b/src/quantem/atoms/measurements.py @@ -16,6 +16,7 @@ __all__ = [ "misorientation", "segment_grains", + "fill_labels", "strain_from_deformation", "bond_angles", "convex_hull_distance", @@ -106,6 +107,44 @@ def segment_grains( return out +def fill_labels( + labels: NDArray, neighbor_index: NDArray, edge_mask: NDArray, member: NDArray +) -> NDArray: + """Assign unlabelled member sites to the majority label of their neighbors. + + Parameters + ---------- + labels : ndarray + ``(N,)`` labels, ``-1`` = unassigned. + neighbor_index : ndarray + ``(N, K)`` neighbor indices (``-1`` = missing). + edge_mask : ndarray + ``(N, K)`` votes are only counted along ``True`` edges. + member : ndarray + ``(N,)`` sites eligible for filling. + + Returns + ------- + ndarray + Updated copy of ``labels``. + """ + labels = np.array(labels, copy=True) + todo = np.where(member & (labels < 0))[0] + if todo.size == 0: + return labels + safe = np.where(neighbor_index >= 0, neighbor_index, 0) + votes = np.where(edge_mask & (neighbor_index >= 0), labels[safe], -1)[todo] + n_max = int(labels.max()) + 2 + counts = np.zeros((todo.size, n_max), dtype=int) + rows = np.repeat(np.arange(todo.size), votes.shape[1]) + valid = votes.ravel() >= 0 + np.add.at(counts, (rows[valid], votes.ravel()[valid]), 1) + best = counts.argmax(axis=1) + has_votes = counts.max(axis=1) > 0 + labels[todo[has_votes]] = best[has_votes] + return labels + + def strain_from_deformation(deformation: NDArray, rotation: NDArray, frame: str = "lab") -> dict: """Small-strain tensor components from a deformation gradient. diff --git a/src/quantem/atoms/show_atoms.py b/src/quantem/atoms/show_atoms.py index 0abf9d41b..ca5523846 100644 --- a/src/quantem/atoms/show_atoms.py +++ b/src/quantem/atoms/show_atoms.py @@ -1,619 +1,9 @@ -"""Interactive 3D viewer for :class:`~quantem.atoms.AtomicModel`. +"""Compatibility shim: the interactive viewer lives in ``quantem.widget``. -:class:`ShowAtoms3D` is a self-contained `anywidget `_ -(no JavaScript build step) that renders all sites of a model in a notebook: - -* drag to rotate, scroll to zoom, shift-drag (or right-drag) to pan, - double-click to reset; -* a **channel** dropdown colors sites by any per-site channel, with a live - histogram whose two handles set the color range; -* **clip** controls cut a slab along the view direction or a model axis so - flat atomic planes can be isolated, and "hide outside range" keeps only sites - whose channel value lies inside the histogram range (e.g. only twin sites); -* the current view is synced back to Python (:attr:`ShowAtoms3D.view_matrix`) - so the same projection can be reproduced with ``model.plot("slab", normal=...)``. - -The widget is optional: it requires the ``anywidget`` package. +``ShowAtoms3D`` moved to the ``quantem.widget`` package +(``quantem.widget.show_atoms3d``) so it shares the widget infrastructure with +the other quantem viewers. This module re-exports it for older imports and +can be removed. """ -from __future__ import annotations - -import json -from typing import Any - -import numpy as np -from numpy.typing import NDArray - -try: - import anywidget - import traitlets -except ImportError as exc: # pragma: no cover - optional dependency - raise ImportError( - "ShowAtoms3D requires the 'anywidget' package: pip install anywidget" - ) from exc - -from quantem.atoms.visualization import _CATEGORICAL_COLORS - -__all__ = ["ShowAtoms3D"] - -_DEFAULT_CMAPS = [ - "viridis", - "turbo", - "inferno", - "magma", - "plasma", - "cividis", - "coolwarm", - "RdBu_r", - "gray", -] - -_ESM = r""" -function mat3mul(a, b) { - const o = new Float64Array(9); - for (let i = 0; i < 3; i++) for (let j = 0; j < 3; j++) { - o[3*i+j] = a[3*i]*b[j] + a[3*i+1]*b[3+j] + a[3*i+2]*b[6+j]; - } - return o; -} -function rotX(t){const c=Math.cos(t),s=Math.sin(t);return [1,0,0,0,c,-s,0,s,c];} -function rotY(t){const c=Math.cos(t),s=Math.sin(t);return [c,0,s,0,1,0,-s,0,c];} -function orthonormalize(m){ - let r=[m[0],m[1],m[2]], u=[m[3],m[4],m[5]]; - const nr=Math.hypot(...r); r=r.map(v=>v/nr); - const d=u[0]*r[0]+u[1]*r[1]+u[2]*r[2]; u=[u[0]-d*r[0],u[1]-d*r[1],u[2]-d*r[2]]; - const nu=Math.hypot(...u); u=u.map(v=>v/nu); - const n=[r[1]*u[2]-r[2]*u[1], r[2]*u[0]-r[0]*u[2], r[0]*u[1]-r[1]*u[0]]; - return [...r,...u,...n]; -} - -function render({ model, el }) { - // ---------------- data ---------------- - const N = model.get("num_sites"); - const pos = new Float32Array(model.get("positions").buffer.slice(0)); - const channelNames = model.get("channel_names"); - const chanData = new Float32Array(model.get("channel_data").buffer.slice(0)); - const categories = model.get("channel_categories"); - const palettes = model.get("channel_palettes"); - const cmapNames = model.get("cmap_names"); - const luts = new Uint8Array(model.get("cmap_luts").buffer.slice(0)); - const radius0 = model.get("model_radius"); - const units = model.get("units"); - - const proj = new Float32Array(N * 3); - const order = new Uint32Array(N); - const depthKey = new Float32Array(N); - const visible = new Uint8Array(N); - let spriteCache = new Map(); - let spriteRadiusPx = -1; - let dragging = false, dragMode = 0, lastX = 0, lastY = 0, moved = false; - let hist = null, histDrag = 0; - - // ---------------- layout ---------------- - el.innerHTML = ""; - const root = document.createElement("div"); root.className = "qa-root"; el.appendChild(root); - const left = document.createElement("div"); left.className = "qa-left"; root.appendChild(left); - const title = document.createElement("div"); title.className = "qa-title"; left.appendChild(title); - const canvas = document.createElement("canvas"); canvas.className = "qa-canvas"; left.appendChild(canvas); - const status = document.createElement("div"); status.className = "qa-status"; left.appendChild(status); - const panel = document.createElement("div"); panel.className = "qa-panel"; root.appendChild(panel); - const ctx = canvas.getContext("2d"); - - function row(label) { - const r = document.createElement("div"); r.className = "qa-row"; - const l = document.createElement("label"); l.textContent = label; r.appendChild(l); - panel.appendChild(r); return r; - } - function select(label, options, key, fmt) { - const r = row(label); const s = document.createElement("select"); - options.forEach(o => { const op = document.createElement("option"); op.value = o; op.textContent = fmt ? fmt(o) : o; s.appendChild(op); }); - s.value = model.get(key); - s.onchange = () => { model.set(key, s.value); model.save_changes(); }; - r.appendChild(s); return s; - } - function slider(label, key, min, max, step, live) { - const r = row(label); const s = document.createElement("input"); s.type = "range"; - s.min = min; s.max = max; s.step = step; s.value = model.get(key); - const v = document.createElement("span"); v.className = "qa-val"; v.textContent = Number(model.get(key)).toFixed(2); - s.oninput = () => { v.textContent = Number(s.value).toFixed(2); model.set(key, Number(s.value)); if (live) draw(); }; - s.onchange = () => { model.set(key, Number(s.value)); model.save_changes(); }; - r.appendChild(s); r.appendChild(v); return { s, v }; - } - function checkbox(label, key) { - const r = row(label); const c = document.createElement("input"); c.type = "checkbox"; c.checked = model.get(key); - c.onchange = () => { model.set(key, c.checked); model.save_changes(); }; - r.appendChild(c); return c; - } - function button(parent, text, fn) { - const b = document.createElement("button"); b.className = "qa-btn"; b.textContent = text; b.onclick = fn; parent.appendChild(b); return b; - } - - const channelSel = select("channel", channelNames, "channel"); - const cmapSel = select("colormap", cmapNames, "cmap"); - const histCanvas = document.createElement("canvas"); histCanvas.className = "qa-hist"; histCanvas.width = 240; histCanvas.height = 96; panel.appendChild(histCanvas); - const hctx = histCanvas.getContext("2d"); - const rangeRow = row("range"); - const vminIn = document.createElement("input"); vminIn.type = "number"; vminIn.className = "qa-num"; vminIn.step = "any"; - const vmaxIn = document.createElement("input"); vmaxIn.type = "number"; vmaxIn.className = "qa-num"; vmaxIn.step = "any"; - rangeRow.appendChild(vminIn); rangeRow.appendChild(vmaxIn); - vminIn.onchange = () => { model.set("vmin", Number(vminIn.value)); model.set("range_auto", false); model.save_changes(); }; - vmaxIn.onchange = () => { model.set("vmax", Number(vmaxIn.value)); model.set("range_auto", false); model.save_changes(); }; - const rr = row(""); button(rr, "auto range", () => { model.set("range_auto", true); model.save_changes(); autoRange(); }); - const hideChk = checkbox("hide outside range", "hide_outside_range"); - const legend = document.createElement("div"); legend.className = "qa-legend"; panel.appendChild(legend); - - const sep1 = document.createElement("div"); sep1.className = "qa-sep"; sep1.textContent = "clip"; panel.appendChild(sep1); - const clipChk = checkbox("enable clip", "clip_enabled"); - const clipAxis = select("axis", ["view", "x", "y", "z"], "clip_axis"); - const clipCenter = slider("center", "clip_center", -radius0, radius0, radius0 / 400, true); - const clipThick = slider("thickness", "clip_thickness", 0.1, 2 * radius0, radius0 / 400, true); - - const sep2 = document.createElement("div"); sep2.className = "qa-sep"; sep2.textContent = "display"; panel.appendChild(sep2); - const sizeSl = slider("marker", "marker_size", 0.05, 2.0, 0.01, true); - const cueSl = slider("depth cue", "depth_cue", 0, 1, 0.01, true); - const edgeChk = checkbox("edges", "show_edges"); - const bgChk = checkbox("dark background", "dark_background"); - const viewRow = row("view"); - button(viewRow, "x", () => setView([0,1,0, 0,0,1, 1,0,0])); - button(viewRow, "y", () => setView([0,0,1, 1,0,0, 0,1,0])); - button(viewRow, "z", () => setView([1,0,0, 0,1,0, 0,0,1])); - button(viewRow, "reset", () => { setView([1,0,0, 0,1,0, 0,0,1]); model.set("zoom", 1.0); model.set("pan", [0, 0]); model.save_changes(); }); - const info = document.createElement("div"); info.className = "qa-info"; panel.appendChild(info); - - function setView(m) { model.set("rotation", Array.from(m)); model.save_changes(); draw(); } - - // ---------------- channel helpers ---------------- - function channelIndex() { const i = channelNames.indexOf(model.get("channel")); return i < 0 ? 0 : i; } - function channelValues() { const i = channelIndex(); return chanData.subarray(i * N, (i + 1) * N); } - function isCategorical() { return categories[model.get("channel")] !== undefined; } - function autoRange() { - const v = channelValues(); const name = model.get("channel"); - if (isCategorical()) { model.set("vmin", -0.5); model.set("vmax", categories[name].length - 0.5); model.save_changes(); return; } - const arr = Array.from(v).filter(Number.isFinite).sort((a, b) => a - b); - if (!arr.length) return; - const lo = arr[Math.floor(0.01 * (arr.length - 1))], hi = arr[Math.floor(0.99 * (arr.length - 1))]; - model.set("vmin", lo); model.set("vmax", hi === lo ? lo + 1e-9 : hi); model.save_changes(); - } - function computeHist() { - const v = channelValues(); const nb = 64; - let lo, hi; - if (isCategorical()) { lo = -1.5; hi = categories[model.get("channel")].length - 0.5; } - else { - let mn = Infinity, mx = -Infinity; - for (let i = 0; i < N; i++) { const x = v[i]; if (Number.isFinite(x)) { if (x < mn) mn = x; if (x > mx) mx = x; } } - lo = mn; hi = mx === mn ? mn + 1e-9 : mx; - } - const counts = new Float64Array(nb); - for (let i = 0; i < N; i++) { const x = v[i]; if (!Number.isFinite(x)) continue; let b = Math.floor((x - lo) / (hi - lo) * nb); if (b >= nb) b = nb - 1; if (b < 0) b = 0; counts[b]++; } - hist = { lo, hi, counts, nb }; - } - function lutColor(t) { - const ci = Math.max(0, cmapNames.indexOf(model.get("cmap"))); - const k = Math.max(0, Math.min(255, Math.round(t * 255))); - const o = (ci * 256 + k) * 3; return [luts[o], luts[o + 1], luts[o + 2]]; - } - function drawHist() { - if (!hist) computeHist(); - const W = histCanvas.width, H = histCanvas.height, pad = 4, hh = H - 22; - const dark = model.get("dark_background"); - hctx.fillStyle = dark ? "#1e1e1e" : "#f4f4f4"; hctx.fillRect(0, 0, W, H); - const vmin = model.get("vmin"), vmax = model.get("vmax"); - const cmax = Math.max(1, ...hist.counts); - const cat = isCategorical(); const name = model.get("channel"); - for (let b = 0; b < hist.nb; b++) { - const x0 = pad + (W - 2 * pad) * b / hist.nb, w = (W - 2 * pad) / hist.nb; - const val = hist.lo + (b + 0.5) / hist.nb * (hist.hi - hist.lo); - let t = (val - vmin) / (vmax - vmin); - let col; - if (cat) { const code = Math.round(val); const p = palettes[name]; col = code < 0 || !p ? [140,140,140] : p.slice(3*(code % (p.length/3)), 3*(code % (p.length/3))+3); } - else col = lutColor(Math.max(0, Math.min(1, t))); - const h = Math.log1p(hist.counts[b]) / Math.log1p(cmax) * hh; - hctx.fillStyle = `rgb(${col[0]},${col[1]},${col[2]})`; - hctx.globalAlpha = (val < vmin || val > vmax) ? 0.3 : 1.0; - hctx.fillRect(x0, pad + hh - h, Math.max(w - 1, 1), h); - } - hctx.globalAlpha = 1; - // colorbar - for (let i = 0; i < W - 2 * pad; i++) { - const val = hist.lo + i / (W - 2 * pad) * (hist.hi - hist.lo); - const t = Math.max(0, Math.min(1, (val - vmin) / (vmax - vmin))); - let col; - if (cat) { const code = Math.round(val); const p = palettes[name]; col = code < 0 || !p ? [140,140,140] : p.slice(3*(code % (p.length/3)), 3*(code % (p.length/3))+3); } - else col = lutColor(t); - hctx.fillStyle = `rgb(${col[0]},${col[1]},${col[2]})`; hctx.fillRect(pad + i, pad + hh + 4, 1, 8); - } - // handles - const xs = [vmin, vmax].map(v => pad + (W - 2 * pad) * (v - hist.lo) / (hist.hi - hist.lo)); - hctx.strokeStyle = dark ? "#fff" : "#000"; hctx.lineWidth = 1.5; - xs.forEach(x => { hctx.beginPath(); hctx.moveTo(x, pad); hctx.lineTo(x, H - 2); hctx.stroke(); }); - vminIn.value = Number(vmin.toPrecision(5)); vmaxIn.value = Number(vmax.toPrecision(5)); - // legend - legend.innerHTML = ""; - if (cat) { - const p = palettes[name] || []; - categories[name].forEach((lab, i) => { - const d = document.createElement("div"); d.className = "qa-legend-item"; - const sw = document.createElement("span"); sw.className = "qa-swatch"; sw.style.background = `rgb(${p[3*i]},${p[3*i+1]},${p[3*i+2]})`; - d.appendChild(sw); d.appendChild(document.createTextNode(`${i}: ${lab}`)); legend.appendChild(d); - }); - const d = document.createElement("div"); d.className = "qa-legend-item"; - const sw = document.createElement("span"); sw.className = "qa-swatch"; sw.style.background = "rgb(140,140,140)"; - d.appendChild(sw); d.appendChild(document.createTextNode("-1: none")); legend.appendChild(d); - } - } - function histPointer(ev) { - const rect = histCanvas.getBoundingClientRect(); const pad = 4, W = histCanvas.width; - const x = (ev.clientX - rect.left) * (W / rect.width); - return hist.lo + (x - pad) / (W - 2 * pad) * (hist.hi - hist.lo); - } - histCanvas.onpointerdown = (ev) => { - if (!hist) computeHist(); - const v = histPointer(ev); const vmin = model.get("vmin"), vmax = model.get("vmax"); - histDrag = Math.abs(v - vmin) < Math.abs(v - vmax) ? 1 : 2; histCanvas.setPointerCapture(ev.pointerId); - }; - histCanvas.onpointermove = (ev) => { - if (!histDrag) return; const v = histPointer(ev); - if (histDrag === 1) model.set("vmin", Math.min(v, model.get("vmax") - 1e-9)); else model.set("vmax", Math.max(v, model.get("vmin") + 1e-9)); - model.set("range_auto", false); drawHist(); draw(); - }; - histCanvas.onpointerup = (ev) => { histDrag = 0; histCanvas.releasePointerCapture(ev.pointerId); model.save_changes(); }; - - // ---------------- sprites ---------------- - function sprite(r, g, b, shade, radiusPx, edges) { - const key = ((r << 16) | (g << 8) | b) * 32 + shade; - let s = spriteCache.get(key); if (s) return s; - const R = Math.max(1, Math.ceil(radiusPx)); const size = 2 * R + 2; - s = document.createElement("canvas"); s.width = size; s.height = size; - const c = s.getContext("2d"); - const f = 1 - 0.85 * shade / 31; - const cx = R + 1, cy = R + 1; - const grad = c.createRadialGradient(cx - 0.35 * R, cy - 0.35 * R, 0.1 * R, cx, cy, R); - grad.addColorStop(0, `rgb(${Math.min(255, r*f+90*f)|0},${Math.min(255, g*f+90*f)|0},${Math.min(255, b*f+90*f)|0})`); - grad.addColorStop(1, `rgb(${(r*f*0.75)|0},${(g*f*0.75)|0},${(b*f*0.75)|0})`); - c.fillStyle = grad; c.beginPath(); c.arc(cx, cy, R, 0, 2 * Math.PI); c.fill(); - if (edges && R > 2) { c.strokeStyle = "rgba(0,0,0,0.6)"; c.lineWidth = Math.max(0.5, R * 0.12); c.stroke(); } - spriteCache.set(key, s); return s; - } - - // ---------------- main draw ---------------- - function draw() { - const W = model.get("canvas_size"), H = W; - if (canvas.width !== W) { canvas.width = W; canvas.height = H; } - const dark = model.get("dark_background"); - ctx.fillStyle = dark ? "#111" : "#fff"; ctx.fillRect(0, 0, W, H); - const m = model.get("rotation"); const zoom = model.get("zoom"); const pan = model.get("pan"); - const scale = zoom * (0.5 * W * 0.92) / radius0; - const cx = W / 2 + pan[0], cy = H / 2 + pan[1]; - const vals = channelValues(); const vmin = model.get("vmin"), vmax = model.get("vmax"); - const cat = isCategorical(); const name = model.get("channel"); const pal = palettes[name]; - const hide = model.get("hide_outside_range"); - const clipOn = model.get("clip_enabled"), clipAx = model.get("clip_axis"); - const cc = model.get("clip_center"), ct = model.get("clip_thickness"); - let count = 0; - let dmin = Infinity, dmax = -Infinity; - for (let i = 0; i < N; i++) { - const x = pos[3*i], y = pos[3*i+1], z = pos[3*i+2]; - const u = m[0]*x + m[1]*y + m[2]*z, v = m[3]*x + m[4]*y + m[5]*z, d = m[6]*x + m[7]*y + m[8]*z; - proj[3*i] = cx + u * scale; proj[3*i+1] = cy - v * scale; proj[3*i+2] = d; - let ok = 1; - if (clipOn) { const c = clipAx === "view" ? d : (clipAx === "x" ? x : (clipAx === "y" ? y : z)); if (Math.abs(c - cc) > ct / 2) ok = 0; } - if (ok && hide) { const val = vals[i]; if (!(val >= vmin && val <= vmax)) ok = 0; } - visible[i] = ok; - if (ok) { if (d < dmin) dmin = d; if (d > dmax) dmax = d; order[count++] = i; } - } - const sub = order.subarray(0, count); - for (let k = 0; k < count; k++) depthKey[sub[k]] = proj[3*sub[k]+2]; - sub.sort((a, b) => depthKey[a] - depthKey[b]); - const radiusPx = Math.max(0.6, model.get("marker_size") * model.get("bond_length") * 0.5 * scale); - if (radiusPx !== spriteRadiusPx) { spriteCache = new Map(); spriteRadiusPx = radiusPx; } - const edges = model.get("show_edges"); const cue = model.get("depth_cue"); - const dr = Math.max(dmax - dmin, 1e-9); - const inv = 1 / Math.max(vmax - vmin, 1e-12); - const ci = Math.max(0, cmapNames.indexOf(model.get("cmap"))); - for (let k = 0; k < count; k++) { - const i = sub[k]; const val = vals[i]; - let r, g, b; - if (cat) { const code = Math.round(val); if (code < 0 || !pal || !Number.isFinite(val)) { r = 140; g = 140; b = 140; } else { const j = 3 * (code % (pal.length / 3)); r = pal[j]; g = pal[j+1]; b = pal[j+2]; } } - else if (!Number.isFinite(val)) { r = 140; g = 140; b = 140; } - else { let t = (val - vmin) * inv; t = t < 0 ? 0 : (t > 1 ? 1 : t); const o = (ci * 256 + Math.round(t * 255)) * 3; r = luts[o]; g = luts[o+1]; b = luts[o+2]; } - const t = (proj[3*i+2] - dmin) / dr; - const shade = Math.round(cue * (1 - t) * 31); - const s = sprite(r, g, b, shade, radiusPx, edges); - ctx.drawImage(s, proj[3*i] - s.width / 2, proj[3*i+1] - s.height / 2); - } - status.textContent = `${count} / ${N} sites shown`; - } - - // ---------------- interaction ---------------- - canvas.oncontextmenu = (e) => e.preventDefault(); - canvas.onpointerdown = (ev) => { - dragging = true; moved = false; lastX = ev.clientX; lastY = ev.clientY; - dragMode = (ev.button === 2 || ev.shiftKey) ? 2 : 1; canvas.setPointerCapture(ev.pointerId); - }; - canvas.onpointermove = (ev) => { - if (!dragging) return; - const dx = ev.clientX - lastX, dy = ev.clientY - lastY; lastX = ev.clientX; lastY = ev.clientY; - if (Math.abs(dx) + Math.abs(dy) > 0) moved = true; - if (dragMode === 2) { const p = model.get("pan"); model.set("pan", [p[0] + dx, p[1] + dy]); } - else { - const m = model.get("rotation"); const k = 0.008; - const rs = mat3mul(rotX(-dy * k), rotY(dx * k)); - model.set("rotation", Array.from(orthonormalize(mat3mul(rs, m)))); - } - draw(); - }; - canvas.onpointerup = (ev) => { - dragging = false; canvas.releasePointerCapture(ev.pointerId); - if (!moved) pick(ev); model.save_changes(); - }; - canvas.onwheel = (ev) => { ev.preventDefault(); const z = model.get("zoom") * Math.exp(-ev.deltaY * 0.0015); model.set("zoom", Math.max(0.05, Math.min(50, z))); draw(); model.save_changes(); }; - canvas.ondblclick = () => { model.set("zoom", 1.0); model.set("pan", [0, 0]); model.save_changes(); draw(); }; - function pick(ev) { - const rect = canvas.getBoundingClientRect(); - const x = (ev.clientX - rect.left) * canvas.width / rect.width, y = (ev.clientY - rect.top) * canvas.height / rect.height; - let best = -1, bd = 1e9; - for (let i = 0; i < N; i++) { if (!visible[i]) continue; const dx = proj[3*i] - x, dy = proj[3*i+1] - y; const d2 = dx*dx + dy*dy - 0.02 * proj[3*i+2]; if (d2 < bd) { bd = d2; best = i; } } - if (best >= 0 && Math.sqrt(bd + 1) < 20) { - model.set("selected", best); model.save_changes(); - const v = channelValues()[best]; const name = model.get("channel"); - const lab = isCategorical() ? (v >= 0 ? categories[name][Math.round(v)] : "none") : v.toPrecision(4); - info.textContent = `site ${best}: ${name} = ${lab} (${pos[3*best].toFixed(2)}, ${pos[3*best+1].toFixed(2)}, ${pos[3*best+2].toFixed(2)}) ${units}`; - } - } - - // ---------------- model listeners ---------------- - model.on("change:channel", () => { channelSel.value = model.get("channel"); hist = null; if (model.get("range_auto")) autoRange(); drawHist(); draw(); }); - model.on("change:cmap", () => { cmapSel.value = model.get("cmap"); spriteCache = new Map(); drawHist(); draw(); }); - ["vmin", "vmax"].forEach(k => model.on("change:" + k, () => { drawHist(); draw(); })); - ["rotation", "zoom", "pan", "marker_size", "depth_cue", "show_edges", "clip_enabled", "clip_axis", "clip_center", "clip_thickness", "hide_outside_range", "canvas_size"].forEach(k => model.on("change:" + k, () => { - if (k === "clip_axis") clipAxis.value = model.get(k); - if (k === "clip_enabled") clipChk.checked = model.get(k); - if (k === "hide_outside_range") hideChk.checked = model.get(k); - if (k === "show_edges") edgeChk.checked = model.get(k); - if (k === "clip_center") { clipCenter.s.value = model.get(k); clipCenter.v.textContent = Number(model.get(k)).toFixed(2); } - if (k === "clip_thickness") { clipThick.s.value = model.get(k); clipThick.v.textContent = Number(model.get(k)).toFixed(2); } - if (k === "marker_size") { sizeSl.s.value = model.get(k); sizeSl.v.textContent = Number(model.get(k)).toFixed(2); } - if (k === "depth_cue") { cueSl.s.value = model.get(k); cueSl.v.textContent = Number(model.get(k)).toFixed(2); } - draw(); - })); - model.on("change:dark_background", () => { bgChk.checked = model.get("dark_background"); root.classList.toggle("qa-dark", model.get("dark_background")); drawHist(); draw(); }); - model.on("change:title", () => { title.textContent = model.get("title"); }); - - title.textContent = model.get("title"); - root.classList.toggle("qa-dark", model.get("dark_background")); - if (model.get("range_auto")) autoRange(); - drawHist(); draw(); -} -export default { render }; -""" - -_CSS = r""" -.qa-root { display: flex; gap: 10px; font-family: system-ui, sans-serif; font-size: 12px; color: #222; background: #fafafa; padding: 6px; border-radius: 6px; } -.qa-root.qa-dark { color: #ddd; background: #222; } -.qa-left { display: flex; flex-direction: column; gap: 4px; } -.qa-title { font-weight: 600; font-size: 13px; } -.qa-canvas { border: 1px solid #888; cursor: grab; touch-action: none; } -.qa-status, .qa-info { font-size: 11px; opacity: 0.8; min-height: 14px; } -.qa-panel { display: flex; flex-direction: column; gap: 4px; width: 250px; } -.qa-row { display: flex; align-items: center; gap: 6px; } -.qa-row label { width: 74px; flex: none; } -.qa-row select { flex: 1; min-width: 0; } -.qa-row input[type=range] { flex: 1; min-width: 0; } -.qa-val { width: 44px; text-align: right; font-variant-numeric: tabular-nums; } -.qa-num { width: 80px; } -.qa-hist { border: 1px solid #888; touch-action: none; cursor: col-resize; } -.qa-sep { margin-top: 6px; font-weight: 600; border-bottom: 1px solid #888; } -.qa-btn { padding: 2px 8px; font-size: 11px; } -.qa-legend { display: flex; flex-wrap: wrap; gap: 4px 10px; } -.qa-legend-item { display: flex; align-items: center; gap: 4px; } -.qa-swatch { width: 10px; height: 10px; border-radius: 50%; border: 1px solid #555; display: inline-block; } -""" - - -def _colormap_luts(names: list[str]) -> bytes: - import matplotlib.pyplot as plt - - out = np.zeros((len(names), 256, 3), dtype=np.uint8) - for i, n in enumerate(names): - rgba = plt.get_cmap(n)(np.linspace(0, 1, 256)) - out[i] = (rgba[:, :3] * 255).round().astype(np.uint8) - return out.tobytes() - - -class ShowAtoms3D(anywidget.AnyWidget): - """Interactive 3D site viewer for an :class:`~quantem.atoms.AtomicModel`. - - Parameters - ---------- - model : AtomicModel - Model to display; all channels are sent to the browser as float32. - channel : str, optional - Initial color channel (default: ``"structure"`` if present). - cmap : str - Initial colormap for continuous channels. - cmaps : sequence of str, optional - Colormaps offered in the dropdown (matplotlib names). - canvas_size : int - Canvas width and height in pixels. - marker_size : float - Marker diameter as a fraction of the bond length. - dark_background : bool - Dark theme. - title : str, optional - Title shown above the canvas (default: model name). - - Attributes - ---------- - rotation : list of float - Row-major ``(3, 3)`` view matrix (rows: right, up, toward-viewer). - view_matrix : ndarray - Same as ``rotation`` as an array; ``view_matrix[2]`` is the viewing - direction usable as ``model.plot("slab", normal=...)``. - selected : int - Index of the last clicked site (``-1`` if none). - """ - - _esm = _ESM - _css = _CSS - - positions = traitlets.Bytes(b"").tag(sync=True) - num_sites = traitlets.Int(0).tag(sync=True) - channel_names = traitlets.List(traitlets.Unicode()).tag(sync=True) - channel_data = traitlets.Bytes(b"").tag(sync=True) - channel_categories = traitlets.Dict().tag(sync=True) - channel_palettes = traitlets.Dict().tag(sync=True) - cmap_names = traitlets.List(traitlets.Unicode()).tag(sync=True) - cmap_luts = traitlets.Bytes(b"").tag(sync=True) - model_radius = traitlets.Float(1.0).tag(sync=True) - bond_length = traitlets.Float(1.0).tag(sync=True) - units = traitlets.Unicode("").tag(sync=True) - title = traitlets.Unicode("").tag(sync=True) - - channel = traitlets.Unicode("").tag(sync=True) - cmap = traitlets.Unicode("viridis").tag(sync=True) - vmin = traitlets.Float(0.0).tag(sync=True) - vmax = traitlets.Float(1.0).tag(sync=True) - range_auto = traitlets.Bool(True).tag(sync=True) - hide_outside_range = traitlets.Bool(False).tag(sync=True) - rotation = traitlets.List(traitlets.Float(), default_value=[1, 0, 0, 0, 1, 0, 0, 0, 1]).tag( - sync=True - ) - zoom = traitlets.Float(1.0).tag(sync=True) - pan = traitlets.List(traitlets.Float(), default_value=[0.0, 0.0]).tag(sync=True) - clip_enabled = traitlets.Bool(False).tag(sync=True) - clip_axis = traitlets.Unicode("view").tag(sync=True) - clip_center = traitlets.Float(0.0).tag(sync=True) - clip_thickness = traitlets.Float(1.0).tag(sync=True) - marker_size = traitlets.Float(0.7).tag(sync=True) - depth_cue = traitlets.Float(0.5).tag(sync=True) - show_edges = traitlets.Bool(True).tag(sync=True) - dark_background = traitlets.Bool(True).tag(sync=True) - canvas_size = traitlets.Int(640).tag(sync=True) - selected = traitlets.Int(-1).tag(sync=True) - - def __init__( - self, - model: Any, - channel: str | None = None, - cmap: str = "viridis", - cmaps: list[str] | None = None, - canvas_size: int = 640, - marker_size: float = 0.7, - dark_background: bool = True, - title: str | None = None, - **kwargs: Any, - ) -> None: - super().__init__(**kwargs) - self._model = model - xyz = np.asarray(model.positions, dtype=np.float32) - center = xyz.mean(0) - xyz = xyz - center - self._center = center - self.positions = np.ascontiguousarray(xyz, dtype=np.float32).tobytes() - self.num_sites = int(xyz.shape[0]) - self.model_radius = float(np.linalg.norm(xyz, axis=1).max()) if xyz.shape[0] else 1.0 - try: - self.bond_length = float(model.bond_length) - except Exception: - self.bond_length = float(self.model_radius / 20) - self.units = str(model.units) - self.title = title if title is not None else str(model.name) - self.cmap_names = list(cmaps or _DEFAULT_CMAPS) - self.cmap_luts = _colormap_luts(self.cmap_names) - self.cmap = cmap if cmap in self.cmap_names else self.cmap_names[0] - self.canvas_size = int(canvas_size) - self.marker_size = float(marker_size) - self.dark_background = bool(dark_background) - self.clip_thickness = float(3 * self.bond_length) - self.update_channels() - if channel is None: - channel = "structure" if "structure" in self.channel_names else self.channel_names[0] - self.channel = channel - - # ------------------------------------------------------------------ # - def update_channels(self) -> None: - """Re-send all channels of the model (call after adding new channels).""" - model = self._model - names = ["x", "y", "z"] + list(model.channels) - data = np.stack([model.get_channel(n) for n in names], axis=0).astype(np.float32) - cats: dict[str, list[str]] = {} - palettes: dict[str, list[int]] = {} - for name, labels in model.categories.items(): - cats[name] = list(labels) - n_cat = max( - len(labels), int(np.nanmax(model.get_channel(name))) + 1 if model.num_sites else 1 - ) - if name == "grain" or n_cat > len(_CATEGORICAL_COLORS): - rng = np.random.default_rng(0) - colors = rng.uniform(0.15, 0.95, (n_cat, 3)) - else: - colors = _CATEGORICAL_COLORS[np.arange(n_cat) % len(_CATEGORICAL_COLORS)] - palettes[name] = (colors * 255).round().astype(int).ravel().tolist() - for name in ("grain",): - if name in model.channels and name not in cats: - codes = model.get_channel(name) - n_cat = int(np.nanmax(codes)) + 1 if codes.size else 1 - cats[name] = [str(i) for i in range(n_cat)] - rng = np.random.default_rng(0) - palettes[name] = ( - (rng.uniform(0.15, 0.95, (max(n_cat, 1), 3)) * 255) - .round() - .astype(int) - .ravel() - .tolist() - ) - self.channel_names = names - self.channel_data = np.ascontiguousarray(data).tobytes() - self.channel_categories = cats - self.channel_palettes = palettes - - @property - def view_matrix(self) -> NDArray: - """``(3, 3)`` current view matrix (rows: right, up, toward-viewer).""" - return np.asarray(self.rotation, dtype=float).reshape(3, 3) - - @property - def view_direction(self) -> NDArray: - """``(3,)`` unit vector pointing from the model toward the viewer.""" - return self.view_matrix[2] - - def set_view(self, normal: str | NDArray, up: str | NDArray | None = None) -> None: - """Look along ``normal`` (as in :func:`quantem.atoms.visualization.view_matrix`).""" - from quantem.atoms.visualization import view_matrix - - self.rotation = view_matrix(normal, up).ravel().tolist() - - def set_range(self, vmin: float, vmax: float) -> None: - """Set the color range and disable auto-ranging.""" - self.range_auto = False - self.vmin, self.vmax = float(vmin), float(vmax) - - def clip( - self, axis: str = "view", center: float = 0.0, thickness: float | None = None - ) -> None: - """Enable slab clipping along ``axis`` (``"view"``, ``"x"``, ``"y"`` or ``"z"``).""" - self.clip_axis = axis - self.clip_center = float(center) - if thickness is not None: - self.clip_thickness = float(thickness) - self.clip_enabled = True - - def slab_kwargs(self) -> dict[str, Any]: - """Keyword arguments reproducing the current view with ``model.plot("slab", ...)``.""" - vm = self.view_matrix - out: dict[str, Any] = {"normal": vm[2], "up": vm[1], "channel": self.channel} - if self.clip_enabled and self.clip_axis == "view": - out.update(offset=self.clip_center, thickness=self.clip_thickness) - if not self.range_auto: - out.update(vmin=self.vmin, vmax=self.vmax) - if self.hide_outside_range: - out["hide"] = None - return out - - def __repr__(self) -> str: - return f"ShowAtoms3D({self.num_sites} sites, channel={self.channel!r}, clip={self.clip_enabled})" - - def _repr_json_state(self) -> str: # pragma: no cover - debugging aid - return json.dumps( - {k: getattr(self, k) for k in ("channel", "cmap", "vmin", "vmax", "rotation", "zoom")} - ) +from quantem.widget.show_atoms3d import ShowAtoms3D as ShowAtoms3D diff --git a/src/quantem/atoms/visualization.py b/src/quantem/atoms/visualization.py index ab31cd75a..4c5a4d637 100644 --- a/src/quantem/atoms/visualization.py +++ b/src/quantem/atoms/visualization.py @@ -127,6 +127,48 @@ def channel_colors( return rgb, colormap, (vmin, vmax), None +def draw_spheres(ax, x, y, rgb, marker_size, edgecolor="k", num_layers: int = 6): + """Draw shaded sphere markers with a stack of offset, brightening discs. + + Parameters + ---------- + ax : Axes + Target axes. + x, y : ndarray + Marker centers (already sorted back-to-front). + rgb : ndarray + ``(N, 3)`` base colors. + marker_size : float + Scatter marker area (points^2) of the full sphere. + edgecolor : str or None + Outline color of the full disc. + num_layers : int + Number of highlight discs stacked toward the upper-left. + """ + radius_pt = np.sqrt(marker_size) / 2.0 + # convert an offset in points to data units using the axes transform + fig = ax.figure + axes_width_pt = ax.get_position().width * fig.get_size_inches()[0] * 72.0 + xlim = ax.get_xlim() if ax.has_data() else (np.min(x), np.max(x)) + data_per_pt = (abs(xlim[1] - xlim[0]) or 1.0) / axes_width_pt + rgb = np.asarray(rgb) + ax.scatter( + x, y, s=marker_size, c=rgb * 0.55, edgecolors=edgecolor, linewidths=0.4 if edgecolor else 0 + ) + for i in range(1, num_layers + 1): + t = i / num_layers + scale = 1.0 - 0.75 * t + shift = 0.28 * t * radius_pt * data_per_pt + color = np.clip(rgb * (0.55 + 0.6 * t) + 0.35 * t**2, 0, 1) + ax.scatter( + x - shift, + y - shift, + s=marker_size * scale**2, + c=color, + edgecolors="none", + ) + + # --------------------------------------------------------------------------- # # Plots # --------------------------------------------------------------------------- # @@ -224,6 +266,7 @@ def plot_slab( depth_cue: float = 0.35, hide: tuple[float, float] | None = None, edgecolor: str | None = "k", + style: str = "spheres", ax=None, figsize: tuple[float, float] = (7, 7), colorbar: bool = True, @@ -254,6 +297,9 @@ def plot_slab( Hide sites whose channel value lies inside this range. edgecolor : str or None Marker edge color. + style : {"spheres", "flat"} + ``"spheres"`` draws shaded spheres (a stack of offset, brightening + discs per site); ``"flat"`` draws plain scatter markers. """ uvd, keep, _ = _project(model, normal, up, offset, thickness, hide, channel) rgb, colormap, (lo, hi), categories = channel_colors(model, channel, cmap, vmin, vmax) @@ -275,14 +321,17 @@ def plot_slab( axes_width_pt = ax.get_position().width * fig.get_size_inches()[0] * 72.0 diameter_pt = 0.9 * model.bond_length * axes_width_pt / extent marker_size = max(diameter_pt**2, 1.0) - sc = ax.scatter( - uvd[:, 1], - uvd[:, 0], - s=marker_size, - c=rgb, - edgecolors=edgecolor, - linewidths=0.3 if edgecolor else 0, - ) + if style == "spheres": + draw_spheres(ax, uvd[:, 1], uvd[:, 0], rgb, marker_size, edgecolor=edgecolor) + else: + ax.scatter( + uvd[:, 1], + uvd[:, 0], + s=marker_size, + c=rgb, + edgecolors=edgecolor, + linewidths=0.3 if edgecolor else 0, + ) ax.set_aspect("equal") ax.invert_yaxis() ax.set_xlabel(f"v [{model.units}]") @@ -305,7 +354,6 @@ def plot_slab( else: sm = plt.cm.ScalarMappable(cmap=colormap, norm=plt.Normalize(lo, hi)) fig.colorbar(sm, ax=ax, fraction=0.04, pad=0.02, label=channel) - del sc return (fig, ax) if returnfig else None diff --git a/tests/atoms/test_atoms.py b/tests/atoms/test_atoms.py index aed14828a..991ff8920 100644 --- a/tests/atoms/test_atoms.py +++ b/tests/atoms/test_atoms.py @@ -236,3 +236,20 @@ def test_atomic_model_channels_and_shapes(): model.set_channel("bar", np.zeros(3)) with pytest.raises(KeyError): model.get_channel("missing") + + +def test_merge_close_sites(): + xyz, _ = make_lattice("fcc", n=5, noise=0.01) + dup = np.vstack([xyz, xyz[:3] + 0.05, xyz[3:4] + np.array([0.02, 0.0, 0.0])]) + model = AtomicModel.from_array(dup) + model.set_channel("intensity", np.r_[np.ones(xyz.shape[0]), 3.0, 3.0, 3.0, 3.0]) + removed = model.merge_close_sites(min_distance=0.3) + assert removed == 4 and model.num_sites == xyz.shape[0] + # merged positions are the intensity-weighted mean and the channel is averaged + model2 = AtomicModel.from_array(dup) + model2.set_channel("intensity", np.r_[np.ones(xyz.shape[0]), 3.0, 3.0, 3.0, 3.0]) + model2.merge_close_sites(min_distance=0.3, weight_channel="intensity") + assert model2["intensity"].max() == pytest.approx(2.5) + model3 = AtomicModel.from_array(dup) + assert model3.merge_close_sites(min_distance=0.3, mode="remove") == 4 + assert model3.merge_close_sites(min_distance=0.3) == 0 From 95b9c46855d5c95f331af6b6ff3ca5ab19add3f1 Mon Sep 17 00:00:00 2001 From: cophus Date: Tue, 15 Sep 2026 17:46:02 -0700 Subject: [PATCH 5/6] updates --- src/quantem/atoms/atomic_model.py | 277 ++++++++++++++++++++++++ src/quantem/atoms/measurements.py | 80 +++++++ src/quantem/atoms/structures.py | 329 +++++++++++++++++++++++++++++ src/quantem/atoms/visualization.py | 25 ++- tests/atoms/test_atoms.py | 70 ++++++ 5 files changed, 774 insertions(+), 7 deletions(-) create mode 100644 src/quantem/atoms/structures.py diff --git a/src/quantem/atoms/atomic_model.py b/src/quantem/atoms/atomic_model.py index 7e3bdc8e6..28cda5ea5 100644 --- a/src/quantem/atoms/atomic_model.py +++ b/src/quantem/atoms/atomic_model.py @@ -931,6 +931,283 @@ def classify_species( self._metadata["species_centers"] = centers.tolist() return labels + # ------------------------------------------------------------------ # + # Multiply twinned particle geometry + # ------------------------------------------------------------------ # + def twin_plane_normals(self, structure: str = "hcp") -> NDArray: + """``(N, 3)`` twin-plane normal at every site from the ``hcp`` template fit. + + The hexagonal ``c`` axis of the fitted HCP template is the normal of + the close-packed plane, which for a site on a coherent twin boundary + is the twin plane. Values are only meaningful at sites classified as + ``structure``. + """ + return self._matches[structure]["rotation"][:, :, 2].copy() + + def fit_icosahedral_centers( + self, + num_centers: int = 2, + min_score: float = 0.6, + max_residual: float = 2.0, + axial_tolerance: float = 0.35, + num_iterations: int = 10, + ) -> dict[str, Any]: + """Locate the centers of two icosahedra that share a 5-fold axis. + + All twin planes of one Mackay icosahedron pass through its center, so + the center is the least-squares intersection of the planes carried by + its twin (``hcp``) sites. Planes that contain the shared axis pass + through both centers and are excluded; the remaining sites are + assigned to the center whose planes they fit best, outliers beyond + ``max_residual`` are dropped, and the fit is iterated. The result + is stored in ``metadata["icosahedral_centers"]``. + + Parameters + ---------- + num_centers : int + Only ``2`` is supported at present. + min_score : float + Minimum ``score_hcp`` of the twin sites used. + max_residual : float + Plane-distance cutoff (calibrated units) for keeping a site. + axial_tolerance : float + Planes whose normal has ``|n . axis|`` below this are treated as + containing the axis and excluded. + num_iterations : int + Outer iterations of the assignment. + + Returns + ------- + dict + ``centers`` (2, 3) relative to the model center, ``axis`` unit + vector from center 0 to center 1, ``separation`` in calibrated + units, ``rms`` plane residual per center, ``num_sites`` per + center, and ``assignment`` (``(N,)`` with ``-1`` = unused). + """ + if num_centers != 2: + raise NotImplementedError("Only two centers are supported.") + if "hcp" not in self._matches: + raise RuntimeError("Match the 'hcp' template first (match_templates(['fcc', 'hcp'])).") + xyz = self.positions - self.center[None, :] + normals = self.twin_plane_normals("hcp") + first_shell = (self.neighbor_distances <= self.first_shell_cutoff).sum(1) + twin = self.get_channel("structure").astype(int) == self._structure_names.index("hcp") + twin &= (self._matches["hcp"]["score"] >= min_score) & (first_shell >= 12) + center0, res = meas.fit_plane_intersection(xyz[twin], normals[twin]) + # the shared axis is the direction most perpendicular to the well-fitting planes + good = np.abs(res) < 1.5 * max_residual + cov = np.einsum("ni,nj->ij", normals[twin][good], normals[twin][good]) + axis = np.linalg.eigh(cov)[1][:, 0] + c1, c2 = center0 - 0.5 * axis, center0 + 0.5 * axis + assign = np.full(self.num_sites, -1) + for _ in range(num_iterations): + non_axial = np.abs(normals @ axis) > axial_tolerance + usable = twin & non_axial + r1 = np.abs(np.einsum("ni,ni->n", normals, xyz - c1)) + r2 = np.abs(np.einsum("ni,ni->n", normals, xyz - c2)) + a1 = usable & (r1 <= r2) & (r1 < max_residual) + a2 = usable & (r2 < r1) & (r2 < max_residual) + if a1.sum() < 10 or a2.sum() < 10: + h = (xyz - center0) @ axis + a1 = usable & (h < 0) + a2 = usable & (h >= 0) + c1n, _ = meas.fit_plane_intersection(xyz[a1], normals[a1]) + c2n, _ = meas.fit_plane_intersection(xyz[a2], normals[a2]) + shift = np.linalg.norm(c1n - c1) + np.linalg.norm(c2n - c2) + c1, c2 = c1n, c2n + axis = (c2 - c1) / np.linalg.norm(c2 - c1) + if shift < 1e-4: + break + assign[a1] = 0 + assign[a2] = 1 + rms = [ + float(np.sqrt(np.mean(np.einsum("ni,ni->n", normals[a], xyz[a] - c) ** 2))) + for a, c in ((a1, c1), (a2, c2)) + ] + result = { + "centers": np.stack([c1, c2]), + "axis": axis, + "separation": float(np.linalg.norm(c2 - c1)), + "rms": np.array(rms), + "num_sites": np.array([int(a1.sum()), int(a2.sum())]), + "assignment": assign, + } + self._metadata["icosahedral_centers"] = { + "centers": result["centers"].tolist(), + "axis": axis.tolist(), + "separation": result["separation"], + } + return result + + def layer_positions( + self, + normal: str | NDArray, + bin_width: float | None = None, + sigma: float | None = None, + min_fraction: float = 0.25, + ) -> NDArray: + """Positions of the atomic layers perpendicular to ``normal``. + + Sites are projected onto ``normal`` (measured from the model center), + the projected density is histogrammed and smoothed, and its peaks are + returned. The result can be passed as ``positions`` to + ``plot("slices", ...)`` to step through the model one layer at a time. + + Parameters + ---------- + normal : str or array + Layer normal (``"x"``, ``"y"``, ``"z"`` or a 3-vector). + bin_width : float, optional + Histogram bin (default: bond length / 40). + sigma : float, optional + Smoothing (default: bond length / 12). + min_fraction : float + Peaks below this fraction of the highest peak are ignored. + + Returns + ------- + ndarray + Layer offsets along ``normal`` relative to the model center. + """ + from quantem.atoms.visualization import view_matrix + + n = view_matrix(normal)[2] + h = (self.positions - self.center[None, :]) @ n + bond = self.bond_length + pos, _, _ = meas.layer_positions( + h, + bin_width=bond / 40.0 if bin_width is None else bin_width, + sigma=bond / 12.0 if sigma is None else sigma, + min_fraction=min_fraction, + ) + return pos + + def explode_grains( + self, + distance: float | None = None, + origin: str | NDArray | None = None, + include_shared: bool = True, + keep_other: bool = False, + structure: str = "fcc", + ) -> "AtomicModel": + """Return a copy with every grain displaced away from its origin. + + Each grain (sector) moves rigidly by ``distance`` along the direction + from ``origin`` to its centroid, so the grains separate and their + boundaries become visible. Twin sites and other boundary sites that + touch several grains are copied into every adjacent grain, displaced + with it, and labelled ``shared`` in the new categorical ``grain`` + channel; sites classified as neither ``structure`` nor a twin are + labelled ``other`` and kept in place when ``keep_other`` is True. + + Parameters + ---------- + distance : float, optional + Displacement per grain (calibrated units); default 2 bond lengths. + origin : {"center", "icosahedral"} or array, optional + Point the grains move away from: the model center (default), the + nearest of the fitted icosahedral centers (after + :meth:`fit_icosahedral_centers`), or an explicit ``(3,)`` point. + include_shared : bool + Copy boundary sites into each adjacent grain. + keep_other : bool + Keep unclassified sites (``other``) at their original positions. + structure : str + Template name of the grains (``"fcc"``). + + Returns + ------- + AtomicModel + New model with channels ``grain`` (categorical), ``source_index`` + (index in this model) and copies of all other channels. + """ + if "grain" not in self.channels: + raise RuntimeError("Call segment_grains() first.") + if distance is None: + distance = 2.0 * self.bond_length + xyz = self.positions + grain = self.get_channel("grain").astype(int) + struct = self.get_channel("structure").astype(int) + num_grains = int(grain.max()) + 1 + centroids = np.array([xyz[grain == g].mean(0) for g in range(num_grains)]) + if origin is None or origin == "center": + origins = np.tile(self.center, (num_grains, 1)) + elif isinstance(origin, str) and origin == "icosahedral": + info = self._metadata.get("icosahedral_centers") + if info is None: + raise RuntimeError( + "Call fit_icosahedral_centers() first for origin='icosahedral'." + ) + cents = np.asarray(info["centers"]) + self.center[None, :] + nearest = np.argmin( + np.linalg.norm(centroids[:, None, :] - cents[None, :, :], axis=2), axis=1 + ) + origins = cents[nearest] + else: + origins = np.tile(np.asarray(origin, dtype=float).reshape(3), (num_grains, 1)) + direction = centroids - origins + norm = np.linalg.norm(direction, axis=1, keepdims=True) + shift = distance * direction / np.where(norm > 1e-9, norm, 1.0) + + table = np.array(self._sites.array, dtype=float) + pos_cols = [self._sites.fields.index(f) for f in _POSITION_FIELDS] + scale = self._sampling[None, :] + rows, labels, source = [], [], [] + # grain members + member = grain >= 0 + idx = np.where(member)[0] + t = table[idx].copy() + t[:, pos_cols] += shift[grain[idx]] / scale + rows.append(t) + labels.append(grain[idx]) + source.append(idx) + # shared boundary sites: copied into each adjacent grain + twin_id = [i for i, n in enumerate(self._structure_names) if n != structure] + boundary = (~member) & ( + np.isin(struct, twin_id) | (struct == self._structure_names.index(structure)) + ) + if include_shared and boundary.any(): + nb_idx = self.neighbor_indices + first = self.neighbor_distances <= self.first_shell_cutoff + nb_grain = np.where(first & (nb_idx >= 0), grain[np.where(nb_idx >= 0, nb_idx, 0)], -1) + for i in np.where(boundary)[0]: + adjacent = np.unique(nb_grain[i][nb_grain[i] >= 0]) + for g in adjacent: + r = table[i].copy() + r[pos_cols] += shift[g] / scale[0] + rows.append(r[None, :]) + labels.append(np.array([num_grains])) + source.append(np.array([i])) + other = (~member) & ~boundary + if keep_other and other.any(): + idx = np.where(other)[0] + rows.append(table[idx]) + labels.append(np.full(idx.size, num_grains + 1)) + source.append(idx) + new_table = np.vstack(rows) + sites = Vector.from_shape( + shape=(), fields=list(self._sites.fields), units=list(self._sites.units), name="sites" + ) + sites[...] = np.ascontiguousarray(new_table) + out = AtomicModel( + sites=sites, + sampling=self._sampling.copy(), + units=self._units, + name=f"{self._name} (exploded)", + metadata={"exploded_distance": float(distance)}, + _token=self._token, + ) + out._categories = {k: v for k, v in self._categories.items() if k != "grain"} + out._structure_names = list(self._structure_names) + out._nn_fit = None if self._nn_fit is None else dict(self._nn_fit) + out.set_channel( + "grain", + np.concatenate(labels), + categories=[str(g) for g in range(num_grains)] + ["shared", "other"], + ) + out.set_channel("source_index", np.concatenate(source)) + return out + # ------------------------------------------------------------------ # # Geometry helpers # ------------------------------------------------------------------ # diff --git a/src/quantem/atoms/measurements.py b/src/quantem/atoms/measurements.py index 2d714c598..27de742f5 100644 --- a/src/quantem/atoms/measurements.py +++ b/src/quantem/atoms/measurements.py @@ -18,6 +18,8 @@ "segment_grains", "fill_labels", "strain_from_deformation", + "fit_plane_intersection", + "layer_positions", "bond_angles", "convex_hull_distance", "sample_volume", @@ -303,3 +305,81 @@ def rotation_to_quaternion(rotation: NDArray) -> NDArray: q = np.roll(q, 1, axis=1) q[q[:, 0] < 0] *= -1 return q + + +def fit_plane_intersection(points: NDArray, normals: NDArray) -> tuple[NDArray, NDArray]: + """Least-squares point closest to a set of planes. + + Each plane passes through ``points[i]`` with unit normal ``normals[i]``; + the returned point minimizes the sum of squared plane distances. + + Returns + ------- + center, residuals : ndarray + ``(3,)`` point and ``(N,)`` signed distances of the point to each plane. + """ + n = np.asarray(normals, dtype=float) + p = np.asarray(points, dtype=float) + a = np.einsum("ni,nj->ij", n, n) + b = np.einsum("ni,ni,nj->j", n, p, n) + center = np.linalg.solve(a + 1e-9 * np.eye(3), b) + return center, np.einsum("ni,ni->n", n, center[None, :] - p) + + +def layer_positions( + heights: NDArray, + bin_width: float, + sigma: float, + min_fraction: float = 0.25, +) -> tuple[NDArray, NDArray, NDArray]: + """Peaks of the site density along one direction (atomic layers). + + Parameters + ---------- + heights : ndarray + ``(N,)`` coordinates of the sites along the direction. + bin_width : float + Histogram bin width. + sigma : float + Gaussian smoothing of the histogram (same units as ``heights``). + min_fraction : float + Peaks lower than this fraction of the highest peak are ignored. + + Returns + ------- + positions, centers, density : ndarray + Peak positions, and the smoothed histogram (bin centers and counts) + for plotting. + """ + from scipy.ndimage import gaussian_filter1d + + h = np.asarray(heights, dtype=float) + edges = np.arange(h.min() - 2 * sigma, h.max() + 2 * sigma + bin_width, bin_width) + counts = np.histogram(h, edges)[0].astype(float) + smooth = gaussian_filter1d(counts, sigma / bin_width) if sigma > 0 else counts + centers = 0.5 * (edges[1:] + edges[:-1]) + inner = smooth[1:-1] + peaks = ( + np.where( + (inner > smooth[:-2]) & (inner >= smooth[2:]) & (inner > min_fraction * smooth.max()) + )[0] + + 1 + ) + # refine each peak with a parabola through its three bins + pos, height = [], [] + for k in peaks: + y0, y1, y2 = smooth[k - 1], smooth[k], smooth[k + 1] + denom = y0 - 2 * y1 + y2 + delta = 0.5 * (y0 - y2) / denom if abs(denom) > 1e-12 else 0.0 + pos.append(centers[k] + delta * bin_width) + height.append(y1) + pos, height = np.asarray(pos), np.asarray(height) + # drop the weaker of any two peaks closer than half the median spacing + if pos.size > 2: + spacing = np.median(np.diff(pos)) + keep = np.ones(pos.size, dtype=bool) + for k in range(1, pos.size): + if pos[k] - pos[k - 1] < 0.5 * spacing: + keep[k if height[k] < height[k - 1] else k - 1] = False + pos = pos[keep] + return pos, centers, smooth diff --git a/src/quantem/atoms/structures.py b/src/quantem/atoms/structures.py new file mode 100644 index 000000000..f052ba789 --- /dev/null +++ b/src/quantem/atoms/structures.py @@ -0,0 +1,329 @@ +"""Ideal nanoparticle structures for comparison with measured atomic models. + +The builders return :class:`~quantem.atoms.AtomicModel` objects in physical +units with two channels: ``shell`` (the shell index of each site, 0 at the +center) and ``sector`` (the tetrahedral sector of multiply twinned particles, +or 0 for single crystals). Distances follow the Mackay convention: sites on +a radial line from the center are spaced by ``bond_length`` and sites within +a shell are 5.15% farther apart, so the 20 tetrahedra of an icosahedron are +slightly distorted FCC. + +Structures +---------- +``icosahedron`` + Mackay icosahedron with ``num_shells`` shells; 20 FCC tetrahedra sharing + one center. +``double_icosahedron`` + Two interpenetrating Mackay icosahedra whose centers are one bond apart + along a common 5-fold axis, related by a mirror through the mid-plane + between the centers (point group D5h). Each half keeps the sites on its + own side of the mid-plane, giving 15 tetrahedral sectors per half plus the + small shared pentagonal bipyramid between the centers. This is the + polyicosahedral motif of the 19-atom double icosahedron, grown shell by + shell. +``attached_icosahedra`` + Two complete Mackay icosahedra of the same size touching at one vertex on + a common 5-fold axis, mirror-related (the geometry of oriented attachment + of two grown particles). +``cuboctahedron`` + FCC cuboctahedron with ``num_shells`` shells (single crystal). +``decahedron`` + Ino decahedron: five FCC tetrahedra sharing a common edge (the 5-fold axis) + with ``num_shells`` shells and no re-entrant Marks facets. +""" + +from __future__ import annotations + +import numpy as np +from numpy.typing import NDArray + +from quantem.atoms.atomic_model import AtomicModel + +__all__ = [ + "icosahedron", + "double_icosahedron", + "cuboctahedron", + "decahedron", + "icosahedron_vertices", +] + +_MACKAY_TANGENTIAL = 1.0515 # tangential / radial spacing of a Mackay icosahedron + + +def icosahedron_vertices() -> tuple[NDArray, NDArray]: + """Unit vertex vectors of an icosahedron and its 20 triangular faces. + + The first vertex points along ``+z`` so that a 5-fold axis is ``z``. + + Returns + ------- + vertices, faces : ndarray + ``(12, 3)`` unit vectors and ``(20, 3)`` vertex indices. + """ + phi = (1.0 + np.sqrt(5.0)) / 2.0 + v = [] + for s1 in (-1, 1): + for s2 in (-1, 1): + v.append([0.0, s1, s2 * phi]) + v.append([s1, s2 * phi, 0.0]) + v.append([s2 * phi, 0.0, s1]) + v = np.array(v) + v /= np.linalg.norm(v, axis=1, keepdims=True) + # rotate so that vertex 0 is along +z + z = v[np.argmax(v[:, 2])] + axis = np.cross(z, [0, 0, 1.0]) + if np.linalg.norm(axis) > 1e-9: + axis /= np.linalg.norm(axis) + ang = np.arccos(np.clip(z @ [0, 0, 1.0], -1, 1)) + k = np.array([[0, -axis[2], axis[1]], [axis[2], 0, -axis[0]], [-axis[1], axis[0], 0]]) + rot = np.eye(3) + np.sin(ang) * k + (1 - np.cos(ang)) * k @ k + v = v @ rot.T + order = np.argsort(-v[:, 2], kind="stable") + v = v[order] + # faces: triples of mutually adjacent vertices (edge length = 1.0515) + edge = np.linalg.norm(v[0] - v[1]) + faces = [] + for i in range(12): + for j in range(i + 1, 12): + for k in range(j + 1, 12): + d = [np.linalg.norm(v[a] - v[b]) for a, b in ((i, j), (j, k), (i, k))] + if all(abs(x - edge) < 1e-6 for x in d): + faces.append((i, j, k)) + return v, np.array(faces) + + +def _sector_sites(num_shells: int) -> tuple[NDArray, NDArray, NDArray]: + """Sites of all 20 tetrahedral sectors of a Mackay icosahedron (unit radial spacing).""" + v, faces = icosahedron_vertices() + pts, shell, sector = [np.zeros((1, 3))], [np.zeros(1, int)], [np.full(1, -1)] + for f, (i, j, k) in enumerate(faces): + for n in range(1, num_shells + 1): + for a in range(n + 1): + for b in range(n + 1 - a): + c = n - a - b + pts.append((a * v[i] + b * v[j] + c * v[k])[None, :]) + shell.append(np.array([n])) + sector.append(np.array([f])) + return np.vstack(pts), np.concatenate(shell), np.concatenate(sector) + + +def _set_sector(model: AtomicModel, sector: NDArray) -> None: + """Store the sector labels as a categorical channel (``-1`` = center site).""" + sector = np.asarray(sector, dtype=int) + num = int(sector.max()) + 1 if sector.size else 1 + model.set_channel("sector", sector, categories=[str(i) for i in range(num)]) + + +def _dedupe(xyz: NDArray, *channels: NDArray, tol: float = 1e-6) -> tuple[NDArray, list[NDArray]]: + key = np.round(xyz / tol).astype(np.int64) + _, first = np.unique(key, axis=0, return_index=True) + first = np.sort(first) + return xyz[first], [c[first] for c in channels] + + +def icosahedron( + num_shells: int, bond_length: float = 1.0, name: str = "icosahedron" +) -> AtomicModel: + """Mackay icosahedron. + + Parameters + ---------- + num_shells : int + Number of shells around the central site (``K``); the particle has + ``(10 K^3 + 15 K^2 + 11 K + 3) / 3`` sites. + bond_length : float + Radial nearest-neighbor spacing in physical units. + name : str + Model name. + """ + xyz, shell, sector = _sector_sites(num_shells) + xyz, (shell, sector) = _dedupe(xyz, shell, sector) + model = AtomicModel.from_array(xyz * bond_length, units="A", name=name) + model.set_channel("shell", shell) + _set_sector(model, sector) + return model + + +def double_icosahedron( + num_shells: int, + bond_length: float = 1.0, + separation: int = 1, + name: str = "double icosahedron", +) -> AtomicModel: + """Two Mackay icosahedra sharing a 5-fold axis, mirror-related (D5h). + + The lower icosahedron is centered at ``(0, 0, -separation / 2)`` (in bond + lengths) and the upper one at ``(0, 0, +separation / 2)``; each half + contributes the sites on its own side of the mid-plane ``z = 0``, and + sites within half a bond of their mirror image are merged onto the plane. + ``separation = 1`` is the polyicosahedral double icosahedron (19 sites for + one shell); larger odd separations place the second center on the 5-fold + axis of the first icosahedron at the position of one of its vertices, so + that the two halves are twins across the mid-plane. Sectors ``0-19`` + belong to the lower half and ``20-39`` to the upper half. + + Parameters + ---------- + num_shells : int + Shells of each icosahedron. + bond_length : float + Radial nearest-neighbor spacing. + separation : int + Distance between the two centers in bond lengths. + name : str + Model name. + """ + xyz, shell, sector = _sector_sites(num_shells) + lower = xyz + np.array([0.0, 0.0, -0.5 * separation]) + keep_lo = lower[:, 2] <= 1e-6 + upper = lower * np.array([1.0, 1.0, -1.0]) # mirror through z = 0 + keep_up = upper[:, 2] >= -1e-6 + pts = np.vstack([lower[keep_lo], upper[keep_up]]) + sh = np.concatenate([shell[keep_lo], shell[keep_up]]) + sec = np.concatenate([sector[keep_lo], sector[keep_up] + 20]) + pts, (sh, sec) = _dedupe(pts, sh, sec) + model = AtomicModel.from_array(pts * bond_length, units="A", name=name) + model.set_channel("shell", sh) + _set_sector(model, sec) + # Mirror images of sites just below the mid-plane lie 0.11, 0.32, 0.53 or + # 0.74 bonds from each other (the Mackay heights i + 0.447 j do not fall on + # the plane); merging every pair closer than 0.75 bonds onto the plane gives + # the shared mid-plane layer of the D5h structure (19 sites for one shell + # at separation 1) with all remaining distances >= 1 bond. + model.merge_close_sites(min_distance=0.75 * bond_length) + model.metadata["centers"] = [ + [0.0, 0.0, -0.5 * separation * bond_length], + [0.0, 0.0, 0.5 * separation * bond_length], + ] + model.metadata["separation"] = int(separation) + return model + + +def cuboctahedron( + num_shells: int, bond_length: float = 1.0, name: str = "cuboctahedron" +) -> AtomicModel: + """FCC cuboctahedron with ``num_shells`` shells around a central site.""" + a = bond_length * np.sqrt(2.0) + n = num_shells + rng = np.arange(-n, n + 1) + ijk = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + basis = np.array([[0, 0, 0], [0, 0.5, 0.5], [0.5, 0, 0.5], [0.5, 0.5, 0]]) + xyz = ((ijk[:, None, :] + basis[None]).reshape(-1, 3)) * a + # shell k is bounded by the cube |x_i| <= k a / 2 and the octahedron |x|+|y|+|z| <= k a + ax = np.abs(xyz) / a + shell = np.rint(np.maximum(2.0 * ax.max(1), ax.sum(1))).astype(int) + keep = shell <= n + model = AtomicModel.from_array(xyz[keep], units="A", name=name) + model.set_channel("shell", shell[keep]) + _set_sector(model, np.zeros(keep.sum(), dtype=int)) + return model + + +def decahedron(num_shells: int, bond_length: float = 1.0, name: str = "decahedron") -> AtomicModel: + """Ino decahedron: five FCC tetrahedra sharing the 5-fold axis ``z``. + + One sector is cut from an FCC crystal as the 70.53 degree wedge between a + (111) and a (11-1) plane meeting along [1-10]; its azimuth about the axis + is stretched to 72 degrees and five copies are placed around the axis. + ``num_shells`` counts the (111) layers from the axis to the surface, and + sites on the shared twin planes and the axis appear once. + + Parameters + ---------- + num_shells : int + Number of (111) layers per sector. + bond_length : float + Nearest-neighbor spacing before the 72/70.53 azimuthal stretch. + name : str + Model name. + """ + a = np.sqrt(2.0) # cubic lattice constant for unit NN spacing + n = num_shells + 2 + rng = np.arange(-n, n + 1) + ijk = np.stack(np.meshgrid(rng, rng, rng, indexing="ij"), -1).reshape(-1, 3) + basis = np.array([[0, 0, 0], [0, 0.5, 0.5], [0.5, 0, 0.5], [0.5, 0.5, 0]]) + fcc = ((ijk[:, None, :] + basis[None]).reshape(-1, 3)) * a + n1 = np.array([1.0, 1.0, 1.0]) / np.sqrt(3) + n2 = np.array([1.0, 1.0, -1.0]) / np.sqrt(3) + d111 = a / np.sqrt(3) + l1 = fcc @ n1 / d111 + l2 = -(fcc @ n2) / d111 + eps = 1e-6 + layer = np.rint(l1 + l2).astype(int) + u = np.array([1.0, -1.0, 0.0]) / np.sqrt(2) + keep = ( + (l1 >= -eps) & (l2 >= -eps) & (layer <= num_shells) & (np.abs(fcc @ u) <= num_shells + eps) + ) + pts = fcc[keep] + layer = layer[keep] + # cylindrical coordinates about the wedge axis u = [1,-1,0]; azimuth from the bisector + bis = n1 - n2 + bis /= np.linalg.norm(bis) # bisector of the wedge, perpendicular to u + e2 = np.cross(u, bis) + z = pts @ u + x = pts @ bis + y = pts @ e2 + rho = np.hypot(x, y) + phi = np.arctan2(y, x) * (72.0 / 70.528779) + out, sh, sec = [], [], [] + for s_idx in range(5): + ang = phi + np.radians(72.0 * s_idx) + out.append(np.stack([rho * np.cos(ang), rho * np.sin(ang), z], 1)) + sh.append(layer) + sec.append(np.full(layer.size, s_idx)) + xyz = np.vstack(out) * bond_length + xyz, (sh, sec) = _dedupe(xyz, np.concatenate(sh), np.concatenate(sec), tol=1e-3 * bond_length) + model = AtomicModel.from_array(xyz, units="A", name=name) + model.set_channel("shell", sh) + _set_sector(model, sec) + return model + + +def attached_icosahedra( + num_shells: int, bond_length: float = 1.0, name: str = "attached icosahedra" +) -> AtomicModel: + """Two complete Mackay icosahedra sharing one vertex site on a 5-fold axis. + + The lower icosahedron is centered at ``z = -num_shells`` bonds and the + upper one, its mirror image, at ``z = +num_shells``; the vertex at the + origin is shared. Sectors ``0-19`` and ``20-39`` label the two particles. + """ + xyz, shell, sector = _sector_sites(num_shells) + lower = xyz + np.array([0.0, 0.0, -float(num_shells)]) + upper = lower * np.array([1.0, 1.0, -1.0]) + pts = np.vstack([lower, upper]) + sh = np.concatenate([shell, shell]) + sec = np.concatenate([sector, sector + 20]) + pts, (sh, sec) = _dedupe(pts, sh, sec, tol=1e-6) + model = AtomicModel.from_array(pts * bond_length, units="A", name=name) + model.set_channel("shell", sh) + _set_sector(model, sec) + model.metadata["centers"] = [ + [0.0, 0.0, -num_shells * bond_length], + [0.0, 0.0, num_shells * bond_length], + ] + return model + + +def growth_steps( + model: AtomicModel, origin: NDArray | None = None, name: str = "growth_step" +) -> NDArray: + """Growth step of every site as its distance from a nucleus in bond lengths. + + Sites are added in order of distance from ``origin`` (default: the first + entry of ``model.metadata["centers"]``, or the model center), rounded up + to whole bond lengths, and the result is stored as channel ``name`` so + that the viewer's growth slider replays an inside-out growth sequence. + + Returns + ------- + ndarray + ``(N,)`` integer growth steps. + """ + if origin is None: + centers = model.metadata.get("centers") + origin = np.asarray(centers[0]) if centers else model.center + r = np.linalg.norm(model.positions - np.asarray(origin, dtype=float)[None, :], axis=1) + steps = np.ceil(r / model.bond_length - 1e-6).astype(int) + model.set_channel(name, steps) + return steps diff --git a/src/quantem/atoms/visualization.py b/src/quantem/atoms/visualization.py index 4c5a4d637..1b4f8e086 100644 --- a/src/quantem/atoms/visualization.py +++ b/src/quantem/atoms/visualization.py @@ -366,6 +366,7 @@ def plot_slices( thickness: float | None = None, start: float | None = None, end: float | None = None, + positions: NDArray | None = None, ncols: int = 3, figsize_per: float = 4.0, returnfig: bool = False, @@ -376,24 +377,34 @@ def plot_slices( Parameters ---------- num_slices : int - Number of slabs. + Number of slabs (ignored when ``positions`` is given). thickness : float, optional Slab thickness; default is the step between slabs. start, end : float, optional Range of slab centers along ``normal`` (relative to the model center); default spans the model. + positions : array, optional + Explicit slab centers along ``normal``, e.g. the atomic layers from + ``model.layer_positions(normal)``. **kwargs Forwarded to :func:`plot_slab`. """ v = view_matrix(normal, kwargs.get("up")) depth = (model.positions - model.center[None, :]) @ v[2] - if start is None: - start = float(depth.min()) + 0.05 * np.ptp(depth) - if end is None: - end = float(depth.max()) - 0.05 * np.ptp(depth) - centers = np.linspace(start, end, num_slices) + if positions is not None: + centers = np.asarray(positions, dtype=float) + num_slices = centers.size + else: + if start is None: + start = float(depth.min()) + 0.05 * np.ptp(depth) + if end is None: + end = float(depth.max()) - 0.05 * np.ptp(depth) + centers = np.linspace(start, end, num_slices) if thickness is None: - thickness = float(centers[1] - centers[0]) if num_slices > 1 else np.ptp(depth) + if positions is not None and num_slices > 1: + thickness = 0.8 * float(np.median(np.diff(centers))) + else: + thickness = float(centers[1] - centers[0]) if num_slices > 1 else np.ptp(depth) nrows = int(np.ceil(num_slices / ncols)) fig, axes = plt.subplots( nrows, ncols, figsize=(figsize_per * ncols, figsize_per * nrows), squeeze=False diff --git a/tests/atoms/test_atoms.py b/tests/atoms/test_atoms.py index 991ff8920..f46de9f3f 100644 --- a/tests/atoms/test_atoms.py +++ b/tests/atoms/test_atoms.py @@ -253,3 +253,73 @@ def test_merge_close_sites(): model3 = AtomicModel.from_array(dup) assert model3.merge_close_sites(min_distance=0.3, mode="remove") == 4 assert model3.merge_close_sites(min_distance=0.3) == 0 + + +def test_ideal_structures(): + from quantem.atoms import structures as st + + assert st.icosahedron(3).num_sites == 147 + assert st.cuboctahedron(3).num_sites == 147 + assert st.double_icosahedron(1).num_sites == 19 + d = st.double_icosahedron(4, bond_length=2.7, separation=5) + xyz = d.positions_native + from scipy.spatial import cKDTree + + assert cKDTree(xyz).query(xyz, k=2)[0][:, 1].min() > 0.99 * 2.7 + # mirror symmetry through z = 0 + mirrored = xyz * np.array([1, 1, -1]) + assert cKDTree(xyz).query(mirrored)[0].max() < 1e-6 + deca = st.decahedron(3) + assert deca.num_sites > 100 and len(np.unique(deca["sector"])) == 5 + + +def test_fit_icosahedral_centers_and_layers(): + from quantem.atoms import structures as st + + model = st.double_icosahedron(7, bond_length=2.75, separation=5) + rng = np.random.default_rng(0) + model.positions_native = model.positions_native + rng.normal(0, 0.05, (model.num_sites, 3)) + model.compute_pdf() + model.find_neighbors(20) + model.match_templates(["fcc", "hcp"], progress=False, device="cpu") + fit = model.fit_icosahedral_centers() + assert abs(fit["separation"] / model.bond_length - 5.0) < 0.3 + assert abs(abs(fit["axis"][2]) - 1.0) < 0.02 + centers = fit["centers"] + model.center + assert np.allclose(np.sort(centers[:, 2]), [-2.5 * 2.75, 2.5 * 2.75], atol=0.6) + layers = model.layer_positions("z") + assert layers.size > 10 + assert np.all(np.diff(layers) > 0.4 * model.bond_length) + + +def test_explode_grains(): + xyz, above = make_twinned_fcc() + model = AtomicModel.from_array(xyz * 7.2, units="voxels") + model.compute_pdf() + model.find_neighbors(20) + model.match_templates(["fcc", "hcp"], progress=False, device="cpu") + model.segment_grains("fcc", angle_threshold=5.0, min_size=20) + ex = model.explode_grains(distance=10.0) + labels = ex["grain"].astype(int) + assert set(np.unique(labels)) <= {0, 1, 2} + assert ex.categories["grain"][2] == "shared" + # each grain moved rigidly by 10 units + for g in (0, 1): + src = ex["source_index"][labels == g].astype(int) + d = ex.positions[labels == g] - model.positions[src] + assert np.allclose(np.linalg.norm(d, axis=1), 10.0, atol=1e-6) + assert (labels == 2).sum() > 0 + assert ex.num_sites > (model["grain"] >= 0).sum() + + +def test_attached_icosahedra_and_growth_steps(): + from quantem.atoms import structures as st + + a = st.attached_icosahedra(3, bond_length=2.0) + assert a.num_sites == 2 * 147 - 1 + steps = st.growth_steps(a) + assert steps.min() == 0 and steps.max() == 9 + d = st.double_icosahedron(4, bond_length=2.0, separation=5) + steps = st.growth_steps(d) + # the upper half only appears once the front crosses the mid-plane at 2.5 bonds + assert steps[d.positions[:, 2] > 0.1].min() >= 3 From 325aa9c6da558fbe12c34243d53b6a712a1a89ec Mon Sep 17 00:00:00 2001 From: cophus Date: Wed, 16 Sep 2026 13:36:02 -0700 Subject: [PATCH 6/6] atomic stuff --- src/quantem/atoms/atomic_model.py | 140 +++++++++++++++--------------- src/quantem/atoms/measurements.py | 50 +++++++++++ tests/atoms/test_atoms.py | 21 +++++ 3 files changed, 140 insertions(+), 71 deletions(-) diff --git a/src/quantem/atoms/atomic_model.py b/src/quantem/atoms/atomic_model.py index 28cda5ea5..2684e6259 100644 --- a/src/quantem/atoms/atomic_model.py +++ b/src/quantem/atoms/atomic_model.py @@ -158,72 +158,6 @@ def from_array( model.set_channel(key, values) return model - @classmethod - def from_mat( - cls, - path: str | Path, - key: str | None = None, - sampling: float | Sequence[float] = 1.0, - units: str = "voxels", - name: str | None = None, - one_based: bool = False, - ) -> "AtomicModel": - """Load coordinates from a MATLAB ``.mat`` file. - - Parameters - ---------- - path : str or Path - File path (v5/v7 or v7.3 HDF5). - key : str, optional - Variable name. If omitted, the first numeric ``(N, 3)`` / ``(3, N)`` - array is used. - sampling, units, name - See :meth:`from_array`. - one_based : bool - Subtract 1 from the coordinates (MATLAB 1-based voxel indices). - """ - path = Path(path) - arrays: dict[str, NDArray] = {} - try: - import scipy.io as sio - - raw = sio.loadmat(path) - arrays = { - k: np.asarray(v) - for k, v in raw.items() - if not k.startswith("__") and isinstance(v, np.ndarray) and v.dtype.kind in "fiu" - } - except NotImplementedError: # v7.3 - import h5py - - with h5py.File(path, "r") as f: - - def _collect(g, prefix=""): - for k, v in g.items(): - if isinstance(v, h5py.Dataset) and v.dtype.kind in "fiu": - arrays[prefix + k] = np.asarray(v[()]).T - elif isinstance(v, h5py.Group): - _collect(v, prefix + k + "/") - - _collect(f) - if key is None: - candidates = [k for k, v in arrays.items() if v.ndim == 2 and 3 in v.shape] - if not candidates: - raise ValueError( - f"No (N, 3) coordinate array found in {path.name}: {list(arrays)}" - ) - key = candidates[0] - xyz = arrays[key].astype(float) - if one_based: - xyz = xyz - 1.0 - return cls.from_array( - xyz, - sampling=sampling, - units=units, - name=name or path.stem, - metadata={"source": str(path), "key": key}, - ) - @classmethod def from_xyz( cls, @@ -578,7 +512,10 @@ def find_neighbors(self, num_neighbors: int = 24, cutoff: float | None = None) - first = dist <= cutoff scale = float(self._sampling.mean()) d_first = np.where(first, dist, np.nan) * scale - with np.errstate(invalid="ignore"): + import warnings + + with warnings.catch_warnings(), np.errstate(invalid="ignore"): + warnings.simplefilter("ignore", RuntimeWarning) self.set_channel("num_neighbors", first.sum(1)) self.set_channel("bond_mean", np.nanmean(d_first, axis=1), self._units) self.set_channel("bond_std", np.nanstd(d_first, axis=1), self._units) @@ -922,13 +859,74 @@ def classify_species( channel: str = "intensity", num_species: int = 2, names: Sequence[str] | None = None, + method: str = "gmm", + min_posterior: float = 0.8, + mask: NDArray | None = None, ) -> NDArray: - """Split a channel (e.g. intensity) into species with 1D k-means; channel ``species``.""" - labels, centers = meas.kmeans_1d(self.get_channel(channel), num_species) + """Assign species from a per-site channel such as the traced intensity. + + With ``method="gmm"`` a one-dimensional Gaussian mixture with + ``num_species`` components is fit by expectation maximization and each + site takes the component with the largest posterior probability; sites + whose largest posterior is below ``min_posterior`` are left unassigned + (code ``-1``, label ``unassigned``) so that ambiguous candidates stay in + the model without being counted as either species. ``method="kmeans"`` + assigns every site to the nearest cluster center. + + Adds channels ``species`` (categorical) and, for the mixture, ``species_posterior``. + The fitted means, widths and fractions are stored in + ``metadata["species_model"]``. + + Parameters + ---------- + channel : str + Channel to split (``"intensity"`` from tracing or volume sampling). + num_species : int + Number of species. + names : sequence of str, optional + Species names in order of increasing channel value. + method : {"gmm", "kmeans"} + Classifier. + min_posterior : float + Minimum posterior probability for an assignment (``"gmm"`` only). + mask : ndarray, optional + ``(N,)`` boolean; only these sites take part in the fit and receive + a species, all others are unassigned. Use it to exclude surface + sites, whose intensities are reduced by the missing neighbors. + + Returns + ------- + ndarray + ``(N,)`` species codes, ``-1`` for unassigned sites. + """ + values = self.get_channel(channel).astype(float) + use = np.ones(self.num_sites, dtype=bool) if mask is None else np.asarray(mask, dtype=bool) + fit_values = np.where(use, values, np.nan) if names is None: names = [f"species_{i}" for i in range(num_species)] - self.set_channel("species", labels, categories=names) - self._metadata["species_centers"] = centers.tolist() + if method == "kmeans": + labels, centers = meas.kmeans_1d(fit_values, num_species) + posterior = np.ones(self.num_sites) + model_info = {"method": "kmeans", "centers": centers.tolist()} + elif method == "gmm": + means, sigmas, weights, resp = meas.gaussian_mixture_1d(fit_values, num_species) + labels = resp.argmax(1) + posterior = resp.max(1) + labels = np.where(posterior >= min_posterior, labels, -1) + model_info = { + "method": "gmm", + "means": means.tolist(), + "sigmas": sigmas.tolist(), + "weights": weights.tolist(), + "min_posterior": float(min_posterior), + } + else: + raise ValueError("method must be 'gmm' or 'kmeans'") + labels = np.where(use, labels, -1) + posterior = np.where(use, posterior, 0.0) + self.set_channel("species", labels, categories=list(names)) + self.set_channel("species_posterior", posterior) + self._metadata["species_model"] = model_info return labels # ------------------------------------------------------------------ # diff --git a/src/quantem/atoms/measurements.py b/src/quantem/atoms/measurements.py index 27de742f5..449f09415 100644 --- a/src/quantem/atoms/measurements.py +++ b/src/quantem/atoms/measurements.py @@ -24,6 +24,7 @@ "convex_hull_distance", "sample_volume", "kmeans_1d", + "gaussian_mixture_1d", "rotation_to_quaternion", ] @@ -297,6 +298,55 @@ def kmeans_1d( return remap[labels], centers[order] +def gaussian_mixture_1d( + values: NDArray, num_components: int = 2, num_iter: int = 200, tol: float = 1e-8 +) -> tuple[NDArray, NDArray, NDArray, NDArray]: + """Fit a one-dimensional Gaussian mixture by expectation maximization. + + Components are initialized from k-means and sorted by increasing mean. + + Returns + ------- + means, sigmas, weights, responsibilities : ndarray + Component parameters ``(K,)`` and the ``(N, K)`` posterior + probabilities of each site (rows sum to 1; non-finite values give + uniform rows). + """ + v = np.asarray(values, dtype=float) + finite = np.isfinite(v) + x = v[finite] + labels, means = kmeans_1d(x, num_components) + sigmas = np.array( + [x[labels == k].std() if np.any(labels == k) else x.std() for k in range(num_components)] + ) + sigmas = np.maximum(sigmas, 1e-6 * (x.std() + 1e-12)) + weights = np.bincount(labels, minlength=num_components) / x.size + prev = -np.inf + for _ in range(num_iter): + log_p = ( + np.log(weights + 1e-300)[None, :] + - 0.5 * ((x[:, None] - means[None, :]) / sigmas[None, :]) ** 2 + - np.log(sigmas[None, :]) + - 0.5 * np.log(2 * np.pi) + ) + log_norm = np.logaddexp.reduce(log_p, axis=1) + resp = np.exp(log_p - log_norm[:, None]) + ll = float(log_norm.sum()) + nk = resp.sum(0) + 1e-12 + means = (resp * x[:, None]).sum(0) / nk + sigmas = np.sqrt((resp * (x[:, None] - means[None, :]) ** 2).sum(0) / nk) + sigmas = np.maximum(sigmas, 1e-6 * (x.std() + 1e-12)) + weights = nk / x.size + if ll - prev < tol * max(1.0, abs(ll)): + break + prev = ll + order = np.argsort(means) + means, sigmas, weights, resp = means[order], sigmas[order], weights[order], resp[:, order] + out = np.full((v.size, num_components), 1.0 / num_components) + out[finite] = resp + return means, sigmas, weights, out + + def rotation_to_quaternion(rotation: NDArray) -> NDArray: """Convert ``(N, 3, 3)`` rotation matrices to ``(N, 4)`` unit quaternions (w, x, y, z).""" from scipy.spatial.transform import Rotation diff --git a/tests/atoms/test_atoms.py b/tests/atoms/test_atoms.py index f46de9f3f..081d02c30 100644 --- a/tests/atoms/test_atoms.py +++ b/tests/atoms/test_atoms.py @@ -323,3 +323,24 @@ def test_attached_icosahedra_and_growth_steps(): steps = st.growth_steps(d) # the upper half only appears once the front crosses the mid-plane at 2.5 bonds assert steps[d.positions[:, 2] > 0.1].min() >= 3 + + +def test_species_mixture_keeps_unassigned(): + rng = np.random.default_rng(0) + xyz = rng.random((4001, 3)) * 30 + model = AtomicModel.from_array(xyz) + # two equal components and one site exactly midway, which is ambiguous by symmetry + intensity = np.r_[rng.normal(1.0, 0.1, 2000), rng.normal(2.0, 0.1, 2000), 1.5] + model.set_channel("intensity", intensity) + labels = model.classify_species(names=["Ni", "Pd"], min_posterior=0.9) + info = model.metadata["species_model"] + assert abs(info["means"][0] - 1.0) < 0.05 and abs(info["means"][1] - 2.0) < 0.05 + assert 0.4 < info["weights"][0] < 0.6 + assert (labels[:2000] == 0).mean() > 0.95 and (labels[2000:4000] == 1).mean() > 0.95 + assert model.categories["species"] == ["Ni", "Pd"] + assert labels[-1] == -1 and model["species_posterior"][-1] < 0.9 + labels = model.classify_species(method="kmeans") + assert labels.min() >= 0 + mask = np.arange(model.num_sites) < 3000 + labels = model.classify_species(names=["Ni", "Pd"], mask=mask) + assert np.all(labels[~mask] == -1) and labels[:2000].max() == 0