diff --git a/.gitignore b/.gitignore index be239dda..e7db765d 100644 --- a/.gitignore +++ b/.gitignore @@ -146,6 +146,8 @@ cython_debug/ # Workspace stuff workspace/ +dev_notes/ +tools/ # Binaries *.tiff @@ -160,4 +162,6 @@ workspace/ # IDE .idea/ .vscode/ -.loglogin \ No newline at end of file +.loglogin + +AGENTS.md \ No newline at end of file diff --git a/docs/source/using_pty_chi/index.rst b/docs/source/using_pty_chi/index.rst index 840391ac..c5cceb2d 100644 --- a/docs/source/using_pty_chi/index.rst +++ b/docs/source/using_pty_chi/index.rst @@ -6,6 +6,7 @@ Using Pty-chi data_structures task + workflows options io initialization_recommendations diff --git a/docs/source/using_pty_chi/workflows.rst b/docs/source/using_pty_chi/workflows.rst new file mode 100644 index 00000000..eb94a85c --- /dev/null +++ b/docs/source/using_pty_chi/workflows.rst @@ -0,0 +1,214 @@ +Workflows +========= + +A :class:`~ptychi.api.task.PtychographyTask` represents one reconstruction run +with fixed initialization and settings. A workflow coordinates multiple tasks, +carrying reconstructed parameters from one task to the next to implement a +larger reconstruction strategy. + + +Multiscan shared-object reconstruction +-------------------------------------- + +:class:`~ptychi.workflows.MultiscanSharedObjectWorkflow` reconstructs several +datasets collected from sub-regions of one object. It creates one task per +dataset so that every scan retains its own probe, positions, OPR weights, and +optimizer state. After each task runs, only its reconstructed object tensor is +copied to the next task in a cyclic sequence. + +Supply one options object and one data array per scan. Every supplied list must +have the same length, and at least one scan is required: + +.. code-block:: python + + import ptychi.api as api + from ptychi.workflows import MultiscanSharedObjectWorkflow + + workflow = MultiscanSharedObjectWorkflow( + [task_options_1, task_options_2], + diffraction_data=[diffraction_data_1, diffraction_data_2], + object_data=[object_guess_1, object_guess_2], + probe_data=[probe_guess_1, probe_guess_2], + probe_position_x_px=[position_x_1, position_x_2], + probe_position_y_px=[position_y_1, position_y_2], + opr_mode_weights_data=[opr_weights_1, None], + valid_pixel_mask=None, + workflow_options=api.MultiscanSharedObjectWorkflowOptions( + num_outer_epochs=20, + num_inner_epochs=1, + ), + ) + workflow.run() + +``opr_mode_weights_data`` and ``valid_pixel_mask`` may each be ``None`` for all +scans, or a list containing an array or ``None`` for each scan. The workflow +sets every task's progress total to ``num_outer_epochs * num_inner_epochs``; +the values originally stored in the caller's task options are not modified. + +All scans must use the same object tensor shape, pixel geometry, and multislice +spacing. They must also use +``object_options.determine_position_origin_coords_by = ObjectPosOriginCoordsMethods.SUPPORT``. +The supplied positions must therefore already be expressed in a common, +SUPPORT-centered coordinate frame. The workflow does not align or stitch scan +coordinates. + +Tasks are available in input order through ``workflow.tasks``. At completion, +the final object is copied to every task, while each task's probe, positions, +OPR weights, preconditioners, and optimizer state remain independent. Inactive +tasks are offloaded to CPU, and all tasks are on CPU when the workflow returns. +A workflow instance cannot be run a second time. + + +Progressive-resolution reconstruction +-------------------------------------- + +:class:`~ptychi.workflows.ProgressiveResolutionWorkflow` starts at reduced +spatial resolution and progressively increases the resolution until it reaches +the resolution of the supplied data. This can provide a useful coarse initial +solution for the subsequent, more expensive resolution levels. + +The downsampling factor at level ``i`` is + +.. math:: + + 2^{N - 1 - i}, + +where ``N`` is the total number of levels and ``i`` starts at zero. Thus, a +three-level workflow uses factors 4, 2, and 1. + +At the first level, the workflow resizes the initial object and probe, divides +the probe positions by the factor, and increases the object pixel size by the +same factor. At every later level, it resizes the reconstructed object and +probe from the previous task, scales the reconstructed probe positions, and +copies the OPR mode weights. The final task therefore uses the full-resolution +data and the original object pixel size. + + +Basic usage +~~~~~~~~~~~ + +Configure the reconstruction algorithm as you would for a regular task, then +provide the number of resolution levels and the number of epochs at each +level: + +.. code-block:: python + + import ptychi.api as api + from ptychi.workflows import ProgressiveResolutionWorkflow + + task_options = api.LSQMLOptions() + task_options.object_options.pixel_size_m = pixel_size_m + task_options.object_options.optimizable = True + task_options.object_options.optimizer = api.Optimizers.SGD + task_options.object_options.step_size = 1 + task_options.probe_options.optimizable = True + task_options.probe_options.optimizer = api.Optimizers.SGD + task_options.probe_options.step_size = 1 + + workflow_options = api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=3, + num_epochs_all_levels=[20, 20, 20], + ) + + workflow = ProgressiveResolutionWorkflow( + task_options, + diffraction_data=diffraction_data, + object_data=object_guess, + probe_data=probe_guess, + probe_position_x_px=positions_px[:, 1], + probe_position_y_px=positions_px[:, 0], + opr_mode_weights_data=opr_mode_weights, + valid_pixel_mask=valid_pixel_mask, + workflow_options=workflow_options, + ) + workflow.run() + +``opr_mode_weights_data`` and ``valid_pixel_mask`` are optional, as they are +for :class:`~ptychi.api.task.PtychographyTask`. The workflow overrides +``task_options.reconstructor_options.num_epochs`` for each task using the +corresponding value in ``num_epochs_all_levels``. It does not modify the +original options object. + + +Far-field and near-field data +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The workflow chooses how to reduce measured data from +``task_options.data_options.free_space_propagation_distance_m``: + +* An infinite distance selects far-field ptychography. Diffraction patterns + are reduced by cropping the low-frequency region in reciprocal space. +* A finite distance selects near-field ptychography. Measured intensity is + real-space data, so it is resized in the same way as the object and probe. + +For far-field data, set ``task_options.data_options.fft_shift`` according to +the layout of the supplied diffraction patterns. Set it to ``True`` when the +DC component is at the center; the dataset will FFT-shift the cropped pattern +to match the forward model, whose DC component is at the top-left corner. Set +it to ``False`` when the supplied data already has DC at the top-left corner. +For near-field data, this option should normally be ``False`` because the +measured intensity is already in real space. + +When a validity mask is provided, the workflow crops it together with +far-field diffraction data. For near-field data it uses nearest-neighbor +resizing so that the mask remains boolean. + + +Accessing results from each level +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Every created task remains available in ``workflow.tasks``, ordered from the +lowest resolution to the full resolution. Use the normal task APIs to retrieve +reconstructed data: + +.. code-block:: python + + for i_level, task in enumerate(workflow.tasks): + object_at_level = task.get_data_to_cpu("object", as_numpy=True) + probe_at_level = task.get_data_to_cpu("probe", as_numpy=True) + + full_resolution_task = workflow.get_full_resolution_task() + reconstructed_object = full_resolution_task.get_data_to_cpu( + "object", as_numpy=True + ) + +:meth:`~ptychi.workflows.ProgressiveResolutionWorkflow.get_full_resolution_task` +is available only after all levels finish successfully. A workflow instance +cannot be run a second time. + + +Memory behavior +~~~~~~~~~~~~~~~ + +The workflow copies all constructor data to CPU memory. If a tensor supplied +by the caller is on an accelerator, the workflow emits a warning and creates a +CPU copy; it does not move or otherwise modify the caller's original tensor. +This means the caller remains responsible for releasing any original GPU +inputs. + +After each resolution level finishes, the workflow calls +:meth:`~ptychi.api.task.PtychographyTask.set_large_tensor_device` to offload +the task's reconstruction parameters, diffraction patterns, optimizer state, +and registered reconstructor buffers to CPU memory. The cached task can still +be inspected through ``workflow.tasks`` without keeping every resolution +level resident on the GPU. + + +API reference +------------- + +.. autoclass:: ptychi.workflows.MultiscanSharedObjectWorkflow + :members: + :show-inheritance: + +.. autoclass:: ptychi.api.options.workflow.MultiscanSharedObjectWorkflowOptions + :members: + :show-inheritance: + +.. autoclass:: ptychi.workflows.ProgressiveResolutionWorkflow + :members: + :show-inheritance: + +.. autoclass:: ptychi.api.options.workflow.ProgressiveResolutionWorkflowOptions + :members: + :show-inheritance: diff --git a/src/ptychi/api/options/__init__.py b/src/ptychi/api/options/__init__.py index 60122a4c..5b5555ba 100644 --- a/src/ptychi/api/options/__init__.py +++ b/src/ptychi/api/options/__init__.py @@ -9,3 +9,4 @@ from .pie import * from .dm import * from .bh import * +from .workflow import * diff --git a/src/ptychi/api/options/data.py b/src/ptychi/api/options/data.py index aaaf27b7..eb1aa6e4 100644 --- a/src/ptychi/api/options/data.py +++ b/src/ptychi/api/options/data.py @@ -28,7 +28,17 @@ class PtychographyDataOptions(base.Options): """The wavelength in meters.""" fft_shift: bool = True - """Whether to FFT-shift the diffraction data.""" + """Whether to FFT-shift the diffraction data when building the dataset. For far-field + ptychography, the forward model does not shift the image after FFT, meaning the + predicted intensity has its DC component at the top left corner. To match the prediction, + measured intensity should be pre-shifted if the DC component of the given diffraction + patterns is at the center. + + However, if the given diffraction patterns are already shifted so that the DC compoenent + is at the top left, or if you are reconstructing near-field ptychography data where + the forward model does not involve Fraunhofer diffraction implemented with FFT, ensure + this option is set to `False` to avoid the erroneous shifting. + """ detector_pixel_size_m: float = 1e-8 """The detector pixel size in meters.""" diff --git a/src/ptychi/api/options/workflow.py b/src/ptychi/api/options/workflow.py new file mode 100644 index 00000000..2eb90f04 --- /dev/null +++ b/src/ptychi/api/options/workflow.py @@ -0,0 +1,61 @@ +# Copyright © 2025 UChicago Argonne, LLC All right reserved +# Full license accessible at https://github.com//AdvancedPhotonSource/pty-chi/blob/main/LICENSE + +from pydantic import model_validator +from pydantic.dataclasses import dataclass + +import ptychi.api.options.base as base + + +__all__ = [ + "WorkflowOptions", + "ProgressiveResolutionWorkflowOptions", + "MultiscanSharedObjectWorkflowOptions", +] + + +@dataclass +class WorkflowOptions(base.Options): + """Base class for options that configure a workflow.""" + + +@dataclass +class ProgressiveResolutionWorkflowOptions(WorkflowOptions): + """Options for a progressive-resolution reconstruction workflow.""" + + num_resolution_levels: int + """The number of resolution levels, including the full-resolution level.""" + + num_epochs_all_levels: list[int] + """The number of reconstruction epochs to run at each resolution level.""" + + @model_validator(mode="after") + def _validate_resolution_levels(self): + if self.num_resolution_levels <= 0: + raise ValueError("`num_resolution_levels` must be greater than 0.") + if len(self.num_epochs_all_levels) != self.num_resolution_levels: + raise ValueError( + "`num_epochs_all_levels` must contain one value for each resolution level." + ) + if any(num_epochs <= 0 for num_epochs in self.num_epochs_all_levels): + raise ValueError("All values in `num_epochs_all_levels` must be greater than 0.") + return self + + +@dataclass +class MultiscanSharedObjectWorkflowOptions(WorkflowOptions): + """Options for a multiscan reconstruction with a shared object.""" + + num_outer_epochs: int + """The number of complete passes over all scans.""" + + num_inner_epochs: int + """The number of epochs to run for each scan in each pass.""" + + @model_validator(mode="after") + def _validate_epochs(self): + if self.num_outer_epochs <= 0: + raise ValueError("`num_outer_epochs` must be greater than 0.") + if self.num_inner_epochs <= 0: + raise ValueError("`num_inner_epochs` must be greater than 0.") + return self diff --git a/src/ptychi/api/task.py b/src/ptychi/api/task.py index ec6c5df6..4391e0a3 100644 --- a/src/ptychi/api/task.py +++ b/src/ptychi/api/task.py @@ -2,6 +2,7 @@ # Full license accessible at https://github.com//AdvancedPhotonSource/pty-chi/blob/main/LICENSE from typing import Literal, Optional, Union, overload +import dataclasses from dataclasses import dataclass from types import TracebackType import random @@ -680,18 +681,18 @@ def set_large_tensor_device( self, device: Literal["cpu", "cuda"] | torch.device | None = None, ) -> None: - """Move large task buffers between CPU and a target device. + """Move task tensors and reconstruction state between devices. This helper is aimed at multi-task workflows where only one task is active on the accelerator at a time. Call it with ``device="cpu"`` to - offload the heavy object/probe/diffraction buffers to system memory, - and call it again (without arguments, or with an explicit device string) - before resuming the task to bring the tensors back to the accelerator. + offload task tensors and optimizer/reconstructor state to system memory, + and call it again before resuming the task to restore that state to the + accelerator. Parameters ---------- device : str | torch.device | None, optional - Target device for the large buffers. If None, tensors are moved back + Target device for task state. If None, tensors are moved back to the current default device. If a string is given, it must be either "cpu" or "cuda". """ @@ -700,21 +701,43 @@ def set_large_tensor_device( device = torch.get_default_device() device = torch.device(device) - if self.reconstructor is None: - raise RuntimeError("Reconstructor is not built yet.") + if self.reconstructor is None or self.dataset is None: + raise RuntimeError("Task is not built yet.") parameter_group = self.reconstructor.parameter_group with torch.no_grad(): - # Move object and probe buffers. - parameter_group.object.to(device) - parameter_group.probe.to(device) + for parameter in parameter_group.get_all_reconstruct_parameters(): + parameter.to(device) + if parameter.optimizer is not None: + for key, value in list(parameter.optimizer.state.items()): + parameter.optimizer.state[key] = utils.move_nested_tensors_to_device( + value, device + ) - # Move diffraction patterns. self.dataset.patterns = self.dataset.patterns.to(device) - # Keep dataset bookkeeping in sync with where patterns live. + self.dataset.move_attributes_to_device(device) self.dataset.save_data_on_device = device.type != "cpu" - # Move intermediate variables in forward model. + buffers = self.reconstructor.reconstructor_buffers + for name in buffers.get_all_names(): + if not hasattr(buffers, name): + continue + buffers.__setattr__( + name, + utils.move_nested_tensors_to_device(getattr(buffers, name), device), + bypass_check=True, + ) + + for name, value in vars(self.reconstructor).items(): + if not name.endswith("_momentum_params") or not dataclasses.is_dataclass(value): + continue + for field in dataclasses.fields(value): + setattr( + value, + field.name, + utils.move_nested_tensors_to_device(getattr(value, field.name), device), + ) + self.reconstructor.forward_model.move_intermediate_variables_to_device(device) if device.type == "cpu": diff --git a/src/ptychi/data_structures/base.py b/src/ptychi/data_structures/base.py index 002e23a3..87d329ce 100644 --- a/src/ptychi/data_structures/base.py +++ b/src/ptychi/data_structures/base.py @@ -86,6 +86,8 @@ class ReconstructParameter(Module): name = None optimizable: bool = True optimization_plan: "api.OptimizationPlan" = None + preconditioner: Optional[Tensor] + update_buffer: Optional[Tensor] optimizer = None step_size_scheduler = None is_dummy = False @@ -151,8 +153,8 @@ def __init__( self.sub_modules = [] self.optimizable_sub_modules = [] self.is_complex = is_complex - self.preconditioner = None - self.update_buffer = None + self.register_buffer("preconditioner", None, persistent=False) + self.register_buffer("update_buffer", None, persistent=False) if is_complex: if data is not None: diff --git a/src/ptychi/workflows/__init__.py b/src/ptychi/workflows/__init__.py new file mode 100644 index 00000000..271479df --- /dev/null +++ b/src/ptychi/workflows/__init__.py @@ -0,0 +1,13 @@ +# Copyright © 2025 UChicago Argonne, LLC All right reserved +# Full license accessible at https://github.com//AdvancedPhotonSource/pty-chi/blob/main/LICENSE + +from .base import BaseWorkflow +from .multiscan_shared_object import MultiscanSharedObjectWorkflow +from .progressive_resolution import ProgressiveResolutionWorkflow + + +__all__ = [ + "BaseWorkflow", + "MultiscanSharedObjectWorkflow", + "ProgressiveResolutionWorkflow", +] diff --git a/src/ptychi/workflows/base.py b/src/ptychi/workflows/base.py new file mode 100644 index 00000000..849140a8 --- /dev/null +++ b/src/ptychi/workflows/base.py @@ -0,0 +1,279 @@ +# Copyright © 2025 UChicago Argonne, LLC All right reserved +# Full license accessible at https://github.com//AdvancedPhotonSource/pty-chi/blob/main/LICENSE + +from abc import ABC, abstractmethod +import copy +import dataclasses +from dataclasses import dataclass +from typing import Optional +import warnings + +import numpy as np +import torch +from numpy import ndarray +from torch import Tensor + +from ptychi.api.options.task import PtychographyTaskOptions +from ptychi.api.options.workflow import WorkflowOptions +from ptychi.api.task import PtychographyTask, TaskArray + + +class _UnsetWorkflowData: + pass + + +_UNSET = _UnsetWorkflowData() + + +@dataclass(frozen=True) +class _WorkflowData: + diffraction_data: TaskArray + object_data: TaskArray + probe_data: TaskArray + probe_position_x_px: TaskArray + probe_position_y_px: TaskArray + opr_mode_weights_data: Optional[TaskArray] + valid_pixel_mask: Optional[TaskArray] + + +@dataclass(frozen=True) +class _WorkflowTensorData: + diffraction_data: Tensor + object_data: Tensor + probe_data: Tensor + probe_position_x_px: Tensor + probe_position_y_px: Tensor + opr_mode_weights_data: Optional[Tensor] + valid_pixel_mask: Optional[Tensor] + + +class BaseWorkflow(ABC): + """Base class for workflows composed of one or more ptychography tasks.""" + + def __init__( + self, + task_options: PtychographyTaskOptions, + *args, + diffraction_data: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + object_data: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + probe_data: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + probe_position_x_px: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + probe_position_y_px: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + opr_mode_weights_data: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + valid_pixel_mask: Optional[TaskArray] | _UnsetWorkflowData = _UNSET, + workflow_options: WorkflowOptions, + **kwargs, + ) -> None: + if not isinstance(task_options, PtychographyTaskOptions): + raise TypeError("`task_options` must be a PtychographyTaskOptions instance.") + if not isinstance(workflow_options, WorkflowOptions): + raise TypeError("`workflow_options` must be a WorkflowOptions instance.") + + self.task_options = task_options + self.workflow_options = workflow_options + self.tasks: list[PtychographyTask] = [] + self._task_args = args + self._task_kwargs = kwargs + + data = self._resolve_workflow_data( + diffraction_data=diffraction_data, + object_data=object_data, + probe_data=probe_data, + probe_position_x_px=probe_position_x_px, + probe_position_y_px=probe_position_y_px, + opr_mode_weights_data=opr_mode_weights_data, + valid_pixel_mask=valid_pixel_mask, + ) + self._warn_for_gpu_data(data) + cpu_data = self._copy_workflow_data_to_cpu(data) + self.diffraction_data = cpu_data.diffraction_data + self.object_data = cpu_data.object_data + self.probe_data = cpu_data.probe_data + self.probe_position_x_px = cpu_data.probe_position_x_px + self.probe_position_y_px = cpu_data.probe_position_y_px + self.opr_mode_weights_data = cpu_data.opr_mode_weights_data + self.valid_pixel_mask = cpu_data.valid_pixel_mask + self.task_options = self._copy_task_options() + + def _resolve_workflow_data( + self, + *, + task_options: Optional[PtychographyTaskOptions] = None, + diffraction_data: Optional[TaskArray] | _UnsetWorkflowData, + object_data: Optional[TaskArray] | _UnsetWorkflowData, + probe_data: Optional[TaskArray] | _UnsetWorkflowData, + probe_position_x_px: Optional[TaskArray] | _UnsetWorkflowData, + probe_position_y_px: Optional[TaskArray] | _UnsetWorkflowData, + opr_mode_weights_data: Optional[TaskArray] | _UnsetWorkflowData, + valid_pixel_mask: Optional[TaskArray] | _UnsetWorkflowData, + ) -> _WorkflowData: + task_options = self.task_options if task_options is None else task_options + return _WorkflowData( + diffraction_data=self._resolve_data_field( + value=diffraction_data, + option_owner=task_options.data_options, + option_field_name="data", + option_path="task_options.data_options.data", + kwarg_name="diffraction_data", + required=True, + ), + object_data=self._resolve_data_field( + value=object_data, + option_owner=task_options.object_options, + option_field_name="initial_guess", + option_path="task_options.object_options.initial_guess", + kwarg_name="object_data", + required=True, + ), + probe_data=self._resolve_data_field( + value=probe_data, + option_owner=task_options.probe_options, + option_field_name="initial_guess", + option_path="task_options.probe_options.initial_guess", + kwarg_name="probe_data", + required=True, + ), + probe_position_x_px=self._resolve_data_field( + value=probe_position_x_px, + option_owner=task_options.probe_position_options, + option_field_name="position_x_px", + option_path="task_options.probe_position_options.position_x_px", + kwarg_name="probe_position_x_px", + required=True, + ), + probe_position_y_px=self._resolve_data_field( + value=probe_position_y_px, + option_owner=task_options.probe_position_options, + option_field_name="position_y_px", + option_path="task_options.probe_position_options.position_y_px", + kwarg_name="probe_position_y_px", + required=True, + ), + opr_mode_weights_data=self._resolve_data_field( + value=opr_mode_weights_data, + option_owner=task_options.opr_mode_weight_options, + option_field_name="initial_weights", + option_path="task_options.opr_mode_weight_options.initial_weights", + kwarg_name="opr_mode_weights_data", + required=False, + ), + valid_pixel_mask=self._resolve_data_field( + value=valid_pixel_mask, + option_owner=task_options.data_options, + option_field_name="valid_pixel_mask", + option_path="task_options.data_options.valid_pixel_mask", + kwarg_name="valid_pixel_mask", + required=False, + ), + ) + + @staticmethod + def _resolve_data_field( + *, + value, + option_owner, + option_field_name: str, + option_path: str, + kwarg_name: str, + required: bool, + ): + option_value = getattr(option_owner, option_field_name) + if value is not _UNSET: + if option_value is not None: + warnings.warn( + f"`{option_path}` is deprecated and was ignored because " + f"`{kwarg_name}` was supplied to the workflow.", + DeprecationWarning, + stacklevel=4, + ) + resolved_value = value + elif option_value is not None: + warnings.warn( + f"Passing workflow data via `{option_path}` is deprecated; pass " + f"`{kwarg_name}` to the workflow instead.", + DeprecationWarning, + stacklevel=4, + ) + resolved_value = option_value + else: + resolved_value = None + + if required and resolved_value is None: + raise ValueError( + f"`{kwarg_name}` is required. Passing it through `{option_path}` is " + "temporarily supported but deprecated." + ) + return resolved_value + + @staticmethod + def _warn_for_gpu_data(data: _WorkflowData) -> None: + gpu_fields = [ + field.name + for field in dataclasses.fields(data) + if isinstance(getattr(data, field.name), Tensor) + and getattr(data, field.name).device.type != "cpu" + ] + if gpu_fields: + warnings.warn( + "Workflow inputs were copied to CPU, but the original tensors remain on " + "their accelerator devices and continue to occupy accelerator memory: " + f"{', '.join(gpu_fields)}.", + UserWarning, + stacklevel=3, + ) + + @staticmethod + def _copy_tensor_to_cpu(data: TaskArray) -> Tensor: + if isinstance(data, Tensor): + return data.detach().to(device="cpu").clone() + if isinstance(data, ndarray): + return torch.from_numpy(np.array(data, copy=True)) + return torch.tensor(data, device="cpu") + + @classmethod + def _copy_optional_tensor_to_cpu(cls, data: Optional[TaskArray]) -> Optional[Tensor]: + if data is None: + return None + return cls._copy_tensor_to_cpu(data) + + @classmethod + def _copy_workflow_data_to_cpu(cls, data: _WorkflowData) -> _WorkflowTensorData: + return _WorkflowTensorData( + diffraction_data=cls._copy_tensor_to_cpu(data.diffraction_data), + object_data=cls._copy_tensor_to_cpu(data.object_data), + probe_data=cls._copy_tensor_to_cpu(data.probe_data), + probe_position_x_px=cls._copy_tensor_to_cpu(data.probe_position_x_px), + probe_position_y_px=cls._copy_tensor_to_cpu(data.probe_position_y_px), + opr_mode_weights_data=cls._copy_optional_tensor_to_cpu( + data.opr_mode_weights_data + ), + valid_pixel_mask=cls._copy_optional_tensor_to_cpu(data.valid_pixel_mask), + ) + + def _copy_task_options( + self, task_options: Optional[PtychographyTaskOptions] = None + ) -> PtychographyTaskOptions: + task_options = self.task_options if task_options is None else task_options + data_fields = ( + task_options.data_options.data, + task_options.data_options.valid_pixel_mask, + task_options.object_options.initial_guess, + task_options.probe_options.initial_guess, + task_options.probe_position_options.position_x_px, + task_options.probe_position_options.position_y_px, + task_options.opr_mode_weight_options.initial_weights, + ) + memo = {id(value): None for value in data_fields if value is not None} + options = copy.deepcopy(task_options, memo) + options.data_options.data = None + options.data_options.valid_pixel_mask = None + options.object_options.initial_guess = None + options.probe_options.initial_guess = None + options.probe_position_options.position_x_px = None + options.probe_position_options.position_y_px = None + options.opr_mode_weight_options.initial_weights = None + return options + + @abstractmethod + def run(self) -> None: + """Run the workflow.""" diff --git a/src/ptychi/workflows/multiscan_shared_object.py b/src/ptychi/workflows/multiscan_shared_object.py new file mode 100644 index 00000000..27a796f9 --- /dev/null +++ b/src/ptychi/workflows/multiscan_shared_object.py @@ -0,0 +1,203 @@ +# Copyright © 2025 UChicago Argonne, LLC All right reserved +# Full license accessible at https://github.com//AdvancedPhotonSource/pty-chi/blob/main/LICENSE + +from typing import Optional + +import ptychi.api as api +from ptychi.api.options.task import PtychographyTaskOptions +from ptychi.api.options.workflow import MultiscanSharedObjectWorkflowOptions +from ptychi.api.task import PtychographyTask, TaskArray +from ptychi.workflows.base import BaseWorkflow, _UNSET, _UnsetWorkflowData + + +class MultiscanSharedObjectWorkflow(BaseWorkflow): + """Reconstruct multiple scans by passing one object between scan-specific tasks.""" + + task_options: list[PtychographyTaskOptions] # type: ignore[assignment] + workflow_options: MultiscanSharedObjectWorkflowOptions + + def __init__( + self, + task_options: list[PtychographyTaskOptions], + *args, + diffraction_data: list[TaskArray] | _UnsetWorkflowData = _UNSET, + object_data: list[TaskArray] | _UnsetWorkflowData = _UNSET, + probe_data: list[TaskArray] | _UnsetWorkflowData = _UNSET, + probe_position_x_px: list[TaskArray] | _UnsetWorkflowData = _UNSET, + probe_position_y_px: list[TaskArray] | _UnsetWorkflowData = _UNSET, + opr_mode_weights_data: ( + list[Optional[TaskArray]] | None | _UnsetWorkflowData + ) = _UNSET, + valid_pixel_mask: list[Optional[TaskArray]] | None | _UnsetWorkflowData = _UNSET, + workflow_options: MultiscanSharedObjectWorkflowOptions, + **kwargs, + ) -> None: + if not isinstance(task_options, list): + raise TypeError("`task_options` must be a list.") + if not task_options: + raise ValueError("`task_options` must contain at least one options object.") + if not all(isinstance(options, PtychographyTaskOptions) for options in task_options): + raise TypeError( + "Every member of `task_options` must be a PtychographyTaskOptions instance." + ) + if not isinstance(workflow_options, MultiscanSharedObjectWorkflowOptions): + raise TypeError( + "`workflow_options` must be a " + "MultiscanSharedObjectWorkflowOptions instance." + ) + + n_tasks = len(task_options) + values_by_name = { + "diffraction_data": self._expand_task_data( + "diffraction_data", diffraction_data, n_tasks + ), + "object_data": self._expand_task_data("object_data", object_data, n_tasks), + "probe_data": self._expand_task_data("probe_data", probe_data, n_tasks), + "probe_position_x_px": self._expand_task_data( + "probe_position_x_px", probe_position_x_px, n_tasks + ), + "probe_position_y_px": self._expand_task_data( + "probe_position_y_px", probe_position_y_px, n_tasks + ), + "opr_mode_weights_data": self._expand_task_data( + "opr_mode_weights_data", opr_mode_weights_data, n_tasks, optional=True + ), + "valid_pixel_mask": self._expand_task_data( + "valid_pixel_mask", valid_pixel_mask, n_tasks, optional=True + ), + } + + copied_options = [] + copied_data = [] + for i_task, options in enumerate(task_options): + data = self._resolve_workflow_data( + task_options=options, + **{name: values[i_task] for name, values in values_by_name.items()}, + ) + self._warn_for_gpu_data(data) + copied_data.append(self._copy_workflow_data_to_cpu(data)) + copied_options.append(self._copy_task_options(options)) + + self.workflow_options = workflow_options + self.task_options = copied_options + self.tasks: list[PtychographyTask] = [] + self._task_args = args + self._task_kwargs = kwargs + self._workflow_task_data = copied_data + self._completed = False + + for field_name in values_by_name: + setattr(self, field_name, [getattr(data, field_name) for data in copied_data]) + + self._validate_shared_object_geometry() + total_epochs = workflow_options.num_outer_epochs * workflow_options.num_inner_epochs + for options in self.task_options: + options.reconstructor_options.num_epochs = total_epochs + + @staticmethod + def _expand_task_data( + name: str, + value, + n_tasks: int, + *, + optional: bool = False, + ) -> list: + if value is _UNSET: + return [_UNSET] * n_tasks + if value is None: + if optional: + return [None] * n_tasks + raise ValueError(f"`{name}` is required.") + if not isinstance(value, list): + raise TypeError(f"`{name}` must be a list.") + if len(value) != n_tasks: + raise ValueError( + f"`{name}` must contain one member for each task " + f"({len(value)} != {n_tasks})." + ) + return value + + def _validate_shared_object_geometry(self) -> None: + if any( + options.object_options.determine_position_origin_coords_by + != api.ObjectPosOriginCoordsMethods.SUPPORT + for options in self.task_options + ): + raise ValueError( + "All task options must set " + "`object_options.determine_position_origin_coords_by` to `SUPPORT`." + ) + + reference_shape = tuple(self._workflow_task_data[0].object_data.shape) + reference_options = self.task_options[0].object_options + for i_task, (data, options) in enumerate( + zip(self._workflow_task_data[1:], self.task_options[1:]), start=1 + ): + object_data = data.object_data + if tuple(object_data.shape) != reference_shape: + raise ValueError( + "All members of `object_data` must have the same shape; " + f"task 0 has {reference_shape} and task {i_task} has " + f"{tuple(object_data.shape)}." + ) + object_options = options.object_options + geometry = ( + object_options.pixel_size_m, + object_options.pixel_size_aspect_ratio, + object_options.slice_spacings_m, + ) + reference_geometry = ( + reference_options.pixel_size_m, + reference_options.pixel_size_aspect_ratio, + reference_options.slice_spacings_m, + ) + if geometry != reference_geometry: + raise ValueError( + "All task options must use matching object pixel and slice geometry." + ) + + def run(self) -> None: + if self.tasks: + raise RuntimeError("This multiscan shared-object workflow has already been run.") + + self._completed = False + self._build_tasks() + for _ in range(self.workflow_options.num_outer_epochs): + for i_task, task in enumerate(self.tasks): + task.build_default_device() + task.build_default_dtype() + task.set_large_tensor_device() + try: + task.run(self.workflow_options.num_inner_epochs) + if len(self.tasks) > 1: + next_task = self.tasks[(i_task + 1) % len(self.tasks)] + next_task.copy_data_from_task(task, params_to_copy=("object",)) + finally: + task.set_large_tensor_device("cpu") + + final_task = self.tasks[-1] + for task in self.tasks[:-1]: + task.copy_data_from_task(final_task, params_to_copy=("object",)) + self._completed = True + + def _build_tasks(self) -> None: + for i_task, options in enumerate(self.task_options): + data = self._workflow_task_data[i_task] + task = PtychographyTask( + options, + *self._task_args, + diffraction_data=data.diffraction_data, + object_data=data.object_data, + probe_data=data.probe_data, + probe_position_x_px=data.probe_position_x_px, + probe_position_y_px=data.probe_position_y_px, + opr_mode_weights_data=data.opr_mode_weights_data, + valid_pixel_mask=data.valid_pixel_mask, + **self._task_kwargs, + ) + self.tasks.append(task) + if i_task > 0 and task.reconstructor is not None: + pbar = getattr(task.reconstructor, "pbar", None) + if pbar is not None: + pbar.disable = True + task.set_large_tensor_device("cpu") diff --git a/src/ptychi/workflows/progressive_resolution.py b/src/ptychi/workflows/progressive_resolution.py new file mode 100644 index 00000000..e762e613 --- /dev/null +++ b/src/ptychi/workflows/progressive_resolution.py @@ -0,0 +1,196 @@ +# Copyright © 2025 UChicago Argonne, LLC All right reserved +# Full license accessible at https://github.com//AdvancedPhotonSource/pty-chi/blob/main/LICENSE + +import math +from typing import Literal + +import torch +import torch.nn.functional as F +from torch import Tensor + +import ptychi.image_proc as image_proc +from ptychi.api.options.workflow import ProgressiveResolutionWorkflowOptions +from ptychi.api.task import PtychographyTask +from ptychi.workflows.base import BaseWorkflow + + +class ProgressiveResolutionWorkflow(BaseWorkflow): + """Reconstruct diffraction data from coarse to full spatial resolution.""" + + workflow_options: ProgressiveResolutionWorkflowOptions + + def run(self) -> None: + if self.tasks: + raise RuntimeError("This progressive-resolution workflow has already been run.") + if not isinstance(self.workflow_options, ProgressiveResolutionWorkflowOptions): + raise TypeError( + "`workflow_options` must be a ProgressiveResolutionWorkflowOptions instance." + ) + + self._validate_spatial_shapes() + self._completed = False + for i_level in range(self.workflow_options.num_resolution_levels): + factor = 2 ** (self.workflow_options.num_resolution_levels - 1 - i_level) + level_data = self._build_level_data(i_level=i_level, factor=factor) + level_options = self._copy_task_options() + level_options.reconstructor_options.num_epochs = ( + self.workflow_options.num_epochs_all_levels[i_level] + ) + level_options.object_options.pixel_size_m = ( + self.task_options.object_options.pixel_size_m * factor + ) + + task = PtychographyTask( + level_options, + *self._task_args, + diffraction_data=self._build_level_diffraction_data(factor), + object_data=level_data["object_data"], + probe_data=level_data["probe_data"], + probe_position_x_px=level_data["probe_position_x_px"], + probe_position_y_px=level_data["probe_position_y_px"], + opr_mode_weights_data=level_data["opr_mode_weights_data"], + valid_pixel_mask=self._build_level_valid_pixel_mask(factor), + **self._task_kwargs, + ) + self.tasks.append(task) + try: + task.run() + finally: + task.set_large_tensor_device("cpu") + + self._completed = True + + def get_full_resolution_task(self) -> PtychographyTask: + """Return the completed task for the full-resolution level.""" + if not getattr(self, "_completed", False): + raise RuntimeError( + "The full-resolution task is not available until the workflow completes." + ) + return self.tasks[-1] + + def _validate_spatial_shapes(self) -> None: + expected_ndim = { + "diffraction_data": 3, + "object_data": 3, + "probe_data": 4, + } + for name, ndim in expected_ndim.items(): + if getattr(self, name).ndim != ndim: + raise ValueError(f"`{name}` must be {ndim}D.") + + if self.valid_pixel_mask is not None: + if self.valid_pixel_mask.ndim != 2: + raise ValueError("`valid_pixel_mask` must be a 2D boolean mask.") + if tuple(self.valid_pixel_mask.shape) != tuple(self.diffraction_data.shape[-2:]): + raise ValueError( + "`valid_pixel_mask.shape` must match the diffraction pattern shape." + ) + + def _build_level_data(self, i_level: int, factor: int) -> dict[str, Tensor | None]: + object_target = self._target_spatial_shape(self.object_data, factor) + probe_target = self._target_spatial_shape(self.probe_data, factor) + if i_level == 0: + return { + "object_data": self._resize_spatial_data(self.object_data, object_target), + "probe_data": self._resize_spatial_data(self.probe_data, probe_target), + "probe_position_x_px": self.probe_position_x_px / factor, + "probe_position_y_px": self.probe_position_y_px / factor, + "opr_mode_weights_data": ( + None + if self.opr_mode_weights_data is None + else self.opr_mode_weights_data.clone() + ), + } + + previous_task = self.tasks[-1] + previous_positions = self._get_task_tensor(previous_task, "probe_positions") + previous_factor = factor * 2 + position_scale = previous_factor / factor + return { + "object_data": self._resize_spatial_data( + self._get_task_tensor(previous_task, "object"), object_target + ), + "probe_data": self._resize_spatial_data( + self._get_task_tensor(previous_task, "probe"), probe_target + ), + "probe_position_x_px": previous_positions[:, 1] * position_scale, + "probe_position_y_px": previous_positions[:, 0] * position_scale, + "opr_mode_weights_data": self._get_task_tensor( + previous_task, "opr_mode_weights" + ), + } + + @staticmethod + def _get_task_tensor( + task: PtychographyTask, + name: Literal["object", "probe", "probe_positions", "opr_mode_weights"], + ) -> Tensor: + data = task.get_data_to_cpu(name) + if not isinstance(data, Tensor): + raise TypeError(f"Expected tensor data for `{name}`.") + return data.detach().cpu().clone() + + @staticmethod + def _target_spatial_shape(data: Tensor, factor: int) -> tuple[int, int]: + height, width = data.shape[-2:] + return ( + max(1, math.floor(height / factor + 0.5)), + max(1, math.floor(width / factor + 0.5)), + ) + + @staticmethod + def _resize_spatial_data(data: Tensor, size: tuple[int, int]) -> Tensor: + if tuple(data.shape[-2:]) == size: + return data.detach().cpu().clone() + + leading_shape = data.shape[:-2] + flattened = data.detach().cpu().reshape(-1, 1, *data.shape[-2:]) + if flattened.is_complex(): + resized = F.interpolate( + flattened.real, size=size, mode="bilinear", align_corners=False + ) + 1j * F.interpolate( + flattened.imag, size=size, mode="bilinear", align_corners=False + ) + else: + if not flattened.is_floating_point(): + flattened = flattened.to(torch.get_default_dtype()) + resized = F.interpolate( + flattened, size=size, mode="bilinear", align_corners=False + ) + return resized.reshape(*leading_shape, *size) + + def _crop_reciprocal_data(self, data: Tensor, factor: int) -> Tensor: + target_size = self._target_spatial_shape(data, factor) + if not self.task_options.data_options.fft_shift: + data = torch.fft.ifftshift(data, dim=(-2, -1)) + data = image_proc.central_crop(data, target_size) + data = torch.fft.fftshift(data, dim=(-2, -1)) + else: + data = image_proc.central_crop(data, target_size) + return data.clone() + + def _build_level_diffraction_data(self, factor: int) -> Tensor: + if math.isfinite( + self.task_options.data_options.free_space_propagation_distance_m + ): + target_size = self._target_spatial_shape(self.diffraction_data, factor) + return self._resize_spatial_data(self.diffraction_data, target_size) + return self._crop_reciprocal_data(self.diffraction_data, factor) + + def _build_level_valid_pixel_mask(self, factor: int) -> Tensor | None: + if self.valid_pixel_mask is None: + return None + if not math.isfinite( + self.task_options.data_options.free_space_propagation_distance_m + ): + return self._crop_reciprocal_data(self.valid_pixel_mask, factor) + + target_size = self._target_spatial_shape(self.valid_pixel_mask, factor) + if tuple(self.valid_pixel_mask.shape[-2:]) == target_size: + return self.valid_pixel_mask.clone() + resized = F.interpolate( + self.valid_pixel_mask[None, None].to(torch.float32), + size=target_size, + mode="nearest", + ) + return resized[0, 0].to(torch.bool) diff --git a/tests/test_2d_ptycho_lsqml_multiscan.py b/tests/test_2d_ptycho_lsqml_multiscan.py index 7feef501..37c21e19 100644 --- a/tests/test_2d_ptycho_lsqml_multiscan.py +++ b/tests/test_2d_ptycho_lsqml_multiscan.py @@ -4,7 +4,7 @@ import torch import ptychi.api as api -from ptychi.api.task import PtychographyTask +from ptychi.workflows import MultiscanSharedObjectWorkflow from ptychi.utils import get_suggested_object_size, get_default_complex_dtype, generate_initial_opr_mode_weights import test_utils as tutils @@ -24,7 +24,7 @@ def test_2d_ptycho_lsqml_multiscan(self): positions_px_1 = positions_px[:500] positions_px_2 = positions_px[500:] - # Create task 1 + # Configure scan 1 options_1 = api.LSQMLOptions() diffraction_data_1 = data1 @@ -49,16 +49,7 @@ def test_2d_ptycho_lsqml_multiscan(self): options_1.reconstructor_options.num_epochs = 8 options_1.reconstructor_options.allow_nondeterministic_algorithms = False - task_1 = PtychographyTask( - options_1, - diffraction_data=diffraction_data_1, - object_data=object_data_1, - probe_data=probe_data_1, - probe_position_x_px=probe_position_x_px_1, - probe_position_y_px=probe_position_y_px_1, - ) - - # Create task 2 + # Configure scan 2 options_2 = api.LSQMLOptions() diffraction_data_2 = data2 @@ -83,28 +74,21 @@ def test_2d_ptycho_lsqml_multiscan(self): options_2.reconstructor_options.num_epochs = 8 options_2.reconstructor_options.allow_nondeterministic_algorithms = False - task_2 = PtychographyTask( - options_2, - diffraction_data=diffraction_data_2, - object_data=object_data_2, - probe_data=probe_data_2, - probe_position_x_px=probe_position_x_px_2, - probe_position_y_px=probe_position_y_px_2, + workflow = MultiscanSharedObjectWorkflow( + [options_1, options_2], + diffraction_data=[diffraction_data_1, diffraction_data_2], + object_data=[object_data_1, object_data_2], + probe_data=[probe_data_1, probe_data_2], + probe_position_x_px=[probe_position_x_px_1, probe_position_x_px_2], + probe_position_y_px=[probe_position_y_px_1, probe_position_y_px_2], + workflow_options=api.MultiscanSharedObjectWorkflowOptions( + num_outer_epochs=options_1.reconstructor_options.num_epochs, + num_inner_epochs=1, + ), ) - - # Disable progress bar for task 2 - task_2.reconstructor.pbar.disable = True - - # Run tasks each for one epoch each time - all_tasks = [task_1, task_2] - for i_epoch in range(options_1.reconstructor_options.num_epochs): - for i_task, task in enumerate(all_tasks): - task.run(1) - # Copy object to next task - i_next_task = (i_task + 1) % len(all_tasks) - all_tasks[i_next_task].copy_data_from_task(task, params_to_copy=("object",)) - - recon = all_tasks[-1].get_data_to_cpu('object', as_numpy=True)[0] + workflow.run() + + recon = workflow.tasks[-1].get_data_to_cpu('object', as_numpy=True)[0] return recon diff --git a/tests/test_large_tensor_offload.py b/tests/test_large_tensor_offload.py index b3b24162..2721b988 100644 --- a/tests/test_large_tensor_offload.py +++ b/tests/test_large_tensor_offload.py @@ -28,7 +28,7 @@ def test_set_large_tensor_device_moves_buffers(self): object_data = object_guess options.object_options.pixel_size_m = 1e-6 - options.object_options.optimizable = False + options.object_options.optimizable = True probe_data = probe_guess options.probe_options.optimizable = False @@ -51,10 +51,25 @@ def test_set_large_tensor_device_moves_buffers(self): fm = task.reconstructor.forward_model indices = torch.arange(2, device="cuda", dtype=torch.long) fm.forward(indices) + task.object.preconditioner = torch.ones(task.object.shape, device="cuda") + task.probe.update_buffer = torch.ones(task.probe.shape, device="cuda") + object_optimizer = task.object.optimizer + object_parameter = object_optimizer.param_groups[0]["params"][0] + object_optimizer.state[object_parameter]["test_buffer"] = torch.ones_like( + object_parameter + ) + reconstructor_buffers = task.reconstructor.reconstructor_buffers assert task.dataset.patterns.device.type == "cuda" assert task.reconstructor.parameter_group.object.tensor.data.device.type == "cuda" assert task.reconstructor.parameter_group.probe.tensor.data.device.type == "cuda" + assert task.probe_positions.data.device.type == "cuda" + assert task.opr_mode_weights.data.device.type == "cuda" + assert task.dataset.valid_pixel_mask.device.type == "cuda" + assert task.object.preconditioner.device.type == "cuda" + assert task.probe.update_buffer.device.type == "cuda" + assert object_optimizer.state[object_parameter]["test_buffer"].device.type == "cuda" + assert reconstructor_buffers.alpha_probe_all_pos.device.type == "cuda" assert fm.intermediate_variables.obj_patches.device.type == "cuda" task.set_large_tensor_device("cpu") @@ -63,6 +78,13 @@ def test_set_large_tensor_device_moves_buffers(self): assert not task.dataset.save_data_on_device assert task.reconstructor.parameter_group.object.tensor.data.device.type == "cpu" assert task.reconstructor.parameter_group.probe.tensor.data.device.type == "cpu" + assert task.probe_positions.data.device.type == "cpu" + assert task.opr_mode_weights.data.device.type == "cpu" + assert task.dataset.valid_pixel_mask.device.type == "cpu" + assert task.object.preconditioner.device.type == "cpu" + assert task.probe.update_buffer.device.type == "cpu" + assert object_optimizer.state[object_parameter]["test_buffer"].device.type == "cpu" + assert reconstructor_buffers.alpha_probe_all_pos.device.type == "cpu" assert fm.intermediate_variables.obj_patches.device.type == "cpu" task.set_large_tensor_device() @@ -71,6 +93,13 @@ def test_set_large_tensor_device_moves_buffers(self): assert task.dataset.save_data_on_device assert task.reconstructor.parameter_group.object.tensor.data.device.type == "cuda" assert task.reconstructor.parameter_group.probe.tensor.data.device.type == "cuda" + assert task.probe_positions.data.device.type == "cuda" + assert task.opr_mode_weights.data.device.type == "cuda" + assert task.dataset.valid_pixel_mask.device.type == "cuda" + assert task.object.preconditioner.device.type == "cuda" + assert task.probe.update_buffer.device.type == "cuda" + assert object_optimizer.state[object_parameter]["test_buffer"].device.type == "cuda" + assert reconstructor_buffers.alpha_probe_all_pos.device.type == "cuda" assert fm.intermediate_variables.obj_patches.device.type == "cuda" diff --git a/tests/test_multiscan_shared_object_workflow.py b/tests/test_multiscan_shared_object_workflow.py new file mode 100644 index 00000000..b2d5bc96 --- /dev/null +++ b/tests/test_multiscan_shared_object_workflow.py @@ -0,0 +1,347 @@ +from types import SimpleNamespace + +import pytest +import torch +from pydantic import ValidationError + +import ptychi.api as api +import ptychi.workflows.multiscan_shared_object as multiscan_module +from ptychi.workflows import MultiscanSharedObjectWorkflow + + +def _task_options(n_tasks=3): + options = [] + for _ in range(n_tasks): + task_options = api.EPIEOptions() + task_options.reconstructor_options.default_device = api.Devices.CPU + task_options.object_options.pixel_size_m = 2.5 + options.append(task_options) + return options + + +def _workflow_data(n_tasks=3): + return { + "diffraction_data": [ + torch.full((2, 5, 5), i_task + 1, dtype=torch.float32, device="cpu") + for i_task in range(n_tasks) + ], + "object_data": [ + torch.full( + (1, 11, 11), + i_task * 10 + 1, + dtype=torch.complex64, + device="cpu", + ) + for i_task in range(n_tasks) + ], + "probe_data": [ + torch.full( + (1, 1, 5, 5), + i_task + 1, + dtype=torch.complex64, + device="cpu", + ) + for i_task in range(n_tasks) + ], + "probe_position_x_px": [ + torch.tensor([-1.0, 1.0], device="cpu") for _ in range(n_tasks) + ], + "probe_position_y_px": [ + torch.tensor([-1.0, 1.0], device="cpu") for _ in range(n_tasks) + ], + } + + +def _workflow_options(outer=2, inner=1): + return api.MultiscanSharedObjectWorkflowOptions( + num_outer_epochs=outer, + num_inner_epochs=inner, + ) + + +class _FakeTask: + instances = [] + events = [] + fail_task = None + + def __init__(self, options, *args, **kwargs): + self.options = options + self.input_data = kwargs + self.index = len(self.__class__.instances) + self.__class__.instances.append(self) + self.reconstructor = SimpleNamespace( + pbar=SimpleNamespace(disable=False), + current_epoch=0, + ) + self.results = { + "object": kwargs["object_data"].clone(), + "probe": kwargs["probe_data"].clone(), + "probe_positions": torch.stack( + [kwargs["probe_position_y_px"], kwargs["probe_position_x_px"]], dim=1 + ), + "opr_mode_weights": kwargs["opr_mode_weights_data"], + } + self.devices = [] + + def build_default_device(self): + self.__class__.events.append(("default_device", self.index)) + + def build_default_dtype(self): + self.__class__.events.append(("default_dtype", self.index)) + + def set_large_tensor_device(self, device=None): + self.devices.append(device) + + def run(self, n_epochs): + self.__class__.events.append(("run", self.index, n_epochs)) + if self.index == self.__class__.fail_task: + raise RuntimeError("task failed") + self.results["object"] = self.results["object"] + self.index + 1 + self.reconstructor.current_epoch += n_epochs + + def copy_data_from_task(self, task, params_to_copy): + self.__class__.events.append(("copy", task.index, self.index, params_to_copy)) + assert params_to_copy == ("object",) + self.results["object"] = task.results["object"].clone() + + def get_data_to_cpu(self, name): + value = self.results[name] + return None if value is None else value.detach().cpu() + + +def _reset_fake_task(): + _FakeTask.instances = [] + _FakeTask.events = [] + _FakeTask.fail_task = None + + +def test_multiscan_options_validate_epoch_counts(): + options = _workflow_options(outer=2, inner=3) + assert options.num_outer_epochs == 2 + assert options.num_inner_epochs == 3 + + with pytest.raises(ValidationError, match="num_outer_epochs"): + _workflow_options(outer=0) + with pytest.raises(ValidationError, match="num_inner_epochs"): + _workflow_options(inner=0) + + +def test_multiscan_requires_nonempty_options_list_and_equal_data_lists(): + with pytest.raises(TypeError, match="must be a list"): + MultiscanSharedObjectWorkflow( + tuple(_task_options(1)), + workflow_options=_workflow_options(), + **_workflow_data(1), + ) + with pytest.raises(ValueError, match="at least one"): + MultiscanSharedObjectWorkflow( + [], + workflow_options=_workflow_options(), + **_workflow_data(0), + ) + + single_scan = MultiscanSharedObjectWorkflow( + _task_options(1), + workflow_options=_workflow_options(), + **_workflow_data(1), + ) + assert len(single_scan.task_options) == 1 + + data = _workflow_data(2) + data["probe_data"] = data["probe_data"][:1] + with pytest.raises(ValueError, match="one member for each task"): + MultiscanSharedObjectWorkflow( + _task_options(2), + workflow_options=_workflow_options(), + **data, + ) + + +def test_multiscan_accepts_optional_whole_none_and_mixed_none_lists(): + data = _workflow_data(2) + data["opr_mode_weights_data"] = None + data["valid_pixel_mask"] = [ + None, + torch.ones((5, 5), dtype=torch.bool, device="cpu"), + ] + workflow = MultiscanSharedObjectWorkflow( + _task_options(2), + workflow_options=_workflow_options(), + **data, + ) + + assert workflow.opr_mode_weights_data == [None, None] + assert workflow.valid_pixel_mask[0] is None + assert torch.equal(workflow.valid_pixel_mask[1], data["valid_pixel_mask"][1]) + + +def test_multiscan_validates_shared_object_geometry_and_support_origin(): + options = _task_options(2) + options[1].object_options.determine_position_origin_coords_by = ( + api.ObjectPosOriginCoordsMethods.POSITIONS + ) + with pytest.raises(ValueError, match="SUPPORT"): + MultiscanSharedObjectWorkflow( + options, + workflow_options=_workflow_options(), + **_workflow_data(2), + ) + + data = _workflow_data(2) + data["object_data"][1] = torch.ones( + (1, 13, 11), dtype=torch.complex64, device="cpu" + ) + with pytest.raises(ValueError, match="same shape"): + MultiscanSharedObjectWorkflow( + _task_options(2), + workflow_options=_workflow_options(), + **data, + ) + + options = _task_options(2) + options[1].object_options.pixel_size_m = 3.0 + with pytest.raises(ValueError, match="pixel and slice geometry"): + MultiscanSharedObjectWorkflow( + options, + workflow_options=_workflow_options(), + **_workflow_data(2), + ) + + +def test_multiscan_runs_ring_then_synchronizes_final_object(monkeypatch): + _reset_fake_task() + monkeypatch.setattr(multiscan_module, "PtychographyTask", _FakeTask) + original_options = _task_options(3) + data = _workflow_data(3) + original_probes = [probe.clone() for probe in data["probe_data"]] + workflow = MultiscanSharedObjectWorkflow( + original_options, + workflow_options=_workflow_options(outer=2, inner=4), + **data, + ) + + assert all(options is not original for options, original in zip(workflow.task_options, original_options)) + assert [options.reconstructor_options.num_epochs for options in workflow.task_options] == [ + 8, + 8, + 8, + ] + assert [options.reconstructor_options.num_epochs for options in original_options] == [ + 100, + 100, + 100, + ] + assert workflow.object_data[0].data_ptr() != data["object_data"][0].data_ptr() + + workflow.run() + + run_events = [event for event in _FakeTask.events if event[0] == "run"] + assert run_events == [ + ("run", 0, 4), + ("run", 1, 4), + ("run", 2, 4), + ("run", 0, 4), + ("run", 1, 4), + ("run", 2, 4), + ] + assert [task.reconstructor.current_epoch for task in workflow.tasks] == [8, 8, 8] + assert not workflow.tasks[0].reconstructor.pbar.disable + assert all(task.reconstructor.pbar.disable for task in workflow.tasks[1:]) + assert all(task.devices == ["cpu", None, "cpu", None, "cpu"] for task in workflow.tasks) + + final_object = workflow.tasks[-1].results["object"] + assert torch.all(final_object == 13) + assert all(torch.equal(task.results["object"], final_object) for task in workflow.tasks) + assert all( + torch.equal(task.results["probe"], original_probe) + for task, original_probe in zip(workflow.tasks, original_probes) + ) + + with pytest.raises(RuntimeError, match="already been run"): + workflow.run() + + +def test_multiscan_offloads_active_task_when_run_fails(monkeypatch): + _reset_fake_task() + _FakeTask.fail_task = 1 + monkeypatch.setattr(multiscan_module, "PtychographyTask", _FakeTask) + workflow = MultiscanSharedObjectWorkflow( + _task_options(3), + workflow_options=_workflow_options(), + **_workflow_data(3), + ) + + with pytest.raises(RuntimeError, match="task failed"): + workflow.run() + + assert workflow.tasks[0].devices[-1] == "cpu" + assert workflow.tasks[1].devices[-1] == "cpu" + assert workflow.tasks[2].devices == ["cpu"] + + +def test_multiscan_deprecated_option_data_fallback_warns(): + options = _task_options(2) + data = _workflow_data(2) + for i_task, task_options in enumerate(options): + task_options.data_options.data = data["diffraction_data"][i_task] + task_options.object_options.initial_guess = data["object_data"][i_task] + task_options.probe_options.initial_guess = data["probe_data"][i_task] + task_options.probe_position_options.position_x_px = data[ + "probe_position_x_px" + ][i_task] + task_options.probe_position_options.position_y_px = data[ + "probe_position_y_px" + ][i_task] + + with pytest.warns(DeprecationWarning): + workflow = MultiscanSharedObjectWorkflow( + options, + workflow_options=_workflow_options(), + ) + + assert len(workflow.diffraction_data) == 2 + assert all(task_options.data_options.data is None for task_options in workflow.task_options) + + +def test_multiscan_runs_real_cpu_tasks(): + options = _task_options(2) + for task_options in options: + task_options.reconstructor_options.batch_size = 2 + task_options.reconstructor_options.allow_nondeterministic_algorithms = False + task_options.object_options.step_size = 0.1 + task_options.probe_options.step_size = 0.1 + workflow = MultiscanSharedObjectWorkflow( + options, + workflow_options=_workflow_options(outer=1, inner=1), + **_workflow_data(2), + ) + + workflow.run() + + final_object = workflow.tasks[-1].get_data_to_cpu("object") + for task in workflow.tasks: + assert task.reconstructor.current_epoch == 1 + assert task.dataset.patterns.device.type == "cpu" + torch.testing.assert_close( + task.get_data_to_cpu("object"), final_object, equal_nan=True + ) + for name in ("object", "probe", "probe_positions", "opr_mode_weights"): + assert task.get_data(name).device.type == "cpu" + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA device required") +def test_multiscan_copies_gpu_inputs_without_modifying_callers(): + data = { + name: [value.cuda() for value in values] + for name, values in _workflow_data(2).items() + } + with pytest.warns(UserWarning, match="original tensors remain"): + workflow = MultiscanSharedObjectWorkflow( + _task_options(2), + workflow_options=_workflow_options(), + **data, + ) + + for name, values in data.items(): + assert all(value.device.type == "cuda" for value in values) + assert all(value.device.type == "cpu" for value in getattr(workflow, name)) diff --git a/tests/test_progressive_resolution_workflow.py b/tests/test_progressive_resolution_workflow.py new file mode 100644 index 00000000..4598bd8f --- /dev/null +++ b/tests/test_progressive_resolution_workflow.py @@ -0,0 +1,415 @@ +import pytest +import torch +from pydantic import ValidationError + +import ptychi.api as api +import ptychi.image_proc as image_proc +import ptychi.workflows.progressive_resolution as progressive_resolution_module +from ptychi.workflows import ProgressiveResolutionWorkflow + + +def _task_options(): + options = api.EPIEOptions() + options.data_options.fft_shift = False + options.reconstructor_options.default_device = api.Devices.CPU + options.object_options.pixel_size_m = 2.5 + return options + + +def _workflow_data(): + return { + "diffraction_data": torch.arange( + 2 * 9 * 11, dtype=torch.float32, device="cpu" + ).reshape(2, 9, 11), + "object_data": torch.ones((1, 17, 19), dtype=torch.complex64, device="cpu"), + "probe_data": torch.ones((1, 1, 9, 11), dtype=torch.complex64, device="cpu"), + "probe_position_x_px": torch.tensor([-4.0, 4.0], device="cpu"), + "probe_position_y_px": torch.tensor([-2.0, 2.0], device="cpu"), + "opr_mode_weights_data": torch.ones((2, 1), device="cpu"), + "valid_pixel_mask": torch.ones((9, 11), dtype=torch.bool, device="cpu"), + } + + +class _MovableParameter: + def __init__(self): + self.devices = [] + self.preconditioner = None + self.update_buffer = None + self.optimizer = None + + def to(self, device): + self.devices.append(torch.device(device)) + return self + + +class _ParameterGroup: + def __init__(self): + self.parameters = [_MovableParameter() for _ in range(4)] + + def get_all_reconstruct_parameters(self): + return self.parameters + + +class _Buffers: + def get_all_names(self): + return [] + + +class _Reconstructor: + def __init__(self): + self.parameter_group = _ParameterGroup() + self.reconstructor_buffers = _Buffers() + + +class _Dataset: + def __init__(self, patterns, valid_pixel_mask): + self.patterns = patterns + self.valid_pixel_mask = valid_pixel_mask + self.devices = [] + + def move_attributes_to_device(self, device=None): + device = torch.device(device) + self.devices.append(device) + if self.valid_pixel_mask is not None: + self.valid_pixel_mask = self.valid_pixel_mask.to(device) + + +class _FakeTask: + instances = [] + fail_level = None + + def __init__(self, options, *args, **kwargs): + self.options = options + self.input_data = kwargs + self.level = len(self.__class__.instances) + self.__class__.instances.append(self) + self.dataset = _Dataset( + kwargs["diffraction_data"], kwargs.get("valid_pixel_mask") + ) + self.reconstructor = _Reconstructor() + self.offload_devices = [] + self.results = { + "object": kwargs["object_data"].clone(), + "probe": kwargs["probe_data"].clone(), + "probe_positions": torch.stack( + [kwargs["probe_position_y_px"], kwargs["probe_position_x_px"]], dim=1 + ), + "opr_mode_weights": kwargs["opr_mode_weights_data"].clone(), + } + + def run(self): + if self.level == self.__class__.fail_level: + raise RuntimeError("level failed") + increment = self.level + 1 + self.results = {name: value + increment for name, value in self.results.items()} + + def get_data_to_cpu(self, name): + return self.results[name].detach().cpu() + + def set_large_tensor_device(self, device=None): + device = torch.device(device) + self.offload_devices.append(device) + self.dataset.patterns = self.dataset.patterns.to(device) + self.dataset.move_attributes_to_device(device) + for parameter in self.reconstructor.parameter_group.get_all_reconstruct_parameters(): + parameter.to(device) + + +def _reset_fake_task(): + _FakeTask.instances = [] + _FakeTask.fail_level = None + + +def test_progressive_resolution_options_validate_levels_and_epochs(): + options = api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1, 2] + ) + assert options.num_resolution_levels == 2 + + with pytest.raises(ValidationError, match="greater than 0"): + api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=0, num_epochs_all_levels=[] + ) + with pytest.raises(ValidationError, match="one value for each"): + api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1] + ) + with pytest.raises(ValidationError, match="greater than 0"): + api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1, 0] + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA device required") +def test_workflow_stores_gpu_inputs_on_cpu(): + data = {name: value.cuda() for name, value in _workflow_data().items()} + with pytest.warns(UserWarning, match="original tensors remain"): + workflow = ProgressiveResolutionWorkflow( + _task_options(), + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=1, num_epochs_all_levels=[1] + ), + **data, + ) + + for name in data: + assert getattr(workflow, name).device.type == "cpu" + assert data[name].device.type == "cuda" + + +def test_workflow_runs_rounded_levels_and_transfers_results(monkeypatch): + _reset_fake_task() + monkeypatch.setattr(progressive_resolution_module, "PtychographyTask", _FakeTask) + data = _workflow_data() + original_options = _task_options() + workflow = ProgressiveResolutionWorkflow( + original_options, + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=3, num_epochs_all_levels=[1, 2, 3] + ), + **data, + ) + + assert workflow.diffraction_data.device.type == "cpu" + assert workflow.diffraction_data.data_ptr() != data["diffraction_data"].data_ptr() + assert workflow.task_options is not original_options + with pytest.raises(RuntimeError, match="not available"): + workflow.get_full_resolution_task() + + workflow.run() + + assert workflow.tasks == _FakeTask.instances + assert workflow.get_full_resolution_task() is workflow.tasks[-1] + assert [task.options.reconstructor_options.num_epochs for task in workflow.tasks] == [ + 1, + 2, + 3, + ] + assert [task.options.object_options.pixel_size_m for task in workflow.tasks] == [ + 10.0, + 5.0, + 2.5, + ] + assert original_options.reconstructor_options.num_epochs == 100 + assert original_options.object_options.pixel_size_m == 2.5 + + assert [tuple(task.input_data["diffraction_data"].shape[-2:]) for task in workflow.tasks] == [ + (2, 3), + (5, 6), + (9, 11), + ] + assert [tuple(task.input_data["object_data"].shape[-2:]) for task in workflow.tasks] == [ + (4, 5), + (9, 10), + (17, 19), + ] + assert [tuple(task.input_data["probe_data"].shape[-2:]) for task in workflow.tasks] == [ + (2, 3), + (5, 6), + (9, 11), + ] + assert tuple(workflow.tasks[-1].input_data["valid_pixel_mask"].shape) == (9, 11) + + for i_level in (1, 2): + previous_positions = workflow.tasks[i_level - 1].results["probe_positions"] + assert torch.equal( + workflow.tasks[i_level].input_data["probe_position_y_px"], + previous_positions[:, 0] * 2, + ) + assert torch.equal( + workflow.tasks[i_level].input_data["probe_position_x_px"], + previous_positions[:, 1] * 2, + ) + assert torch.equal( + workflow.tasks[i_level].input_data["opr_mode_weights_data"], + workflow.tasks[i_level - 1].results["opr_mode_weights"], + ) + + for task in workflow.tasks: + assert task.dataset.patterns.device.type == "cpu" + assert task.dataset.devices == [torch.device("cpu")] + assert task.offload_devices == [torch.device("cpu")] + assert all( + parameter.devices == [torch.device("cpu")] + for parameter in task.reconstructor.parameter_group.parameters + ) + + with pytest.raises(RuntimeError, match="already been run"): + workflow.run() + + +def test_reciprocal_crop_respects_fft_shift(): + options = _task_options() + options.data_options.fft_shift = True + workflow = ProgressiveResolutionWorkflow( + options, + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1, 1] + ), + **_workflow_data(), + ) + data = workflow.diffraction_data + expected_size = (5, 6) + + cropped_centered_data = workflow._build_level_diffraction_data(factor=2) + expected = image_proc.central_crop(data, expected_size) + assert torch.equal(cropped_centered_data, expected) + + workflow.task_options.data_options.fft_shift = False + cropped_raw_data = workflow._crop_reciprocal_data(data, factor=2) + expected = torch.fft.fftshift( + image_proc.central_crop( + torch.fft.ifftshift(data, dim=(-2, -1)), expected_size + ), + dim=(-2, -1), + ) + assert torch.equal(cropped_raw_data, expected) + + +def test_near_field_data_and_mask_are_resized_in_real_space(): + options = _task_options() + options.data_options.free_space_propagation_distance_m = 0.1 + data = _workflow_data() + data["valid_pixel_mask"] = ( + torch.arange(9 * 11, device="cpu").reshape(9, 11) % 3 != 0 + ) + workflow = ProgressiveResolutionWorkflow( + options, + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1, 1] + ), + **data, + ) + expected_size = (5, 6) + + resized_data = workflow._build_level_diffraction_data(factor=2) + expected_data = torch.nn.functional.interpolate( + data["diffraction_data"][:, None], + size=expected_size, + mode="bilinear", + align_corners=False, + )[:, 0] + assert torch.equal(resized_data, expected_data) + + resized_mask = workflow._build_level_valid_pixel_mask(factor=2) + expected_mask = torch.nn.functional.interpolate( + data["valid_pixel_mask"][None, None].to(torch.float32), + size=expected_size, + mode="nearest", + )[0, 0].to(torch.bool) + assert resized_mask is not None + assert resized_mask.dtype == torch.bool + assert torch.equal(resized_mask, expected_mask) + + assert torch.equal( + workflow._build_level_diffraction_data(factor=1), + data["diffraction_data"], + ) + assert torch.equal( + workflow._build_level_valid_pixel_mask(factor=1), + data["valid_pixel_mask"], + ) + + +def test_failed_level_is_offloaded(monkeypatch): + _reset_fake_task() + _FakeTask.fail_level = 1 + monkeypatch.setattr(progressive_resolution_module, "PtychographyTask", _FakeTask) + workflow = ProgressiveResolutionWorkflow( + _task_options(), + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=3, num_epochs_all_levels=[1, 1, 1] + ), + **_workflow_data(), + ) + + with pytest.raises(RuntimeError, match="level failed"): + workflow.run() + + assert len(workflow.tasks) == 2 + assert all(task.offload_devices == [torch.device("cpu")] for task in workflow.tasks) + with pytest.raises(RuntimeError, match="not available"): + workflow.get_full_resolution_task() + + +def test_deprecated_option_data_fallback_warns(): + options = _task_options() + data = _workflow_data() + options.data_options.data = data["diffraction_data"] + options.object_options.initial_guess = data["object_data"] + options.probe_options.initial_guess = data["probe_data"] + options.probe_position_options.position_x_px = data["probe_position_x_px"] + options.probe_position_options.position_y_px = data["probe_position_y_px"] + + with pytest.warns(DeprecationWarning): + workflow = ProgressiveResolutionWorkflow( + options, + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=1, num_epochs_all_levels=[1] + ), + ) + + assert workflow.diffraction_data.device.type == "cpu" + + +def test_multilevel_workflow_runs_real_tasks(): + options = _task_options() + options.reconstructor_options.batch_size = 2 + options.reconstructor_options.allow_nondeterministic_algorithms = False + options.object_options.step_size = 0.1 + options.probe_options.step_size = 0.1 + data = { + "diffraction_data": torch.ones((2, 5, 5), device="cpu"), + "object_data": torch.ones((1, 11, 11), dtype=torch.complex64, device="cpu"), + "probe_data": torch.ones((1, 1, 5, 5), dtype=torch.complex64, device="cpu"), + "probe_position_x_px": torch.tensor([-1.0, 1.0], device="cpu"), + "probe_position_y_px": torch.tensor([-1.0, 1.0], device="cpu"), + } + workflow = ProgressiveResolutionWorkflow( + options, + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1, 1] + ), + **data, + ) + + workflow.run() + + task = workflow.get_full_resolution_task() + assert len(workflow.tasks) == 2 + assert workflow.tasks[0].dataset.patterns.shape[-2:] == (3, 3) + assert task.reconstructor.current_epoch == 1 + assert task.dataset.patterns.device.type == "cpu" + assert task.get_data_to_cpu("object").shape == data["object_data"].shape + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA device required") +def test_multilevel_workflow_offloads_real_gpu_tasks(): + options = _task_options() + options.data_options.save_data_on_device = True + options.reconstructor_options.default_device = api.Devices.GPU + options.reconstructor_options.batch_size = 2 + options.object_options.step_size = 0.1 + options.probe_options.step_size = 0.1 + data = { + "diffraction_data": torch.ones((2, 5, 5), device="cpu"), + "object_data": torch.ones((1, 11, 11), dtype=torch.complex64, device="cpu"), + "probe_data": torch.ones((1, 1, 5, 5), dtype=torch.complex64, device="cpu"), + "probe_position_x_px": torch.tensor([-1.0, 1.0], device="cpu"), + "probe_position_y_px": torch.tensor([-1.0, 1.0], device="cpu"), + } + workflow = ProgressiveResolutionWorkflow( + options, + workflow_options=api.ProgressiveResolutionWorkflowOptions( + num_resolution_levels=2, num_epochs_all_levels=[1, 1] + ), + **data, + ) + + workflow.run() + + for task in workflow.tasks: + assert task.dataset.patterns.device.type == "cpu" + assert task.dataset.valid_pixel_mask.device.type == "cpu" + for name in ("object", "probe", "probe_positions", "opr_mode_weights"): + assert task.get_data(name).device.type == "cpu"