From f84e4418438c41289b9fe57ee4a72416ff97929d Mon Sep 17 00:00:00 2001 From: Abe Levitan Date: Thu, 5 Aug 2021 16:21:58 -0400 Subject: [PATCH] Standardize the way that data is copied when models are created --- CDTools/__init__.py | 8 + CDTools/models/bragg_2d_ptycho.py | 133 ++++---- CDTools/models/fancy_ptycho.py | 431 ++++++++++++------------- CDTools/models/multislice_2d_ptycho.py | 147 ++++----- CDTools/models/rpi.py | 53 +-- CDTools/models/simple_ptycho.py | 17 +- 6 files changed, 383 insertions(+), 406 deletions(-) diff --git a/CDTools/__init__.py b/CDTools/__init__.py index df67147..258a104 100644 --- a/CDTools/__init__.py +++ b/CDTools/__init__.py @@ -1,3 +1,11 @@ +# This is needed to allow us to use torch.tensor in the module without it +# constantly complaining. +import warnings +warnings.filterwarnings("ignore", + message='To copy construct from a tensor, ') + +__all__ = ['tools', 'datasets', 'models'] + from CDTools import tools from CDTools import datasets from CDTools import models diff --git a/CDTools/models/bragg_2d_ptycho.py b/CDTools/models/bragg_2d_ptycho.py index 43ce500..37188cd 100644 --- a/CDTools/models/bragg_2d_ptycho.py +++ b/CDTools/models/bragg_2d_ptycho.py @@ -57,14 +57,13 @@ class Bragg2DPtycho(CDIModel): def __init__(self, wavelength, detector_geometry, probe_basis, probe_guess, obj_guess, detector_slice=None, - min_translation = t.Tensor([0,0]), - median_propagation = t.Tensor(data=[0]), - background = None, translation_offsets=None, mask=None, - weights = None, translation_scale = 1, saturation=None, - probe_support = None, obj_support=None, oversampling=1, + min_translation=t.tensor([0, 0], dtype=t.float32), + median_propagation=t.tensor(0, dtype=t.float32), + background=None, translation_offsets=None, mask=None, + weights=None, translation_scale=1, saturation=None, + probe_support=None, oversampling=1, propagate_probe=True, correct_tilt=True, lens=False): - # We need the detector geometry # We need the probe basis (but in this case, we don't need the surface # normal because it comes implied by the probe basis @@ -73,94 +72,90 @@ class Bragg2DPtycho(CDIModel): # The median propagation should be needed as well # translation_offsets can stay 2D for now # propagate_probe and correct_tilt are important! - - super(Bragg2DPtycho,self).__init__() - self.wavelength = t.Tensor([wavelength]) + + super(Bragg2DPtycho, self).__init__() + self.wavelength = t.tensor(wavelength) self.detector_geometry = copy(detector_geometry) det_geo = self.detector_geometry if hasattr(det_geo, 'distance'): - det_geo['distance'] = t.Tensor(det_geo['distance']) + det_geo['distance'] = t.tensor(det_geo['distance']) if hasattr(det_geo, 'basis'): - det_geo['basis'] = t.Tensor(det_geo['basis']) + det_geo['basis'] = t.tensor(det_geo['basis']) if hasattr(det_geo, 'corner'): - det_geo['corner'] = t.Tensor(det_geo['corner']) + det_geo['corner'] = t.tensor(det_geo['corner']) - self.min_translation = t.Tensor(min_translation) - self.median_propagation = median_propagation + self.min_translation = t.tensor(min_translation) + self.median_propagation = t.tensor(median_propagation) - self.probe_basis = t.Tensor(probe_basis) - self.detector_slice = detector_slice + self.probe_basis = t.tensor(probe_basis) + self.detector_slice = copy(detector_slice) # calculate the surface normal from the probe basis surface_normal = np.cross(np.array(probe_basis)[:,1], np.array(probe_basis)[:,0]) surface_normal /= np.linalg.norm(surface_normal) - self.surface_normal = t.Tensor(surface_normal) + self.surface_normal = t.tensor(surface_normal) self.saturation = saturation if mask is None: self.mask = mask else: - self.mask = t.BoolTensor(mask) - + self.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() > 3: - self.probe_norm = 1 * t.max(t.abs(probe_guess[0]).to(t.float32)) + if probe_guess.dim() > 2: + self.probe_norm = 1 * t.max(t.abs(probe_guess[0])) else: - self.probe_norm = 1 * t.max(t.abs(probe_guess).to(t.float32)) + self.probe_norm = 1 * t.max(t.abs(probe_guess)) # Not strictly necessary but otherwise it will return # a probe with the stuff outside of the support unchanged after # optimization if probe_support is not None: + self.probe_support = t.tensor(probe_support, dtype=t.bool) probe_guess = probe_guess * probe_support # This seems dumb, but otherwise it winds up with a mixture # of negative and positive zeros and it's super annoying when # you look at the phase map probe_guess[probe_guess == 0] = 0 - - self.probe = t.nn.Parameter(probe_guess.to(t.complex64) - / self.probe_norm) - - self.obj = t.nn.Parameter(obj_guess.to(t.complex64)) - + else: + self.probe_support = t.ones(self.probe[0].shape, dtype=t.bool) + + self.probe = t.nn.Parameter(probe_guess / self.probe_norm) + self.obj = t.nn.Parameter(obj_guess) + if background is None: if detector_slice is not None: - background = 1e-6 * t.ones(self.probe[0][self.detector_slice].shape) + background = 1e-6 * t.ones( + self.probe[0][self.detector_slice].shape, + dtype=t.float32) else: - background = 1e-6 * t.ones(self.probe[0].shape) + background = 1e-6 * t.ones(self.probe[0].shape, + dtype=t.float32) - - self.background = t.nn.Parameter(t.as_tensor(background,dtype=t.float32)) + self.background = t.nn.Parameter(background) if weights is None: self.weights = None else: - self.weights = t.nn.Parameter(t.as_tensor(weights,dtype=t.float32)) - + # No incoherent + unstable here yet + self.weights = t.nn.Parameter(t.tensor(weights, + dtype=t.float32)) + if translation_offsets is None: self.translation_offsets = None else: - self.translation_offsets = t.nn.Parameter( - t.as_tensor(translation_offsets,dtype=t.float32) / - translation_scale) + t_o = t.tensor(translation_offsets, dtype=t.float32) + t_o = t_o / translation_scale + self.translation_offsets = t.nn.Parameter(t_o) self.translation_scale = translation_scale - if probe_support is not None: - self.probe_support = probe_support - else: - self.probe_support = t.ones(self.probe[0].shape,dtype=t.bool) - - if obj_support is not None: - self.obj_support = obj_support - self.obj.data = self.obj * obj_support - else: - self.obj_support = t.ones_like(self.obj, dtype=t.bool) - - self.oversampling = oversampling self.propagate_probe = propagate_probe @@ -174,7 +169,7 @@ class Bragg2DPtycho(CDIModel): self.k_map, self.intensity_map = \ tools.propagators.generate_high_NA_k_intensity_map( self.probe_basis, - self.detector_geometry['basis']/ oversampling, + self.detector_geometry['basis'] / oversampling, probe_shape, self.detector_geometry['distance'], self.wavelength,dtype=t.float32, @@ -183,21 +178,21 @@ class Bragg2DPtycho(CDIModel): self.k_map = None self.intensity_map = None - self.prop_dir = t.Tensor([0,0,1]).to(dtype=t.float32) + self.prop_dir = t.tensor([0, 0, 1], dtype=t.float32) # This propagator should be able to be multiplied by the propagation # distance each time to get a propagator - self.universal_propagator = t.angle(ggasp(self.probe.shape[1:], - self.probe_basis, self.wavelength, - t.Tensor([0,0,self.wavelength/(2*np.pi)]), - propagation_vector=self.prop_dir, - dtype=t.complex64, - propagate_along_offset=True)) + self.universal_propagator = t.angle(ggasp( + self.probe.shape[1:], + self.probe_basis, self.wavelength, + t.tensor([0, 0, self.wavelength/(2*np.pi)], dtype=t.float32), + propagation_vector=self.prop_dir, + dtype=t.complex64, + propagate_along_offset=True)) - @classmethod - def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=True, propagate_probe=True,correct_tilt=True, lens=False, opt_for_fft=False): + def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, scattering_mode=None, oversampling=1, auto_center=True, propagate_probe=True,correct_tilt=True, lens=False, opt_for_fft=False): wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] @@ -262,11 +257,8 @@ class Bragg2DPtycho(CDIModel): # equal the input vector with a trailing 0, so we can do the # projection with a pseudoinverse and removing the last column - projector = np.linalg.pinv(mat)[:,:3] - - + projector = np.linalg.pinv(mat)[:, :3] probe_basis = t.Tensor(np.dot(projector, ew_basis)) - # Now we need a much better way to handle the translations here # than translations_to_pixel @@ -324,17 +316,6 @@ class Bragg2DPtycho(CDIModel): else: probe_support = None; - if restrict_obj != -1: - ro = restrict_obj - os = np.array(obj_size) - ps = np.array(probe_shape) - obj_support = t.zeros(obj.shape,dtype=t.bool) - obj_support[ps[0]//2-ro:os[0]+ro-ps[0]//2, - ps[1]//2-ro:os[1]+ro-ps[1]//2] = 1 - else: - obj_support = None - - # Here we need to implement a simple condition to choose whether # to propagate the probe or not if not( propagate_probe is True or propagate_probe is False): @@ -353,7 +334,6 @@ class Bragg2DPtycho(CDIModel): translation_scale=translation_scale, saturation=saturation, probe_support=probe_support, - obj_support=obj_support, oversampling=oversampling, propagate_probe=propagate_probe, correct_tilt=correct_tilt, @@ -385,7 +365,7 @@ class Bragg2DPtycho(CDIModel): prs[j] = tools.propagators.near_field(prs[j], propagator) exit_waves = self.probe_norm * tools.interactions.ptycho_2D_sinc( - prs, self.obj_support * self.obj,pix_trans, + prs, self.obj,pix_trans, shift_probe=True, multiple_modes=True) return exit_waves @@ -443,7 +423,6 @@ class Bragg2DPtycho(CDIModel): self.probe_basis = self.probe_basis.to(*args,**kwargs) self.probe_norm = self.probe_norm.to(*args,**kwargs) self.probe_support = self.probe_support.to(*args,**kwargs) - self.obj_support = self.obj_support.to(*args,**kwargs) self.surface_normal = self.surface_normal.to(*args, **kwargs) self.prop_dir = self.prop_dir.to(*args, **kwargs) self.universal_propagator = self.universal_propagator.to(*args,**kwargs) diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py index a0c9a95..84b4b48 100644 --- a/CDTools/models/fancy_ptycho.py +++ b/CDTools/models/fancy_ptycho.py @@ -12,64 +12,68 @@ from copy import copy __all__ = ['FancyPtycho'] + class FancyPtycho(CDIModel): def __init__(self, wavelength, detector_geometry, probe_basis, probe_guess, obj_guess, detector_slice=None, - surface_normal=np.array([0.,0.,1.]), - min_translation = t.Tensor([0,0]), - background = None, translation_offsets=None, mask=None, - weights = None, translation_scale = 1, saturation=None, - probe_support = None, obj_support=None, oversampling=1, - loss='amplitude mse',units='um'): - - super(FancyPtycho,self).__init__() - self.wavelength = t.Tensor([wavelength]) + surface_normal=t.tensor([0., 0., 1.], dtype=t.float32), + min_translation=t.tensor([0, 0], dtype=t.float32), + background=None, translation_offsets=None, mask=None, + weights=None, translation_scale=1, saturation=None, + probe_support=None, oversampling=1, + loss='amplitude mse', units='um'): + + super(FancyPtycho, self).__init__() + self.wavelength = t.tensor(wavelength) self.detector_geometry = copy(detector_geometry) det_geo = self.detector_geometry if hasattr(det_geo, 'distance'): - det_geo['distance'] = t.Tensor(det_geo['distance']) + det_geo['distance'] = t.tensor(det_geo['distance']) if hasattr(det_geo, 'basis'): - det_geo['basis'] = t.Tensor(det_geo['basis']) + det_geo['basis'] = t.tensor(det_geo['basis']) if hasattr(det_geo, 'corner'): - det_geo['corner'] = t.Tensor(det_geo['corner']) + det_geo['corner'] = t.tensor(det_geo['corner']) - self.min_translation = t.Tensor(min_translation) + self.min_translation = t.tensor(min_translation) + + self.probe_basis = t.tensor(probe_basis) + self.detector_slice = copy(detector_slice) + self.surface_normal = t.tensor(surface_normal) - self.probe_basis = t.Tensor(probe_basis) - self.detector_slice = detector_slice - self.surface_normal = t.Tensor(surface_normal) - self.saturation = saturation self.units = units - + if mask is None: self.mask = mask else: - self.mask = t.BoolTensor(mask) - + self.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: - self.probe_norm = 1 * t.max(t.abs(probe_guess[0].to(t.complex64))) + self.probe_norm = 1 * t.max(t.abs(probe_guess[0])) else: - self.probe_norm = 1 * t.max(t.abs(probe_guess.to(t.complex64))) - - self.probe = t.nn.Parameter(probe_guess.to(t.complex64) - / self.probe_norm) - - self.obj = t.nn.Parameter(obj_guess.to(t.complex64)) - + self.probe_norm = 1 * t.max(t.abs(probe_guess)) + + self.probe = t.nn.Parameter(probe_guess / self.probe_norm) + self.obj = t.nn.Parameter(obj_guess) + if background is None: if detector_slice is not None: - background = 1e-6 * t.ones(self.probe[0][self.detector_slice].shape) + background = 1e-6 * t.ones( + self.probe[0][self.detector_slice].shape, + dtype=t.float32) else: - background = 1e-6 * t.ones(self.probe[0].shape) + background = 1e-6 * t.ones(self.probe[0].shape, + dtype=t.float32) - - self.background = t.nn.Parameter(t.Tensor(background).to(t.float32)) + self.background = t.nn.Parameter(background) if weights is None: self.weights = None @@ -78,53 +82,42 @@ class FancyPtycho(CDIModel): # 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).to(t.float32)) + self.weights = t.nn.Parameter(t.tensor(weights, + dtype=t.float32)) else: - # Now this is a matrix of weights, so we - if type(weights) == type(t.zeros(1)): - self.weights = t.nn.Parameter(weights.to(t.complex64)) - else: - # There is a good chance that this doesn't work - self.weights = t.nn.Parameter(t.Tensor(weights).to(t.complex64)) + # 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: - self.translation_offsets = t.nn.Parameter(t.Tensor(translation_offsets).to(t.float32)/ translation_scale) + t_o = t.tensor(translation_offsets, dtype=t.float32) + t_o = t_o / translation_scale + self.translation_offsets = t.nn.Parameter(t_o) self.translation_scale = translation_scale if probe_support is not None: self.probe_support = probe_support else: - self.probe_support = t.ones_like(self.probe[0]) - - if obj_support is not None: - self.obj_support = obj_support - self.obj.data = self.obj * obj_support - else: - self.obj_support = t.ones_like(self.obj) + self.probe_support = t.ones_like(self.probe[0], dtype=t.bool) self.oversampling = oversampling - # Here we set the appropriate loss function - if loss.lower().strip() == 'amplitude mse'\ - or loss.lower().strip() == 'amplitude_mse': + 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': + 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, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um', polarized=False): - - # if polarized=True, this function takes care of the dataset iterator by taking into account the polarizer and analyzer components - # however, it drops them afterwards and treats the dataset as if it's not polarized + def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale=1, saturation=None, probe_support_radius=None, propagation_distance=None, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um'): wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] @@ -134,56 +127,53 @@ class FancyPtycho(CDIModel): # always do this on the cpu get_as_args = dataset.get_as_args dataset.get_as(device='cpu') - if not polarized: - (indices, translations), patterns = dataset[:] - else: - (indices, translations, polarizer, analyzer), patterns = dataset[:] - dataset.get_as(*get_as_args[0],**get_as_args[1]) + + # 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]) # Set to none to avoid issues with things outside the detector if auto_center: - center = tools.image_processing.centroid(t.sum(patterns,dim=0)) + center = tools.image_processing.centroid(t.sum(patterns, dim=0)) else: center = None - + # Then, generate the probe geometry from the dataset ewg = tools.initializers.exit_wave_geometry - probe_basis, probe_shape, det_slice = ewg(det_basis, - det_shape, - wavelength, - distance, - center=center, - padding=padding, - opt_for_fft=opt_for_fft, - oversampling=oversampling) - + probe_basis, probe_shape, det_slice = ewg(det_basis, + det_shape, + wavelength, + distance, + center=center, + padding=padding, + opt_for_fft=opt_for_fft, + 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.]) - + 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.]) + 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.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 = 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(probe_basis, translations, surface_normal=surface_normal) obj_size, min_translation = tools.initializers.calc_object_setup(probe_shape, pix_translations, padding=200) - if hasattr(dataset, 'background') and dataset.background is not None: background = t.sqrt(dataset.background) else: @@ -195,18 +185,17 @@ class FancyPtycho(CDIModel): else: probe = tools.initializers.gaussian_probe(dataset, probe_basis, probe_shape, probe_size, propagation_distance=propagation_distance) - # 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)] - probe = t.stack([probe,] + probe_stack) - #probe = t.stack([tools.propagators.far_field(probe),] + probe_stack) + probe_stack = [0.01 * probe_max * t.rand(probe.shape, dtype=probe.dtype) for i in range(n_modes - 1)] + probe = t.stack([probe, ] + probe_stack) + # probe = t.stack([tools.propagators.far_field(probe),] + probe_stack) + + obj = t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5)) - obj = t.exp(1j*randomize_ang * (t.rand(obj_size)-0.5)) - det_geo = dataset.detector_geometry - translation_offsets = 0 * (t.rand((len(dataset),2)) - 0.5) + 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: @@ -214,103 +203,92 @@ class FancyPtycho(CDIModel): 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) + + 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 + 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, '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]] + 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= 2: Ws = self.weights.detach().cpu().numpy() - rhos_out = np.matmul(np.swapaxes(Ws,1,2), Ws.conj()) + 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 tidy_probes(self, normalization=1, normalize=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) + self.probe.data = t.as_tensor( + ortho_probes, + device=self.probe.device, + dtype=self.probe.dtype) return # This is for the unified mode case - + # 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. - + rhos = self.get_rhos() - overall_rho = np.mean(rhos,axis=0) + overall_rho = np.mean(rhos, axis=0) probe = self.probe.detach().cpu().numpy() - ortho_probes, A = analysis.orthogonalize_probes(probe, - density_matrix=overall_rho, - keep_transform=True, - normalize=normalize) + 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 = 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, v = sla.eigh(rho) w = w[::-1][:dm_rank] - v = v[:,::-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())) - + 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) - - self.probe.data = t.as_tensor(ortho_probes, - device=self.probe.device,dtype=self.probe.dtype) + self.weights.data = t.as_tensor( + new_Ws, dtype=self.weights.dtype, device=self.weights.device) + + self.probe.data = t.as_tensor( + ortho_probes, device=self.probe.device, dtype=self.probe.dtype) - def plot_wavefront_variation(self, dataset,fig=None,mode='amplitude',**kwargs): + 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) + 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) + 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) + + 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) + + 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.probe_basis, nanomap_colorbar_title='Total Probe Intensity',cmap=cmap,**kwargs), + + p.plot_nanomap_with_images(self.corrected_translations(dataset), get_probes, values=values, fig=fig, units=self.units, basis=self.probe_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, 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, 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, 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 Amplitudes (scroll to view modes)', - lambda self, fig: p.plot_amplitude(self.probe, fig=fig, basis=self.probe_basis,units=self.units)), + lambda self, fig: p.plot_amplitude(self.probe, fig=fig, basis=self.probe_basis, units=self.units)), ('Basis Probe Phases (scroll to view modes)', - lambda self, fig: p.plot_phase(self.probe, fig=fig, basis=self.probe_basis,units=self.units)), + lambda self, fig: p.plot_phase(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), + 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, fig=fig, basis=self.probe_basis,units=self.units)), + 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, fig=fig, basis=self.probe_basis, units=self.units)), ('Object Phase', - lambda self, fig: p.plot_phase(self.obj, fig=fig, basis=self.probe_basis,units=self.units)), + lambda self, fig: p.plot_phase(self.obj, fig=fig, basis=self.probe_basis, units=self.units)), ('Corrected Translations', - lambda self, fig, dataset: p.plot_translations(self.corrected_translations(dataset), fig=fig,units=self.units)), + lambda self, fig, dataset: p.plot_translations(self.corrected_translations(dataset), fig=fig, units=self.units)), ('Background', lambda self, fig: plt.figure(fig.number) and plt.imshow(self.background.detach().cpu().numpy()**2)) ] - + def save_results(self, dataset): basis = self.probe_basis.detach().cpu().numpy() translations = self.corrected_translations(dataset).detach().cpu().numpy() @@ -556,8 +539,8 @@ class FancyPtycho(CDIModel): obj = self.obj.detach().cpu().numpy() background = self.background.detach().cpu().numpy()**2 weights = self.weights.detach().cpu().numpy() - - return {'basis':basis, 'translation':translations, - 'probe':probe,'obj':obj, - 'background':background, - 'weights':weights} + + return {'basis': basis, 'translation': translations, + 'probe': probe, 'obj': obj, + 'background': background, + 'weights': weights} diff --git a/CDTools/models/multislice_2d_ptycho.py b/CDTools/models/multislice_2d_ptycho.py index a996f0b..2843e3d 100644 --- a/CDTools/models/multislice_2d_ptycho.py +++ b/CDTools/models/multislice_2d_ptycho.py @@ -18,21 +18,27 @@ class Multislice2DPtycho(CDIModel): @property def probe(self): - return t.complex(self.probe_real,self.probe_imag) + return t.complex(self.probe_real, self.probe_imag) @property def obj(self): - return t.complex(self.obj_real,self.obj_imag) - - def __init__(self, wavelength, detector_geometry, + return t.complex(self.obj_real, self.obj_imag) + + def __init__(self, + wavelength, + detector_geometry, probe_basis, probe_guess, obj_guess, dz, nz, detector_slice=None, - surface_normal=np.array([0.,0.,1.]), - min_translation = t.Tensor([0,0]), - background = None, translation_offsets=None, mask=None, - weights = None, translation_scale = 1, saturation=None, - probe_support = None, + surface_normal=np.array([0., 0., 1.]), + min_translation=t.Tensor([0, 0]), + background=None, + translation_offsets=None, + mask=None, + weights=None, + translation_scale=1, + saturation=None, + probe_support=None, oversampling=1, bandlimit=None, subpixel=True, @@ -40,69 +46,71 @@ class Multislice2DPtycho(CDIModel): fourier_probe=False, prevent_aliasing=True, phase_only=False, - units='um'): - - super(Multislice2DPtycho,self).__init__() - self.wavelength = t.Tensor([wavelength]) + units='um', + ): + + super(Multislice2DPtycho, self).__init__() + self.wavelength = t.tensor(wavelength) self.detector_geometry = copy(detector_geometry) self.dz = dz self.nz = nz det_geo = self.detector_geometry if hasattr(det_geo, 'distance'): - det_geo['distance'] = t.Tensor(det_geo['distance']) + det_geo['distance'] = t.tensor(det_geo['distance']) if hasattr(det_geo, 'basis'): - det_geo['basis'] = t.Tensor(det_geo['basis']) + det_geo['basis'] = t.tensor(det_geo['basis']) if hasattr(det_geo, 'corner'): - det_geo['corner'] = t.Tensor(det_geo['corner']) + det_geo['corner'] = t.tensor(det_geo['corner']) - self.min_translation = t.Tensor(min_translation) + self.min_translation = t.tensor(min_translation) + + self.probe_basis = t.tensor(probe_basis) + self.detector_slice = copy(detector_slice) + self.surface_normal = t.tensor(surface_normal) - self.probe_basis = t.Tensor(probe_basis) - self.detector_slice = detector_slice - self.surface_normal = t.Tensor(surface_normal) - self.saturation = saturation self.subpixel = subpixel self.exponentiate_obj = exponentiate_obj self.fourier_probe = fourier_probe self.units = units - self.phase_only=phase_only - self.prevent_aliasing=prevent_aliasing - + self.phase_only = phase_only + self.prevent_aliasing = prevent_aliasing + if mask is None: self.mask = mask else: - self.mask = t.BoolTensor(mask) - + self.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() > 3: - self.probe_norm = 1 * t.max(t.abs(probe_guess[0]).to(t.float32)) + self.probe_norm = 1 * t.max(t.abs(probe_guess[0])) else: - self.probe_norm = 1 * t.max(t.abs(probe_guess).to(t.float32)) + self.probe_norm = 1 * t.max(t.abs(probe_guess)) - pg = probe_guess.to(t.complex64)/self.probe_norm + pg = probe_guess / self.probe_norm self.probe_real = t.nn.Parameter(pg.real) self.probe_imag = t.nn.Parameter(pg.imag) - og = obj_guess.to(t.complex64) - self.obj_real = t.nn.Parameter(og.real) - self.obj_imag = t.nn.Parameter(og.imag) - - + self.obj_real = t.nn.Parameter(obj_guess.real) + self.obj_imag = t.nn.Parameter(obj_guess.imag) + #self.probe = t.nn.Parameter(probe_guess.to(t.complex64) # / self.probe_norm) - #self.obj = t.nn.Parameter(obj_guess.to(t.complex64)) if background is None: if detector_slice is not None: - background = 1e-6 * t.ones(self.probe[0][self.detector_slice].shape) + background = 1e-6 * t.ones( + self.probe[0][self.detector_slice].shape, dtype=t.float32) else: - background = 1e-6 * t.ones(self.probe[0].shape) + background = 1e-6 * t.ones(self.probe[0].shape, + dtype=t.float32) - - self.background = t.nn.Parameter(t.as_tensor(background,dtype=t.float32)) + self.background = t.nn.Parameter(background) if weights is None: self.weights = None @@ -111,47 +119,43 @@ class Multislice2DPtycho(CDIModel): # 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.as_tensor(weights, - dtype=t.float32)) + self.weights = t.nn.Parameter(t.tensor(weights, + dtype=t.float32)) else: - # Now this is a matrix of weights, so we - self.weights = t.nn.Parameter(t.as_tensor(weights, - dtype=t.complex64)) - + # 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: - self.translation_offsets = t.nn.Parameter(t.as_tensor(translation_offsets,dtype=t.float32)/ translation_scale) + t_o = t.tensor(translation_offsets, dtype=t.float32) + t_o = t_o / translation_scale + self.translation_offsets = t.nn.Parameter(t_o) self.translation_scale = translation_scale if probe_support is not None: - self.probe_support = t.as_tensor(probe_support,dtype=t.bool) + self.probe_support = t.tensor(probe_support, dtype=t.bool) else: - self.probe_support = None#t.ones_like(self.probe,dtype=t.bool)#None - + self.probe_support = None + self.oversampling = oversampling - spacing = np.linalg.norm(self.probe_basis,axis=0) + spacing = np.linalg.norm(self.probe_basis, axis=0) shape = np.array(self.probe.shape[1:]) if prevent_aliasing: shape *= 2 spacing /= 2 - + self.bandlimit = bandlimit + self.as_prop = tools.propagators.generate_angular_spectrum_propagator(shape, spacing, self.wavelength, self.dz, self.bandlimit) - self.as_prop = tools.propagators.generate_angular_spectrum_propagator(shape, spacing, self.wavelength, self.dz, bandlimit=1/np.sqrt(2))#self.bandlimit) - #plt.imshow(t.abs(self.as_prop)) - #plt.figure() - #plt.imshow(t.abs(t.fft.fftshift(t.fft.ifft2(self.as_prop)))) - #plt.show() - #exit() - @classmethod - def from_dataset(cls, dataset, dz, nz, probe_convergence_semiangle, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, propagation_distance=None, scattering_mode=None, oversampling=1, auto_center=True, bandlimit=None, replicate_slice=False, subpixel=True, exponentiate_obj=True, units='um', fourier_probe=False, phase_only=False, prevent_aliasing=True, probe_support_radius=None): - + def from_dataset(cls, dataset, dz, nz, probe_convergence_semiangle, padding=0, n_modes=1, dm_rank=None, translation_scale=1, saturation=None, propagation_distance=None, scattering_mode=None, oversampling=1, auto_center=True, bandlimit=None, replicate_slice=False, subpixel=True, exponentiate_obj=True, units='um', fourier_probe=False, phase_only=False, prevent_aliasing=True, probe_support_radius=None): + wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] det_shape = dataset[0][1].shape @@ -161,24 +165,24 @@ class Multislice2DPtycho(CDIModel): get_as_args = dataset.get_as_args dataset.get_as(device='cpu') (indices, translations), patterns = dataset[:] - dataset.get_as(*get_as_args[0],**get_as_args[1]) + dataset.get_as(*get_as_args[0], **get_as_args[1]) # Set to none to avoid issues with things outside the detector if auto_center: - center = tools.image_processing.centroid(t.sum(patterns,dim=0)) + center = tools.image_processing.centroid(t.sum(patterns, dim=0)) else: center = None - + # Then, generate the probe geometry from the dataset ewg = tools.initializers.exit_wave_geometry - probe_basis, probe_shape, det_slice = ewg(det_basis, - det_shape, - wavelength, - distance, - center=center, - padding=padding, - opt_for_fft=False, - oversampling=oversampling) + probe_basis, probe_shape, det_slice = ewg(det_basis, + det_shape, + wavelength, + distance, + center=center, + padding=padding, + opt_for_fft=False, + oversampling=oversampling) if hasattr(dataset, 'sample_info') and \ @@ -227,7 +231,6 @@ class Multislice2DPtycho(CDIModel): if fourier_probe: probe = tools.propagators.far_field(probe) - # Consider a different start if exponentiate_obj: obj = t.zeros(obj_size, dtype=t.complex64) diff --git a/CDTools/models/rpi.py b/CDTools/models/rpi.py index 60ed3b4..ae0602c 100644 --- a/CDTools/models/rpi.py +++ b/CDTools/models/rpi.py @@ -51,67 +51,70 @@ class RPI(CDIModel): return t.complex(self.obj_real, self.obj_imag) def __init__(self, wavelength, detector_geometry, probe_basis, - probe, obj_guess, detector_slice=None, - background = None, mask=None, saturation=None, + probe, obj_guess, detector_slice=None, + background=None, mask=None, saturation=None, obj_support=None, oversampling=1): - super(RPI,self).__init__() + super(RPI, self).__init__() - self.wavelength = t.Tensor([wavelength]) + self.wavelength = t.tensor(wavelength) self.detector_geometry = copy(detector_geometry) det_geo = self.detector_geometry if hasattr(det_geo, 'distance'): - det_geo['distance'] = t.Tensor(det_geo['distance']) + det_geo['distance'] = t.tensor(det_geo['distance']) if hasattr(det_geo, 'basis'): - det_geo['basis'] = t.Tensor(det_geo['basis']) + det_geo['basis'] = t.tensor(det_geo['basis']) if hasattr(det_geo, 'corner'): - det_geo['corner'] = t.Tensor(det_geo['corner']) + det_geo['corner'] = t.tensor(det_geo['corner']) - - self.probe_basis = t.Tensor(probe_basis) - - scale_factor = t.Tensor([probe.shape[-1]/obj_guess.shape[-1], + self.probe_basis = t.tensor(probe_basis) + + scale_factor = t.tensor([probe.shape[-1]/obj_guess.shape[-1], probe.shape[-2]/obj_guess.shape[-2]]) self.obj_basis = self.probe_basis / scale_factor self.detector_slice = detector_slice # Maybe something to include in a bit - #self.surface_normal = t.Tensor(surface_normal) + # self.surface_normal = t.tensor(surface_normal) self.saturation = saturation if mask is None: self.mask = mask else: - self.mask = t.BoolTensor(mask) + self.mask = t.tensor(mask, dtype=t.bool) - self.probe = probe.to(t.complex64) + self.probe = t.tensor(probe, dtype=t.complex64) if obj_guess.dim() == 2: - obj_guess = obj_guess[None,:,:] - + obj_guess = obj_guess[None, :, :] - self.obj_real = t.nn.Parameter(obj_guess.real.to(t.float32)) - self.obj_imag = t.nn.Parameter(obj_guess.imag.to(t.float32)) + obj_guess = t.tensor(obj_guess, dtype=t.complex64) + + self.obj_real = t.nn.Parameter(obj_guess.real) + self.obj_imag = t.nn.Parameter(obj_guess.imag) # Wait for LBFGS to be updated for complex-valued parameters - #self.obj = t.nn.Parameter(obj_guess.to(t.float32)) - + # self.obj = t.nn.Parameter(obj_guess.to(t.float32)) + if background is None: if detector_slice is not None: - background = 1e-6 * t.ones(self.probe[0][self.detector_slice].shape[:-1]) + background = 1e-6 * t.ones( + self.probe[0][self.detector_slice].shape, + dtype=t.float32) else: - background = 1e-6 * t.ones(self.probe[0].shape[:-1]) + background = 1e-6 * t.ones(self.probe[0].shape, + dtype=t.float32) - self.background = t.Tensor(background).to(t.float32) + self.background = t.tensor(background, dtype=t.float32) if obj_support is not None: self.obj_support = obj_support - self.obj.data = self.obj * obj_support[None,...] + self.obj.data = self.obj * obj_support[None, ...] else: - self.obj_support = t.ones_like(self.obj[0,...]) + self.obj_support = t.ones_like(self.obj[0, ...]) self.oversampling = oversampling diff --git a/CDTools/models/simple_ptycho.py b/CDTools/models/simple_ptycho.py index f306068..234bbe2 100644 --- a/CDTools/models/simple_ptycho.py +++ b/CDTools/models/simple_ptycho.py @@ -22,7 +22,7 @@ class SimplePtycho(CDIModel): surface_normal=np.array([0.,0.,1.]), mask=None): super(SimplePtycho,self).__init__() - self.wavelength = t.tensor([wavelength]) + self.wavelength = t.tensor(wavelength) self.detector_geometry = copy(detector_geometry) det_geo = self.detector_geometry if hasattr(det_geo, 'distance'): @@ -35,23 +35,24 @@ class SimplePtycho(CDIModel): self.min_translation = t.tensor(min_translation) self.probe_basis = t.tensor(probe_basis) - self.detector_slice = detector_slice + self.detector_slice = copy(detector_slice) self.surface_normal = t.tensor(surface_normal) if mask is None: self.mask = None else: - self.mask = t.BoolTensor(mask) + self.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 - self.probe_norm = t.max(t.abs(probe_guess.to(t.complex64))) - - self.probe = t.nn.Parameter(probe_guess.to(t.complex64) - / self.probe_norm) - self.obj = t.nn.Parameter(obj_guess.to(t.complex64)) + self.probe_norm = t.max(t.abs(probe_guess)) + self.probe = t.nn.Parameter(probe_guess / self.probe_norm) + self.obj = t.nn.Parameter(obj_guess) @classmethod