diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py index e0bb8ae..6bbf2bf 100644 --- a/CDTools/models/fancy_ptycho.py +++ b/CDTools/models/fancy_ptycho.py @@ -24,18 +24,19 @@ class FancyPtycho(CDIModel): background=None, translation_offsets=None, mask=None, weights=None, translation_scale=1, saturation=None, probe_support=None, oversampling=1, - loss='amplitude mse', units='um'): + loss='amplitude mse', units='um', + simulate_probe_translation=False): 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']) - if hasattr(det_geo, 'basis'): - det_geo['basis'] = t.tensor(det_geo['basis']) - if hasattr(det_geo, 'corner'): - det_geo['corner'] = t.tensor(det_geo['corner']) + if 'distance' in det_geo: + det_geo['distance'] = t.tensor(det_geo['distance'], dtype=t.float32) + if 'basis' in det_geo: + det_geo['basis'] = t.tensor(det_geo['basis'], dtype=t.float32) + if 'corner' in det_geo: + det_geo['corner'] = t.tensor(det_geo['corner'], dtype=t.float32) self.min_translation = t.tensor(min_translation) @@ -105,6 +106,14 @@ class FancyPtycho(CDIModel): self.oversampling = oversampling + self.simulate_probe_translation = simulate_probe_translation + if simulate_probe_translation: + Is = t.arange(self.probe.shape[-2], dtype=t.float32) + Js = t.arange(self.probe.shape[-1], dtype=t.float32) + Is, Js = t.meshgrid(Is/t.max(Is), Js/t.max(Js)) + self.I_phase = 2 * np.pi* Is + self.J_phase = 2 * np.pi* Js + # Here we set the appropriate loss function if (loss.lower().strip() == 'amplitude mse' or loss.lower().strip() == 'amplitude_mse'): @@ -117,7 +126,7 @@ class FancyPtycho(CDIModel): @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, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um'): + 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', simulate_probe_translation=False): wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] @@ -244,7 +253,8 @@ class FancyPtycho(CDIModel): saturation=saturation, probe_support=probe_support, oversampling=oversampling, - loss=loss, units=units) + loss=loss, units=units, + simulate_probe_translation=simulate_probe_translation) def interaction(self, index, translations, *args): @@ -280,12 +290,29 @@ class FancyPtycho(CDIModel): # 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 = t.tensordot( + # translations, + # t.as_tensor(self.detector_geometry['basis'], + # dtype=t.float32), + # dims=1) + det_pix_trans = tools.interactions.translations_to_pixel( + self.detector_geometry['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,:,:] + + # 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( prs, self.obj, pix_trans, shift_probe=True, multiple_modes=True) - return exit_waves @@ -315,16 +342,20 @@ class FancyPtycho(CDIModel): self.wavelength = self.wavelength.to(*args, **kwargs) # move the detector geometry too det_geo = self.detector_geometry - if hasattr(det_geo, 'distance'): + if 'distance' in det_geo: det_geo['distance'] = det_geo['distance'].to(*args, **kwargs) - if hasattr(det_geo, 'basis'): + if 'basis' in det_geo: det_geo['basis'] = det_geo['basis'].to(*args, **kwargs) - if hasattr(det_geo, 'corner'): + if 'corner' in det_geo: det_geo['corner'] = det_geo['corner'].to(*args, **kwargs) if self.mask is not None: self.mask = self.mask.to(*args, **kwargs) + if self.simulate_probe_translation: + self.I_phase = self.I_phase.to(*args, **kwargs) + self.J_phase = self.J_phase.to(*args, **kwargs) + self.min_translation = self.min_translation.to(*args, **kwargs) self.probe_basis = self.probe_basis.to(*args, **kwargs) self.probe_norm = self.probe_norm.to(*args, **kwargs) @@ -359,7 +390,7 @@ class FancyPtycho(CDIModel): indices, translations = args_list # Then we simulate the results - data = self.forward(indices, translations) + data = self.forward(indices, translations).detach() # And finally, we make the dataset return Ptycho2DDataset( diff --git a/CDTools/models/rpi.py b/CDTools/models/rpi.py index ae0602c..6bd2e66 100644 --- a/CDTools/models/rpi.py +++ b/CDTools/models/rpi.py @@ -4,6 +4,7 @@ from CDTools import tools from CDTools.tools import plotting as p from CDTools.tools.interactions import RPI_interaction from CDTools.tools import initializers +from scipy.ndimage.morphology import binary_dilation import numpy as np from copy import copy @@ -49,11 +50,18 @@ class RPI(CDIModel): @property def obj(self): return t.complex(self.obj_real, self.obj_imag) + + @property + def weights(self): + ws = t.complex(self.weights_real, self.weights_imag) + return ws / 10# / self.obj_real.size().numel() + + def __init__(self, wavelength, detector_geometry, probe_basis, probe, obj_guess, detector_slice=None, background=None, mask=None, saturation=None, - obj_support=None, oversampling=1): + obj_support=None, oversampling=1, weight_matrix=False): super(RPI, self).__init__() @@ -72,7 +80,7 @@ class RPI(CDIModel): 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.obj_basis = self.probe_basis * scale_factor self.detector_slice = detector_slice # Maybe something to include in a bit @@ -84,10 +92,10 @@ class RPI(CDIModel): self.mask = mask else: self.mask = t.tensor(mask, dtype=t.bool) - + self.probe = t.tensor(probe, dtype=t.complex64) - + if obj_guess.dim() == 2: obj_guess = obj_guess[None, :, :] @@ -96,6 +104,13 @@ class RPI(CDIModel): self.obj_real = t.nn.Parameter(obj_guess.real) self.obj_imag = t.nn.Parameter(obj_guess.imag) + self.weights_real = t.nn.Parameter(t.eye(probe.shape[0])* 10)# * self.obj_real.size().numel()) + self.weights_imag = t.nn.Parameter(t.zeros(probe.shape[0])) + + if not weight_matrix: + self.weights_real.requires_grad=False + self.weights_imag.requires_grad=False + # Wait for LBFGS to be updated for complex-valued parameters # self.obj = t.nn.Parameter(obj_guess.to(t.float32)) @@ -120,7 +135,7 @@ class RPI(CDIModel): @classmethod - def from_dataset(cls, dataset, probe, obj_size=None, background=None, mask=None, padding=0, n_modes=1, saturation=None, scattering_mode=None, oversampling=1, auto_center=False, initialization='random', opt_for_fft=False): + def from_dataset(cls, dataset, probe, obj_size=None, background=None, mask=None, padding=0, n_modes=1, saturation=None, scattering_mode=None, oversampling=1, auto_center=False, initialization='random', opt_for_fft=False, weight_matrix=False, probe_threshold=0): wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] @@ -211,16 +226,24 @@ class RPI(CDIModel): else: raise KeyError('Initialization "' + str(initialization) + \ '" invalid - use "spectral" or "random"') - - - # Maybe put something here to initialize an object support based on - # a probe threshold? - obj_support=None + + probe_intensity = t.sqrt(t.sum(t.abs(probe)**2,axis=0)) + probe_fft = tools.propagators.far_field(probe_intensity) + pad0l = (probe.shape[-2] - obj_size[-2])//2 + pad0r = probe.shape[-2] - obj_size[-2] - pad0l + pad1l = (probe.shape[-1] - obj_size[-1])//2 + pad1r = probe.shape[-1] - obj_size[-1] - pad1l + probe_lr_fft = probe_fft[pad0l:-pad0r,pad1l:-pad1r] + probe_lr = t.abs(tools.propagators.inverse_far_field(probe_lr_fft)) + + obj_support = probe_lr > t.max(probe_lr) * probe_threshold + obj_support = t.as_tensor(binary_dilation(obj_support)) return cls(wavelength, det_geo, probe_basis, probe, obj_guess, detector_slice=det_slice, background=background, mask=mask, saturation=saturation, - obj_support=obj_support, oversampling=oversampling) + obj_support=obj_support, oversampling=oversampling, + weight_matrix=weight_matrix) def random_init(self, pattern): @@ -251,8 +274,12 @@ class RPI(CDIModel): all_exit_waves = [] + + # Mix the probes with the weight matrix + prs = t.sum(self.weights[..., None, None] * self.probe, axis=-3) + for i in range(self.probe.shape[0]): - pr = self.probe[i] + pr = prs[i] # Here we have a 3D probe (one single mode) # and a 4D object (multiple modes mixing incoherently) exit_waves = RPI_interaction(pr, @@ -275,7 +302,6 @@ class RPI(CDIModel): output = output.unsqueeze(0).repeat(1,len(index),1,1,1) except TypeError: pass - return output @@ -336,7 +362,7 @@ class RPI(CDIModel): plot_list = [ ('Root Sum Squared Amplitude of all Probes', lambda self, fig: p.plot_amplitude( - np.sqrt(np.sum((t.abs(self.probe)**2).cpu().numpy(),axis=0)), + np.sqrt(np.sum((t.abs(t.sum(self.weights[..., None, None].detach() * self.probe, axis=-3))**2).cpu().numpy(),axis=0)), fig=fig, basis=self.probe_basis)), ('Dominant Object Amplitude', lambda self, fig: p.plot_amplitude(self.obj[0], fig=fig,