From a990ec647896f1bd1431c688948525a6531339f3 Mon Sep 17 00:00:00 2001 From: Chang Liu Date: Wed, 30 Sep 2026 17:19:27 +0200 Subject: [PATCH 1/2] Remove pcPIE.py but keep the same interface, functions already included in mPIE --- PtyLab/Engines/__init__.py | 14 +- PtyLab/Engines/mPIE.py | 392 ++++++++++++++++++++-- tests/Engines/test_remove_pcPIE_engine.py | 43 +++ 3 files changed, 418 insertions(+), 31 deletions(-) create mode 100644 tests/Engines/test_remove_pcPIE_engine.py diff --git a/PtyLab/Engines/__init__.py b/PtyLab/Engines/__init__.py index 2f87ce9..17ed5a0 100644 --- a/PtyLab/Engines/__init__.py +++ b/PtyLab/Engines/__init__.py @@ -14,16 +14,24 @@ "mqNewton", "multiPIE", "OPR", - "pcPIE", "qNewton", "zPIE", ] from .e3PIE import e3PIE from .ePIE import ePIE -from .mPIE import mPIE +from .mPIE import mPIE, pcPIE from .mqNewton import mqNewton from .multiPIE import multiPIE from .OPR import OPR -from .pcPIE import pcPIE from .qNewton import qNewton from .zPIE import zPIE + + +import warnings + +warnings.warn( + "`pcPIE` is deprecated. Use `mPIE` with " + "`params.positionCorrectionSwitch = True` instead.", + DeprecationWarning, + stacklevel=2, +) \ No newline at end of file diff --git a/PtyLab/Engines/mPIE.py b/PtyLab/Engines/mPIE.py index 1238d90..a2a93c0 100644 --- a/PtyLab/Engines/mPIE.py +++ b/PtyLab/Engines/mPIE.py @@ -26,6 +26,77 @@ class mPIE(BaseEngine): + r""" + Momentum-accelerated ptychographic iterative engine (mPIE). + + mPIE extends the standard ePIE reconstruction by combining + regularized PIE (rPIE) object/probe updates with momentum + acceleration.[^maiden2017] + + As in ePIE, the exit-wave correction at scan position $j$ is + + $$ + \Delta\Psi_j = \Psi'_j - \Psi_j + $$ + + For conventional ptychography, the object update is regularized as + + $$ + O'_j = O_j + \beta_O \frac{P^*}{\alpha_O P_{\max} + (1-\alpha_O)|P|^2}\Delta\Psi_j + $$ + + where $P_{\max}=\max |P|^2$, $\beta_O$ controls the object update + step size, and $\alpha_O$ controls the spatial regularization. + + The probe is updated analogously: + + $$ + P' = P + \beta_P \frac{O_j^*}{\alpha_P O_{\max} + (1-\alpha_P)|O_j|^2}\Delta\Psi_j + $$ + + where $O_{\max}=\max |O_j|^2$, and $\alpha_P$ and $\beta_P$ control + probe regularization and update strength. + + Compared with ePIE, these denominators retain stronger updates in + moderately illuminated regions while suppressing unstable updates where + the corresponding probe or object intensity is small. + + Momentum acceleration is periodically applied to both object and probe. + In the PtyLab implementation, the momentum state is updated as + + $$ + M^{(n)} = G^{(n)} + \eta M^{(n-1)} + $$ + + followed by + + $$ + X^{(n+1)} = X^{(n)} - \gamma M^{(n)} + $$ + + where $X$ denotes the object or probe, $\eta$ is `frictionM`, and + $\gamma$ is `feedbackM`. The current implementation applies these + momentum updates stochastically during the scan-position loop. + + The default mPIE parameters are `betaObject = 0.25`, + `betaProbe = 0.25`, `alphaObject = 0.1`, `alphaProbe = 0.1`, + `feedbackM = 0.3`, and `frictionM = 0.7`. + + [^maiden2017]: A. M. Maiden, D. Johnson, and P. Li, + "Further improvements to the ptychographical iterative engine," + Optica 4, 736-745 (2017). + https://doi.org/10.1364/OPTICA.4.000736 + + Attributes: + keepPatches (bool): + If enabled, store the reconstructed object patch associated with each + scan position for debugging or detailed analysis. This option can + require substantial additional memory. + See Also: + `ePIE` + Baseline ePIE reconstruction without rPIE regularization or + momentum acceleration. + """ def __init__( self, reconstruction: Reconstruction, @@ -33,11 +104,32 @@ def __init__( params: Params, monitor: Monitor, ): - # This contains reconstruction parameters that are specific to the reconstruction - # but not necessarily to ePIE reconstruction + """ + Initialize the mPIE reconstruction engine. + + Shared reconstruction state is initialized through `BaseEngine`, followed + by the mPIE-specific reconstruction parameters and momentum buffers. + + Momentum acceleration is enabled through + `params.momentumAcceleration`, allowing shared BaseEngine operations such + as modal orthogonalization to keep the corresponding momentum and buffer + arrays consistent with the reconstructed object and probe. + + Args: + reconstruction (Reconstruction): + Reconstruction state containing the current object, probe, and + geometry. + experimentalData (ExperimentalData): + Experimental diffraction data and acquisition parameters. + params (Params): + Shared reconstruction parameters and constraint settings. + monitor (Monitor): + Monitor used for reconstruction visualization and progress + reporting. + """ super().__init__(reconstruction, experimentalData, params, monitor) self.logger = logging.getLogger("mPIE") - self.logger.info("Sucesfully created mPIE mPIE_engine") + self.logger.info("Successfully created mPIE engine") self.logger.info("Wavelength attribute: %s", self.reconstruction.wavelength) # initialize mPIE Params self.initializeReconstructionParams() @@ -46,10 +138,12 @@ def __init__( @property def keepPatches(self): - """Wether or not to keep track of the individual object update patches. + """ + Whether to store the reconstructed object patch for every scan position. - This strongly increases the amount of memory required, only use when absolutely required. + Enable with `engine.keepPatches = True` and disable with `engine.keepPatches = False`. + This option is intended for debugging or detailed analysis and may require a large amount of additional memory. """ return hasattr(self, "patches") @@ -72,8 +166,20 @@ def keepPatches(self, keep_them): def initializeReconstructionParams(self): """ - Set parameters that are specific to the mPIE settings. - :return: + Initialize mPIE-specific reconstruction parameters and momentum state. + + The default mPIE parameters are: + + - `betaObject = 0.25`: object update step size. + - `betaProbe = 0.25`: probe update step size. + - `alphaObject = 0.1`: object-update regularization parameter. + - `alphaProbe = 0.1`: probe-update regularization parameter. + - `feedbackM = 0.3`: momentum feedback strength. + - `frictionM = 0.7`: momentum memory coefficient. + - `numIterations = 50`: number of reconstruction iterations. + + Object and probe momentum arrays are initialized together with corresponding + buffers that store the reconstruction state used by the momentum updates. """ # self.eswUpdate = self.reconstruction.esw.copy() self.betaProbe = 0.25 @@ -94,8 +200,70 @@ def initializeReconstructionParams(self): self.reconstruction.probeWindow = np.abs(self.reconstruction.probe) def reconstruct(self, experimentalData=None, reconstruction=None, vis_after_each_iteration=None): - """Reconstruct object. If experimentalData is given, it replaces the current data. Idem for reconstruction.""" - + r""" + Run the mPIE reconstruction to completion. + + The reconstruction follows the standard ptychographic position loop. + At each scan position $j$, the current object patch and probe form the + exit surface wave + + $$ + \Psi_j = O_j P + $$ + + which is propagated to the detector plane and constrained by the measured + diffraction intensity through `intensityProjection()`. The corresponding + exit-wave correction is + + $$ + \Delta\Psi_j = \Psi'_j - \Psi_j + $$ + + The object and probe are then updated using the mPIE/rPIE update rules + implemented by `objectPatchUpdate()` and `probeUpdate()`. + + If `params.objectTVregSwitch` is enabled, the TV-regularized object update + `objectPatchUpdate_TV()` is used every `params.objectTVfreq` iterations. + The native mPIE object update is retained and an additional TV + regularization term is added with strength controlled by + `params.objectTVregStepSize`. + + If `params.weigh_probe_updates_by_intensity` is enabled, the probe update + is scaled by the relative intensity of the current diffraction frame. + + If `params.positionCorrectionSwitch` is enabled, scan-position correction + is applied after the object and probe updates. + + Momentum acceleration is applied stochastically during the scan-position + loop. In the current implementation, each position update has + approximately a 5% probability of triggering `objectMomentumUpdate()` and + `probeMomentumUpdate()`. + + If `keepPatches` is enabled, the reconstructed patch associated with each + scan position is additionally stored in `self.patches` for diagnostics or + further analysis. + + After all scan positions in an iteration have been processed, + `getErrorMetrics()` evaluates the reconstruction error, + `applyConstraints()` applies the enabled reconstruction constraints, and + `showReconstruction()` updates the reconstruction monitor. + + Args: + experimentalData (ExperimentalData, optional): + Experimental dataset to use for the reconstruction. If provided, + it replaces the currently attached experimental data. + reconstruction (Reconstruction, optional): + Reconstruction state to optimize. If provided, it replaces the + currently attached reconstruction object. + vis_after_each_iteration (callable, optional): + Callback executed after each reconstruction iteration as + `vis_after_each_iteration(loop, reconstruction)`. + + Notes: + The object and probe momentum buffers are reset after + `_prepareReconstruction()` so that they remain synchronized with any + initialization changes applied before reconstruction starts. + """ self.changeExperimentalData(experimentalData) self.changeOptimizable(reconstruction) @@ -144,12 +312,11 @@ def reconstruct(self, experimentalData=None, reconstruction=None, vis_after_each else: object_patch = self.objectPatchUpdate(objectPatch, DELTA) + self.reconstruction.object[..., sy, sx] = object_patch if self.keepPatches: self.patches[positionIndex, ..., sy, sx] = asNumpyArray( abs(object_patch) ** 2 ) - else: - self.reconstruction.object[..., sy, sx] = object_patch # probe update weight = 1 @@ -161,7 +328,7 @@ def reconstruct(self, experimentalData=None, reconstruction=None, vis_after_each # self.reconstruction.push_probe_update(self.reconstruction.probe, positionIndex, self.experimentalData.ptychogram.shape[0]) if self.params.positionCorrectionSwitch: - shifter = self.positionCorrection( + self.positionCorrection( objectPatch, positionIndex, sy, sx ) # self.pbar_pos.write(f'Corr: {shifter[0]*1e6:.2f} um x {shifter[1]*1e6:.2f} um') @@ -194,9 +361,38 @@ def reconstruct(self, experimentalData=None, reconstruction=None, vis_after_each # todo clearMemory implementation def objectMomentumUpdate(self): - """ - momentum update object, save updated objectMomentum and objectBuffer. - :return: + r""" + Apply the mPIE momentum update to the reconstructed object. + + The change in the object since the previous momentum update is estimated + from the stored object buffer: + + $$ + G_O^{(n)} = O_{\mathrm{buf}}^{(n)} - O^{(n)} + $$ + + The object momentum is updated according to + + $$ + M_O^{(n)} = G_O^{(n)} + \eta M_O^{(n-1)} + $$ + + where $\eta$ is `frictionM`. + + The accumulated momentum is then fed back into the object estimate: + + $$ + O^{(n+1)} = O^{(n)} - \gamma M_O^{(n)} + $$ + + where $\gamma$ is `feedbackM`. + + After the momentum correction, `objectBuffer` is updated with the current + object estimate for the next momentum step. + + Notes: + This update is triggered stochastically from `reconstruct()` rather + than after every scan-position update. """ gradient = self.reconstruction.objectBuffer - self.reconstruction.object self.reconstruction.objectMomentum = ( @@ -209,9 +405,12 @@ def objectMomentumUpdate(self): self.reconstruction.objectBuffer = self.reconstruction.object.copy() def probeMomentumUpdate(self): - """ - momentum update probe, save updated probeMomentum and probeBuffer. - :return: + r""" + Apply the mPIE momentum update to the reconstructed probe. Similar to `objectMomentumUpdate()`. + + See Also: + `objectMomentumUpdate` + Equivalent momentum update applied to the reconstructed object. """ gradient = self.reconstruction.probeBuffer - self.reconstruction.probe self.reconstruction.probeMomentum = ( @@ -224,11 +423,56 @@ def probeMomentumUpdate(self): self.reconstruction.probeBuffer = self.reconstruction.probe.copy() def objectPatchUpdate(self, objectPatch: np.ndarray, DELTA: np.ndarray): - """ - Todo add docstring - :param objectPatch: - :param DELTA: - :return: + r""" + Update the object patch using the regularized mPIE/rPIE object-update rule. + + The probe intensity is first evaluated and its maximum value is used as a + global normalization scale: + + $$ + P_{\max} = \max_{x,y}\sum |P(x,y)|^2 + $$ + + For conventional ptychography, the probe weighting is + + $$ + W_P =\frac{P^*}{\alpha_O P_{\max} + (1-\alpha_O)|P|^2} + $$ + + and the object patch is updated according to + + $$ + O'_j =O_j + \beta_O \sum W_P\Delta\Psi_j + $$ + + where $\Delta\Psi_j$ is the exit-wave correction, $\beta_O$ is + `betaObject`, and $\alpha_O$ is `alphaObject`. + + The parameter `alphaObject` controls the balance between global + normalization by the maximum probe intensity and local normalization by + the spatially varying probe intensity. + + For Fourier ptychography (`operationMode == "FPM"`), an additional + probe-amplitude weighting is applied: + + $$ + W_P^{\mathrm{FPM}} =\frac{|P|}{P_{\max}}\frac{P^*}{\alpha_O P_{\max} + (1-\alpha_O)|P|^2} + $$ + + In the multidimensional PtyLab representation, the object correction is + summed over the probe-mode axis before being added to the current object + patch. + + Args: + objectPatch (ndarray): + Current object patch at the active scan position. + DELTA (ndarray): + Exit-wave correction + `reconstruction.eswUpdate - reconstruction.esw`. + + Returns: + ndarray: + Updated object patch. """ # find out which array module to use, numpy or cupy (or other...) xp = getArrayModule(objectPatch) @@ -251,11 +495,58 @@ def objectPatchUpdate(self, objectPatch: np.ndarray, DELTA: np.ndarray): ) def probeUpdate(self, objectPatch: np.ndarray, DELTA: np.ndarray, weight: float): - """ - Todo add docstring - :param objectPatch: - :param DELTA: - :return: + r""" + Update the probe using the regularized mPIE/rPIE probe-update rule. + + The object intensity is first evaluated and its maximum value is used as + a global normalization scale: + + $$ + O_{\max} = \max_{x,y}\sum |O_j(x,y)|^2 + $$ + + The object weighting is then calculated as + + $$ + W_O = \frac{O_j^*}{\alpha_P O_{\max} + (1-\alpha_P)|O_j|^2} + $$ + + and the probe is updated according to + + $$ + P' = P + w\beta_P\sum W_O\Delta\Psi_j + $$ + + where $\Delta\Psi_j$ is the exit-wave correction, $\beta_P$ is + `betaProbe`, $\alpha_P$ is `alphaProbe`, and $w$ is the optional + intensity-dependent update weight. + + The parameter `alphaProbe` controls the balance between global + normalization by the maximum object intensity and local normalization by + the spatially varying object intensity. + + By default, $w=1$. If `params.weigh_probe_updates_by_intensity` is + enabled in `reconstruct()`, $w$ is set to the relative intensity of the + current diffraction frame. + + In the current multidimensional PtyLab representation, the probe + correction is summed over axis `1` before being added to the current + probe estimate. + + Args: + objectPatch (ndarray): + Current object patch at the active scan position. + DELTA (ndarray): + Exit-wave correction + `reconstruction.eswUpdate - reconstruction.esw`. + weight (float): + Multiplicative weight applied to the probe update. Typically `1`, + or the relative intensity of the current diffraction frame when + intensity-weighted probe updates are enabled. + + Returns: + ndarray: + Updated probe. """ # find out which array module to use, numpy or cupy (or other...) xp = getArrayModule(objectPatch) @@ -268,3 +559,48 @@ def probeUpdate(self, objectPatch: np.ndarray, DELTA: np.ndarray, weight: float) frac * DELTA, axis=1, keepdims=True ) return r + + +class pcPIE(mPIE): + """ + Backward-compatible wrapper for :class:`mPIE`. + + Position correction is now provided by `mPIE` through + `params.positionCorrectionSwitch`. This class is retained only for + compatibility with existing code using `Engines.pcPIE`. + """ + + def __init__( + self, + reconstruction: Reconstruction, + experimentalData: ExperimentalData, + params: Params, + monitor: Monitor, + ): + super().__init__( + reconstruction, + experimentalData, + params, + monitor, + ) + + self.name = "pcPIE" + self.logger = logging.getLogger("pcPIE") + + @property + def betaM(self): + """Deprecated alias for `feedbackM`.""" + return self.feedbackM + + @betaM.setter + def betaM(self, value): + self.feedbackM = value + + @property + def stepM(self): + """Deprecated alias for `frictionM`.""" + return self.frictionM + + @stepM.setter + def stepM(self, value): + self.frictionM = value \ No newline at end of file diff --git a/tests/Engines/test_remove_pcPIE_engine.py b/tests/Engines/test_remove_pcPIE_engine.py new file mode 100644 index 0000000..c60e5a8 --- /dev/null +++ b/tests/Engines/test_remove_pcPIE_engine.py @@ -0,0 +1,43 @@ +import pytest + +from PtyLab import Engines +from PtyLab.Engines.mPIE import mPIE, pcPIE + + +@pytest.fixture +def engine(): + """Create a wrapper instance without initializing reconstruction data.""" + return object.__new__(pcPIE) + + +def test_pcpie_is_exported(): + """Engines.pcPIE should remain available for backward compatibility.""" + assert hasattr(Engines, "pcPIE") + assert Engines.pcPIE is pcPIE + + +def test_pcpie_inherits_mpie(): + """pcPIE should reuse the mPIE implementation.""" + assert issubclass(pcPIE, mPIE) + + +@pytest.mark.parametrize( + "legacy_name,current_name,initial_value,updated_value", + [ + ("betaM", "feedbackM", 0.3, 0.5), + ("stepM", "frictionM", 0.7, 0.8), + ], + ids=["betaM-feedbackM", "stepM-frictionM"], +) +def test_legacy_momentum_parameter_aliases( + engine, legacy_name, current_name, initial_value, updated_value +): + """Legacy pcPIE momentum names should map to the mPIE parameters.""" + setattr(engine, current_name, initial_value) + + # Old pcPIE names should read from the new mPIE attributes + assert getattr(engine, legacy_name) == initial_value + + # Setting the old names should update the new attributes + setattr(engine, legacy_name, updated_value) + assert getattr(engine, current_name) == updated_value From 81d1127092c63351fe0c5cd7c9f7a950371795c2 Mon Sep 17 00:00:00 2001 From: Chang Liu Date: Wed, 30 Sep 2026 17:22:24 +0200 Subject: [PATCH 2/2] delete pcPIE.py --- PtyLab/Engines/pcPIE.py | 202 ---------------------------------------- 1 file changed, 202 deletions(-) delete mode 100644 PtyLab/Engines/pcPIE.py diff --git a/PtyLab/Engines/pcPIE.py b/PtyLab/Engines/pcPIE.py deleted file mode 100644 index 0802a0e..0000000 --- a/PtyLab/Engines/pcPIE.py +++ /dev/null @@ -1,202 +0,0 @@ -import numpy as np -from matplotlib import pyplot as plt - -from PtyLab.utils import gpuUtils - -try: - import cupy as cp -except ImportError: - # print("Cupy not available, will not be able to run GPU based computation") - # Still define the name, we'll take care of it later but in this way it's still possible - # to see that gPIE exists for example. - cp = None - -import logging -import sys - -import tqdm - -from PtyLab.Engines.BaseEngine import BaseEngine -from PtyLab.ExperimentalData.ExperimentalData import ExperimentalData -from PtyLab.Monitor.Monitor import Monitor -from PtyLab.Params.Params import Params - -# PtyLab imports -from PtyLab.Reconstruction.Reconstruction import Reconstruction -from PtyLab.utils.gpuUtils import asNumpyArray, getArrayModule -from PtyLab.utils.utils import fft2c, ifft2c - - -class pcPIE(BaseEngine): - def __init__( - self, - reconstruction: Reconstruction, - experimentalData: ExperimentalData, - params: Params, - monitor: Monitor, - ): - # This contains reconstruction parameters that are specific to the reconstruction - # but not necessarily to ePIE reconstruction - super().__init__(reconstruction, experimentalData, params, monitor) - self.logger = logging.getLogger("pcPIE") - self.logger.info("Successfully created pcPIE pcPIE_engine") - self.logger.info("Wavelength attribute: %s", self.reconstruction.wavelength) - # initialize pcPIE Params - self.initializeReconstructionParams() - # initialize momentum - self.reconstruction.initializeObjectMomentum() - self.reconstruction.initializeProbeMomentum() - # set object and probe buffers - self.reconstruction.objectBuffer = self.reconstruction.object.copy() - self.reconstruction.probeBuffer = self.reconstruction.probe.copy() - - self.params.momentumAcceleration = True - - def initializeReconstructionParams(self): - """ - Set parameters that are specific to the pcPIE settings. - :return: - """ - # these are same as mPIE - # self.eswUpdate = self.reconstruction.esw.copy() - self.betaProbe = 0.25 - self.betaObject = 0.25 - self.alphaProbe = 0.1 # probe regularization - self.alphaObject = 0.1 # object regularization - self.betaM = 0.3 # feedback - self.stepM = 0.7 # friction - # self.probeWindow = np.abs(self.reconstruction.probe) - self.numIterations = 50 - - def reconstruct(self): - self._prepareReconstruction() - - # actual reconstruction ePIE_engine - - self.pbar = tqdm.trange( - self.numIterations, desc="pcPIE", file=sys.stdout, leave=True - ) # in order to change description to the tqdm progress bar - for loop in self.pbar: - # set position order - self.setPositionOrder() - - for positionLoop, positionIndex in enumerate(self.positionIndices): - # get object patch - row, col = self.reconstruction.positions[positionIndex] - sy = slice(row, row + self.reconstruction.Np) - sx = slice(col, col + self.reconstruction.Np) - # note that object patch has size of probe array - objectPatch = self.reconstruction.object[..., sy, sx].copy() - - # make exit surface wave - self.reconstruction.esw = objectPatch * self.reconstruction.probe - - # propagate to camera, intensityProjection, propagate back to object - self.intensityProjection(positionIndex) - - # difference term - DELTA = self.reconstruction.eswUpdate - self.reconstruction.esw - - # object update - self.reconstruction.object[..., sy, sx] = self.objectPatchUpdate( - objectPatch, DELTA - ) - - # probe update - self.reconstruction.probe = self.probeUpdate(objectPatch, DELTA) - if self.params.positionCorrectionSwitch: - self.positionCorrection(objectPatch, positionIndex, sy, sx) - - # momentum updates - if np.random.rand(1) > 0.95: - self.objectMomentumUpdate() - self.probeMomentumUpdate() - - # get error metric - self.getErrorMetrics() - - # apply Constraints - self.applyConstraints(loop) - - # show reconstruction - self.showReconstruction(loop) - - # todo clearMemory implementation - - if self.params.gpuFlag: - self.logger.info("switch to cpu") - self._move_data_to_cpu() - self.params.gpuFlag = 0 - - def objectMomentumUpdate(self): - """ - momentum update object, save updated objectMomentum and objectBuffer. - :return: - """ - gradient = self.reconstruction.objectBuffer - self.reconstruction.object - self.reconstruction.objectMomentum = ( - gradient + self.stepM * self.reconstruction.objectMomentum - ) - self.reconstruction.object = ( - self.reconstruction.object - self.betaM * self.reconstruction.objectMomentum - ) - self.reconstruction.objectBuffer = self.reconstruction.object.copy() - - def probeMomentumUpdate(self): - """ - momentum update probe, save updated probeMomentum and probeBuffer. - :return: - """ - gradient = self.reconstruction.probeBuffer - self.reconstruction.probe - self.reconstruction.probeMomentum = ( - gradient + self.stepM * self.reconstruction.probeMomentum - ) - self.reconstruction.probe = ( - self.reconstruction.probe - self.betaM * self.reconstruction.probeMomentum - ) - self.reconstruction.probeBuffer = self.reconstruction.probe.copy() - - def objectPatchUpdate(self, objectPatch: np.ndarray, DELTA: np.ndarray): - """ - Todo add docstring - :param objectPatch: - :param DELTA: - :return: - """ - # find out which array module to use, numpy or cupy (or other...) - xp = getArrayModule(objectPatch) - absP2 = xp.abs(self.reconstruction.probe) ** 2 - Pmax = xp.max(xp.sum(absP2, axis=(0, 1, 2, 3)), axis=(-1, -2)) - if self.experimentalData.operationMode == "FPM": - frac = ( - abs(self.reconstruction.probe) - / Pmax - * self.reconstruction.probe.conj() - / (self.alphaObject * Pmax + (1 - self.alphaObject) * absP2) - ) - else: - frac = self.reconstruction.probe.conj() / ( - self.alphaObject * Pmax + (1 - self.alphaObject) * absP2 - ) - return objectPatch + self.betaObject * xp.sum( - frac * DELTA, axis=2, keepdims=True - ) - - def probeUpdate(self, objectPatch: np.ndarray, DELTA: np.ndarray): - """ - Todo add docstring - :param objectPatch: - :param DELTA: - :return: - """ - # find out which array module to use, numpy or cupy (or other...) - xp = getArrayModule(objectPatch) - absO2 = xp.abs(objectPatch) ** 2 - Omax = xp.max(xp.sum(absO2, axis=(0, 1, 2, 3)), axis=(-1, -2)) - frac = objectPatch.conj() / ( - self.alphaProbe * Omax + (1 - self.alphaProbe) * absO2 - ) - r = self.reconstruction.probe + self.betaProbe * xp.sum( - frac * DELTA, axis=1, keepdims=True - ) - return r