diff --git a/src/quantem/core/datastructures/dataset4dstem.py b/src/quantem/core/datastructures/dataset4dstem.py index 004db427..67bdc279 100644 --- a/src/quantem/core/datastructures/dataset4dstem.py +++ b/src/quantem/core/datastructures/dataset4dstem.py @@ -1,3 +1,4 @@ +from os import PathLike from typing import Any, Self import matplotlib.pyplot as plt @@ -97,15 +98,15 @@ def __init__( self._virtual_detectors = {} # Store detector information for regeneration @classmethod - def from_file(cls, file_path: str, file_type: str) -> "Dataset4dstem": + def from_file(cls, file_path: str | PathLike, file_type: str | None = None) -> "Dataset4dstem": """ Create a new Dataset4dstem from a file. Parameters ---------- - file_path : str + file_path : str | PathLike Path to the data file - file_type : str + file_type : str | None The type of file reader needed. See rosettasciio for supported formats https://hyperspy.org/rosettasciio/supported_formats/index.html diff --git a/src/quantem/core/fitting/__init__.py b/src/quantem/core/fitting/__init__.py new file mode 100644 index 00000000..1bb92620 --- /dev/null +++ b/src/quantem/core/fitting/__init__.py @@ -0,0 +1,21 @@ +from quantem.core.fitting.background import DCBackground as DCBackground +from quantem.core.fitting.background import GaussianBackground as GaussianBackground +from quantem.core.fitting.base import Component as Component +from quantem.core.fitting.base import Model as Model +from quantem.core.fitting.base import ModelContext as ModelContext +from quantem.core.fitting.base import OriginND as OriginND +from quantem.core.fitting.base import Parameter as Parameter +from quantem.core.fitting.diffraction import DiskTemplate as DiskTemplate +from quantem.core.fitting.diffraction import SyntheticDiskLattice as SyntheticDiskLattice + +__all__ = [ + "Component", + "DCBackground", + "DiskTemplate", + "GaussianBackground", + "Model", + "ModelContext", + "OriginND", + "Parameter", + "SyntheticDiskLattice", +] diff --git a/src/quantem/core/fitting/background.py b/src/quantem/core/fitting/background.py new file mode 100644 index 00000000..000d4a01 --- /dev/null +++ b/src/quantem/core/fitting/background.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +from typing import Any, Sequence + +import torch +from torch import nn + +from quantem.core.fitting.base import OriginND, RenderComponent, RenderContext + + +class DCBackground(RenderComponent): + def __init__( + self, + *, + intensity: float | int | Sequence[float | int | None] = 0.0, + name: str = "dc_background", + constraint_params: dict[str, Any] | None = None, + ): + """ + Build a constant background component. + + Notes + ----- + Validity is enforced via hard constraints/parameter bounds. Forward + intentionally avoids hard clamps for gradient flow. + """ + super().__init__() + self.name = str(name) + intensity_init, intensity_lo, intensity_hi = self.parse_bounded_init( + intensity, name="intensity" + ) + self.intensity_raw = nn.Parameter(torch.tensor(intensity_init, dtype=torch.float32)) + bounded_lo = 0.0 if intensity_lo is None else max(float(intensity_lo), 0.0) + self.register_parameter_bounds("intensity_raw", bounded_lo, intensity_hi) + if constraint_params is not None: + self.apply_constraint_params(constraint_params, strict=True) + self._enforce_parameter_bounds() + + def forward(self, ctx: RenderContext) -> torch.Tensor: + """ + Render constant background from raw trainable intensity. + + Notes + ----- + Validity is enforced via hard constraints/parameter bounds, not via + forward-time hard clamps. + """ + inten = self.intensity_raw.to(device=ctx.device, dtype=ctx.dtype) + return torch.ones(ctx.shape, device=ctx.device, dtype=ctx.dtype) * inten + + +class GaussianBackground(RenderComponent): # TODO this should be N dimensional by default + def __init__( + self, + *, + sigma: float | int | Sequence[float | int | None] = (40.0, 5.0, None), + intensity: float | int | Sequence[float | int | None] = 0.0, + origin: OriginND | None = None, + origin_key: str = "origin", + name: str = "gaussian_background", + constraint_params: dict[str, Any] | None = None, + ): + """ + Build a Gaussian background component centered at origin. + + Notes + ----- + ``sigma_raw`` and ``intensity_raw`` validity is enforced via hard + constraints/parameter bounds. Forward intentionally avoids hard clamps + for gradient flow. + """ + super().__init__() + self.name = str(name) + self.origin = origin + self.origin_key = str(origin_key) + sigma_init, sigma_lo, sigma_hi = self.parse_bounded_init(sigma, name="sigma") + intensity_init, intensity_lo, intensity_hi = self.parse_bounded_init( + intensity, name="intensity" + ) + self.sigma_raw = nn.Parameter(torch.tensor(sigma_init, dtype=torch.float32)) + sigma_bounded_lo = 1e-6 if sigma_lo is None else max(float(sigma_lo), 1e-6) + self.register_parameter_bounds("sigma_raw", sigma_bounded_lo, sigma_hi) + self.intensity_raw = nn.Parameter(torch.tensor(intensity_init, dtype=torch.float32)) + intensity_bounded_lo = 0.0 if intensity_lo is None else max(float(intensity_lo), 0.0) + self.register_parameter_bounds("intensity_raw", intensity_bounded_lo, intensity_hi) + if constraint_params is not None: + self.apply_constraint_params(constraint_params, strict=True) + self._enforce_parameter_bounds() + + def set_origin(self, origin: OriginND) -> None: + self.origin = origin + + def forward(self, ctx: RenderContext) -> torch.Tensor: + """ + Render Gaussian background from raw trainable parameters. + + Notes + ----- + Validity is enforced via hard constraints/parameter bounds, not via + forward-time hard clamps. + """ + if self.origin is None: + raise RuntimeError("GaussianBackground requires an OriginND instance.") + + rr = torch.arange(ctx.shape[0], device=ctx.device, dtype=ctx.dtype)[:, None] + cc = torch.arange(ctx.shape[1], device=ctx.device, dtype=ctx.dtype)[None, :] + r0, c0 = self.origin.coords[0], self.origin.coords[1] + + sigma = self.sigma_raw.to(device=ctx.device, dtype=ctx.dtype) + inten = self.intensity_raw.to(device=ctx.device, dtype=ctx.dtype) + r2 = (rr - r0) ** 2 + (cc - c0) ** 2 + return inten * torch.exp(-0.5 * r2 / (sigma * sigma)) diff --git a/src/quantem/core/fitting/base.py b/src/quantem/core/fitting/base.py new file mode 100644 index 00000000..74ce1f0f --- /dev/null +++ b/src/quantem/core/fitting/base.py @@ -0,0 +1,822 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Literal, Self, Sequence, cast + +import numpy as np +import torch +from torch import nn +from tqdm import tqdm + +from quantem.core.ml.optimizer_mixin import ( + OptimizerMixin, + OptimizerParams, + OptimizerType, + SchedulerType, +) + + +def parse_bounded_init( + value: float | int | Sequence[float | int | None], *, name: str +) -> tuple[float, float | None, float | None]: + """ + Parse a scalar or bounded initializer specification. + + Parameters + ---------- + value : float | int | Sequence[float | int | None] + Accepted forms: + - ``x`` -> init ``x`` with no bounds. + - ``(x0, delta)`` -> init ``x0`` with bounds ``[x0-|delta|, x0+|delta|]``. + - ``(x0, lo, hi)`` -> init ``x0`` with explicit bounds. + name : str + Parameter name used in error messages. + + Returns + ------- + tuple[float, float | None, float | None] + Parsed ``(init, lo, hi)``. + + Raises + ------ + ValueError + If the sequence form is invalid, contains required ``None`` entries, + has invalid ordering, or ``init`` lies outside explicit bounds. + """ + if not isinstance(value, (list, tuple, np.ndarray)): + x = float(cast(float | int, value)) + return x, None, None + + seq = list(value) + if len(seq) == 0: + raise ValueError(f"{name} cannot be empty.") + if seq[0] is None: + raise ValueError(f"{name} initial value cannot be None.") + x0 = float(cast(float | int, seq[0])) + + if len(seq) == 1: + return x0, None, None + if len(seq) == 2: + if seq[1] is None: + raise ValueError(f"{name} delta cannot be None.") + delta = abs(float(cast(float | int, seq[1]))) + return x0, x0 - delta, x0 + delta + if len(seq) == 3: + if seq[1] is None or seq[2] is None: + raise ValueError(f"{name} bounds cannot contain None.") + lo = float(cast(float | int, seq[1])) + hi = float(cast(float | int, seq[2])) + if lo > hi: + raise ValueError(f"{name} has invalid bounds: lo ({lo}) > hi ({hi}).") + if x0 < lo or x0 > hi: + raise ValueError(f"{name} initial value {x0} is outside bounds [{lo}, {hi}].") + return x0, lo, hi + + raise ValueError(f"{name} must be scalar, (x0, delta), or (x0, lo, hi).") + + +@dataclass +class RenderContext: + shape: tuple[int, ...] + device: torch.device + dtype: torch.dtype + mask: torch.Tensor | None = None + fields: dict[str, Any] = field(default_factory=dict) + + +class OriginND(nn.Module): + def __init__(self, *, ndim: int, init: Sequence[float]): + super().__init__() + if int(ndim) <= 0: + raise ValueError("ndim must be >= 1.") + if len(init) != int(ndim): + raise ValueError("init length must match ndim.") + self.ndim = int(ndim) + self.coords = nn.Parameter(torch.as_tensor(init, dtype=torch.float32).reshape(self.ndim)) + + +class RenderComponent(nn.Module): + DEFAULT_HARD_CONSTRAINTS: dict[str, Any] = {} + DEFAULT_SOFT_CONSTRAINTS: dict[str, Any] = {} + + def __init__(self) -> None: + super().__init__() + self.hard_constraints: dict[str, Any] = dict(self.DEFAULT_HARD_CONSTRAINTS) + self.soft_constraints: dict[str, Any] = dict(self.DEFAULT_SOFT_CONSTRAINTS) + self.parameter_bounds: dict[str, tuple[float | None, float | None]] = {} + + @staticmethod + def parse_bounded_init( + value: float | int | Sequence[float | int | None], *, name: str + ) -> tuple[float, float | None, float | None]: + """ + Parse bounded initializer forms into ``(init, lo, hi)``. + + Parameters + ---------- + value : float | int | Sequence[float | int | None] + Scalar, ``(x0, delta)``, or ``(x0, lo, hi)``. + name : str + Parameter name used in error messages. + + Returns + ------- + tuple[float, float | None, float | None] + Parsed ``(init, lo, hi)``. + """ + return parse_bounded_init(value, name=name) + + def register_parameter_bounds( + self, parameter_name: str, lo: float | None, hi: float | None + ) -> None: + """ + Register hard bounds for a trainable parameter. + + Parameters + ---------- + parameter_name : str + Name of an ``nn.Parameter`` attribute on this component. + lo : float | None + Lower bound, or ``None`` for unbounded lower side. + hi : float | None + Upper bound, or ``None`` for unbounded upper side. + + Returns + ------- + None + + Raises + ------ + ValueError + If ``lo > hi``. + """ + if lo is not None and hi is not None and float(lo) > float(hi): + raise ValueError(f"Invalid bounds for {parameter_name}: lo ({lo}) > hi ({hi}).") + self.parameter_bounds[str(parameter_name)] = ( + None if lo is None else float(lo), + None if hi is None else float(hi), + ) + + def _enforce_parameter_bounds(self) -> None: + """ + Clamp registered parameters in-place to configured bounds. + + Returns + ------- + None + + Raises + ------ + AttributeError + If a registered parameter attribute is missing. + TypeError + If a registered attribute is not an ``nn.Parameter``. + """ + if not self.parameter_bounds: + return + with torch.no_grad(): + for param_name, (lo, hi) in self.parameter_bounds.items(): + if not hasattr(self, param_name): + raise AttributeError( + f"Parameter '{param_name}' is not an attribute of {self.__class__.__name__}." + ) + param = getattr(self, param_name) + if not isinstance(param, nn.Parameter): + raise TypeError( + f"Attribute '{param_name}' on {self.__class__.__name__} is not an nn.Parameter." + ) + if lo is None and hi is None: + continue + if lo is None: + assert hi is not None + param.clamp_(max=float(hi)) + elif hi is None: + assert lo is not None + param.clamp_(min=float(lo)) + else: + param.clamp_(min=float(lo), max=float(hi)) + + def _set_constraints( + self, + current: dict[str, Any], + defaults: dict[str, Any], + constraints: dict[str, Any], + *, + strict: bool, + ) -> None: + if strict: + unknown = [k for k in constraints if k not in defaults] + if unknown: + keys = ", ".join(str(k) for k in unknown) + raise KeyError(f"Unknown constraint keys: {keys}") + current.update(constraints) + + def set_hard_constraints(self, constraints: dict[str, Any], strict: bool = True) -> None: + self._set_constraints( + self.hard_constraints, self.DEFAULT_HARD_CONSTRAINTS, constraints, strict=strict + ) + + def set_soft_constraints(self, constraints: dict[str, Any], strict: bool = True) -> None: + self._set_constraints( + self.soft_constraints, self.DEFAULT_SOFT_CONSTRAINTS, constraints, strict=strict + ) + + def apply_constraint_params(self, params: dict[str, Any], strict: bool = True) -> None: + if not isinstance(params, dict): + raise TypeError("constraint params must be a dict.") + if "hard" in params or "soft" in params: + hard = params.get("hard") + soft = params.get("soft") + if hard is not None: + if not isinstance(hard, dict): + raise TypeError("constraint params 'hard' value must be a dict.") + self.set_hard_constraints(hard, strict=strict) + if soft is not None: + if not isinstance(soft, dict): + raise TypeError("constraint params 'soft' value must be a dict.") + self.set_soft_constraints(soft, strict=strict) + return + + hard_updates: dict[str, Any] = {} + soft_updates: dict[str, Any] = {} + unknown: dict[str, Any] = {} + for k, v in params.items(): + if k in self.DEFAULT_HARD_CONSTRAINTS: + hard_updates[k] = v + elif k in self.DEFAULT_SOFT_CONSTRAINTS: + soft_updates[k] = v + else: + unknown[k] = v + + if unknown and strict: + keys = ", ".join(str(k) for k in unknown.keys()) + raise KeyError(f"Unknown constraint keys for {self.__class__.__name__}: {keys}") + if unknown: + soft_updates.update(unknown) + if hard_updates: + self.set_hard_constraints(hard_updates, strict=strict) + if soft_updates: + self.set_soft_constraints(soft_updates, strict=strict) + + def effective_soft_constraints(self, params: dict[str, Any] | None = None) -> dict[str, Any]: + effective = dict(self.soft_constraints) + if isinstance(params, dict): + effective.update(params) + return effective + + def enforce_hard_constraints(self, ctx: RenderContext) -> None: + self._enforce_parameter_bounds() + + def forward(self, ctx: RenderContext) -> torch.Tensor: + raise NotImplementedError + + def constraint_loss( + self, ctx: RenderContext, params: dict[str, Any] | None = None + ) -> torch.Tensor: + return torch.zeros((), device=ctx.device, dtype=ctx.dtype) + + +class AdditiveRenderModel(nn.Module): + def __init__(self, *, origin: nn.Module, components: list[RenderComponent]): + super().__init__() + self.origin = origin + self.components = nn.ModuleList(components) + + def forward(self, ctx: RenderContext) -> torch.Tensor: + if len(self.components) == 0: + return torch.zeros(ctx.shape, device=ctx.device, dtype=ctx.dtype) + out = self.components[0](ctx) + for component in self.components[1:]: + out = out + component(ctx) + return out + + def _component_constraint_name(self, component: RenderComponent, idx: int) -> str: + name = getattr(component, "name", None) + if isinstance(name, str) and name: + return name + class_name = component.__class__.__name__ + if class_name: + return class_name + return f"component_{idx}" + + def apply_constraint_params( + self, constraint_params: dict[str, Any], strict: bool = True + ) -> None: + if not isinstance(constraint_params, dict): + raise TypeError("constraint_params must be a dict.") + source = constraint_params.get("components") + component_map = source if isinstance(source, dict) else constraint_params + for target, params in component_map.items(): + if not isinstance(params, dict): + if strict: + raise TypeError(f"Constraint params for '{target}' must be a dict.") + continue + target_str = str(target) + name_matches: list[RenderComponent] = [] + class_matches: list[RenderComponent] = [] + for idx, module in enumerate(self.components): + component = cast(RenderComponent, module) + if self._component_constraint_name(component, idx) == target_str: + name_matches.append(component) + if component.__class__.__name__ == target_str: + class_matches.append(component) + targets = name_matches if name_matches else class_matches + if not targets: + if strict: + raise KeyError(f"No matching component for constraint target '{target_str}'.") + continue + for component in targets: + component.apply_constraint_params(params, strict=strict) + + def apply_hard_constraints(self, ctx: RenderContext) -> None: + for module in self.components: + component = cast(RenderComponent, module) + component.enforce_hard_constraints(ctx) + + def total_constraint_loss(self, ctx: RenderContext) -> torch.Tensor: + loss = torch.zeros((), device=ctx.device, dtype=ctx.dtype) + for module in self.components: + component = cast(RenderComponent, module) + loss = loss + component.constraint_loss(ctx) + return loss + + +@dataclass +class FitResult: + losses: list[float] + lrs: list[float] + final_loss: float + num_steps: int + metrics: dict[str, list[float]] = field(default_factory=dict) + + +class FitBase(OptimizerMixin): + DEFAULT_LR = 1e-2 + DEFAULT_OPTIMIZER_TYPE = "adam" + + def __init__(self): + super().__init__() + # Core wiring + self.loss_fn = torch.nn.MSELoss(reduction="mean") + self.model: AdditiveRenderModel | None = None + self.ctx: RenderContext | None = None + + # State/checkpoints + self.state_initialized: dict[str, torch.Tensor] | None = None + + # Histories/results + self.fit_history: dict[str, FitResult] = {} + + def get_optimization_parameters(self) -> Any: + if self.model is None: + return [] + return [p for p in self.model.parameters() if p.requires_grad] + + @property + def state_current(self) -> dict[str, torch.Tensor] | None: + if self.model is None: + return None + return self._get_model_state_dict_copy() + + @property + def render_initialized(self) -> np.ndarray: + if self.state_initialized is None: + raise RuntimeError("initialized state is unavailable. Call .define_model(...) first.") + return self._render_state_array(self.state_initialized) + + @property + def render_current(self) -> np.ndarray: + if self.model is None or self.ctx is None: + raise RuntimeError("Call .define_model(...) first.") + return self.model(self.ctx).detach().cpu().numpy() + + def reset( + self, + reset_to: Literal["initialized"] = "initialized", + ) -> Self: + if reset_to != "initialized": + raise ValueError("FitBase.reset only supports reset_to='initialized'.") + if self.state_initialized is None: + raise RuntimeError("initialized state is unavailable. Call .define_model(...) first.") + self._load_model_state_dict_copy(self.state_initialized) + self._clear_fit_history_all() + return self + + def set_component_trainable( + self, component_name: str, enabled: bool, rebuild_optimizer: bool = True + ) -> None: + """ + Enable or disable optimization for all parameters in one component. + + Parameters + ---------- + component_name : str + Resolved component name. + enabled : bool + If ``True``, mark component parameters trainable. + rebuild_optimizer : bool, optional + If ``True``, rebuild optimizer param groups after toggling. + + Returns + ------- + None + + Raises + ------ + RuntimeError + If the model is not defined. + KeyError + If ``component_name`` is unknown. + + Notes + ----- + When rebuilding, the optimizer is reconstructed from stored optimizer + parameters if available, otherwise inferred from the current optimizer + type and learning rate, else defaults. Scheduler state is cleared + predictably by setting scheduler type to ``"none"``. + """ + component = self._resolve_component_by_name(component_name) + for _, param in component.named_parameters(recurse=True): + param.requires_grad_(bool(enabled)) + if rebuild_optimizer: + self._rebuild_optimizer_after_trainability_change() + + def set_parameter_trainable( + self, + component_name: str, + parameter_name: str, + enabled: bool, + rebuild_optimizer: bool = True, + ) -> None: + """ + Enable or disable optimization for one component parameter. + + Parameters + ---------- + component_name : str + Resolved component name. + parameter_name : str + Parameter name from ``component.named_parameters()``. + enabled : bool + If ``True``, mark parameter trainable. + rebuild_optimizer : bool, optional + If ``True``, rebuild optimizer param groups after toggling. + + Returns + ------- + None + + Raises + ------ + RuntimeError + If the model is not defined. + KeyError + If ``component_name`` or ``parameter_name`` is unknown. + + Notes + ----- + When rebuilding, scheduler state is cleared by setting scheduler type + to ``"none"``. + """ + component = self._resolve_component_by_name(component_name) + params = dict(component.named_parameters(recurse=True)) + if parameter_name not in params: + known = ", ".join(sorted(params.keys())) + raise KeyError( + f"Parameter '{parameter_name}' not found in component '{component_name}'. " + f"Known parameters: {known}" + ) + params[parameter_name].requires_grad_(bool(enabled)) + if rebuild_optimizer: + self._rebuild_optimizer_after_trainability_change() + + def set_parameters_trainable( + self, + component_name: str, + parameter_names: list[str], + enabled: bool, + rebuild_optimizer: bool = True, + ) -> None: + """ + Enable or disable optimization for multiple component parameters. + + Parameters + ---------- + component_name : str + Resolved component name. + parameter_names : list[str] + Parameter names from ``component.named_parameters()``. + enabled : bool + If ``True``, mark parameters trainable. + rebuild_optimizer : bool, optional + If ``True``, rebuild optimizer param groups after toggling. + + Returns + ------- + None + + Raises + ------ + RuntimeError + If the model is not defined. + KeyError + If any parameter name is unknown. + """ + component = self._resolve_component_by_name(component_name) + params = dict(component.named_parameters(recurse=True)) + missing = [name for name in parameter_names if name not in params] + if missing: + known = ", ".join(sorted(params.keys())) + raise KeyError( + f"Unknown parameters for component '{component_name}': {', '.join(missing)}. " + f"Known parameters: {known}" + ) + for name in parameter_names: + params[name].requires_grad_(bool(enabled)) + if rebuild_optimizer: + self._rebuild_optimizer_after_trainability_change() + + def get_component_trainable(self, component_name: str) -> dict[str, bool]: + """ + Return trainability flags for one component's parameters. + + Parameters + ---------- + component_name : str + Resolved component name. + + Returns + ------- + dict[str, bool] + Mapping of parameter name to ``requires_grad``. + + Raises + ------ + RuntimeError + If the model is not defined. + KeyError + If ``component_name`` is unknown. + """ + component = self._resolve_component_by_name(component_name) + return {name: bool(param.requires_grad) for name, param in component.named_parameters()} + + def fit_render( + self, + *, + target: torch.Tensor, + n_steps: int, + constraint_weight: float = 1.0, + constraint_params: dict[str, Any] | None = None, + optimizer_params: OptimizerType | dict | None = None, + scheduler_params: SchedulerType | dict | None = None, + progress: bool = False, + run_key: str = "default", + **kwargs: Any, + ) -> FitResult: + """ + Fit model parameters to a target render. + + Parameters + ---------- + target : torch.Tensor + Target tensor to fit. + n_steps : int + Number of optimization steps. + constraint_weight : float, optional + Multiplier applied to the summed soft-constraint loss. + constraint_params : dict[str, Any] | None, optional + Optional constraint updates applied once to matching components before + optimization starts. If ``None``, existing component constraints are reused. + optimizer_params : dict | None, optional + Optimizer configuration override for this call. + scheduler_params : dict | None, optional + Scheduler configuration override for this call. + progress : bool, optional + If ``True``, display a progress bar. + run_key : str, optional + History key used to store/append fit metrics. + **kwargs : Any + Forwarded to internal forward/loss hooks. + + Returns + ------- + FitResult + Fit history and final loss metadata for this run key. + + Raises + ------ + RuntimeError + If model/context are undefined. + + Notes + ----- + Hard constraints are applied after each optimizer step. + """ + if self.model is None or self.ctx is None: + raise RuntimeError("Model and context are not defined for fitting.") + if constraint_params is not None: + self.model.apply_constraint_params(constraint_params, strict=True) + + optimizer_rebuilt = False + if optimizer_params is not None: + self.set_optimizer(optimizer_params) + optimizer_rebuilt = True + elif self.optimizer is None: + if self.optimizer_params: + self.set_optimizer(self.optimizer_params) + else: + self.set_optimizer( + { + "type": getattr(self, "DEFAULT_OPTIMIZER_TYPE", "adamw"), + "lr": float(getattr(self, "DEFAULT_LR", self.DEFAULT_LR)), + } + ) + optimizer_rebuilt = True + + n_steps = int(n_steps) + if scheduler_params is not None: + self.set_scheduler(scheduler_params, num_iter=n_steps) + elif self.scheduler is None and self.scheduler_params: + self.set_scheduler(self.scheduler_params, num_iter=n_steps) + elif optimizer_rebuilt and self.scheduler is not None and self.optimizer is not None: + self.scheduler.optimizer = self.optimizer + + pbar = tqdm(range(n_steps), desc="Fit render", disable=not progress) + + losses: list[float] = [] + lrs: list[float] = [] + for _ in pbar: + self.zero_optimizer_grad() + pred = self._forward_for_fit(target=target, **kwargs) + data_loss = self._fidelity_loss(pred, target, **kwargs) + constraint_loss = self._constraint_loss(pred, target, **kwargs) + total_loss = data_loss + constraint_weight * constraint_loss + total_loss.backward() + self.step_optimizer() + if self.model is None or self.ctx is None: + raise RuntimeError("Model and context are not defined for fitting.") + self.model.apply_hard_constraints(self.ctx) + total_loss_value = float(total_loss.detach().cpu()) + self.step_scheduler(total_loss_value) + losses.append(total_loss_value) + lrs.append(float(self.get_current_lr())) + + key = str(run_key) + if key in self.fit_history: + prev = self.fit_history[key] + prev.losses.extend(losses) + prev.lrs.extend(lrs) + prev.final_loss = prev.losses[-1] if prev.losses else float("nan") + prev.num_steps = len(prev.losses) + result = prev + else: + result = FitResult( + losses=losses, + lrs=lrs, + final_loss=(losses[-1] if losses else float("nan")), + num_steps=n_steps, + ) + self.fit_history[key] = result + return result + + def _iter_named_components(self) -> list[tuple[str, RenderComponent]]: + """ + Return canonical component names paired with components. + + Returns + ------- + list[tuple[str, RenderComponent]] + ``(name, component)`` entries using the model's canonical naming + rule. Names fall back to class-name/index behavior when ``.name`` is + missing. + + Raises + ------ + RuntimeError + If the model is not defined. + """ + if self.model is None: + raise RuntimeError("Call .define_model(...) first.") + entries: list[tuple[str, RenderComponent]] = [] + for idx, module in enumerate(self.model.components): + component = cast(RenderComponent, module) + name = self.model._component_constraint_name(component, idx) + entries.append((name, component)) + return entries + + def get_component_names(self) -> list[str]: + """ + Return canonical component names. + + Returns + ------- + list[str] + Canonical component names. + """ + return [name for name, _ in self._iter_named_components()] + + def _resolve_component_by_name(self, component_name: str) -> RenderComponent: + target = str(component_name) + for resolved_name, component in self._iter_named_components(): + if resolved_name == target: + return component + known = ", ".join(self.get_component_names()) + raise KeyError(f"Component not found: {target}. Known components: {known}") + + def _infer_optimizer_rebuild_params(self) -> dict[str, Any]: + if self.optimizer_params: + op = self.optimizer_params + if isinstance(op, OptimizerParams.NoneOptimizer): + return {"type": "none"} + out: dict[str, Any] = dict(op.params()) + out["type"] = op._name + return out + if self.optimizer is not None: + opt_type: str | type[torch.optim.Optimizer] + if isinstance(self.optimizer, torch.optim.AdamW): + opt_type = "adamw" + elif isinstance(self.optimizer, torch.optim.Adam): + opt_type = "adam" + elif isinstance(self.optimizer, torch.optim.SGD): + opt_type = "sgd" + else: + opt_type = type(self.optimizer) + lr = float( + self.optimizer.param_groups[0].get( + "lr", getattr(self, "DEFAULT_LR", self.DEFAULT_LR) + ) + ) + return {"type": opt_type, "lr": lr} + return { + "type": getattr(self, "DEFAULT_OPTIMIZER_TYPE", self.DEFAULT_OPTIMIZER_TYPE), + "lr": float(getattr(self, "DEFAULT_LR", self.DEFAULT_LR)), + } + + def _rebuild_optimizer_after_trainability_change(self) -> None: + if self.model is None: + raise RuntimeError("Call .define_model(...) first.") + rebuild_params = self._infer_optimizer_rebuild_params() + self.set_optimizer(rebuild_params) + self.set_scheduler({"type": "none"}) + + def _clone_state_dict(self, state: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + return {k: v.detach().clone() for k, v in state.items()} + + def _get_model_state_dict_copy(self) -> dict[str, torch.Tensor]: + if self.model is None: + raise RuntimeError("Call .define_model(...) first.") + return self._clone_state_dict(self.model.state_dict()) + + def _load_model_state_dict_copy(self, state: dict[str, torch.Tensor]) -> None: + if self.model is None: + raise RuntimeError("Call .define_model(...) first.") + self.model.load_state_dict(self._clone_state_dict(state), strict=True) + + def _clear_fit_history_all(self) -> None: + self.fit_history.clear() + + def _clear_fit_history_run(self, run_key: str) -> None: + self.fit_history.pop(str(run_key), None) + + def _render_state_array(self, state: dict[str, torch.Tensor]) -> np.ndarray: + if self.model is None or self.ctx is None: + raise RuntimeError("Call .define_model(...) first.") + live = self._get_model_state_dict_copy() + try: + self._load_model_state_dict_copy(state) + arr = self.model(self.ctx).detach().cpu().numpy() + finally: + self._load_model_state_dict_copy(live) + return arr + + def _forward_for_fit(self, *, target: torch.Tensor, **kwargs: Any) -> torch.Tensor: + if self.model is None or self.ctx is None: + raise RuntimeError("Model and context are not defined for fitting.") + return self.model(self.ctx) + + def _fidelity_loss( + self, pred: torch.Tensor, target: torch.Tensor, **kwargs: Any + ) -> torch.Tensor: + if self.ctx is not None and self.ctx.mask is not None: + # TODO -- use loss modules (currently implemented in tomo branch) + # and update them to allow for masking at module level + diff = (pred - target) * self.ctx.mask + denom = torch.clamp(torch.sum(self.ctx.mask), min=1.0) + return torch.sum(diff * diff) / denom + return self.loss_fn(pred, target) + + def _constraint_loss( + self, + pred: torch.Tensor, + target: torch.Tensor, + **kwargs: Any, + ) -> torch.Tensor: + if self.model is None or self.ctx is None: + raise RuntimeError("Model and context are not defined for fitting.") + return self.model.total_constraint_loss(self.ctx) + + +Component = RenderComponent +ModelContext = RenderContext +Model = AdditiveRenderModel +Parameter = nn.Parameter diff --git a/src/quantem/core/fitting/diffraction.py b/src/quantem/core/fitting/diffraction.py new file mode 100644 index 00000000..30934bf5 --- /dev/null +++ b/src/quantem/core/fitting/diffraction.py @@ -0,0 +1,581 @@ +from __future__ import annotations + +from typing import Any, Iterable, Sequence, cast + +import numpy as np +import torch +import torch.nn.functional as F +from torch import nn + +from quantem.core.fitting.base import OriginND, RenderComponent, RenderContext + + +def _splat_patch( + out: torch.Tensor, + *, + r0: torch.Tensor, + c0: torch.Tensor, + patch_vals: torch.Tensor, + dr: torch.Tensor, + dc: torch.Tensor, + scale: torch.Tensor, +) -> None: + h, w = out.shape + r = r0 + dr + c = c0 + dc + + r_base = torch.floor(r) + c_base = torch.floor(c) + fr = r - r_base + fc = c - c_base + r0i = r_base.to(torch.long) + c0i = c_base.to(torch.long) + + w00 = (1.0 - fr) * (1.0 - fc) + w01 = (1.0 - fr) * fc + w10 = fr * (1.0 - fc) + w11 = fr * fc + v = patch_vals * scale + + def put(rr: torch.Tensor, cc: torch.Tensor, ww: torch.Tensor) -> None: + keep = (rr >= 0) & (rr < h) & (cc >= 0) & (cc < w) + if torch.any(keep): + out.index_put_((rr[keep], cc[keep]), v[keep] * ww[keep], accumulate=True) + + put(r0i, c0i, w00) + put(r0i, c0i + 1, w01) + put(r0i + 1, c0i, w10) + put(r0i + 1, c0i + 1, w11) + + +class DiskTemplate(RenderComponent): + DEFAULT_HARD_CONSTRAINTS: dict[str, bool] = { + "force_center": False, + "force_positive": True, + } + DEFAULT_SOFT_CONSTRAINTS: dict[str, float] = {"tv_weight": 0.0} + + def __init__( + self, + *, + name: str, + array: np.ndarray, + refine_all_pixels: bool = False, + normalize: str = "none", + origin: OriginND | None = None, + origin_key: str = "origin", + intensity: float | Sequence[float] = 1.0, + constraint_params: dict[str, Any] | None = None, + ): + """ + Build a disk template renderer centered at the shared origin. + + Parameters + ---------- + intensity : float | Sequence[float], optional + Trainable scalar amplitude applied to the rendered template. + Accepts ``x``, ``(x0, delta)``, or ``(x0, lo, hi)``. + + Returns + ------- + None + + Raises + ------ + ValueError + If ``array`` is not 2D or if ``normalize`` is unsupported. + + Notes + ----- + ``template_raw`` controls template shape and ``intensity_raw`` controls + center-disk amplitude. + """ + super().__init__() + self.name = str(name) + self.refine_all_pixels = bool(refine_all_pixels) + self.origin = origin + self.origin_key = str(origin_key) + intensity_init, intensity_lo, intensity_hi = self.parse_bounded_init( + intensity, name="intensity" + ) + self.intensity_raw = nn.Parameter(torch.tensor(intensity_init, dtype=torch.float32)) + if intensity_lo is not None or intensity_hi is not None: + self.register_parameter_bounds("intensity_raw", intensity_lo, intensity_hi) + + a = np.asarray(array, dtype=np.float32) + if a.ndim != 2: + raise ValueError("DiskTemplate.array must be 2D.") + if normalize == "max": + s = float(np.max(a)) + if s > 0.0: + a = a / s + elif normalize == "mean": + s = float(np.mean(a)) + if s != 0.0: + a = a / s + elif normalize != "none": + raise ValueError("normalize must be one of: 'none', 'max', 'mean'.") + + template = torch.as_tensor(a, dtype=torch.float32) + self.template_raw = nn.Parameter(template.clone(), requires_grad=self.refine_all_pixels) + + ht, wt = int(template.shape[0]), int(template.shape[1]) + rr, cc = np.mgrid[0:ht, 0:wt] + rr = rr.astype(np.float32) - (ht - 1) * 0.5 + cc = cc.astype(np.float32) - (wt - 1) * 0.5 + self.register_buffer("dr", torch.as_tensor(rr.ravel(), dtype=torch.float32)) + self.register_buffer("dc", torch.as_tensor(cc.ravel(), dtype=torch.float32)) + if constraint_params is not None: + self.apply_constraint_params(constraint_params, strict=True) + if bool(self.hard_constraints.get("force_positive", False)): + self._enforce_positivity() + + @classmethod + def from_array( + cls, + *, + name: str, + array: np.ndarray, + refine_all_pixels: bool = False, + normalize: str = "none", + origin: OriginND | None = None, + origin_key: str = "origin", + intensity: float | Sequence[float] = 1.0, + constraint_params: dict[str, Any] | None = None, + ) -> "DiskTemplate": + return cls( + name=name, + array=array, + refine_all_pixels=refine_all_pixels, + normalize=normalize, + origin=origin, + origin_key=origin_key, + intensity=intensity, + constraint_params=constraint_params, + ) + + def set_origin(self, origin: OriginND) -> None: + self.origin = origin + + def set_intensity(self, value: float | int) -> None: + """Assign ``intensity_raw`` in-place.""" + with torch.no_grad(): + self.intensity_raw.copy_(torch.as_tensor(float(value), dtype=self.intensity_raw.dtype)) + + def patch_values(self) -> torch.Tensor: + return self.template_raw.reshape(-1) + + def patch_offsets(self) -> tuple[torch.Tensor, torch.Tensor]: + return cast(torch.Tensor, self.dr), cast(torch.Tensor, self.dc) + + def add_patch( + self, out: torch.Tensor, *, r0: torch.Tensor, c0: torch.Tensor, scale: torch.Tensor + ) -> None: + vals = self.patch_values().to(device=out.device, dtype=out.dtype) + dr = cast(torch.Tensor, self.dr).to(device=out.device, dtype=out.dtype) + dc = cast(torch.Tensor, self.dc).to(device=out.device, dtype=out.dtype) + _splat_patch(out, r0=r0, c0=c0, patch_vals=vals, dr=dr, dc=dc, scale=scale) + + def forward(self, ctx: RenderContext) -> torch.Tensor: + """ + Render template at origin with scalar amplitude ``intensity_raw``. + + Parameters + ---------- + ctx : RenderContext + Rendering context. + + Returns + ------- + torch.Tensor + Rendered center disk image. + """ + out = torch.zeros(ctx.shape, device=ctx.device, dtype=ctx.dtype) + if self.origin is None: + raise RuntimeError("DiskTemplate.forward() requires an OriginND instance.") + r0, c0 = self.origin.coords[0], self.origin.coords[1] + scale = self.intensity_raw.to(device=ctx.device, dtype=ctx.dtype) + self.add_patch(out, r0=r0, c0=c0, scale=scale) + return out + + def _center_disk(self) -> None: + with torch.no_grad(): + template = self.template_raw + h, w = int(template.shape[0]), int(template.shape[1]) + weights = torch.clamp(template, min=0.0) + mass = torch.sum(weights) + if float(mass.detach().cpu()) <= 1e-12: + return + rr = torch.arange(h, device=template.device, dtype=template.dtype)[:, None] + cc = torch.arange(w, device=template.device, dtype=template.dtype)[None, :] + com_r = torch.sum(weights * rr) / mass + com_c = torch.sum(weights * cc) / mass + target_r = torch.as_tensor((h - 1) * 0.5, device=template.device, dtype=template.dtype) + target_c = torch.as_tensor((w - 1) * 0.5, device=template.device, dtype=template.dtype) + shift_r = target_r - com_r + shift_c = target_c - com_c + denom_h = max(h - 1, 1) + denom_w = max(w - 1, 1) + ty = -2.0 * shift_r / float(denom_h) + tx = -2.0 * shift_c / float(denom_w) + theta = torch.as_tensor( + [[1.0, 0.0, tx], [0.0, 1.0, ty]], + device=template.device, + dtype=template.dtype, + )[None, ...] + src = template[None, None, :, :] + grid = F.affine_grid(theta, [1, 1, h, w], align_corners=True) + shifted = F.grid_sample( + src, + grid, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + )[0, 0] + self.template_raw.copy_(shifted) + + def _enforce_positivity(self) -> None: + with torch.no_grad(): + self.template_raw.clamp_(min=0.0) + self.intensity_raw.clamp_(min=0.0) + + def enforce_hard_constraints(self, ctx: RenderContext) -> None: + if bool(self.hard_constraints.get("force_center", False)): + self._center_disk() + if bool(self.hard_constraints.get("force_positive", False)): + self._enforce_positivity() + super().enforce_hard_constraints(ctx) + + def constraint_loss( + self, ctx: RenderContext, params: dict[str, object] | None = None + ) -> torch.Tensor: + cfg = self.effective_soft_constraints(cast(dict[str, object] | None, params)) + tv_weight = float(cfg.get("tv_weight", 0.0)) + if tv_weight <= 0.0: + return torch.zeros((), device=ctx.device, dtype=ctx.dtype) + template = self.template_raw.to(device=ctx.device, dtype=ctx.dtype) + tv_r = ( + torch.mean(torch.abs(template[1:, :] - template[:-1, :])) + if template.shape[0] > 1 + else torch.zeros((), device=ctx.device, dtype=ctx.dtype) + ) + tv_c = ( + torch.mean(torch.abs(template[:, 1:] - template[:, :-1])) + if template.shape[1] > 1 + else torch.zeros((), device=ctx.device, dtype=ctx.dtype) + ) + return torch.as_tensor(tv_weight, device=ctx.device, dtype=ctx.dtype) * (tv_r + tv_c) + + +class SyntheticDiskLattice(RenderComponent): + DEFAULT_HARD_CONSTRAINTS: dict[str, bool] = { + "force_positive_intensity": True, + } + + def __init__( + self, + *, + name: str, + disk: DiskTemplate, + u_row: float | Sequence[float], + u_col: float | Sequence[float], + v_row: float | Sequence[float], + v_col: float | Sequence[float], + u_max: int = 0, + v_max: int = 0, + intensity_0: float | Sequence[float] = 0.0, + intensity_row: float | Sequence[float] = 0.0, + intensity_col: float | Sequence[float] = 0.0, + intensity_row_row: float | Sequence[float] = 0.0, + intensity_col_col: float | Sequence[float] = 0.0, + intensity_row_col: float | Sequence[float] = 0.0, + per_disk_intensity: bool = False, + per_disk_slopes: bool = True, + max_intensity_order: int | None = None, + default_pattern_intensity_order: int | None = None, + center_intensity_0: float | Sequence[float] | None = None, + exclude_indices: Iterable[tuple[int, int]] | None = None, + boundary_px: float = 0.0, + origin: OriginND | None = None, + origin_key: str = "origin", + constraint_params: dict[str, Any] | None = None, + ): + """ + Build a synthetic disk lattice renderer. + + Parameters + ---------- + u_row, u_col, v_row, v_col : float | Sequence[float | int | None] + Lattice basis parameters. Accept ``x``, ``(x0, delta)``, or + ``(x0, lo, hi)``. + intensity_0 : float | Sequence[float], optional + Baseline intensity for included lattice disks. Accepts ``x``, + ``(x0, delta)``, or ``(x0, lo, hi)``. + intensity_row, intensity_col, intensity_row_row, intensity_col_col, intensity_row_col : + float | Sequence[float | int | None] + Intensity polynomial parameters. Accept ``x``, ``(x0, delta)``, or + ``(x0, lo, hi)``. + center_intensity_0 : float | Sequence[float] | None, optional + Optional center-disk baseline. Accepts ``x``, ``(x0, delta)``, or + ``(x0, lo, hi)`` and routes by center ownership rules. + exclude_indices : Iterable[tuple[int, int]] | None, optional + Lattice indices excluded from rendering. By default, ``(0, 0)`` is + excluded. To include center explicitly, pass ``exclude_indices`` that + does not contain ``(0, 0)``. + + Returns + ------- + None + + Raises + ------ + ValueError + If ``center_intensity_0`` is provided with center included while + ``per_disk_intensity=False``. + + Notes + ----- + Center-intensity routing is explicit: + - Center excluded (default): ``center_intensity_0`` maps to + ``disk.intensity_raw``. + - Center included with ``per_disk_intensity=True``: + ``center_intensity_0`` maps to lattice center ``i0_raw`` entry. + In this case ``disk.intensity_raw`` is set to ``0`` to avoid + duplicate center ownership when disk is rendered standalone. + """ + super().__init__() + self.name = str(name) + self.disk = disk + self.origin = origin + self.origin_key = str(origin_key) + self.per_disk_intensity = bool(per_disk_intensity) + self.u_max = int(u_max) + self.v_max = int(v_max) + self.boundary_px = float(boundary_px) + + if max_intensity_order is None: + max_intensity_order = 1 if bool(per_disk_slopes) else 0 + self.max_intensity_order = int(max_intensity_order) + if self.max_intensity_order < 0 or self.max_intensity_order > 2: + raise ValueError("max_intensity_order must be 0, 1, or 2.") + + if default_pattern_intensity_order is None: + default_pattern_intensity_order = self.max_intensity_order + self.default_pattern_intensity_order = int(default_pattern_intensity_order) + + u_row_init, u_row_lo, u_row_hi = self.parse_bounded_init(u_row, name="u_row") + u_col_init, u_col_lo, u_col_hi = self.parse_bounded_init(u_col, name="u_col") + v_row_init, v_row_lo, v_row_hi = self.parse_bounded_init(v_row, name="v_row") + v_col_init, v_col_lo, v_col_hi = self.parse_bounded_init(v_col, name="v_col") + self.u_row = nn.Parameter(torch.tensor(u_row_init, dtype=torch.float32)) + if u_row_lo is not None or u_row_hi is not None: + self.register_parameter_bounds("u_row", u_row_lo, u_row_hi) + self.u_col = nn.Parameter(torch.tensor(u_col_init, dtype=torch.float32)) + if u_col_lo is not None or u_col_hi is not None: + self.register_parameter_bounds("u_col", u_col_lo, u_col_hi) + self.v_row = nn.Parameter(torch.tensor(v_row_init, dtype=torch.float32)) + if v_row_lo is not None or v_row_hi is not None: + self.register_parameter_bounds("v_row", v_row_lo, v_row_hi) + self.v_col = nn.Parameter(torch.tensor(v_col_init, dtype=torch.float32)) + if v_col_lo is not None or v_col_hi is not None: + self.register_parameter_bounds("v_col", v_col_lo, v_col_hi) + + exclude = {(0, 0)} if exclude_indices is None else set(exclude_indices) + center_included = (0, 0) not in exclude + if center_intensity_0 is not None and center_included and not self.per_disk_intensity: + raise ValueError( + "center_intensity_0 with center included requires per_disk_intensity=True, " + "or exclude (0,0) and use DiskTemplate intensity ownership." + ) + uv: list[tuple[int, int]] = [] + for u in range(-self.u_max, self.u_max + 1): + for v in range(-self.v_max, self.v_max + 1): + if (u, v) not in exclude: + uv.append((u, v)) + uv_t = ( + torch.as_tensor(uv, dtype=torch.long) if uv else torch.zeros((0, 2), dtype=torch.long) + ) + self.register_buffer("uv_indices", uv_t) + + n_uv = int(uv_t.shape[0]) + i0_init, i0_lo, i0_hi = self.parse_bounded_init(intensity_0, name="intensity_0") + if center_intensity_0 is None: + i0_center, i0_center_lo, i0_center_hi = i0_init, None, None + else: + i0_center, i0_center_lo, i0_center_hi = self.parse_bounded_init( + center_intensity_0, name="center_intensity_0" + ) + if center_intensity_0 is not None and not center_included: + self.disk.set_intensity(float(i0_center)) + if i0_center_lo is not None or i0_center_hi is not None: + self.disk.register_parameter_bounds("intensity_raw", i0_center_lo, i0_center_hi) + ir_init, ir_lo, ir_hi = self.parse_bounded_init(intensity_row, name="intensity_row") + ic_init, ic_lo, ic_hi = self.parse_bounded_init(intensity_col, name="intensity_col") + irr_init, irr_lo, irr_hi = self.parse_bounded_init( + intensity_row_row, name="intensity_row_row" + ) + icc_init, icc_lo, icc_hi = self.parse_bounded_init( + intensity_col_col, name="intensity_col_col" + ) + irc_init, irc_lo, irc_hi = self.parse_bounded_init( + intensity_row_col, name="intensity_row_col" + ) + self._center_i0_bounds: tuple[float | None, float | None] | None = None + self._center_i0_index: int | None = None + + if self.per_disk_intensity: + i0_values = torch.full((n_uv,), float(i0_init), dtype=torch.float32) + if n_uv > 0: + center_mask = (uv_t[:, 0] == 0) & (uv_t[:, 1] == 0) + i0_values[center_mask] = float(i0_center) + if center_intensity_0 is not None and bool(torch.any(center_mask)): + self._center_i0_index = int(torch.nonzero(center_mask, as_tuple=False)[0, 0]) + self._center_i0_bounds = (i0_center_lo, i0_center_hi) + self.i0_raw = nn.Parameter(i0_values) + if i0_lo is not None or i0_hi is not None: + self.register_parameter_bounds("i0_raw", i0_lo, i0_hi) + if center_intensity_0 is not None and center_included: + self.disk.set_intensity(0.0) + if self.max_intensity_order >= 1: + self.ir = nn.Parameter(torch.full((n_uv,), float(ir_init), dtype=torch.float32)) + self.ic = nn.Parameter(torch.full((n_uv,), float(ic_init), dtype=torch.float32)) + if ir_lo is not None or ir_hi is not None: + self.register_parameter_bounds("ir", ir_lo, ir_hi) + if ic_lo is not None or ic_hi is not None: + self.register_parameter_bounds("ic", ic_lo, ic_hi) + else: + self.ir = None + self.ic = None + if self.max_intensity_order >= 2: + self.irr = nn.Parameter(torch.full((n_uv,), float(irr_init), dtype=torch.float32)) + self.icc = nn.Parameter(torch.full((n_uv,), float(icc_init), dtype=torch.float32)) + self.irc = nn.Parameter(torch.full((n_uv,), float(irc_init), dtype=torch.float32)) + if irr_lo is not None or irr_hi is not None: + self.register_parameter_bounds("irr", irr_lo, irr_hi) + if icc_lo is not None or icc_hi is not None: + self.register_parameter_bounds("icc", icc_lo, icc_hi) + if irc_lo is not None or irc_hi is not None: + self.register_parameter_bounds("irc", irc_lo, irc_hi) + else: + self.irr = None + self.icc = None + self.irc = None + else: + self.i0_raw = nn.Parameter(torch.tensor(i0_init, dtype=torch.float32)) + if i0_lo is not None or i0_hi is not None: + self.register_parameter_bounds("i0_raw", i0_lo, i0_hi) + self.ir = nn.Parameter(torch.tensor(ir_init, dtype=torch.float32)) + if ir_lo is not None or ir_hi is not None: + self.register_parameter_bounds("ir", ir_lo, ir_hi) + self.ic = nn.Parameter(torch.tensor(ic_init, dtype=torch.float32)) + if ic_lo is not None or ic_hi is not None: + self.register_parameter_bounds("ic", ic_lo, ic_hi) + self.irr = nn.Parameter(torch.tensor(irr_init, dtype=torch.float32)) + if irr_lo is not None or irr_hi is not None: + self.register_parameter_bounds("irr", irr_lo, irr_hi) + self.icc = nn.Parameter(torch.tensor(icc_init, dtype=torch.float32)) + if icc_lo is not None or icc_hi is not None: + self.register_parameter_bounds("icc", icc_lo, icc_hi) + self.irc = nn.Parameter(torch.tensor(irc_init, dtype=torch.float32)) + if irc_lo is not None or irc_hi is not None: + self.register_parameter_bounds("irc", irc_lo, irc_hi) + if constraint_params is not None: + self.apply_constraint_params(constraint_params, strict=True) + if bool(self.hard_constraints.get("force_positive_intensity", False)): + self._enforce_positive_intensity_params() + + def set_origin(self, origin: OriginND) -> None: + self.origin = origin + + def _enforce_positive_intensity_params(self) -> None: + """ + Project base intensity parameter(s) to nonnegative values. + + Notes + ----- + Positivity is enforced as a hard projection after optimizer steps. + The forward path intentionally avoids clamp-based dead gradients. + Only ``i0_raw`` is projected; slope terms remain unconstrained. + """ + with torch.no_grad(): + self.i0_raw.clamp_(min=0.0) + + def enforce_hard_constraints(self, ctx: RenderContext) -> None: + if bool(self.hard_constraints.get("force_positive_intensity", False)): + self._enforce_positive_intensity_params() + if self._center_i0_bounds is not None and self._center_i0_index is not None: + with torch.no_grad(): + lo, hi = self._center_i0_bounds + idx = self._center_i0_index + if lo is not None: + self.i0_raw[idx].clamp_(min=float(lo)) + if hi is not None: + self.i0_raw[idx].clamp_(max=float(hi)) + super().enforce_hard_constraints(ctx) + + def forward(self, ctx: RenderContext) -> torch.Tensor: + if self.origin is None: + raise RuntimeError("SyntheticDiskLattice requires an OriginND instance.") + + out = torch.zeros(ctx.shape, device=ctx.device, dtype=ctx.dtype) + uv_indices = cast(torch.Tensor, self.uv_indices) + if torch.numel(uv_indices) == 0: + return out + + uv = torch.as_tensor(uv_indices, device=ctx.device) + u = uv[:, 0].to(dtype=ctx.dtype) + v = uv[:, 1].to(dtype=ctx.dtype) + r0, c0 = self.origin.coords[0], self.origin.coords[1] + centers_r = r0 + u * self.u_row + v * self.v_row + centers_c = c0 + u * self.u_col + v * self.v_col + + b = torch.as_tensor(self.boundary_px, device=ctx.device, dtype=ctx.dtype) + keep = (centers_r >= b) & (centers_r <= (ctx.shape[0] - 1) - b) + keep = keep & (centers_c >= b) & (centers_c <= (ctx.shape[1] - 1) - b) + keep_idx = torch.nonzero(keep, as_tuple=False).reshape(-1) + if keep_idx.numel() == 0: + return out + + active_order = int( + ctx.fields.get( + "lattice_intensity_order_override", self.default_pattern_intensity_order + ) + ) + active_order = max(0, min(active_order, self.max_intensity_order)) + + dr, dc = self.disk.patch_offsets() + dr = dr.to(device=ctx.device, dtype=ctx.dtype) + dc = dc.to(device=ctx.device, dtype=ctx.dtype) + dr2 = dr * dr + dc2 = dc * dc + drdc = dr * dc + + for j in keep_idx: + rr0 = centers_r[j] + cc0 = centers_c[j] + + if self.per_disk_intensity: + inten = self.i0_raw[j] + if active_order >= 1 and self.ir is not None and self.ic is not None: + inten = inten + self.ir[j] * dr + self.ic[j] * dc + if ( + active_order >= 2 + and self.irr is not None + and self.icc is not None + and self.irc is not None + ): + inten = inten + self.irr[j] * dr2 + self.icc[j] * dc2 + self.irc[j] * drdc + else: + inten = self.i0_raw + if active_order >= 1: + assert self.ir is not None and self.ic is not None + inten = inten + self.ir * rr0 + self.ic * cc0 + if active_order >= 2: + assert self.irr is not None and self.icc is not None and self.irc is not None + inten = ( + inten + self.irr * rr0 * rr0 + self.icc * cc0 * cc0 + self.irc * rr0 * cc0 + ) + + self.disk.add_patch(out, r0=rr0, c0=cc0, scale=inten) + + return out diff --git a/src/quantem/diffraction/__init__.py b/src/quantem/diffraction/__init__.py index e69de29b..4367682b 100644 --- a/src/quantem/diffraction/__init__.py +++ b/src/quantem/diffraction/__init__.py @@ -0,0 +1 @@ +from quantem.diffraction.model_fitting import ModelDiffraction as ModelDiffraction diff --git a/src/quantem/diffraction/model_fitting.py b/src/quantem/diffraction/model_fitting.py new file mode 100644 index 00000000..41e6b8c6 --- /dev/null +++ b/src/quantem/diffraction/model_fitting.py @@ -0,0 +1,535 @@ +from __future__ import annotations + +from typing import Any, Literal, Sequence, cast + +import numpy as np +import torch +from scipy.ndimage import shift as ndi_shift +from scipy.signal.windows import tukey + +from quantem.core.datastructures import Dataset2d, Dataset3d, Dataset4d, Dataset4dstem +from quantem.core.fitting.base import ( + AdditiveRenderModel, + FitBase, + OriginND, + RenderComponent, + RenderContext, +) +from quantem.core.fitting.diffraction import DiskTemplate, SyntheticDiskLattice +from quantem.core.io.serialize import AutoSerialize +from quantem.core.ml.optimizer_mixin import OptimizerType, SchedulerType +from quantem.core.utils.imaging_utils import cross_correlation_shift +from quantem.diffraction.model_fitting_visualizations import ModelDiffractionVisualizations + + +def _parse_init(value: float | int | Sequence[float | int | None], *, name: str) -> float: + if isinstance(value, (list, tuple, np.ndarray)): + if len(value) == 0: + raise ValueError(f"{name} cannot be empty.") + if value[0] is None: + raise ValueError(f"{name} initial value cannot be None.") + return float(value[0]) + return float(cast(float | int, value)) + + +class ModelDiffraction(ModelDiffractionVisualizations, FitBase, AutoSerialize): + _token = object() + DEFAULT_LR = 5e-2 + DEFAULT_OPTIMIZER_TYPE = "adam" + + def __init__(self, dataset: Any, _token: object | None = None): + if _token is not self._token: + raise RuntimeError("Use ModelDiffraction.from_dataset() or .from_file().") + AutoSerialize.__init__(self) + FitBase.__init__(self) + + # Dataset/input references + self.dataset = dataset + self.image_ref: np.ndarray | None = None + self.preprocess_shifts: np.ndarray | None = None + self.index_shape: tuple[int, ...] | None = None + self.target_mean: torch.Tensor | None = None + + # Diffraction-specific state/checkpoints + self.state_mean_refined: dict[str, torch.Tensor] | None = None + self.mean_refined: bool = False + + # Misc metadata + self.metadata: dict[str, Any] = {} + + @classmethod + def from_dataset( + cls, dataset: Dataset2d | Dataset3d | Dataset4d | Dataset4dstem | Any + ) -> "ModelDiffraction": + if isinstance(dataset, (Dataset2d, Dataset3d, Dataset4d, Dataset4dstem)): + return cls(dataset=dataset, _token=cls._token) + raise TypeError( + "from_dataset expects a Dataset2d, Dataset3d, Dataset4d, or Dataset4dstem instance." + ) + + @property + def components(self) -> torch.nn.ModuleList: + if self.model is None: + raise RuntimeError("Call .define_model(...) first.") + return self.model.components + + def get_component(self, name: str) -> RenderComponent: + """ + Return a live model component by resolved name. + + Parameters + ---------- + name : str + Resolved component name. + + Returns + ------- + RenderComponent + The live component object. + + Raises + ------ + RuntimeError + If the model is not defined. + KeyError + If no component matches ``name``. + """ + return self._resolve_component_by_name(name) + + def get_rendered_component(self, name: str) -> np.ndarray: + """ + Render a component and return a NumPy array. + + Parameters + ---------- + name : str + Resolved component name. + + Returns + ------- + np.ndarray + Rendered component image. + + Raises + ------ + RuntimeError + If model/context are not defined. + KeyError + If no component matches ``name``. + """ + if self.ctx is None: + raise RuntimeError("Call .define_model(...) first.") + ctx = self.ctx + component = self._resolve_component_by_name(name) + rendered = component(ctx) + return rendered.detach().cpu().numpy() + + def get_rendered_disk_template(self, name: str | None = None) -> np.ndarray: + """ + Return a DiskTemplate patch as a numpy array--not rendered onto the full frame. + + Parameters + ---------- + name : str | None, optional + DiskTemplate component name. If omitted, requires exactly one DiskTemplate. + + Returns + ------- + np.ndarray + Template-sized array from ``template_raw``. + + Raises + ------ + RuntimeError + If model/context are not defined, no DiskTemplate exists, or multiple + DiskTemplates exist when ``name`` is omitted. + TypeError + If a named component exists but is not a DiskTemplate. + """ + if self.ctx is None or self.model is None: + raise RuntimeError("Call .define_model(...) first.") + + if name is not None: + component = self._resolve_component_by_name(name) + if not isinstance(component, DiskTemplate): + raise TypeError(f"Component '{name}' is not a DiskTemplate.") + return component.template_raw.detach().cpu().numpy() + matches = [m for m in self.model.components if isinstance(m, DiskTemplate)] + if len(matches) == 0: + raise RuntimeError("No DiskTemplate components found.") + if len(matches) > 1: + raise RuntimeError("Multiple DiskTemplate components found; pass name explicitly.") + disk = cast(DiskTemplate, matches[0]) + return disk.template_raw.detach().cpu().numpy() + + def set_disk_template_trainable( + self, enabled: bool, name: str | None = None, rebuild_optimizer: bool = True + ) -> None: + """ + Toggle DiskTemplate ``template_raw`` trainability. + + Parameters + ---------- + enabled : bool + If ``True``, enable optimization of ``template_raw``. + name : str | None, optional + DiskTemplate component name. If ``None``, applies to all DiskTemplate + components in the current model. + rebuild_optimizer : bool, optional + If ``True``, rebuild optimizer param groups after toggling. + + Returns + ------- + None + + Raises + ------ + KeyError + If ``name`` does not match any component. + RuntimeError + If model is not defined or no DiskTemplate components are found. + TypeError + If ``name`` resolves to a non-DiskTemplate component. + + Notes + ----- + This toggles only ``template_raw.requires_grad``. Other DiskTemplate + parameters (for example ``intensity_raw``) are unchanged. When + ``rebuild_optimizer=True``, optimizer param groups are rebuilt to match + current ``requires_grad`` flags. + """ + if self.model is None: + raise RuntimeError("Call .define_model(...) first.") + + if name is not None: + component = self._resolve_component_by_name(name) + if not isinstance(component, DiskTemplate): + raise TypeError(f"Component '{name}' is not a DiskTemplate.") + self.set_parameter_trainable( + name, + "template_raw", + enabled=enabled, + rebuild_optimizer=rebuild_optimizer, + ) + return + + disk_names = [ + component_name + for component_name, component in self._iter_named_components() + if isinstance(component, DiskTemplate) + ] + if len(disk_names) == 0: + raise RuntimeError("No DiskTemplate components found.") + + for disk_name in disk_names: + self.set_parameter_trainable( + disk_name, + "template_raw", + enabled=enabled, + rebuild_optimizer=False, + ) + if rebuild_optimizer: + self._rebuild_optimizer_after_trainability_change() + + def get_component_constraints(self, name: str) -> dict[str, dict[str, Any]]: + component = self._resolve_component_by_name(name) + return { + "hard": dict(component.hard_constraints), + "soft": dict(component.soft_constraints), + } + + def get_overlay_coordinates(self) -> tuple[np.ndarray, np.ndarray]: + """ + Return origin and lattice disk-center coordinates for overlay plotting. + + Parameters + ---------- + None + + Returns + ------- + origin_rc : np.ndarray + Origin coordinate array with shape ``(2,)`` as ``(row, col)``. + disk_centers_rc : np.ndarray + Disk-center array with shape ``(N, 2)`` as ``(row, col)``. + + Raises + ------ + RuntimeError + If model/context are not defined. + + Notes + ----- + Coordinates are computed from current model parameters without mutating state. + Boundary filtering matches ``SyntheticDiskLattice.forward`` behavior. + """ + if self.model is None or self.ctx is None: + raise RuntimeError("Call .define_model(...) first.") + + with torch.no_grad(): + origin = cast(OriginND, self.model.origin) + origin_rc = origin.coords[:2].detach().cpu().numpy().astype(np.float32, copy=False) + + centers: list[np.ndarray] = [] + for module in self.model.components: + component = cast(RenderComponent, module) + if not isinstance(component, SyntheticDiskLattice): + continue + if component.origin is None: + continue + uv_indices = cast(torch.Tensor, component.uv_indices) + if torch.numel(uv_indices) == 0: + continue + + uv = torch.as_tensor(uv_indices, device=self.ctx.device) + u = uv[:, 0].to(dtype=self.ctx.dtype) + v = uv[:, 1].to(dtype=self.ctx.dtype) + r0, c0 = component.origin.coords[0], component.origin.coords[1] + centers_r = r0 + u * component.u_row + v * component.v_row + centers_c = c0 + u * component.u_col + v * component.v_col + + b = torch.as_tensor( + component.boundary_px, device=self.ctx.device, dtype=self.ctx.dtype + ) + keep = (centers_r >= b) & (centers_r <= (self.ctx.shape[0] - 1) - b) + keep = keep & (centers_c >= b) & (centers_c <= (self.ctx.shape[1] - 1) - b) + if torch.any(keep): + rc = torch.stack((centers_r[keep], centers_c[keep]), dim=1) + centers.append(rc.detach().cpu().numpy().astype(np.float32, copy=False)) + + if centers: + disk_centers_rc = np.concatenate(centers, axis=0) + else: + disk_centers_rc = np.zeros((0, 2), dtype=np.float32) + + return origin_rc, disk_centers_rc + + def preprocess( + self, + *, + align: bool = False, + edge_blend: float = 8.0, + upsample_factor: int = 32, + max_shift: float | None = None, + shift_order: int = 1, + ) -> "ModelDiffraction": + arr = np.asarray(self.dataset.array) + if arr.ndim < 2: + raise ValueError("dataset.array must have at least 2 dimensions.") + h, w = arr.shape[-2], arr.shape[-1] + self.index_shape = tuple(arr.shape[:-2]) + + stack = arr.reshape((-1, h, w)).astype(np.float32, copy=False) + n = stack.shape[0] + if not align or n <= 1: + self.image_ref = np.mean(stack, axis=0) + self.preprocess_shifts = None + return self + + alpha_r = 0.0 if edge_blend <= 0 else min(1.0, 2.0 * float(edge_blend) / float(h)) + alpha_c = 0.0 if edge_blend <= 0 else min(1.0, 2.0 * float(edge_blend) / float(w)) + window = tukey(h, alpha=alpha_r)[:, None] * tukey(w, alpha=alpha_c)[None, :] + window = window.astype(np.float32, copy=False) + + shifts = np.zeros((n, 2), dtype=np.float32) + fft_ref = np.fft.fft2(window * stack[0]) + for i in range(1, n): + fft_i = np.fft.fft2(window * stack[i]) + drc, fft_shift = cross_correlation_shift( + fft_ref, + fft_i, + upsample_factor=int(upsample_factor), + max_shift=max_shift, + fft_input=True, + fft_output=True, + return_shifted_image=True, + ) + if not isinstance(drc, (list, tuple, np.ndarray)) or len(drc) < 2: + raise RuntimeError("cross_correlation_shift returned an invalid shift vector.") + shifts[i, 0] = float(drc[0]) + shifts[i, 1] = float(drc[1]) + fft_ref = fft_ref * (i / (i + 1)) + fft_shift / (i + 1) + + shifts -= np.mean(shifts, axis=0, keepdims=True) + aligned = np.empty_like(stack, dtype=np.float32) + for i in range(n): + aligned[i] = ndi_shift( + stack[i], + shift=(float(shifts[i, 0]), float(shifts[i, 1])), + order=int(shift_order), + mode="nearest", + prefilter=False, + ) + + self.image_ref = np.mean(aligned, axis=0) + self.preprocess_shifts = shifts.reshape(self.index_shape + (2,)) + return self + + def define_model( + self, + *, + origin_row: float | Sequence[float], + origin_col: float | Sequence[float], + components: list[RenderComponent], + device: torch.device | str | None = None, + dtype: torch.dtype | None = None, + mask: np.ndarray | torch.Tensor | None = None, + origin_key: str = "origin", + ) -> "ModelDiffraction": + if self.image_ref is None: + self.preprocess() + if self.image_ref is None: + raise RuntimeError("image_ref not available.") + + h, w = int(self.image_ref.shape[0]), int(self.image_ref.shape[1]) + dev = torch.device(device) if device is not None else torch.device("cpu") + dt = dtype if dtype is not None else torch.float32 + + mask_t = None + if mask is not None: + mask_t = ( + mask.to(device=dev, dtype=dt) + if torch.is_tensor(mask) + else torch.as_tensor(mask, device=dev, dtype=dt) + ) + if tuple(mask_t.shape) != (h, w): + raise ValueError("mask must have shape (H, W).") + + origin = OriginND( + ndim=2, + init=[ + _parse_init(origin_row, name="origin_row"), + _parse_init(origin_col, name="origin_col"), + ], + ) + origin._quantem_origin_key = str(origin_key) # type: ignore[attr-defined] + + for component in components: + if hasattr(component, "set_origin"): + component.set_origin(origin) # type: ignore[misc] + elif hasattr(component, "origin") and getattr(component, "origin") is None: + component.origin = origin # type: ignore[attr-defined] + + self.model = AdditiveRenderModel(origin=origin, components=list(components)).to( + device=dev, dtype=dt + ) + self.ctx = RenderContext(shape=(h, w), device=dev, dtype=dt, mask=mask_t, fields={}) + self.target_mean = torch.as_tensor(self.image_ref, device=dev, dtype=dt) + + s0 = self._get_model_state_dict_copy() + self.state_initialized = s0 + self.state_mean_refined = None + self.mean_refined = False + self._clear_fit_history_all() + self.remove_optimizer() + return self + + def fit_mean_diffraction_pattern( + self, + *, + n_steps: int = 200, + reset: bool | Literal["initialized", "mean_refined"] = False, + optimizer_params: OptimizerType | dict | None = None, + scheduler_params: SchedulerType | dict | None = None, + constraint_weight: float = 1.0, + constraint_params: dict[str, Any] | None = None, + progress: bool = True, + ) -> "ModelDiffraction": + """ + Fit the mean diffraction pattern. + + Parameters + ---------- + n_steps : int, optional + Number of optimization steps. + reset : bool | Literal["initialized", "mean_refined"], optional + Reset behavior before fitting. + optimizer_params : dict | None, optional + Optimizer override for this fit call. + scheduler_params : dict | None, optional + Scheduler override for this fit call. + constraint_weight : float, optional + Global multiplier for soft-constraint loss. + constraint_params : dict[str, Any] | None, optional + Optional constraint updates applied once to components before fitting. + If ``None``, previously assigned constraints are reused. + progress : bool, optional + If ``True``, show progress bar. + + Returns + ------- + ModelDiffraction + Self, with updated fit state and history. + + Raises + ------ + RuntimeError + If model/context/target are not defined. + ValueError + If ``reset`` has an unsupported value. + + Notes + ----- + Constraint assignments persist on components across fit calls. + """ + if self.model is None or self.ctx is None or self.target_mean is None: + raise RuntimeError("Call .define_model(...) first.") + if reset is True: + self.reset("initialized") + elif isinstance(reset, str): + if reset not in ("initialized", "mean_refined"): + raise ValueError("reset must be False, True, 'initialized', or 'mean_refined'.") + self.reset(reset_to=cast(Literal["initialized", "mean_refined"], reset)) + elif reset not in (False,): + raise ValueError("reset must be False, True, 'initialized', or 'mean_refined'.") + + self.fit_render( + target=self.target_mean, + n_steps=int(n_steps), + constraint_weight=float(constraint_weight), + constraint_params=constraint_params, + optimizer_params=optimizer_params, + scheduler_params=scheduler_params, + progress=bool(progress), + run_key="mean", + ) + + s_fit = self._get_model_state_dict_copy() + self.state_mean_refined = self._clone_state_dict(s_fit) + self.mean_refined = True + return self + + def reset( + self, + reset_to: Literal["initialized", "mean_refined"] = "mean_refined", + ) -> "ModelDiffraction": + if reset_to == "initialized": + state = self.state_initialized + if state is None: + raise RuntimeError( + "initialized state is unavailable. Call .define_model(...) first." + ) + self._clear_fit_history_all() + elif reset_to == "mean_refined": + state = self.state_mean_refined + if state is None: + raise RuntimeError( + "mean_refined state is unavailable. Run .fit_mean_diffraction_pattern(...) first." + ) + mean_hist = self.fit_history.get("mean") + self._clear_fit_history_all() + if mean_hist is not None: + self.fit_history["mean"] = mean_hist + else: + raise ValueError("reset_to must be 'initialized' or 'mean_refined'.") + + self._load_model_state_dict_copy(state) + return self + + @property + def render_mean_refined(self) -> np.ndarray: + if self.state_mean_refined is None: + raise RuntimeError( + "mean_refined state is unavailable. Run .fit_mean_diffraction_pattern(...) first." + ) + return self._render_state_array(self.state_mean_refined) diff --git a/src/quantem/diffraction/model_fitting_visualizations.py b/src/quantem/diffraction/model_fitting_visualizations.py new file mode 100644 index 00000000..6dfb4f62 --- /dev/null +++ b/src/quantem/diffraction/model_fitting_visualizations.py @@ -0,0 +1,428 @@ +from typing import TYPE_CHECKING, Any, Literal, cast + +import numpy as np +from matplotlib import gridspec +from matplotlib import pyplot as plt + +from quantem.core import config +from quantem.core.visualization import show_2d + +if TYPE_CHECKING: + from quantem.diffraction.model_fitting import ModelDiffraction + + +class ModelDiffractionVisualizations: + def _plot_overlays( + self, + ax: Any, + origin_rc: np.ndarray, + disk_centers_rc: np.ndarray, + *, + overlay_origin: bool = True, + overlay_disks: bool = True, + origin_marker_kwargs: dict[str, Any] | None = None, + disk_marker_kwargs: dict[str, Any] | None = None, + ) -> None: + """ + Plot origin and disk-center overlays on an axis. + + Parameters + ---------- + ax : Any + Matplotlib axis receiving overlays. + origin_rc : np.ndarray + Origin coordinate as ``(row, col)``. + disk_centers_rc : np.ndarray + Disk centers as ``(N, 2)`` in ``(row, col)`` order. + overlay_origin : bool, optional + If ``True``, plot origin marker. + overlay_disks : bool, optional + If ``True``, plot disk-center markers. + origin_marker_kwargs : dict[str, Any] | None, optional + Matplotlib kwargs merged onto origin marker defaults. + disk_marker_kwargs : dict[str, Any] | None, optional + Matplotlib kwargs merged onto disk marker defaults. + + Returns + ------- + None + """ + colors = config.get("viz.colors.set") + if overlay_origin and origin_rc.shape == (2,): + kw_origin = { + "marker": "+", + "color": colors[0], + "markersize": 10, + "markeredgewidth": 3, + "linestyle": "None", + } + if origin_marker_kwargs is not None: + kw_origin.update(origin_marker_kwargs) + ax.plot(float(origin_rc[1]), float(origin_rc[0]), **kw_origin) + + if overlay_disks and disk_centers_rc.ndim == 2 and disk_centers_rc.shape[0] > 0: + kw_disks = { + "marker": "x", + "color": colors[1], + "markersize": 5, + "markeredgewidth": 2.0, + "linestyle": "None", + } + if disk_marker_kwargs is not None: + kw_disks.update(disk_marker_kwargs) + ax.plot(disk_centers_rc[:, 1], disk_centers_rc[:, 0], **kw_disks) + + def plot_losses( + self, figax: tuple[Any, Any] | None = None, plot_lrs: bool = True + ) -> tuple[Any, Any]: + md = cast("ModelDiffraction", self) + colors = config.get("viz.colors.set") + loss_color = "k" + lr_color = colors[8] + + if figax is None: + fig, ax = plt.subplots() + else: + fig, ax = figax + + mean_hist = md.fit_history.get("mean") + losses = np.asarray([] if mean_hist is None else mean_hist.losses, dtype=np.float64) + if losses.size == 0: + ax.text( + 0.5, + 0.5, + "No fit history available", + ha="center", + va="center", + transform=ax.transAxes, + ) + ax.set_xlabel("Iterations") + ax.set_ylabel("Loss") + if figax is None: + plt.tight_layout() + plt.show() + return fig, ax + + iters = np.arange(losses.size) + lines: list[Any] = [] + lines.extend(ax.semilogy(iters, losses, c=loss_color, lw=2, label="loss")) + ax.set_xlabel("Iterations") + ax.set_ylabel("Loss", color=loss_color) + ax.tick_params(axis="y", which="both", colors=loss_color) + ax.spines["left"].set_color(loss_color) + ax.set_xbound(-2, max(1, int(iters.max())) + 2) + + lrs = np.asarray([] if mean_hist is None else mean_hist.lrs, dtype=np.float64) + if plot_lrs and lrs.size > 0: + if lrs.size == losses.size and not np.allclose(lrs, lrs[0]): + ax_lr = ax.twinx() + ax.set_zorder(2) + ax_lr.set_zorder(1) + ax.patch.set_visible(False) + ax_lr.spines["left"].set_visible(False) + lines.extend( + ax_lr.semilogy(np.arange(lrs.size), lrs, c=lr_color, lw=2, ls="--", label="LR") + ) + ax_lr.set_ylabel("LR", color=lr_color) + ax_lr.tick_params(axis="y", which="both", colors=lr_color) + ax_lr.spines["right"].set_color(lr_color) + else: + ax.set_title(f"LR: {float(lrs[-1]):.2e}", fontsize=10) + + labels = [line.get_label() for line in lines] + if len(labels) > 1: + ax.legend(lines, labels, loc="upper right") + + if figax is None: + plt.tight_layout() + plt.show() + return fig, ax + + def visualize( + self, + *, + power: float = 0.25, + cbar: bool = False, + axsize: tuple[int, int] = (6, 6), + overlay: bool = True, + overlay_origin: bool = True, + overlay_disks: bool = True, + overlay_on: Literal["model", "both"] = "model", + origin_marker_kwargs: dict[str, Any] | None = None, + disk_marker_kwargs: dict[str, Any] | None = None, + ) -> tuple[Any, Any]: + """ + Visualize fit losses with reference/model image panels. + + Parameters + ---------- + power : float, optional + Power-law display scaling. + cbar : bool, optional + If ``True``, draw colorbars. + axsize : tuple[int, int], optional + Axis size passed through to ``show_2d``. + overlay : bool, optional + If ``True``, draw coordinate overlays. + overlay_origin : bool, optional + If ``True``, include origin marker in overlays. + overlay_disks : bool, optional + If ``True``, include disk-center markers in overlays. + overlay_on : {"model", "both"}, optional + Which image panel(s) receive overlays. + origin_marker_kwargs : dict[str, Any] | None, optional + Marker kwargs override for origin marker. + disk_marker_kwargs : dict[str, Any] | None, optional + Marker kwargs override for disk-center markers. + + Returns + ------- + tuple[Any, Any] + ``(fig, axs)`` for further editing. + + Raises + ------ + RuntimeError + If model/context are not defined. + ValueError + If ``overlay_on`` is invalid. + """ + md = cast("ModelDiffraction", self) + + if md.image_ref is None: + md.preprocess() + if md.image_ref is None or md.model is None or md.ctx is None: + raise RuntimeError("Call .define_model(...) first.") + + fig = plt.figure(figsize=(12, 7)) + gs = gridspec.GridSpec(2, 1, height_ratios=[1, 2], hspace=0.3) + ax_top = fig.add_subplot(gs[0]) + md.plot_losses(figax=(fig, ax_top), plot_lrs=True) + + ref = np.asarray(md.image_ref, dtype=np.float32) + pred = md.render_current + refp = ref if power == 1.0 else np.maximum(ref, 0.0) ** float(power) + predp = pred if power == 1.0 else np.maximum(pred, 0.0) ** float(power) + vmin = float(min(refp.min(), predp.min())) + vmax = float(max(refp.max(), predp.max())) + + gs_bot = gridspec.GridSpecFromSubplotSpec(1, 2, subplot_spec=gs[1], wspace=0.15) + axs = np.array( + [fig.add_subplot(gs_bot[0, 0]), fig.add_subplot(gs_bot[0, 1])], dtype=object + ) + show_2d( + [refp, predp], + figax=(fig, axs), + title=["image_ref", "model"], + cmap=config.get("viz.cmap"), + cbar=bool(cbar), + returnfig=False, + axsize=axsize, + vmin=vmin, + vmax=vmax, + ) + + if overlay: + if overlay_on not in ("model", "both"): + raise ValueError("overlay_on must be 'model' or 'both'.") + origin_rc, disk_centers_rc = md.get_overlay_coordinates() + axes = [axs[1]] if overlay_on == "model" else [axs[0], axs[1]] + for ax in axes: + self._plot_overlays( + ax, + origin_rc, + disk_centers_rc, + overlay_origin=overlay_origin, + overlay_disks=overlay_disks, + origin_marker_kwargs=origin_marker_kwargs, + disk_marker_kwargs=disk_marker_kwargs, + ) + + mean_hist = md.fit_history.get("mean") + if mean_hist is not None and len(mean_hist.losses) > 0: + fig.suptitle( + f"Final loss: {mean_hist.losses[-1]:.3e} | Iters: {len(mean_hist.losses)}", + fontsize=13, + y=0.98, + ) + plt.show() + return fig, axs + + def plot_mean_model( + self, + *, + power: float = 0.25, + returnfig: bool = False, + axsize: tuple[int, int] = (6, 6), + overlay: bool = True, + overlay_origin: bool = True, + overlay_disks: bool = True, + overlay_on: Literal["model", "both"] = "model", + origin_marker_kwargs: dict[str, Any] | None = None, + disk_marker_kwargs: dict[str, Any] | None = None, + **_: Any, + ) -> tuple[Any, Any] | None: + """ + Plot reference and model mean diffraction images. + + Parameters + ---------- + power : float, optional + Power-law display scaling. + returnfig : bool, optional + If ``True``, return ``(fig, ax)``. + axsize : tuple[int, int], optional + Axis size passed through to ``show_2d``. + overlay : bool, optional + If ``True``, draw coordinate overlays. + overlay_origin : bool, optional + If ``True``, include origin marker in overlays. + overlay_disks : bool, optional + If ``True``, include disk-center markers in overlays. + overlay_on : {"model", "both"}, optional + Which image panel(s) receive overlays. + origin_marker_kwargs : dict[str, Any] | None, optional + Marker kwargs override for origin marker. + disk_marker_kwargs : dict[str, Any] | None, optional + Marker kwargs override for disk-center markers. + **_ : Any + Ignored extra kwargs for backward compatibility. + + Returns + ------- + tuple[Any, Any] | None + Figure/axes tuple when ``returnfig=True``; otherwise ``None``. + + Raises + ------ + RuntimeError + If model/context are not defined. + ValueError + If ``overlay_on`` is invalid. + """ + md = cast("ModelDiffraction", self) + if md.image_ref is None: + md.preprocess() + if md.image_ref is None or md.model is None or md.ctx is None: + raise RuntimeError("Call .define_model(...) first.") + + ref = np.asarray(md.image_ref, dtype=np.float32) + pred = md.render_current + + refp = ref if power == 1.0 else np.maximum(ref, 0.0) ** float(power) + predp = pred if power == 1.0 else np.maximum(pred, 0.0) ** float(power) + vmin = float(min(refp.min(), predp.min())) + vmax = float(max(refp.max(), predp.max())) + + fig, ax = show_2d( + [refp, predp], + title=["image_ref", "model"], + cmap=config.get("viz.cmap"), + cbar=False, + returnfig=True, + axsize=axsize, + vmin=vmin, + vmax=vmax, + ) + if overlay: + if overlay_on not in ("model", "both"): + raise ValueError("overlay_on must be 'model' or 'both'.") + origin_rc, disk_centers_rc = md.get_overlay_coordinates() + axs_arr = np.asarray(ax, dtype=object).reshape(-1) + axes = [axs_arr[1]] if overlay_on == "model" else [axs_arr[0], axs_arr[1]] + for a in axes: + self._plot_overlays( + a, + origin_rc, + disk_centers_rc, + overlay_origin=overlay_origin, + overlay_disks=overlay_disks, + origin_marker_kwargs=origin_marker_kwargs, + disk_marker_kwargs=disk_marker_kwargs, + ) + if returnfig: + return fig, ax + return None + + def visualize_components( + self, + components: str | list[str], + *, + power: float = 0.25, + cbar: bool = False, + axsize: tuple[int, int] = (6, 6), + returnfig: bool = False, + overlay: bool = True, + overlay_origin: bool = True, + overlay_disks: bool = True, + origin_marker_kwargs: dict[str, Any] | None = None, + disk_marker_kwargs: dict[str, Any] | None = None, + ) -> tuple[Any, Any] | None: + """ + Render and display a summed component image. + + Parameters + ---------- + components : str | list[str] + Component name or list of component names. Multiple names are + composited by summing component renders. + power : float, optional + Power-law display scaling. + cbar : bool, optional + If ``True``, draw colorbars. + axsize : tuple[int, int], optional + Axis size passed through to ``show_2d``. + returnfig : bool, optional + If ``True``, return ``(fig, ax)``. + overlay : bool, optional + If ``True``, draw coordinate overlays on all component panels. + overlay_origin : bool, optional + If ``True``, include origin marker in overlays. + overlay_disks : bool, optional + If ``True``, include disk-center markers in overlays. + origin_marker_kwargs : dict[str, Any] | None, optional + Marker kwargs override for origin marker. + disk_marker_kwargs : dict[str, Any] | None, optional + Marker kwargs override for disk-center markers. + + Returns + ------- + tuple[Any, Any] | None + Figure/axes tuple when ``returnfig=True``; otherwise ``None``. + """ + md = cast("ModelDiffraction", self) + names = [components] if isinstance(components, str) else list(components) + if len(names) == 0: + raise ValueError("components must contain at least one component name.") + + rendered = [ + np.asarray(md.get_rendered_component(name), dtype=np.float32) for name in names + ] + summed = np.sum(np.stack(rendered, axis=0), axis=0) + summed_scaled = summed if power == 1.0 else np.maximum(summed, 0.0) ** float(power) + title = names[0] if len(names) == 1 else " + ".join(names) + + fig, ax = show_2d( + summed_scaled, + title=title, + cmap=config.get("viz.cmap"), + cbar=bool(cbar), + returnfig=True, + axsize=axsize, + ) + + if overlay: + origin_rc, disk_centers_rc = md.get_overlay_coordinates() + self._plot_overlays( + ax, + origin_rc, + disk_centers_rc, + overlay_origin=overlay_origin, + overlay_disks=overlay_disks, + origin_marker_kwargs=origin_marker_kwargs, + disk_marker_kwargs=disk_marker_kwargs, + ) + + if returnfig: + return fig, ax + return None