diff --git a/CDTools/datasets/base.py b/CDTools/datasets/base.py index 469a978..f4d3067 100644 --- a/CDTools/datasets/base.py +++ b/CDTools/datasets/base.py @@ -121,7 +121,7 @@ class CDataset(torchdata.Dataset): self.mask = self.mask.to(*args,**mask_kwargs) if self.background is not None: self.background = self.background.to(*args,**kwargs) - + def get_as(self, *args, **kwargs): """Sets the dataset to return data on the given device and dtype diff --git a/CDTools/models/base.py b/CDTools/models/base.py index 146422d..306abd9 100644 --- a/CDTools/models/base.py +++ b/CDTools/models/base.py @@ -156,7 +156,6 @@ class CDIModel(t.nn.Module): sim_patterns = self.forward(*inp) if hasattr(self, 'mask'): - loss = self.loss(pats,sim_patterns, mask=self.mask) else: loss = self.loss(pats,sim_patterns) @@ -266,7 +265,7 @@ class CDIModel(t.nn.Module): calculation_width=calculation_width) - def LBFGS_optimize(self, iterations, dataset, batch_size=None, + def LBFGS_optimize(self, iterations, dataset, lr=0.1,history_size=2, subset=None, regularization_factor=None, thread=True, calculation_width=10): @@ -276,6 +275,9 @@ class CDIModel(t.nn.Module): situations or geometries it can be shockingly efficient. Like all the other optimization routines, it is defined as a generator function which yields the average loss each epoch. + + Note: There is no batch size, because it is a usually a bad idea to use + LBFGS on anything but all the data at onece Parameters ---------- @@ -283,8 +285,6 @@ class CDIModel(t.nn.Module): How many epochs of the algorithm to run dataset : CDataset The dataset to reconstruct against - batch_size : int - Optional, the size of the minibatches to use lr : float Optional, the learning rate to use history_size : int @@ -303,12 +303,9 @@ class CDIModel(t.nn.Module): subset = [subset] dataset = torchdata.Subset(dataset, subset) - # Make a dataloader - if batch_size is not None: - data_loader = torchdata.DataLoader(dataset, batch_size=batch_size, - shuffle=True) - else: - data_loader = torchdata.DataLoader(dataset, batch_size=len(dataset)) + # Make a dataloader. This basically does nothing but load all the + # data at once + data_loader = torchdata.DataLoader(dataset, batch_size=len(dataset)) # Define the optimizer diff --git a/CDTools/models/multislice_2d_ptycho.py b/CDTools/models/multislice_2d_ptycho.py index e7958b4..7ef3297 100644 --- a/CDTools/models/multislice_2d_ptycho.py +++ b/CDTools/models/multislice_2d_ptycho.py @@ -25,8 +25,9 @@ class Multislice2DPtycho(CDIModel): weights = None, translation_scale = 1, saturation=None, #probe_support = None, probe_fourier_support=None, - obj_support=None, oversampling=1, - bandlimit=4/5): + oversampling=1, + bandlimit=4/5, + subpixel=True): super(Multislice2DPtycho,self).__init__() self.wavelength = t.Tensor([wavelength]) @@ -48,6 +49,7 @@ class Multislice2DPtycho(CDIModel): self.surface_normal = t.Tensor(surface_normal) self.saturation = saturation + self.subpixel = subpixel if mask is None: self.mask = mask @@ -88,22 +90,12 @@ class Multislice2DPtycho(CDIModel): self.translation_scale = translation_scale self.probe_fourier_support = t.Tensor(probe_fourier_support).to(t.float32) + # In case real-space-support gets added back #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 - if self.obj.dim() == 3: - self.obj.data = self.obj * obj_support - elif self.obj.dim() == 4: - self.obj.data = self.obj * obj_support[None,...] - else: - if self.obj.dim() == 3: - self.obj_support = t.ones_like(self.obj) - elif self.obj.dim() == 4: - self.obj_support = t.ones_like(self.obj[0]) self.oversampling = oversampling spacing = np.linalg.norm(self.probe_basis,axis=0) @@ -115,7 +107,7 @@ class Multislice2DPtycho(CDIModel): @classmethod - def from_dataset(cls, dataset, dz, nz, probe_convergence_radius, probe_size=None, 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, bandlimit=4/5, replicate_slice=False): + def from_dataset(cls, dataset, dz, nz, probe_convergence_radius, probe_size=None, 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, bandlimit=4/5, replicate_slice=False, subpixel=True): wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] @@ -219,18 +211,6 @@ class Multislice2DPtycho(CDIModel): else: probe_support = None; - if restrict_obj != -1: - ro = restrict_obj - os = np.array(obj_size) - ps = np.array(probe_shape) - if replicate_slice: - obj_support = t.zeros_like(obj.to(dtype=t.float32)) - else: - obj_support = t.zeros_like(obj[0].to(dtype=t.float32)) - 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 probe_support = t.zeros_like(probe[0].to(dtype=t.float32)) @@ -252,9 +232,9 @@ class Multislice2DPtycho(CDIModel): saturation=saturation, #probe_support=probe_support, probe_fourier_support=probe_support, - obj_support=obj_support, oversampling=oversampling, - bandlimit=bandlimit) + bandlimit=bandlimit, + subpixel=subpixel) def interaction(self, index, translations): @@ -265,71 +245,52 @@ class Multislice2DPtycho(CDIModel): if self.translation_offsets is not None: pix_trans += self.translation_scale * self.translation_offsets[index] - if len(pix_trans.shape) == 1: - pix_trans = [pix_trans] - index = [index] - strip_first_index = True - else: - strip_first_index = False - # For a Fourier-space probe prs = tools.propagators.inverse_far_field(self.probe*self.probe_fourier_support[None,:,:]) # Here is where the mixing would happen, if it happened - all_exit_waves = [] - for i in range(self.probe.shape[0]): - pr = prs[i] - - #exit_waves = pr - #print(self.probe_norm) - #for i in range(self.nz): - - exit_waves = self.probe_norm * pr - for i in range(self.nz): - # If only one object slice - if self.obj.dim() == 3: - - #exit_wave = tools.interactions.ptycho_2D_sinc( - # exit_wave, - # self.obj_support*cmath.cexpi(self.obj/self.nz), - # pix_trans, shift_probe=True) - exit_waves = tools.interactions.ptycho_2D_round( - exit_waves, - self.obj_support*cmath.cexpi(self.obj/self.nz), - trans) - - elif self.obj.dim() == 4: - # If separate slices - #exit_wave = tools.interactions.ptycho_2D_sinc( - # exit_wave, - # self.obj_support*cmath.cexpi(self.obj[i]/self.nz), - # trans, shift_probe=True) - # I see, it needs to know that if exit_wave is the - # same shape as the translations, it should be broadcast - # along that dimension - exit_waves = tools.interactions.ptycho_2D_round( - exit_waves, - self.obj_support*cmath.cexpi(self.obj[i]/self.nz), - pix_trans) - - exit_waves = tools.propagators.near_field( - exit_waves,self.as_prop) - - - if exit_waves.dim() == 4: - # If the index is a list and not a single index - exit_waves = self.weights[index][:,None,None,None] * exit_waves + exit_waves = self.probe_norm * prs + for i in range(self.nz): + # If only one object slice + if self.obj.dim() == 3: + if i == 0 and self.subpixel: + # We only need to apply the subpixel shift to the first + # slice, because it shifts the probe + exit_waves = tools.interactions.ptycho_2D_sinc( + exit_waves, cmath.cexpi(self.obj/self.nz), + pix_trans, shift_probe=True, + multiple_modes=True) else: - # If the index a single index - exit_waves = self.weights[index] * exit_waves + exit_waves = tools.interactions.ptycho_2D_round( + exit_waves,cmath.cexpi(self.obj/self.nz), + pix_trans, multiple_modes=True) + + elif self.obj.dim() == 4: + # If separate slices + if i == 0 and self.subpixel: + exit_waves = tools.interactions.ptycho_2D_sinc( + exit_waves, cmath.cexpi(self.obj[i]/self.nz), + pix_trans, shift_probe=True, + multiple_modes=True) + else: + exit_waves = tools.interactions.ptycho_2D_round( + exit_waves, cmath.cexpi(self.obj[i]/self.nz), + pix_trans, multiple_modes=True) - if strip_first_index: - exit_waves = exit_waves[0,...] + exit_waves = tools.propagators.near_field( + exit_waves,self.as_prop) + + + if exit_waves.dim() == 5: + # If the index is a list and not a single index + exit_waves = self.weights[index][...,None,None,None,None] * exit_waves + else: + # If the index a single index + exit_waves = self.weights[index] * exit_waves - all_exit_waves.append(exit_waves) - return t.stack(all_exit_waves) + return exit_waves def forward_propagator(self, wavefields): @@ -375,7 +336,7 @@ class Multislice2DPtycho(CDIModel): self.probe_norm = self.probe_norm.to(*args,**kwargs) #self.probe_support = self.probe_support.to(*args,**kwargs) self.probe_fourier_support = self.probe_fourier_support.to(*args,**kwargs) - self.obj_support = self.obj_support.to(*args,**kwargs) + self.surface_normal = self.surface_normal.to(*args, **kwargs) self.as_prop = self.as_prop.to(*args, **kwargs) diff --git a/CDTools/tools/interactions/interactions.py b/CDTools/tools/interactions/interactions.py index c68e947..9d6a882 100644 --- a/CDTools/tools/interactions/interactions.py +++ b/CDTools/tools/interactions/interactions.py @@ -226,7 +226,7 @@ def project_translations_to_sample(sample_basis, translations): -def ptycho_2D_round(probe, obj, translations): +def ptycho_2D_round(probe, obj, translations, multiple_modes=False): """Returns a stack of exit waves without accounting for subpixel shifts This function returns a collection of exit waves, with the first @@ -234,15 +234,25 @@ def ptycho_2D_round(probe, obj, translations): corresponding to the detector. The exit waves are calculated by shifting the probe by the rounded value of the translation + If multiple_modes is set to False, any additional dimensions in the + ptycho_2D_round function will be assumed to correspond to the translation + index. If multiple_modes is set to true, the (-4th) dimension of the probe + will always be assumed to be defining a set of (P) incoherently mixing + modes to be broadcast all translation indices. If any additional dimensions + closer to the start exist, they will be assumed to be translation indices + + Parameters ---------- probe : torch.Tensor - A (P)xMxL probe function for the exit waves + A (P)xMxLx2 probe function to illuminate the object object : torch.Tensor The object function to be probed translations : torch.Tensor The (N)x2 array of (i,j) translations to simulate - + multuple_modes : bool + Default False, whether to assume the probe contains multiple modes + Returns ------- exit_waves : torch.Tensor @@ -260,11 +270,17 @@ def ptycho_2D_round(probe, obj, translations): tr[1]:tr[1]+probe.shape[-2]] for tr in integer_translations]) - if single_translation: - return cmult(probe,selection)[0] + if multiple_modes: + # if the probe dimension is 4, then this hasn't yet been broadcast + # over the translation dimensions + output = cmult(probe,selections[:,None,:,:,:]) else: - return cmult(probe,selections) + output = cmult(probe,selections) + if single_translation: + return output[0] + else: + return output @@ -363,9 +379,7 @@ def ptycho_2D_linear(probe, obj, translations, shift_probe=True): return t.stack(exit_waves) - - -def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10): +def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multiple_modes=True): """Returns a stack of exit waves accounting for subpixel shifts This function returns a collection of exit waves, with the first @@ -376,25 +390,32 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10): in Fourier space) If shift_probe is True, it applies the subpixel shift to the probe, - otherwise the subpixel shift is applied to the object + otherwise the subpixel shift is applied to the object [not yet implemented] + + If multiple_modes is set to False, any additional dimensions in the + ptycho_2D_round function will be assumed to correspond to the translation + index. If multiple_modes is set to true, the (-4th) dimension of the probe + will always be assumed to be defining a set of (P) incoherently mixing + modes to be broadcast all translation indices. If any additional dimensions + closer to the start exist, they will be assumed to be translation indices Parameters ---------- probe : torch.Tensor - An MxL probe function for the exit waves + An (P)xMxLx2 probe function for the exit waves object : torch.Tensor The object function to be probed translations : torch.Tensor - The Nx2 array of translations to simulate + The (N)x2 array of translations to simulate shift_probe : bool Default True, Whether to subpixel shift the probe or object - padding : int - Default 10, if shifting the object, the padding to apply to the object to avoid circular shift effects + multuple_modes : bool + Default False, whether to assume the probe contains multiple modes Returns ------- exit_waves : torch.Tensor - An NxMxL tensor of the calculated exit waves + An (N)x(P)xMxLx2 tensor of the calculated exit waves """ single_translation = False if translations.dim() == 1: @@ -406,35 +427,47 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10): integer_translations = t.floor(translations) subpixel_translations = translations - integer_translations integer_translations = integer_translations.to(dtype=t.int32) + + selections = t.stack([obj[tr[0]:tr[0]+probe.shape[-3], + tr[1]:tr[1]+probe.shape[-2]] + for tr in integer_translations]) exit_waves = [] if shift_probe: - i = t.arange(probe.shape[0]) - probe.shape[0]//2 - j = t.arange(probe.shape[1]) - probe.shape[1]//2 + i = t.arange(probe.shape[-3],device=probe.device,dtype=probe.dtype) \ + - probe.shape[-3]//2 + j = t.arange(probe.shape[-2],device=probe.device,dtype=probe.dtype) \ + - probe.shape[-2]//2 I,J = t.meshgrid(i,j) - I = 2 * np.pi * I.to(t.float32) / probe.shape[0] - J = 2 * np.pi * J.to(t.float32) / probe.shape[1] - I = I.to(dtype=probe.dtype,device=probe.device) - J = J.to(dtype=probe.dtype,device=probe.device) - - for tr, sp in zip(integer_translations, - subpixel_translations): - fft_probe = fftshift(t.fft(probe, 2)) - shifted_fft_probe = cmult(fft_probe, expi(-sp[0]*I - sp[1]*J)) - shifted_probe = t.ifft(ifftshift(shifted_fft_probe),2) + I = 2 * np.pi * I / probe.shape[-3] + J = 2 * np.pi * J / probe.shape[-2] - obj_slice = obj[tr[0]:tr[0]+probe.shape[0], - tr[1]:tr[1]+probe.shape[1]] + phase_masks = expi(-subpixel_translations[:,0,None,None]*I + -subpixel_translations[:,1,None,None]*J) + fft_probe = fftshift(t.fft(probe, 2)) + if multiple_modes: + # if the probe dimension is 4, then this hasn't yet been broadcast + # over the translation dimensions + shifted_fft_probe = cmult(fft_probe,phase_masks[:,None,:,:,:]) + else: + shifted_fft_probe = cmult(fft_probe,phase_masks) - exit_waves.append(cmult(shifted_probe, obj_slice)) + shifted_probe = t.ifft(ifftshift(shifted_fft_probe),2) + + if multiple_modes: + # if the probe dimension is 4, then this hasn't yet been broadcast + # over the translation dimensions + output = cmult(shifted_probe,selections[:,None,:,:,:]) + else: + output = cmult(shifted_probe,selections) else: raise NotImplementedError('Object shift not yet implemented') if single_translation: - return exit_waves[0] + return output[0] else: - return t.stack(exit_waves) + return output def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, padding=10): diff --git a/CDTools/tools/measurements/measurements.py b/CDTools/tools/measurements/measurements.py index 9d3fba4..9871341 100644 --- a/CDTools/tools/measurements/measurements.py +++ b/CDTools/tools/measurements/measurements.py @@ -169,10 +169,10 @@ def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=Non the wavefields. If a detector slice is given, the returned array will only include that slice from the simulated wavefronts. - The first index is the set of incoherently adding patterns, and - the second index is the index of the diffraction pattern to measure. - The next two indices index the wavefield. The final index is the complex - index. + The (-4th) index is the set of incoherently adding patterns, and any + indexes further to the front correspond to the set of diffraction patterns + to meaasure. The (-3rd) and (-2nd) indices are the wavefield, and the final + index is the complex index Parameters ---------- @@ -188,10 +188,10 @@ def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=Non Returns ------- sim_patterns : torch.Tensor - A real JXMxN array storing the incoherently summed intensities + A real LXMxN array storing the incoherently summed intensities """ - output = t.sum(cmath.cabssq(wavefields),dim=0) + output = t.sum(cmath.cabssq(wavefields),dim=-3) # Now we apply oversampling if oversampling != 1: