diff --git a/CDTools/models/__init__.py b/CDTools/models/__init__.py index 7befb04..4e7a76e 100644 --- a/CDTools/models/__init__.py +++ b/CDTools/models/__init__.py @@ -123,19 +123,44 @@ class CDIModel(t.nn.Module): yield loss - def Adam_optimize(self, iterations, dataset, batch_size=15, lr=0.005): + def Adam_optimize(self, iterations, dataset, batch_size=15, lr=0.005, schedule=False): # Make a dataloader data_loader = torchdata.DataLoader(dataset, batch_size=batch_size, shuffle=True) - # Define the optimizer optimizer = t.optim.Adam(self.parameters(), lr = lr) + + + # Define the scheduler + if schedule: + scheduler = t.optim.ReduceLROnPlateau(optimizer, factor=0.2) + else: + scheduler = None + + return self.AD_optimize(iterations, data_loader, optimizer, scheduler=scheduler) + + + def LBFGS_optimize(self, iterations, dataset, batch_size=None, + lr=0.1,history_size=2): + + # 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) + + + # Define the optimizer + optimizer = t.optim.LBFGS(self.parameters(), + lr = lr, history_size=history_size) return self.AD_optimize(iterations, data_loader, optimizer) - + from CDTools.models.simple_ptycho import SimplePtycho from CDTools.models.fancy_ptycho import FancyPtycho +from CDTools.models.incoherent_ptycho import IncoherentPtycho diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py new file mode 100644 index 0000000..87664c0 --- /dev/null +++ b/CDTools/models/fancy_ptycho.py @@ -0,0 +1,227 @@ +from __future__ import division, print_function, absolute_import + +import torch as t +from CDTools.models import CDIModel +from CDTools import tools +from CDTools.tools import cmath +import numpy as np +from copy import copy + + +class FancyPtycho(CDIModel): + + def __init__(self, wavelength, detector_geometry, + probe_basis, detector_slice, + probe_guess, obj_guess, min_translation = t.Tensor([0,0]), + background = None, translation_offsets=None, mask=None, + weights = None, translation_scale = 1, saturation=None, + probe_support = None): + + 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']) + + self.min_translation = t.Tensor(min_translation) + + self.probe_basis = t.Tensor(probe_basis) + self.detector_slice = detector_slice + + self.saturation = saturation + + if mask is None: + self.mask = mask + else: + self.mask = t.ByteTensor(mask) + + # We rescale the probe here so it learns at the same rate as the + # object + if probe_guess.dim() > 3: + self.probe_norm = t.max(tools.cmath.cabs(probe_guess[0].to(t.float32))) + else: + self.probe_norm = t.max(tools.cmath.cabs(probe_guess.to(t.float32))) + + self.probe = t.nn.Parameter(probe_guess.to(t.float32) + / self.probe_norm) + + self.obj = t.nn.Parameter(obj_guess.to(t.float32)) + + if background is None: + background = 1e-6 * t.ones(self.probe[0][self.detector_slice].shape[:-1]) + + + self.background = t.nn.Parameter(t.Tensor(background).to(t.float32)) + + if weights is None: + self.weights = None + else: + self.weights = t.nn.Parameter(t.Tensor(weights).to(t.float32)) + + 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) + + 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]) + + + @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): + 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') + (indices, translations), patterns = dataset[:] + dataset.get_as(*get_as_args[0],**get_as_args[1]) + + # Set to none to avoid issues with things outside the detector + center = tools.image_processing.centroid(t.sum(patterns,dim=0)) + + # 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) + + # Next generate the object geometry from the probe geometry and + # the translations + pix_translations = tools.interactions.translations_to_pixel(probe_basis, translations) + + obj_size, min_translation = tools.initializers.calc_object_setup(probe_shape, pix_translations, padding=50) + + # Finally, initialize the probe and object using this information + if probe_size is None: + probe = tools.initializers.SHARP_style_probe(dataset, probe_shape, det_slice) + else: + probe = tools.initializers.gaussian_probe(dataset, probe_basis, probe_shape, probe_size) + + + # Now we initialize all the subdominant probe modes + probe_max = t.max(cmath.cabs(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) + + obj = tools.cmath.expi(randomize_ang * (t.rand(obj_size)-0.5)) + + det_geo = dataset.detector_geometry + + translation_offsets = 0 * (t.rand((len(dataset),2)) - 0.5) + + weights = t.ones(len(dataset)) + + if hasattr(dataset, 'mask') and dataset.mask is not None: + mask = dataset.mask.to(t.uint8) + else: + mask = None + + if probe_support_radius is not None: + probe_support = t.zeros_like(probe[0].to(dtype=t.float32)) + p_cent = np.array(probe.shape[1:3]).astype(int) // 2 + psr = int(probe_support_radius) + probe_support[p_cent[0]-psr:p_cent[0]+psr, + p_cent[1]-psr:p_cent[1]+psr] = 1 + else: + probe_support = t.ones_like(probe[0].to(dtype=t.float32)) + + + + #return cls(wavelength, det_geo, probe_basis, det_slice, probe, obj, min_translation=min_translation, translation_offsets = translation_offsets) + return cls(wavelength, det_geo, probe_basis, det_slice, probe, obj, min_translation=min_translation, translation_offsets = translation_offsets, weights=weights, mask=mask, translation_scale=translation_scale, saturation=saturation, probe_support=probe_support) + + + def interaction(self, index, translations): + pix_trans = tools.interactions.translations_to_pixel(self.probe_basis, + translations) + pix_trans -= self.min_translation + + if self.translation_offsets is not None: + pix_trans += self.translation_scale * self.translation_offsets[index] + + all_exit_waves = [] + for i in range(self.probe.shape[0]): + pr = self.probe[i] * self.probe_support + #exit_waves = self.probe_norm * tools.interactions.ptycho_2D_round(self.probe[i], + # self.obj, + # pix_trans) + exit_waves = self.probe_norm * tools.interactions.ptycho_2D_sinc(pr, + self.obj, + pix_trans, + shift_probe=True) + exit_waves = exit_waves * self.probe_support[...,:,:] + + + if exit_waves.dim() == 4: + exit_waves = self.weights[index][:,None,None,None] * exit_waves + else: + exit_waves = self.weights[index] * exit_waves + + all_exit_waves.append(exit_waves) + + return t.stack(all_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, + detector_slice=self.detector_slice, + measurement=tools.measurements.incoherent_sum, + saturation=self.saturation ) + + + def loss(self, sim_data, real_data, mask=None): + return tools.losses.amplitude_mse(real_data, sim_data, mask=mask) + + + def to(self, *args, **kwargs): + super(FancyPtycho, self).to(*args, **kwargs) + self.wavelength = self.wavelength.to(*args,**kwargs) + # move the detector geometry too + det_geo = self.detector_geometry + if hasattr(det_geo, 'distance'): + det_geo['distance'] = det_geo['distance'].to(*args,**kwargs) + if hasattr(det_geo, 'basis'): + det_geo['basis'] = det_geo['basis'].to(*args,**kwargs) + if hasattr(det_geo, 'corner'): + det_geo['corner'] = det_geo['corner'].to(*args,**kwargs) + + if self.mask is not None: + self.mask = self.mask.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) + self.probe_support = self.probe_support.to(*args,**kwargs) + + + def sim_to_dataset(self, args_list): + pass + + diff --git a/CDTools/models/incoherent_ptycho.py b/CDTools/models/incoherent_ptycho.py new file mode 100644 index 0000000..2d405cc --- /dev/null +++ b/CDTools/models/incoherent_ptycho.py @@ -0,0 +1,188 @@ +from __future__ import division, print_function, absolute_import + +import torch as t +from CDTools.models import CDIModel +from CDTools import tools +from copy import copy +import numpy as np + + +class IncoherentPtycho(CDIModel): + + def __init__(self, wavelength, detector_geometry, + probe_basis, detector_slice, + probe_guess, obj_guess, min_translation = t.Tensor([0,0]), + translation_offsets=None, + background = None, mask=None, weights = None): + + super(IncoherentPtycho,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']) + + self.min_translation = t.Tensor(min_translation) + + self.probe_basis = t.Tensor(probe_basis) + self.detector_slice = detector_slice + + + if mask is None: + self.mask = mask + else: + self.mask = t.ByteTensor(mask) + + # We rescale the probe here so it learns at the same rate as the + # object + probe_norm = t.max(tools.cmath.cabs(probe_guess[0].to(t.float32))) + + self.probe = t.nn.Parameter(probe_guess.to(t.float32)/probe_norm) + + self.probe_norm = float(probe_norm.numpy()) + + self.obj = t.nn.Parameter(obj_guess.to(t.float32)) + + if background is None: + background = 1e-6 * t.ones(self.probe[(np.s_[0],)+self.detector_slice].shape[:-1]) + self.background = t.nn.Parameter(t.Tensor(background).to(t.float32)) + + if weights is None: + self.weights = None + else: + self.weights = t.nn.Parameter(t.Tensor(weights).to(t.float32)) + + if translation_offsets is None: + self.translation_offsets = None + else: + self.translation_offsets = t.nn.Parameter(t.Tensor(translation_offsets).to(t.float32)) + + + + @classmethod + def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0): + 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') + (indices, translations), patterns = dataset[:] + dataset.get_as(*get_as_args[0],**get_as_args[1]) + + # Set to none to avoid issues with things outside the detector + center = tools.image_processing.centroid(t.sum(patterns,dim=0)) + + # 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) + + # Next generate the object geometry from the probe geometry and + # the translations + pix_translations = tools.interactions.translations_to_pixel(probe_basis, translations) + + obj_size, min_translation = tools.initializers.calc_object_setup(probe_shape, pix_translations, padding=20) + + # Finally, initialize the probe and object using this information + if probe_size is None: + probe = tools.initializers.SHARP_style_probe(dataset, probe_shape, det_slice) + else: + probe = tools.initializers.gaussian_probe(dataset, probe_basis, probe_shape, probe_size) + + translation_offsets = 0 * (t.rand((len(dataset),2)) - 0.5) + + # For incoherent probe mixing + probe = t.stack((probe,0.05*t.rand(probe.shape).to(probe.dtype))) + + obj = tools.cmath.expi(randomize_ang * (t.rand(obj_size)-0.5)) + + det_geo = dataset.detector_geometry + + weights = t.ones(len(dataset)) + + if hasattr(dataset, 'mask') and dataset.mask is not None: + mask = dataset.mask.to(t.uint8) + else: + mask = None + + return cls(wavelength, det_geo, probe_basis, det_slice, probe, obj, min_translation=min_translation, translation_offsets=translation_offsets, weights=weights, mask=mask) + + + def interaction(self, index, translations): + pix_trans = tools.interactions.translations_to_pixel(self.probe_basis, + translations) + # The 10x term is to condition the translation offsets + pix_trans -= self.min_translation + pix_trans = pix_trans + self.translation_offsets[index] + + all_exit_waves = [] + for i in range(self.probe.shape[0]): + exit_waves = self.probe_norm * tools.interactions.ptycho_2D_sinc(self.probe[i], + self.obj, + pix_trans, + shift_probe=True) + + if exit_waves.dim() == 4: + exit_waves = self.weights[index][:,None,None,None] * exit_waves + else: + exit_waves = self.weights[index] * exit_waves + + all_exit_waves.append(exit_waves) + + return t.stack(all_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, + detector_slice=self.detector_slice, + measurement=tools.measurements.incoherent_sum) + + + def loss(self, sim_data, real_data, mask=None): + return tools.losses.amplitude_mse(real_data, sim_data, mask=mask) + + + def to(self, *args, **kwargs): + super(IncoherentPtycho, self).to(*args, **kwargs) + self.wavelength = self.wavelength.to(*args,**kwargs) + # move the detector geometry too + det_geo = self.detector_geometry + if hasattr(det_geo, 'distance'): + det_geo['distance'] = det_geo['distance'].to(*args,**kwargs) + if hasattr(det_geo, 'basis'): + det_geo['basis'] = det_geo['basis'].to(*args,**kwargs) + if hasattr(det_geo, 'corner'): + det_geo['corner'] = det_geo['corner'].to(*args,**kwargs) + + if self.mask is not None: + self.mask = self.mask.to(*args, **kwargs) + + + self.min_translation = self.min_translation.to(*args,**kwargs) + self.probe_basis = self.probe_basis.to(*args,**kwargs) + + + def sim_to_dataset(self, args_list): + pass + + diff --git a/CDTools/scripts/synthesize.py b/CDTools/scripts/synthesize.py new file mode 100644 index 0000000..03e87dd --- /dev/null +++ b/CDTools/scripts/synthesize.py @@ -0,0 +1,178 @@ +from __future__ import division, print_function, absolute_import +import numpy as np +import torch as t +from matplotlib import pyplot as plt +import pickle +import argparse +from scipy import fftpack + +from CDTools.tools import cmath, plotting +from CDTools.tools import image_processing as ip + + +def standardize(probe, obj, obj_slice=None): + # First, we normalize the probe intensity to a fixed value. + # Should this be the maximum or the integrated intensity? I think + # probably the integrated intensity. We set the average per-[ixel + # intensity in the probe to be one + normalization = np.sqrt(np.sum(np.abs(probe)**2) / len(probe.ravel())) + probe = cmath.complex_to_torch(probe / normalization) + obj = cmath.complex_to_torch(obj * normalization) + + + # Default slice of the object to use for alignment, etc. + if obj_slice is None: + obj_slice = np.s_[(obj.shape[0]//8)*3:(obj.shape[0]//8)*5, + (obj.shape[1]//8)*3:(obj.shape[1]//8)*5] + + + + # Now we get rid of the probe's phase ramp + # Currently disabled + #center_freq = ip.centroid_sq(cmath.fftshift(t.fft(probe,2)),comp=True) + #center_freq -= (t.tensor(probe.shape[:-1]) // 2).to(t.float32) + #center_freq /= t.tensor(probe.shape[:-1]).to(t.float32) + + + + #Is, Js = np.mgrid[:probe.shape[0],:probe.shape[1]] + #probe_phase_ramp = cmath.expi(2*np.pi * + # (center_freq[0] * t.tensor(Is).to(t.float32) + + # center_freq[1] * t.tensor(Js).to(t.float32))) + #probe = cmath.cmult(probe, cmath.cconj(probe_phase_ramp)) + #Is, Js = np.mgrid[:obj.shape[0],:obj.shape[1]] + #obj_phase_ramp = cmath.expi(2*np.pi * + # (center_freq[0] * t.tensor(Is).to(t.float32) + + # center_freq[1] * t.tensor(Js).to(t.float32))) + #obj = cmath.cmult(obj, obj_phase_ramp) + + + # Then, we set them to consistent absolute phases + probe_angle = cmath.cphase(t.sum(probe,dim=(0,1))) + obj_angle = cmath.cphase(t.sum(obj[obj_slice],dim=(0,1))) + + probe = cmath.cmult(probe, cmath.expi(-probe_angle)) + obj = cmath.cmult(obj, cmath.expi(-obj_angle)) + + return probe, obj + + +def synthesize_reconstructions(probes, objects, use_probe=False, obj_slice=None): + + if obj_slice is None: + obj_slice = np.s_[(objects[0].shape[0]//8)*3:(objects[0].shape[0]//8)*5, + (objects[0].shape[1]//8)*3:(objects[0].shape[1]//8)*5] + + + synth_probe, synth_obj = standardize(probes[0], objects[0]) + obj_stack = [cmath.torch_to_complex(synth_obj)] + for i, (probe, obj) in enumerate(zip(probes[1:],objects[1:])): + probe, obj = standardize(probe, obj) + probe = probe[0] + print(i) + #plt.imshow(np.angle(cmath.torch_to_complex(obj[obj_slice]))) + #plt.show() + + if use_probe: + shift = ip.find_shift(synth_probe,probe, resolution=50) + else: + shift = ip.find_shift(synth_obj[obj_slice],obj[obj_slice], resolution=50) + + + obj = ip.sinc_subpixel_shift(obj,np.array(shift)) + probe = ip.sinc_subpixel_shift(probe,tuple(shift)) + #obj = t.roll(obj,tuple(int(s) for s in shift),dims=(0,1)) + #probe = t.roll(probe,tuple(int(s) for s in shift),dims=(0,1)) + + synth_probe += probe + synth_obj += obj + obj_stack.append(cmath.torch_to_complex(obj)) + + + + # If there only was one image + try: + i + except: + i = -1 + + synth_probe = cmath.torch_to_complex(synth_probe) + synth_obj = cmath.torch_to_complex(synth_obj) + return synth_probe/(i+2), synth_obj/(i+2), obj_stack + + + +def calc_prtf(synth_obj, objects, basis, obj_slice=None): + if obj_slice is None: + obj_slice = np.s_[(objects[0].shape[0]//8)*3:(objects[0].shape[0]//8)*5, + (objects[0].shape[1]//8)*3:(objects[0].shape[1]//8)*5] + + synth_obj = cmath.complex_to_torch(synth_obj[obj_slice]) + + synth_fft = cmath.cabssq(cmath.fftshift(t.fft(synth_obj,2))).numpy() + + prtfs = [] + for obj in objects: + obj = cmath.complex_to_torch(obj[obj_slice]) + single_fft = cmath.cabssq(cmath.fftshift(t.fft(obj,2))).numpy() + + + di = np.linalg.norm(basis[:,0]) + dj = np.linalg.norm(basis[:,1]) + + i_freqs = fftpack.fftshift(fftpack.fftfreq(single_fft.shape[0],d=di)) + j_freqs = fftpack.fftshift(fftpack.fftfreq(single_fft.shape[1],d=dj)) + + Js,Is = np.meshgrid(j_freqs,i_freqs) + Is = Is - np.mean(Is) + Js = Js - np.mean(Js) + Rs = np.sqrt(Is**2+Js**2) + + single_ints, bins = np.histogram(Rs,bins=100,weights=single_fft) + synth_ints, bins = np.histogram(Rs,bins=100,weights=synth_fft) + + prtfs.append(synth_ints/single_ints) + + return bins[:-1], np.mean(prtfs,axis=0) + + + +def make_argparser(): + parser = argparse.ArgumentParser(description=__doc__) + + parser.add_argument('file', help='The reconstruction file to calculate metrics for') + parser.add_argument('--use-probe', '-up', action='store_true', help='Use the probe instead of the object to align the reconstructions') + return parser + + +if __name__ == '__main__': + + args = make_argparser().parse_args() + + with open(args.file, 'rb') as f: + dataset = pickle.load(f) + + synth_probe, synth_obj, aligned_objs = synthesize_reconstructions( + dataset['probe'], dataset['obj'], args.use_probe) + + freqs, prtf = calc_prtf(synth_obj, aligned_objs, dataset['basis']) + + print(np.linalg.norm(dataset['basis'],axis=0)) + plotting.plot_phase(dataset['probe'][0][0],basis=1e6*dataset['basis']) + plotting.plot_amplitude(dataset['probe'][0][0],basis=1e6*dataset['basis']) + plotting.plot_colorized(dataset['probe'][0][0],basis=1e6*dataset['basis']) + #plotting.plot_amplitude(synth_obj[400:750,450:850],basis=1e6*dataset['basis']) + #plotting.plot_colorized(synth_obj[400:750,450:850],basis=1e6*dataset['basis']) + #plotting.plot_phase(synth_obj[400:750,450:850],basis=1e6*dataset['basis']) + plotting.plot_amplitude(synth_obj[::-1,::-1][450:900,325:775],basis=1e6*dataset['basis']) + plotting.plot_phase(synth_obj[::-1,::-1][450:900,325:775],basis=1e6*dataset['basis']) + plotting.plot_colorized(synth_obj[::-1,::-1][450:900,325:775],basis=1e6*dataset['basis']) + + plt.figure() + real_translations = dataset['basis'].dot(dataset['translation'][0].transpose()) + real_translations -= np.min(real_translations,axis=1)[:,None] + plt.plot(real_translations[0]*1e6,real_translations[1]*1e6,'k.') + plt.plot(real_translations[0]*1e6,real_translations[1]*1e6,'b-',linewidth=0.5) + plt.figure() + plt.plot(freqs*1e-6, prtf) + plt.show() diff --git a/CDTools/tools/__init__.py b/CDTools/tools/__init__.py index dbc7ff8..ec25428 100644 --- a/CDTools/tools/__init__.py +++ b/CDTools/tools/__init__.py @@ -10,3 +10,4 @@ from CDTools.tools import projectors from CDTools.tools import interactions from CDTools.tools import propagators from CDTools.tools import measurements +from CDTools.tools import analysis diff --git a/CDTools/tools/analysis.py b/CDTools/tools/analysis.py new file mode 100644 index 0000000..9f87a92 --- /dev/null +++ b/CDTools/tools/analysis.py @@ -0,0 +1,56 @@ +from __future__ import division, print_function + +import torch as t +import numpy as np +from CDTools.tools import cmath + + + +def orthogonalize_probes(probes): + """Orthogonalizes a set of incoherently mixing probes + + The strategy is to define a reduced orthogonal basis that spans + all of the retrieved probes, and then build the density matrix + defined by the probes in that basis. After diagonalization, the + eigenvectors can be recast into the original basis and returned + + Args: + probes (t.Tensor) : n x (image) size tensor, a stack of probes + + Returns: + (t.Tensor) : n x (image) size tensor, a stack of probes + """ + + try: + probes = cmath.torch_to_complex(probes.detach().cpu()) + except: + pass + + bases = [] + coefficients = np.zeros((probes.shape[0],probes.shape[0]), dtype=np.complex64) + for i, probe in enumerate(probes): + ortho_probe = np.copy(probe) + for j, basis in enumerate(bases): + coefficients[j,i] = np.sum(basis.conj()*ortho_probe) + ortho_probe -= basis * coefficients[i,j] + + + coefficients[i,i] = np.sqrt(np.sum(np.abs(ortho_probe)**2)) + bases.append(ortho_probe / coefficients[i,i]) + + density_mat = np.conj(coefficients).transpose().dot(coefficients) + eigvals, eigvecs = np.linalg.eigh(density_mat) + + ortho_probes = [] + for i in range(len(eigvals)): + coefficients = np.sqrt(eigvals[i]) * eigvecs[:,i] + probe = np.zeros(bases[0].shape, dtype=np.complex64) + for coefficient, basis in zip(coefficients, bases): + probe += basis * coefficient + ortho_probes.append(probe) + + + return cmath.complex_to_torch(np.stack(ortho_probes)) + + + diff --git a/CDTools/tools/image_processing.py b/CDTools/tools/image_processing.py index 291cc8d..44f4eb8 100644 --- a/CDTools/tools/image_processing.py +++ b/CDTools/tools/image_processing.py @@ -57,6 +57,33 @@ def centroid_sq(im, dims=2, comp=False): return centroid(im_sq, dims=dims) +def sinc_subpixel_shift(im, shift): + """Performs a subpixel shift with sinc interpolation on the given tensor + + The subpixel shift is done circularly via a multiplication with a linear + phase mask in Fourier space. + + Args: + im (torch.Tensor) : A complex-valued tensor to perform the subpixel shift on + shift (array_like) : A length-2 array_like object describing the shift to perform, in pixels + + Returns: + (torch.Tensor) : The subpixel shifted tensor + """ + i = t.arange(im.shape[0]) - im.shape[0]//2 + j = t.arange(im.shape[1]) - im.shape[1]//2 + I,J = t.meshgrid(i,j) + I = 2 * np.pi * I.to(t.float32) / im.shape[0] + J = 2 * np.pi * J.to(t.float32) / im.shape[1] + I = I.to(dtype=im.dtype,device=im.device) + J = J.to(dtype=im.dtype,device=im.device) + + fft_im = cmath.fftshift(t.fft(im, 2)) + shifted_fft_im = cmath.cmult(fft_im, cmath.expi(-shift[0]*I - shift[1]*J)) + return t.ifft(cmath.ifftshift(shifted_fft_im),2) + + + def find_subpixel_shift(im1, im2, search_around=(0,0), resolution=10): """Calculates the subpixel shift between two images by maximizing the autocorrelation @@ -91,7 +118,8 @@ def find_subpixel_shift(im1, im2, search_around=(0,0), resolution=10): # Not sure if this is more or less stable than just the correlation # maximum - requires some testing cor = t.ifft(cor_fft / cmath.cabs(cor_fft)[:,:,None],2) - + + # Now, I need to shift the array to pull out a contiguous window # around the correlation maximum try: @@ -149,7 +177,7 @@ def find_pixel_shift(im1, im2): # Not sure if this is more or less stable than just the correlation # maximum - requires some testing cor = cmath.cabs(t.ifft(cor_fft / cmath.cabs(cor_fft)[:,:,None],2)) - #cor = cmath.cabs(t.ifft(cor_fft,2)) + sh = t.tensor(cor.shape).to(device=im1.device) cormax = t.tensor([t.argmax(cor) // sh[1], diff --git a/CDTools/tools/initializers.py b/CDTools/tools/initializers.py index 80cab00..d5ad876 100644 --- a/CDTools/tools/initializers.py +++ b/CDTools/tools/initializers.py @@ -56,6 +56,7 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None, if opt_for_fft: full_shape = t.Tensor([next_fast_len(dim) for dim in full_shape]).to(t.int32) + # Then, generate a slice that pops the actual detector from the full # detector shape full_center = full_shape // 2 @@ -64,9 +65,10 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None, int(full_center[1]-center[1]): int(full_center[1]-center[1]+det_shape[1])] + # Finally, generate the basis for the exit wave in real space # I believe this calculation is incorrect for non-rectangular - # detectors, because the real space basis shoud be related to the + # detectors, because the real space basis should be related to the # dual of the original basis. Leaving this for now since # non-rectangular detectors are not a pressing concern. basis_dirs = det_basis / t.norm(det_basis, dim=0) @@ -75,7 +77,8 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None, # Finally, convert the shape back to a torch.Size full_shape = t.Size([dim for dim in full_shape]) - + + return real_space_basis, full_shape, det_slice @@ -245,6 +248,10 @@ def SHARP_style_probe(dataset, shape, det_slice): # Now we remove the central pixel center = np.array(probe_guess.shape) // 2 + + + # I had to remove this because it put some intensity outside of + # the detector region that caused issues probe_guess[center[0], center[1]]=np.mean([ probe_guess[center[0]-1, center[1]], diff --git a/CDTools/tools/interactions.py b/CDTools/tools/interactions.py index 0e62433..c70d6dd 100644 --- a/CDTools/tools/interactions.py +++ b/CDTools/tools/interactions.py @@ -178,7 +178,7 @@ def ptycho_2D_linear(probe, obj, translations, shift_probe=True): -#TODO: Implement a sinc-interpolated shift using a fourier space shifting op + def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10): """Returns a stack of exit waves accounting for subpixel shifts diff --git a/CDTools/tools/measurements.py b/CDTools/tools/measurements.py index d5add69..1c83ab0 100644 --- a/CDTools/tools/measurements.py +++ b/CDTools/tools/measurements.py @@ -12,7 +12,7 @@ import numpy as np __all__ = ['intensity', 'incoherent sum', 'quadratic_background'] -def intensity(wavefield, detector_slice=None, epsilon=1e-7): +def intensity(wavefield, detector_slice=None, epsilon=1e-7, saturation=None): """Returns the intensity of a wavefield The intensity is defined as the magnitude squared of the @@ -22,20 +22,25 @@ def intensity(wavefield, detector_slice=None, epsilon=1e-7): Args: wavefield (torch.Tensor) : A JxMxNx2 stack of complex wavefields detector_slice (slice) : Optional, a slice or tuple of slices defining a section of the simulation to return + saturation (float) : Optional, a maximum saturation value to clamp the resulting intensities to Returns: torch.Tensor : A real MxN array storing the wavefield's intensities """ if detector_slice is None: - return cmath.cabssq(wavefield) + epsilon + output = cmath.cabssq(wavefield) + epsilon else: if wavefield.dim() == 3: - return cmath.cabssq(wavefield[detector_slice]) + epsilon + output = cmath.cabssq(wavefield[detector_slice]) + epsilon else: - return cmath.cabssq(wavefield[(np.s_[:],) + detector_slice]) + epsilon + output = cmath.cabssq(wavefield[(np.s_[:],) + detector_slice]) + epsilon + if saturation is None: + return output + else: + return t.clamp(output,0,saturation) -def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7): +def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=None): """Returns the incoherent sum of the intensities of the wavefields The intensity is defined as the sum of the magnitudes squared of @@ -50,22 +55,27 @@ def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7): Args: wavefields (torch.Tensor) : A JxLxMxNx2 stack of complex wavefields detector_slice (slice) : Optional, a slice or tuple of slices defining a section of the simulation to return + saturation (float) : Optional, a maximum saturation value to clamp the resulting intensities to Returns: torch.Tensor : A real MxN array storing the incoherently summed intensities """ # This syntax just adds an axis to the slice to preserve the J direction if detector_slice is None: - return t.sum(cmath.cabssq(wavefields),dim=-3) + epsilon + output = t.sum(cmath.cabssq(wavefields),dim=-3) + epsilon else: if wavefields.dim() == 4: - return t.sum(cmath.cabssq(wavefields[(np.s_[:],)+detector_slice]),dim=-3) + epsilon + output = t.sum(cmath.cabssq(wavefields[(np.s_[:],)+detector_slice]),dim=0) + epsilon else: - return t.sum(cmath.cabssq(wavefields[(np.s_[:],np.s_[:])+detector_slice]),dim=-3) + epsilon + output = t.sum(cmath.cabssq(wavefields[(np.s_[:],np.s_[:])+detector_slice]),dim=0) + epsilon + if saturation is None: + return output + else: + return t.clamp(output,0,saturation) -def quadratic_background(wavefield, background, detector_slice=None, measurement=intensity, epsilon=1e-7): +def quadratic_background(wavefield, background, detector_slice=None, measurement=intensity, epsilon=1e-7, saturation=None): """Returns the intensity of a wavefield plus a background The intensity is calculated via the given measurment function @@ -78,13 +88,18 @@ def quadratic_background(wavefield, background, detector_slice=None, measurement background (torch.Tensor) : An tensor storing the square root of the detector background detector_slice (slice) : Optional, a slice or tuple of slices defining a section of the simulation to return measurement (function) : Optional, the measurement function to use. The default is measurements.intensity + saturation (float) : Optional, a maximum saturation value to clamp the resulting intensities to Returns: torch.Tensor : A real MxN array storing the wavefield's intensities """ if detector_slice is None: - return measurement(wavefield, epsilon=epsilon) + background**2 + output = measurement(wavefield, epsilon=epsilon) + background**2 else: - return measurement(wavefield, detector_slice, epsilon=epsilon) \ + output = measurement(wavefield, detector_slice, epsilon=epsilon) \ + background**2 + if saturation is None: + return output + else: + return t.clamp(output,0,saturation) diff --git a/CDTools/tools/plotting.py b/CDTools/tools/plotting.py index 130c05d..c9fbb47 100644 --- a/CDTools/tools/plotting.py +++ b/CDTools/tools/plotting.py @@ -64,7 +64,7 @@ def plot_amplitude(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), ** if fig is None: fig = plt.figure() ax = fig.add_subplot(111, **kwargs) - basis_norm = np.linalg.norm(basis, axis = -1) + basis_norm = np.linalg.norm(basis, axis = 0) if isinstance(im, t.Tensor): absolute = cmath.cabs(im).detach().cpu().numpy() else: @@ -93,7 +93,7 @@ def plot_phase(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), **kwa phase = cmath.cphase(im).detach().cpu().numpy() else: phase = np.angle(im) - basis_norm = np.linalg.norm(basis, axis = -1) + basis_norm = np.linalg.norm(basis, axis = 0) try: plt.imshow(phase, cmap = 'twilight', extent = [0, phase.shape[-1]*basis_norm[1], 0, phase.shape[-2]*basis_norm[0]]) except: plt.imshow(phase, cmap = 'hsv', extent = [0, phase.shape[-1]*basis_norm[1], 0, phase.shape[-2]*basis_norm[0]]) plt.colorbar() @@ -119,7 +119,7 @@ def plot_colorized(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), * ax = fig.add_subplot(111, **kwargs) if isinstance(im, t.Tensor): im = cmath.torch_to_complex(im.detach().cpu()) - basis_norm = np.linalg.norm(basis, axis = -1) + basis_norm = np.linalg.norm(basis, axis = 0) colorized = colorize(im) plt.imshow(colorized, extent = [0, im.shape[-1]*basis_norm[1], 0, im.shape[-2]*basis_norm[0]]) return fig diff --git a/README.md b/README.md index 1d1ca85..843ccf7 100644 --- a/README.md +++ b/README.md @@ -1 +1,11 @@ # CDTools + +Outline: + +Description of toolbox + +Note about authors and availability + +Example usage + +Available methods \ No newline at end of file diff --git a/examples/gold_ball_ptycho.py b/examples/gold_ball_ptycho.py new file mode 100644 index 0000000..911e7c2 --- /dev/null +++ b/examples/gold_ball_ptycho.py @@ -0,0 +1,50 @@ +from __future__ import division, print_function, absolute_import + +import CDTools +from CDTools.tools.plotting import * +from CDTools.tools.cmath import * +from CDTools.tools import interactions +import h5py +import numpy as np +from matplotlib import pyplot as plt + +filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' +#filename = '/media/Data Bank/CSX_3_19/Processed_CXIs/115195_p.cxi' + +with h5py.File(filename,'r') as f: + dataset = CDTools.datasets.Ptycho_2D_Dataset.from_cxi(f) + darks = np.array(f['entry_1/instrument_1/detector_1/data_dark']) + +old_patterns = dataset.patterns.clone() +dataset.patterns -= t.tensor(np.nanmean(darks,axis=0)) +dataset.patterns = t.clamp(dataset.patterns,min=0) + +model = CDTools.models.FancyPtycho.from_dataset(dataset,n_modes=3,randomize_ang=0.1*np.pi) +dataset.patterns = old_patterns + +# default is CPU with 32-bit floats +model.to(device='cuda') +dataset.get_as(device='cuda') + +#model.translation_offsets.requires_grad = False + +for i, loss in enumerate(model.Adam_optimize(30, dataset, batch_size=100)): + print(i,loss) + +for i, loss in enumerate(model.Adam_optimize(30, dataset, batch_size=100, lr=0.001)): + print(i,loss) + +for i, loss in enumerate(model.Adam_optimize(50, dataset, batch_size=100, lr=0.0001)): + print(i,loss) + + +# Show some figures of merit +plot_amplitude(model.probe[0], basis=model.probe_basis.cpu()*1e6) +plot_phase(model.probe[0], basis=model.probe_basis.cpu()*1e6) +plot_amplitude(model.obj, basis=model.probe_basis.cpu()*1e6) +plot_phase(model.obj, basis=model.probe_basis.cpu()*1e6) +translations = (interactions.translations_to_pixel(model.probe_basis.cpu(), dataset.translations) + model.translation_offsets.detach().cpu()).numpy() +plt.figure() +plt.plot(translations[:,1],translations[:,0],'k-',linewidth=0.5) +plt.plot(translations[:,1],translations[:,0],'b.') +plt.show() diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index 7e65f43..54daa27 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -9,12 +9,13 @@ import numpy as np filename = '../../../Projects/CSX_3_19/cxis/processed/114429_p.cxi' filename = '../../../Projects/CSX_3_19/cxis/processed/115145_p.cxi' - +filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' with h5py.File(filename,'r') as f: dataset = CDTools.datasets.Ptycho_2D_Dataset.from_cxi(f) + model = CDTools.models.SimplePtycho.from_dataset(dataset) @@ -24,9 +25,13 @@ model.to(device='cuda') #dataset.to(device='cuda') dataset.get_as(device='cuda') -for loss in model.Adam_optimize(100, dataset): +#model.probe.requires_grad = False +for loss in model.Adam_optimize(10, dataset): print(loss) - + +#for loss in model.Adam_optimize(20, dataset, lr=0.0005): +# print(loss) + from matplotlib import pyplot as plt plot_amplitude(model.probe)