From edc94ec09bf5bb559dba6dcb8225db6dfc2a8885 Mon Sep 17 00:00:00 2001 From: Abe Levitan Date: Thu, 18 Apr 2024 09:28:03 +0200 Subject: [PATCH] Introducing new multislice model --- src/cdtools/models/__init__.py | 1 + src/cdtools/models/fancy_ptycho.py | 97 +-- src/cdtools/models/multislice_ptycho.py | 916 ++++++++++++++++++++++++ src/cdtools/tools/analysis/analysis.py | 57 ++ 4 files changed, 1030 insertions(+), 41 deletions(-) create mode 100644 src/cdtools/models/multislice_ptycho.py diff --git a/src/cdtools/models/__init__.py b/src/cdtools/models/__init__.py index 0bc6689..9a10e91 100644 --- a/src/cdtools/models/__init__.py +++ b/src/cdtools/models/__init__.py @@ -31,6 +31,7 @@ from cdtools.models.polarized_fancy_ptycho import PolarizedFancyPtycho from cdtools.models.polarization_swept_ptycho import PolarizationSweptPtycho from cdtools.models.bragg_2d_ptycho import Bragg2DPtycho from cdtools.models.multislice_2d_ptycho import Multislice2DPtycho +from cdtools.models.multislice_ptycho import MultislicePtycho from cdtools.models.rpi import RPI from cdtools.models.multimode_rpi import MultimodeRPI from cdtools.models.time_resolved_ptycho_calibration import TimeResolvedPtychoCalibration diff --git a/src/cdtools/models/fancy_ptycho.py b/src/cdtools/models/fancy_ptycho.py index 6f3fbd9..252eda8 100644 --- a/src/cdtools/models/fancy_ptycho.py +++ b/src/cdtools/models/fancy_ptycho.py @@ -296,12 +296,12 @@ class FancyPtycho(CDIModel): if n_obj_modes != 1: obj = t.stack([obj,] + [0.05*t.ones_like(obj),]*(n_obj_modes-1)) + pfc = (probe_fourier_crop if probe_fourier_crop else 0) if obj_view_crop is None: - obj_view_crop = (min(probe.shape[-2], probe.shape[-1]) // 2 - + probe_fourier_crop) + obj_view_crop = min(probe.shape[-2], probe.shape[-1]) // 2 + pfc if obj_view_crop < 0: - obj_view_crop += (min(probe.shape[-2], probe.shape[-1]) // 2 - + probe_fourier_crop) + obj_view_crop += min(probe.shape[-2], probe.shape[-1]) // 2 + pfc + obj_view_crop += obj_padding det_geo = dataset.detector_geometry @@ -436,7 +436,8 @@ class FancyPtycho(CDIModel): prs = tools.propagators.far_field(prs) prs = t.nn.functional.pad(prs, padding) prs = tools.propagators.inverse_far_field(prs) - + + # Now we actually do the interaction, using the sinc subpixel # translation model as per usual exit_waves = self.probe_norm * tools.interactions.ptycho_2D_sinc( @@ -556,14 +557,14 @@ class FancyPtycho(CDIModel): self.probe.data = centered_probe.to(device=self.probe.data.device) - def tidy_probes(self, normalization=1, normalize=False): + def tidy_probes(self, normalize=False, tidy_each_frame=False): """Tidies up the probes What we want to do here is use all the information on all the probes to calculate a natural basis for the experiment, and update all the density matrices to operate in that updated basis """ - + # First we treat the purely incoherent case # I don't love this pattern of using an if statement with a return @@ -579,50 +580,64 @@ class FancyPtycho(CDIModel): return # This is for the unified mode case + + # We concatenate all the weight matrices + all_weights = t.cat(t.unbind(self.weights.detach().cpu(), dim=0), dim=0) - # Note to future: We could probably do this more cleanly with an - # SVD directly on the Ws matrix, instead of an eigendecomposition - # of the rho matrix. + # We use that to calculate the density matrix of the full experiment, + # normalized by number of exposures + overall_rho = (t.mm(all_weights.transpose(0,1), all_weights.conj()) + / self.weights.shape[0]) - rhos = self.get_rhos() - overall_rho = np.mean(rhos, axis=0) + # We generate the orthogonal probes based on this full-experiment + # density matrix. We also keep the transform matrix A probe = self.probe.detach().cpu().numpy() ortho_probes, A = analysis.orthogonalize_probes( probe, density_matrix=overall_rho, keep_transform=True, normalize=normalize) - Aconj = A.conj() - Atrans = np.transpose(A) - new_rhos = np.matmul(Atrans, np.matmul(rhos, Aconj)) - new_rhos /= normalization - ortho_probes *= np.sqrt(normalization) - - dm_rank = self.weights.shape[1] - - new_Ws = [] - for rho in new_rhos: - # These are returned from smallest to largest - we want to keep - # the largest ones - w, v = sla.eigh(rho) - w = w[::-1][:dm_rank] - v = v[:, ::-1][:, :dm_rank] - # For situations where the rank of the density matrix is not - # full in reality, but we keep more modes around than needed, - # some ws can go negative due to numerical error! This is - # extremely rare, but comon enough to cause crashes occasionally - # when there are thousands of individual matrices to transform - # every time this is called. - w = np.maximum(w, 0) - - new_Ws.append(np.dot(np.diag(np.sqrt(w)), v.transpose())) - - new_Ws = np.array(new_Ws) - - self.weights.data = t.as_tensor( - new_Ws, dtype=self.weights.dtype, device=self.weights.device) + # We apply A to the weight matrices to update them along with the + # probes + new_weights = t.matmul( + t.as_tensor(A).transpose(0,1), + self.weights.detach().cpu().transpose(-2,-1)).transpose(-2,-1) self.probe.data = t.as_tensor( ortho_probes, device=self.probe.device, dtype=self.probe.dtype) + + self.weights.data = new_weights.to(device=self.weights.device, + dtype=self.weights.dtype) + + # At this point, we now have a new set of basis probes, which are the + # eigenbasis for the full-experiment density matrix, and we have + # re-expressed all the shot-to-shot weight matrices in that basis. + # But, the shot-to-shot probes (self.weights self.probe) + # are still exactly the same as they were before. + # + # Oftentimes, we also want the shot-to-shot weights to be re-expressed + # so that the shot-to-shot probes are the eigenbasis for each + # individual shot's density matrix. That's what we do below. + # + if tidy_each_frame: + + # TODO: I need to check that this really works + + dm_rank = self.weights.shape[-2] + + for idx in range(self.weights.shape[0]): + weights = self.weights.data[idx].detach().cpu() + rho = t.mm(weights.transpose(0,1), weights.conj()) + + ortho_probes, A = analysis.orthogonalize_probes( + self.probe.detach().cpu(), density_matrix=rho, + normalize=False, keep_transform=True) + + + new_weights = t.linalg.pinv(A)[:dm_rank, :].conj() + + self.weights.data[idx] = new_weights.to( + device=self.weights.device, + dtype=self.weights.dtype) def plot_wavefront_variation(self, dataset, fig=None, mode='amplitude', **kwargs): diff --git a/src/cdtools/models/multislice_ptycho.py b/src/cdtools/models/multislice_ptycho.py new file mode 100644 index 0000000..9a1c660 --- /dev/null +++ b/src/cdtools/models/multislice_ptycho.py @@ -0,0 +1,916 @@ +import torch as t +from cdtools.models import CDIModel +from cdtools.datasets import Ptycho2DDataset +from cdtools import tools +from cdtools.tools import plotting as p +from cdtools.tools import analysis +from matplotlib import pyplot as plt +from datetime import datetime +import numpy as np +from scipy import linalg as sla +from copy import copy + +__all__ = ['MultislicePtycho'] + +class MultislicePtycho(CDIModel): + + def __init__(self, + wavelength, + detector_geometry, + obj_basis, + probe_guess, + obj_guess, + interslice_propagator, + surface_normal=t.tensor([0., 0., 1.], dtype=t.float32), + min_translation=t.tensor([0, 0], dtype=t.float32), + background=None, + probe_basis=None, + translation_offsets=None, + mask=None, + weights=None, + translation_scale=1, + saturation=None, + probe_support=None, + oversampling=1, + fourier_probe=False, + loss='amplitude mse', + units='um', + simulate_probe_translation=False, + simulate_finite_pixels=False, + dtype=t.float32, + exponentiate_obj=False, + obj_view_crop=0 + ): + + super(MultislicePtycho, self).__init__() + self.register_buffer('wavelength', + t.tensor(wavelength, dtype=dtype)) + self.store_detector_geometry(detector_geometry, + dtype=dtype) + + self.register_buffer('min_translation', + t.tensor(min_translation, dtype=dtype)) + + self.register_buffer('obj_basis', + t.tensor(obj_basis, dtype=dtype)) + + self.register_buffer('exponentiate_obj', + t.tensor(exponentiate_obj, dtype=bool)) + + self.register_buffer('interslice_propagator', + t.tensor(interslice_propagator, dtype=t.complex64)) + + if probe_basis is None: + self.register_buffer('probe_basis', + t.tensor(obj_basis, dtype=dtype)) + else: + self.register_buffer('probe_basis', + t.tensor(probe_basis, dtype=dtype)) + + self.register_buffer('surface_normal', + t.tensor(surface_normal, dtype=dtype)) + + if saturation is None: + self.saturation = None + else: + self.register_buffer('saturation', + t.tensor(saturation, dtype=dtype)) + + self.register_buffer('fourier_probe', + t.tensor(fourier_probe, dtype=bool)) + + # Not sure how to make this a buffer... + self.units = units + + if mask is None: + self.mask = None + else: + self.register_buffer('mask', + t.tensor(mask, dtype=t.bool)) + + probe_guess = t.tensor(probe_guess, dtype=t.complex64) + obj_guess = t.tensor(obj_guess, dtype=t.complex64) + + # We rescale the probe here so it learns at the same rate as the + # object + if probe_guess.dim() > 2: + probe_norm = 1 * t.max(t.abs(probe_guess[0])) + else: + probe_norm = 1 * t.max(t.abs(probe_guess)) + self.register_buffer('probe_norm', probe_norm.to(dtype)) + + self.probe = t.nn.Parameter(probe_guess / self.probe_norm) + self.obj = t.nn.Parameter(obj_guess) + + + self.obj_view_slice = np.s_[obj_view_crop:-obj_view_crop, + obj_view_crop:-obj_view_crop] + + # TODO: perhaps not working anymore for fourier cropped probes + if background is None: + raise NotImplementedError('Issues with this due to probe fourier padding') + shape = [s//oversampling for s in self.probe[0]] + background = 1e-6 * t.ones(shape, dtype=t.float32) + + self.background = t.nn.Parameter(background) + + if weights is None: + self.weights = None + else: + # We now need to distinguish between real-valued per-image + # weights and complex-valued per-mode weight matrices + if len(weights.shape) == 1: + # This is if it's just a list of numbers + self.weights = t.nn.Parameter(t.tensor(weights, + dtype=t.float32)) + else: + # Now this is a matrix of weights, so it needs to be complex + self.weights = t.nn.Parameter(t.tensor(weights, + dtype=t.complex64)) + + if translation_offsets is None: + self.translation_offsets = None + else: + t_o = t.tensor(translation_offsets, dtype=t.float32) + t_o = t_o / translation_scale + self.translation_offsets = t.nn.Parameter(t_o) + + self.register_buffer('translation_scale', + t.tensor(translation_scale, dtype=dtype)) + + if probe_support is None: + probe_support = t.ones_like(self.probe[0], dtype=t.bool) + self.register_buffer('probe_support', + t.tensor(probe_support, dtype=t.bool)) + self.probe.data *= self.probe_support + + self.register_buffer('oversampling', + t.tensor(oversampling, dtype=int)) + + self.register_buffer('simulate_probe_translation', + t.tensor(simulate_probe_translation, dtype=bool)) + + if simulate_probe_translation: + Is = t.arange(self.probe.shape[-2], dtype=dtype) + Js = t.arange(self.probe.shape[-1], dtype=dtype) + Is, Js = t.meshgrid(Is/t.max(Is), Js/t.max(Js)) + + I_phase = 2 * np.pi* Is * self.oversampling + J_phase = 2 * np.pi* Js * self.oversampling + self.register_buffer('I_phase', I_phase) + self.register_buffer('J_phase', J_phase) + + #from matplotlib import pyplot as plt + #p.plot_real( + # self.obj[(np.s_[:],) + self.obj_view_slice], + # fig=fig, + # basis=self.obj_basis, + # units=self.units) + #plt.show() + + self.register_buffer('simulate_finite_pixels', + t.tensor(simulate_finite_pixels, dtype=bool)) + + # Here we set the appropriate loss function + if (loss.lower().strip() == 'amplitude mse' + or loss.lower().strip() == 'amplitude_mse'): + self.loss = tools.losses.amplitude_mse + elif (loss.lower().strip() == 'poisson nll' + or loss.lower().strip() == 'poisson_nll'): + self.loss = tools.losses.poisson_nll + else: + raise KeyError('Specified loss function not supported') + + + @classmethod + def from_dataset(cls, + dataset, + dz, + nz, + probe_size=None, + randomize_ang=0, + n_modes=1, + n_obj_modes=1, + dm_rank=None, + translation_scale=1, + saturation=None, + probe_support_radius=None, + probe_fourier_crop=None, + propagator_fourier_crop=None, + propagation_distance=None, + scattering_mode=None, + oversampling=1, + fourier_probe=False, + loss='amplitude mse', + units='um', + simulate_probe_translation=False, + simulate_finite_pixels=False, + obj_view_crop=None, + obj_padding=200, + exponentiate_obj=False, + ): + + wavelength = dataset.wavelength + det_basis = dataset.detector_geometry['basis'] + det_shape = dataset[0][1].shape + distance = dataset.detector_geometry['distance'] + + # always do this on the cpu + get_as_args = dataset.get_as_args + dataset.get_as(device='cpu') + + # We include the *extras to make this work even with datasets, like + # polarization dependent datasets, that might toss out extra inputs + (indices, translations, *extras), patterns = dataset[:] + + dataset.get_as(*get_as_args[0], **get_as_args[1]) + + # Then, generate the probe geometry from the dataset + ewg = tools.initializers.exit_wave_geometry + obj_basis = ewg( + det_basis, + det_shape, + wavelength, + distance, + oversampling=oversampling, + ) + + if hasattr(dataset, 'sample_info') and \ + dataset.sample_info is not None and \ + 'orientation' in dataset.sample_info: + surface_normal = dataset.sample_info['orientation'][2] + else: + surface_normal = np.array([0., 0., 1.]) + + # If this information is supplied when the function is called, + # then we override the information in the .cxi file + if scattering_mode in {'t', 'transmission'}: + surface_normal = np.array([0., 0., 1.]) + elif scattering_mode in {'r', 'reflection'}: + outgoing_dir = np.cross(det_basis[:, 0], det_basis[:, 1]) + outgoing_dir /= np.linalg.norm(outgoing_dir) + surface_normal = outgoing_dir + np.array([0., 0., 1.]) + surface_normal /= -np.linalg.norm(surface_normal) + + # Next generate the object geometry from the probe geometry and + # the translations + + pix_translations = tools.interactions.translations_to_pixel( + obj_basis, + translations, + surface_normal=surface_normal, + ) + + obj_size, min_translation = tools.initializers.calc_object_setup( + [s * oversampling for s in det_shape], + pix_translations, + padding=obj_padding, + ) + + # Finally, initialize the probe and object using this information + if probe_size is None: + probe = tools.initializers.SHARP_style_probe( + dataset, + propagation_distance=propagation_distance, + oversampling=oversampling, + ) + else: + probe = tools.initializers.gaussian_probe( + dataset, + obj_basis, + probe_shape, + probe_size, + propagation_distance=propagation_distance, + ) + + if hasattr(dataset, 'background') and dataset.background is not None: + background = t.sqrt(dataset.background) + else: + background = 1e-6 * t.ones( + dataset.patterns.shape[-2:], dtype=t.float32) + + if probe_fourier_crop is not None: + probe = tools.propagators.far_field(probe) + probe = probe[probe_fourier_crop : probe.shape[-2] + - probe_fourier_crop, + probe_fourier_crop : probe.shape[-1] + - probe_fourier_crop] + probe = tools.propagators.inverse_far_field(probe) + # TODO: This may fail with oversampling != 1 + scale_factor = np.array(det_shape) / np.array(probe.shape) + probe_basis = obj_basis * scale_factor[None,:] + else: + probe_basis = obj_basis.clone() + + # Now we initialize all the subdominant probe modes + probe_max = t.max(t.abs(probe)) + probe_stack = [0.01 * probe_max * t.rand(probe.shape, dtype=probe.dtype) for i in range(n_modes - 1)] + + # For a Fourier space probe + if fourier_probe: + probe = tools.propagators.far_field(probe) + + probe = t.stack([probe, ] + probe_stack) + + # If exponentiate_obj, we're basically recovering the transmission + # matrix T - so, we don't exponentiate the initialization + if not exponentiate_obj: + obj = t.stack([t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5)) + for idx in range(nz)]) + else: + obj = t.stack([randomize_ang * (t.rand(obj_size)-0.5) + for idx in range(nz)]) + + pfc = (probe_fourier_crop if probe_fourier_crop else 0) + if obj_view_crop is None: + obj_view_crop = min(probe.shape[-2], probe.shape[-1]) // 2 + pfc + if obj_view_crop < 0: + obj_view_crop += min(probe.shape[-2], probe.shape[-1]) // 2 + pfc + + obj_view_crop += obj_padding + + det_geo = dataset.detector_geometry + + translation_offsets = 0 * (t.rand((len(dataset), 2)) - 0.5) + + if dm_rank is not None and dm_rank != 0: + if dm_rank > n_modes: + raise KeyError('Density matrix rank cannot be greater than the number of modes. Use dm_rank = -1 to use a full rank matrix.') + elif dm_rank == -1: + # dm_rank == -1 is defined to mean full-rank + dm_rank = n_modes + + Ws = t.zeros(len(dataset), dm_rank, n_modes, dtype=t.complex64) + # Start with as close to the identity matrix as possible, + # cutting of when we hit the specified maximum rank + for i in range(0, dm_rank): + Ws[:, i, i] = 1 + else: + # dm_rank == None or dm_rank = 0 triggers a special case where + # a standard incoherent multi-mode model is used. This is the + # default, because it is so common. + # In this case, we define a set of weights which only has one index + Ws = t.ones(len(dataset)) + + if hasattr(dataset, 'intensities') and dataset.intensities is not None: + Ws *= (dataset.intensities.to(dtype=Ws.dtype)[:,...] + / t.mean(dataset.intensities)) + + if hasattr(dataset, 'mask') and dataset.mask is not None: + mask = dataset.mask.to(t.bool) + else: + mask = None + + if probe_support_radius is not None: + probe_support = t.zeros(probe[0].shape, dtype=t.bool) + xs, ys = np.mgrid[:probe.shape[-2], :probe.shape[-1]] + xs = xs - np.mean(xs) + ys = ys - np.mean(ys) + Rs = np.sqrt(xs**2 + ys**2) + + probe_support[Rs < probe_support_radius] = 1 + probe = probe * probe_support[None, :, :] + + else: + probe_support = None + + # Here we define the inter-slice propagator + # TODO: this will probably fail for oversampling != 1 and + # for any non-square spacing (I didn't check if this was the + # correct ordering) + spacing = [np.abs(obj_basis[0,1]), np.abs(obj_basis[1,0])] + + interslice_propagator = \ + tools.propagators.generate_angular_spectrum_propagator( + det_shape, spacing, wavelength, dz) + + if ((propagator_fourier_crop is not None) + and (propagator_fourier_crop != 0)): + interslice_propagator = t.fft.fftshift(interslice_propagator) + interslice_propagator[:propagator_fourier_crop,:] = 0 + interslice_propagator[:,:propagator_fourier_crop] = 0 + interslice_propagator[-propagator_fourier_crop:,:] = 0 + interslice_propagator[:,-propagator_fourier_crop:] = 0 + interslice_propagator = t.fft.ifftshift(interslice_propagator) + + + return cls( + wavelength, + det_geo, + obj_basis, + probe, + obj, + interslice_propagator, + surface_normal=surface_normal, + min_translation=min_translation, + translation_offsets=translation_offsets, + weights=Ws, mask=mask, background=background, + translation_scale=translation_scale, + saturation=saturation, + probe_basis=probe_basis, + probe_support=probe_support, + fourier_probe=fourier_probe, + oversampling=oversampling, + loss=loss, units=units, + simulate_probe_translation=simulate_probe_translation, + simulate_finite_pixels=simulate_finite_pixels, + exponentiate_obj=exponentiate_obj, + obj_view_crop=obj_view_crop, + ) + + + def interaction(self, index, translations, *args): + + # The *args is included so that this can work even when given, say, + # a polarized ptycho dataset that might spit out more inputs. + + # Step 1 is to convert the translations for each position into a + # value in pixels + pix_trans = tools.interactions.translations_to_pixel( + self.obj_basis, + translations, + surface_normal=self.surface_normal) + pix_trans -= self.min_translation + # We then add on any recovered translation offset, if they exist + if self.translation_offsets is not None: + pix_trans += (self.translation_scale * + self.translation_offsets[index]) + + # This restricts the basis probes to stay within the probe support + basis_prs = self.probe * self.probe_support[..., :, :] + + # For a Fourier-space probe, we take an IFT + if self.fourier_probe: + basis_prs = tools.propagators.inverse_far_field(basis_prs) + + # Now we construct the probes for each shot from the basis probes + if self.weights is not None: + Ws = self.weights[index] + else: + try: + Ws = t.ones(len(index)) # I'm positive this introduced a bug + except: + Ws = 1 + + if self.weights is None or len(self.weights[0].shape) == 0: + # If a purely stable coherent illumination is defined + prs = Ws[..., None, None, None] * basis_prs + else: + # If a frame-by-frame weight matrix is defined + # This takes the dot product of all the weight matrices with + # the probes. The output has dimensions of translation, then + # coherent mode index, then x,y, and then complex index + # Maybe this can be done with a matmul now? + prs = t.sum(Ws[..., None, None] * basis_prs, axis=-3) + + if self.simulate_probe_translation: + det_pix_trans = tools.interactions.translations_to_pixel( + self.det_basis, + translations, + surface_normal=self.surface_normal) + + probe_masks = t.exp(1j* (det_pix_trans[:,0,None,None] * + self.I_phase[None,...] + + det_pix_trans[:,1,None,None] * + self.J_phase[None,...])) + prs = prs * probe_masks[...,None,:,:] + + + # We automatically rescale the probe to match the background size, + # which allows us to do stuff like let the object be super-resolution, + # while restricting the probe to the detector resolution but still + # doing an explicit real-space limitation of the probe + padding = [self.oversampling * self.background.shape[-2] - prs.shape[-2], + self.oversampling * self.background.shape[-1] - prs.shape[-1]] + + if any([p != 0 for p in padding]): # For probe_fourier_crop != 0. + padding = [padding[-1]//2, padding[-1]-padding[-1]//2, + padding[-2]//2, padding[-2]-padding[-2]//2] + prs = tools.propagators.far_field(prs) + prs = t.nn.functional.pad(prs, padding) + prs = tools.propagators.inverse_far_field(prs) + + # Now we actually do the interaction, using the sinc subpixel + # translation model as per usual + exit_waves = self.probe_norm * prs + + # Exponentiate the obj if we have to + obj = (self.obj if not self.exponentiate_obj + else t.exp(1j * self.obj)) + + for idx in range(self.obj.shape[0]): + # Interact with the object + exit_waves = tools.interactions.ptycho_2D_sinc( + exit_waves, obj[idx], pix_trans, + shift_probe=True, multiple_modes=True) + + # For all but the last slice + if idx <= (self.obj.shape[0] - 1): + exit_waves = tools.propagators.near_field( + exit_waves, self.interslice_propagator) + + return exit_waves + + + + def forward_propagator(self, wavefields): + return tools.propagators.far_field(wavefields) + + + def backward_propagator(self, wavefields): + return tools.propagators.inverse_far_field(wavefields) + + + def measurement(self, wavefields): + return tools.measurements.quadratic_background( + wavefields, + self.background, + measurement=tools.measurements.incoherent_sum, + saturation=self.saturation, + oversampling=int(self.oversampling), + simulate_finite_pixels=self.simulate_finite_pixels, + ) + + + # Note: No "loss" function is defined here, because it is added + # dynamically during object creation in __init__ + + def sim_to_dataset(self, args_list, calculation_width=None): + # In the future, potentially add more control + # over what metadata is saved (names, etc.) + + # First, I need to gather all the relevant data + # that needs to be added to the dataset + entry_info = {'program_name': 'cdtools', + 'instrument_n': 'Simulated Data', + 'start_time': datetime.now()} + + surface_normal = self.surface_normal.detach().cpu().numpy() + xsurfacevec = np.cross(np.array([0., 1., 0.]), surface_normal) + xsurfacevec /= np.linalg.norm(xsurfacevec) + ysurfacevec = np.cross(surface_normal, xsurfacevec) + ysurfacevec /= np.linalg.norm(ysurfacevec) + orientation = np.array([xsurfacevec, ysurfacevec, surface_normal]) + + sample_info = {'description': 'A simulated sample', + 'orientation': orientation} + + + mask = self.mask + wavelength = self.wavelength + indices, translations = args_list + + data = [] + len(indices) + if calculation_width is None: + calculation_width = len(indices) + index_chunks = [indices[i:i + calculation_width] + for i in range(0, len(indices), + calculation_width)] + translation_chunks = [translations[i:i + calculation_width] + for i in range(0, len(indices), + calculation_width)] + + + # Then we simulate the results + data = [self.forward(idx, trans).detach() + for idx, trans in zip(index_chunks, translation_chunks)] + + data = t.cat(data, dim=0) + # And finally, we make the dataset + return Ptycho2DDataset( + translations, data, + entry_info=entry_info, + sample_info=sample_info, + wavelength=wavelength, + detector_geometry=self.get_detector_geometry(), + mask=mask) + + + def corrected_translations(self, dataset): + translations = dataset.translations.to( + dtype=t.float32, device=self.probe.device) + if (hasattr(self, 'translation_offsets') and + self.translation_offsets is not None): + t_offset = tools.interactions.pixel_to_translations( + self.obj_basis, + self.translation_offsets * self.translation_scale, + surface_normal=self.surface_normal) + return translations + t_offset + else: + return translations + + + def get_rhos(self): + # If this is the general unified mode model + if self.weights.dim() >= 2: + Ws = self.weights.detach().cpu().numpy() + rhos_out = np.matmul(np.swapaxes(Ws, 1, 2), Ws.conj()) + return rhos_out + # This is the purely incoherent case + else: + return np.array([np.eye(self.probe.shape[0])]*self.weights.shape[0], + dtype=np.complex64) + + def center_probes(self, iterations=4): + """Centers the probes + + Note that this does not compensate for the centering by adjusting + the object, so it's a good idea to reset the object after centering + the probes + """ + centered_probe = tools.image_processing.center( + self.probe.data.cpu(), iterations=iterations) + self.probe.data = centered_probe.to(device=self.probe.data.device) + + + def tidy_probes(self, normalize=False, tidy_each_frame=False): + """Tidies up the probes + + What we want to do here is use all the information on all the probes + to calculate a natural basis for the experiment, and update all the + density matrices to operate in that updated basis + """ + + # First we treat the purely incoherent case + + # I don't love this pattern of using an if statement with a return + # to catch this case, but because it's so much simpler than the + # unified mode case I think it's appropriate + if self.weights.dim() == 1: + probe = self.probe.detach().cpu().numpy() + ortho_probes = analysis.orthogonalize_probes(probe) + self.probe.data = t.as_tensor( + ortho_probes, + device=self.probe.device, + dtype=self.probe.dtype) + return + + # This is for the unified mode case + + # We concatenate all the weight matrices + all_weights = t.cat(t.unbind(self.weights.detach().cpu(), dim=0), dim=0) + + # We use that to calculate the density matrix of the full experiment, + # normalized by number of exposures + overall_rho = (t.mm(all_weights.transpose(0,1), all_weights.conj()) + / self.weights.shape[0]) + + # We generate the orthogonal probes based on this full-experiment + # density matrix. We also keep the transform matrix A + probe = self.probe.detach().cpu().numpy() + ortho_probes, A = analysis.orthogonalize_probes( + probe, density_matrix=overall_rho, + keep_transform=True, normalize=normalize) + + # We apply A to the weight matrices to update them along with the + # probes + new_weights = t.matmul( + t.as_tensor(A).transpose(0,1), + self.weights.detach().cpu().transpose(-2,-1)).transpose(-2,-1) + + self.probe.data = t.as_tensor( + ortho_probes, device=self.probe.device, dtype=self.probe.dtype) + + self.weights.data = new_weights.to(device=self.weights.device, + dtype=self.weights.dtype) + + # At this point, we now have a new set of basis probes, which are the + # eigenbasis for the full-experiment density matrix, and we have + # re-expressed all the shot-to-shot weight matrices in that basis. + # But, the shot-to-shot probes (self.weights self.probe) + # are still exactly the same as they were before. + # + # Oftentimes, we also want the shot-to-shot weights to be re-expressed + # so that the shot-to-shot probes are the eigenbasis for each + # individual shot's density matrix. That's what we do below. + # + if tidy_each_frame: + + # TODO: I need to check that this really works + + dm_rank = self.weights.shape[-2] + + for idx in range(self.weights.shape[0]): + weights = self.weights.data[idx].detach().cpu() + rho = t.mm(weights.transpose(0,1), weights.conj()) + + ortho_probes, A = analysis.orthogonalize_probes( + self.probe.detach().cpu(), density_matrix=rho, + normalize=False, keep_transform=True) + + + new_weights = t.linalg.pinv(A)[:dm_rank, :].conj() + + self.weights.data[idx] = new_weights.to( + device=self.weights.device, + dtype=self.weights.dtype) + + + def plot_wavefront_variation(self, dataset, fig=None, mode='amplitude', **kwargs): + def get_probes(idx): + basis_prs = self.probe * self.probe_support[..., :, :] + prs = t.sum(self.weights[idx, :, :, None, None] * basis_prs, + axis=-4) + ortho_probes = analysis.orthogonalize_probes(prs) + + if mode.lower() == 'amplitude': + return np.abs(ortho_probes.detach().cpu().numpy()) + if mode.lower() == 'root_sum_intensity': + return np.sum(np.abs(ortho_probes.detach().cpu().numpy())**2, + axis=0) + if mode.lower() == 'phase': + return np.angle(ortho_probes.detach().cpu().numpy()) + + probe_matrix = np.zeros([self.probe.shape[0]]*2, + dtype=np.complex64) + np_probes = self.probe.detach().cpu().numpy() + for i in range(probe_matrix.shape[0]): + for j in range(probe_matrix.shape[0]): + probe_matrix[i,j] = np.sum(np_probes[i]*np_probes[j].conj()) + + weights = self.weights.detach().cpu().numpy() + + probe_intensities = np.sum(np.tensordot(weights, probe_matrix, axes=1) + * weights.conj(), axis=2) + + # Imaginary part is already essentially zero up to rounding error + probe_intensities = np.real(probe_intensities) + + values = np.sum(probe_intensities, axis=1) + if mode.lower() == 'amplitude' or mode.lower() == 'root_sum_intensity': + cmap = 'viridis' + else: + cmap = 'twilight' + + p.plot_nanomap_with_images( + self.corrected_translations(dataset), + get_probes, + values=values, + fig=fig, + units=self.units, + basis=self.obj_basis, + nanomap_colorbar_title='Total Probe Intensity', + cmap=cmap, + **kwargs), + + + plot_list = [ + ('', + lambda self, fig, dataset: self.plot_wavefront_variation( + dataset, + fig=fig, + mode='root_sum_intensity', + image_title='Root Summed Probe Intensities', + image_colorbar_title='Square Root of Intensity'), + lambda self: len(self.weights.shape) >= 2), + ('', + lambda self, fig, dataset: self.plot_wavefront_variation( + dataset, + fig=fig, + mode='amplitude', + image_title='Probe Amplitudes (scroll to view modes)', + image_colorbar_title='Probe Amplitude'), + lambda self: len(self.weights.shape) >= 2), + ('', + lambda self, fig, dataset: self.plot_wavefront_variation( + dataset, + fig=fig, + mode='phase', + image_title='Probe Phases (scroll to view modes)', + image_colorbar_title='Probe Phase'), + lambda self: len(self.weights.shape) >= 2), + ('Basis Probe Fourier Space Amplitudes', + lambda self, fig: p.plot_amplitude( + (self.probe if self.fourier_probe + else tools.propagators.inverse_far_field(self.probe)), + fig=fig)), + ('Basis Probe Fourier Space Phases', + lambda self, fig: p.plot_phase( + (self.probe if self.fourier_probe + else tools.propagators.inverse_far_field(self.probe)) + , fig=fig)), + ('Basis Probe Real Space Amplitudes', + lambda self, fig: p.plot_amplitude( + (self.probe if not self.fourier_probe + else tools.propagators.inverse_far_field(self.probe)), + fig=fig, + basis=self.probe_basis, + units=self.units)), + ('Basis Probe Real Space Phases', + lambda self, fig: p.plot_phase( + (self.probe if not self.fourier_probe + else tools.propagators.inverse_far_field(self.probe)), + fig=fig, + basis=self.probe_basis, + units=self.units)), + ('Average Density Matrix Amplitudes', + lambda self, fig: p.plot_amplitude( + np.nanmean(np.abs(self.get_rhos()), axis=0), + fig=fig), + lambda self: len(self.weights.shape) >= 2), + ('% Power in Top Mode (only accurate after tidy_probes)', + lambda self, fig, dataset: p.plot_nanomap( + self.corrected_translations(dataset), + analysis.calc_top_mode_fraction(self.get_rhos()), + fig=fig, + units=self.units), + lambda self: len(self.weights.shape) >= 2), + ('Object Amplitude', + lambda self, fig: p.plot_amplitude( + self.obj[(np.s_[:],) + self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + lambda self: not self.exponentiate_obj), + ('Object (T) Imaginary Part', + lambda self, fig: p.plot_imag( + self.obj[(np.s_[:],) + self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + lambda self: self.exponentiate_obj), + ('Object Phase', + lambda self, fig: p.plot_phase( + self.obj[(np.s_[:],) + self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + lambda self: not self.exponentiate_obj), + ('Object (T) Real Part', + lambda self, fig: p.plot_real( + self.obj[(np.s_[:],) + self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units, + cmap='cividis'), + lambda self: self.exponentiate_obj), + ('Object Product Amplitude', + lambda self, fig: p.plot_amplitude( + t.prod(self.obj, dim=0)[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + lambda self: not self.exponentiate_obj), + ('Object (T) Sum Imaginary Part', + lambda self, fig: p.plot_imag( + t.sum(self.obj, dim=0)[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + lambda self: self.exponentiate_obj), + ('Object Product Phase', + lambda self, fig: p.plot_phase( + t.prod(self.obj, dim=0)[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + lambda self: not self.exponentiate_obj), + ('Object (T) Sum Real Part', + lambda self, fig: p.plot_real( + t.sum(self.obj, dim=0)[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units, + cmap='cividis'), + lambda self: self.exponentiate_obj), + ('Corrected Translations', + lambda self, fig, dataset: p.plot_translations(self.corrected_translations(dataset), fig=fig, units=self.units)), + ('Background', + lambda self, fig: p.plot_amplitude(self.background**2, fig=fig)) + ] + + + def save_results(self, dataset): + # This will save out everything needed to recreate the object + # in the same state, but it's not the best formatted. For example, + # "background" stores the square root of the background, etc. + base_results = super().save_results() + + # We also save out the main results in a more readable format + obj_basis = self.obj_basis.detach().cpu().numpy() + probe_basis = self.probe_basis.detach().cpu().numpy() + translations=self.corrected_translations(dataset).detach().cpu().numpy() + original_translations = dataset.translations.detach().cpu().numpy() + probe = self.probe.detach().cpu().numpy() + probe = probe * self.probe_norm.detach().cpu().numpy() + obj = self.obj.detach().cpu().numpy() + background = self.background.detach().cpu().numpy()**2 + weights = self.weights.detach().cpu().numpy() + oversampling = self.oversampling.cpu().numpy() + wavelength = self.wavelength.cpu().numpy() + + results = { + 'obj_basis': obj_basis, + 'probe_basis': probe_basis, + 'translations': translations, + 'original_translations': original_translations, + 'probe': probe, + 'obj': obj, + 'background': background, + 'oversampling': oversampling, + 'weights': weights, + 'wavelength': wavelength, + } + + return {**base_results, **results} diff --git a/src/cdtools/tools/analysis/analysis.py b/src/cdtools/tools/analysis/analysis.py index d65aa94..187a6c8 100644 --- a/src/cdtools/tools/analysis/analysis.py +++ b/src/cdtools/tools/analysis/analysis.py @@ -21,6 +21,62 @@ __all__ = ['orthogonalize_probes', 'standardize', 'synthesize_reconstructions', 'standardize_reconstruction_set'] +def orthogonalize_probes_t( + probes, + density_matrix=None, + keep_transform=False, + normalize=False, + n_probe_dims=2, +): + """ Orthogonalizes a set of incoherently mixing probes + + TODO: actually make this, and replace ortho_probes with a fully + pytorch-based function + + Any set of probe modes defines a density matrix (a.k.a mutual coherence + function) that is the ultimate description of the state of the light + field. This function takes any set of probe modes - not necessarily + orthogonalized - and returns an orthogonalized set of probe modes. + Formally, it returns the eigenbasis of the density matrix, ordered + from largest to smallest eigenvalue. + + If normalize is set to True, then it will return the normalized + eigenbasis. Otherwise, it will return a scaled version of the eigenbasis, + so that the returned probes can be used directly for multi-mode + ptychography. + + If a density matrix is explicitly given, it will instead + consider the problem of extracting the eigenbasis of the matrix + probes * denstity_matrix * probes^dagger, where probes is the + column matrix of the given probe functions. This latter problem arises + in the generalization of the probe mixing model, and reduces to the + simpler case when the density matrix is equal to the identity matrix + + If the parameter "keep_transform" is set, the function will additionally + return the matrix A such that A * ortho_probes^dagger = probes^dagger + TODO: is the above right, or are ortho_probes and probes flipped? + + Parameters + ---------- + probes : array + An l x () complex array representing a stack of probes + density_matrix : array + An optional l x l density matrix further elaborating on the state + keep_transform : bool + Default False, whether to return the map from probes to ortho_probes + normalize : bool + Default False, whether to normalize the probe modes + n_probe_dims : int + Default 2, the number of trailing dimensions defining each probe + + Returns + ------- + ortho_probes: array + An l x () complex array representing a stack of probes + """ + pass + + def orthogonalize_probes(probes, density_matrix=None, keep_transform=False, normalize=False): """Orthogonalizes a set of incoherently mixing probes @@ -40,6 +96,7 @@ def orthogonalize_probes(probes, density_matrix=None, keep_transform=False, norm If the parameter "keep_transform" is set, the function will additionally return the matrix A such that A * ortho_probes^dagger = probes^dagger + TODO: is the above right, or are ortho_probes and probes flipped? If the parameter "normalize" is False (as is the default), the variation in intensities in the probe modes will be kept in the probe modes, as is