diff --git a/CDTools/models/base.py b/CDTools/models/base.py index 63d3aa7..5ceca97 100644 --- a/CDTools/models/base.py +++ b/CDTools/models/base.py @@ -163,12 +163,19 @@ class CDIModel(t.nn.Module): exit() sim_patterns = self.forward(*inp) + + #sim_patterns.retain_grad() if hasattr(self, 'mask'): loss = self.loss(pats,sim_patterns, mask=self.mask) else: loss = self.loss(pats,sim_patterns) - loss.backward() + loss.backward()#retain_variables=True) + #plt.figure() + #plt.imshow(sim_patterns.grad.cpu()[0]) + #plt.colorbar() + #plt.show() + total_loss += loss.detach() #print('probe grad') @@ -346,10 +353,10 @@ class CDIModel(t.nn.Module): # Define the optimizer - #optimizer = t.optim.LBFGS(self.parameters(), - # lr = lr, history_size=history_size) - optimizer = MyLBFGS(self.parameters(), - lr = lr, history_size=history_size) + optimizer = t.optim.LBFGS(self.parameters(), + lr = lr, history_size=history_size) + #optimizer = MyLBFGS(self.parameters(), + # lr = lr, history_size=history_size) return self.AD_optimize(iterations, data_loader, optimizer, regularization_factor=regularization_factor, diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py index 42ed20c..89a3750 100644 --- a/CDTools/models/fancy_ptycho.py +++ b/CDTools/models/fancy_ptycho.py @@ -66,9 +66,9 @@ class FancyPtycho(CDIModel): if background is None: if detector_slice is not None: - background = 1e-6 * t.ones(self.probe[0][self.detector_slice]) + background = 1e-6 * t.ones(self.probe[0][self.detector_slice].shape) else: - background = 1e-6 * t.ones(self.probe[0]) + background = 1e-6 * t.ones(self.probe[0].shape) self.background = t.nn.Parameter(t.Tensor(background).to(t.float32)) @@ -229,7 +229,7 @@ class FancyPtycho(CDIModel): mask = None if probe_support_radius is not None: - probe_support = t.zeros_like(probe[0]) + 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) @@ -239,7 +239,7 @@ class FancyPtycho(CDIModel): probe = probe * probe_support[None,:,:] else: - probe_support = None; + probe_support = None if restrict_obj != -1: ro = restrict_obj diff --git a/CDTools/models/multislice_2d_ptycho.py b/CDTools/models/multislice_2d_ptycho.py index 8768cc9..e79e2b3 100644 --- a/CDTools/models/multislice_2d_ptycho.py +++ b/CDTools/models/multislice_2d_ptycho.py @@ -14,6 +14,14 @@ __all__ = ['Multislice2DPtycho'] class Multislice2DPtycho(CDIModel): + @property + def probe(self): + 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, probe_basis, probe_guess, obj_guess, dz, nz, @@ -27,6 +35,7 @@ class Multislice2DPtycho(CDIModel): bandlimit=None, subpixel=True, exponentiate_obj=True, + low_res_obj=False, fourier_probe=False, prevent_aliasing=True, phase_only=False, @@ -58,6 +67,7 @@ class Multislice2DPtycho(CDIModel): self.units = units self.phase_only=phase_only self.prevent_aliasing=prevent_aliasing + self.low_res_obj = low_res_obj if mask is None: self.mask = mask @@ -70,11 +80,21 @@ class Multislice2DPtycho(CDIModel): self.probe_norm = 1 * t.max(t.abs(probe_guess[0]).to(t.float32)) else: self.probe_norm = 1 * t.max(t.abs(probe_guess).to(t.float32)) - - self.probe = t.nn.Parameter(probe_guess.to(t.complex64) - / self.probe_norm) + + pg = probe_guess.to(t.complex64)/self.probe_norm + self.probe_real = t.nn.Parameter(pg.real) + self.probe_imag = t.nn.Parameter(pg.imag) + #self.probe = t.complex(self.probe_real,self.probe_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 = t.complex(self.obj_real,self.obj_imag) - self.obj = t.nn.Parameter(obj_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)) if background is None: if detector_slice is not None: @@ -126,7 +146,7 @@ class Multislice2DPtycho(CDIModel): @classmethod - def from_dataset(cls, dataset, dz, nz, probe_convergence_semiangle, probe_size=None, 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): + 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, low_res_obj=False): wavelength = dataset.wavelength det_basis = dataset.detector_geometry['basis'] @@ -178,8 +198,14 @@ class Multislice2DPtycho(CDIModel): # 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) + if low_res_obj: # obj in half normal resolution + pix_translations /= 2 - obj_size, min_translation = tools.initializers.calc_object_setup(probe_shape, pix_translations, padding=200) + + obj_size, min_translation = tools.initializers.calc_object_setup(probe_shape, pix_translations, padding=100) + + if low_res_obj: + obj_size, min_translation = tools.initializers.calc_object_setup(t.as_tensor(probe_shape)//2, pix_translations, padding=100) if hasattr(dataset, 'background') and dataset.background is not None: background = t.sqrt(dataset.background) @@ -187,7 +213,8 @@ class Multislice2DPtycho(CDIModel): background = None # Finally, initialize the probe and object using this information - probe = tools.initializers.STEM_style_probe(dataset, probe_shape, det_slice, probe_convergence_semiangle, propagation_distance=propagation_distance, oversampling=oversampling) + #probe = tools.initializers.STEM_style_probe(dataset, probe_shape, det_slice, probe_convergence_semiangle, propagation_distance=propagation_distance, oversampling=oversampling) + probe = tools.initializers.SHARP_style_probe(dataset, probe_shape, det_slice, propagation_distance=propagation_distance, oversampling=oversampling) # Now we initialize all the subdominant probe modes probe_max = t.max(t.abs(probe)) @@ -239,17 +266,25 @@ class Multislice2DPtycho(CDIModel): else: mask = None - # probe_support = t.zeros(probe[0].shape,dtype=t.bool) - # xs, ys = np.mgrid[:probe.shape[-2],:probe.shape[-1]] - # xs = xs - np.mean(xs) - # ys = ys - np.mean(ys) - # Rs = np.sqrt(xs**2 + ys**2) + if probe_support_radius is not None: + probe_support = t.zeros(probe[0].shape, dtype=t.bool) + xs, ys = np.mgrid[:probe.shape[-2],:probe.shape[-1]] + xs = xs - np.mean(xs) + ys = ys - np.mean(ys) + Rs = np.sqrt(xs**2 + ys**2) + + probe_support[Rs= 1: + # plt.imshow(t.abs(tools.propagators.far_field( + # exit_waves[0,0].detach()).cpu())) + # plt.colorbar() + # plt.show() if i < self.nz-1: #on all but the last iteration exit_waves = tools.propagators.near_field( exit_waves,self.as_prop) diff --git a/CDTools/tools/initializers/initializers.py b/CDTools/tools/initializers/initializers.py index 2bf04c7..b61b5d5 100644 --- a/CDTools/tools/initializers/initializers.py +++ b/CDTools/tools/initializers/initializers.py @@ -335,15 +335,15 @@ def SHARP_style_probe(dataset, shape, det_slice, propagation_distance=None, over probe_guess = inverse_far_field(probe_fft).numpy() # Now we remove the central pixel - center = np.array(probe_guess.shape) // 2 + #center = np.array(probe_guess.shape) // 2 # I'm always unsure whether to use this modification: - probe_guess[center[0], center[1]]=np.mean([ - probe_guess[center[0]-1, center[1]], - probe_guess[center[0]+1, center[1]], - probe_guess[center[0], center[1]-1], - probe_guess[center[0], center[1]+1]]) + #probe_guess[center[0], center[1]]=np.mean([ + # probe_guess[center[0]-1, center[1]], + # probe_guess[center[0]+1, center[1]], + # probe_guess[center[0], center[1]-1], + # probe_guess[center[0], center[1]+1]]) probe_guess = t.as_tensor(probe_guess, dtype=t.complex64) @@ -420,8 +420,10 @@ def STEM_style_probe(dataset, shape, det_slice, convergence_semiangle, propagati # Fourier space (simulated on a larger stage in real space). That factor # is defined by oversampling. - probe_basis = dataset.detector_geometry['basis'] / oversampling - + probe_basis = (t.as_tensor(dataset.detector_geometry['basis'], + dtype=t.float32) + / oversampling) + mean_im = t.mean(dataset.patterns,dim=0) center = image_processing.centroid(mean_im) diff --git a/CDTools/tools/propagators/propagators.py b/CDTools/tools/propagators/propagators.py index 5bb8f3d..2d517ab 100644 --- a/CDTools/tools/propagators/propagators.py +++ b/CDTools/tools/propagators/propagators.py @@ -369,10 +369,14 @@ def generate_angular_spectrum_propagator(shape, spacing, wavelength, z, *args, r offset = t.zeros([3], dtype=basis.dtype) offset[2] = z + # And we call the generalized function! propagator = generate_generalized_angular_spectrum_propagator(shape, basis, - wavelength, offset, + wavelength, offset, propagate_along_offset=remove_z_phase, **kwargs) + if z < 0: + propagator = t.conj(propagator) + # Bandlimiting is not implemented in the generalized function, because it # has a less clear meaning in that setting, so we apply it here instead @@ -508,6 +512,7 @@ def generate_generalized_angular_spectrum_propagator(shape, basis, wavelength, o if propagation_vector is not None: perpendicular_dir *= t.sign(t.dot(perpendicular_dir,propagation_vector)) else: + pass perpendicular_dir *= t.sign(t.dot(perpendicular_dir,offset_vector)) # Then, if we have a propagation vector, we shift the in-plane diff --git a/README.md b/README.md index 711e646..6ada9df 100644 --- a/README.md +++ b/README.md @@ -6,21 +6,21 @@ CDTools is a python library for ptychography and CDI reconstructions, using an A # imports from matplotlib import pyplot as plt from CDTools.datasets import Ptycho_2D_Dataset -from CDTools.models import SimplePtycho +from CDTools.models import FancyPtycho # Load the file dataset = Ptycho_2D_Dataset.from_cxi('ptycho_data.cxi') # Generate a model from the data -model = SimplePtycho.from_dataset(dataset) +model = FancyPtycho.from_dataset(dataset) # Run a reconstruction -for i, loss in enumerate(model.Adam_optimize(10, dataset)): - print(i, loss) +for loss in model.Adam_optimize(10, dataset): + print(model.report()) # And look at the results! -model.inspect(dataset) -model.compare(dataset) +model.inspect(dataset) # See the reconstructed object, probe, etc. +model.compare(dataset) # See how the simulated and measured patterns compare plt.show() ``` diff --git a/conda_requirements.txt b/conda_requirements.txt index 186c07b..4726895 100644 --- a/conda_requirements.txt +++ b/conda_requirements.txt @@ -2,7 +2,7 @@ numpy>=1.0 scipy>=1.0 matplotlib>=2.0 python-dateutil -pytorch>=1.8.0 +pytorch>=1.9.0 h5py>=2.1 pytest sphinx diff --git a/converters/NSLS2_HXN_hdf5_to_CXI.py b/converters/NSLS2_HXN_hdf5_to_CXI.py deleted file mode 100755 index ed800f4..0000000 --- a/converters/NSLS2_HXN_hdf5_to_CXI.py +++ /dev/null @@ -1,131 +0,0 @@ -""" -Purpose: Convert NSLSII HXN hdf5 files to CXI files for analysis with CDTools. -Author: David Rower -Date: December 2019 -""" - -import numpy as np -import pickle -import h5py -import os -import CDTools -from CDTools.tools import data as cdtdata -from matplotlib import pyplot as plt -from scipy.spatial.transform import Rotation -from datetime import datetime - -def create_cxi_from_NSLS2_HXN_2DFly(data_dir, save_str, scan_number, - wavelength, theta, ROI_corner_xy, metadata): - """Converts NSLS2 HXN 2D Fly scan data (from pickle and hdf5) to CXI format - - Assumes scan files will live in data_dir with naming convention - pickle: /scan_.pickle, - hdf5: /scan_.hdf5, - and will create the file /scan_.cxi. - - Parameters - ---------- - data_dir : str - Input data directory - save_str : str - Output data name - scan_number : int - A scan index number - theta : float - Rotation angle of sample in HXN convention, in degrees - ROI_corner_xy : np.array - 1x2 array containing x, y corner of detector ROI - metadata : dict - Contains metadata relevant to the experiment - """ - - ## Load in pickle and hdf5 files - scan_str = "scan_" + scan_number - - # Load pickle (includes useful data about scan not in .hdf5 file) - with open(os.path.join(data_dir, scan_str+".pickle"), 'rb') as f: - scan_pickle = pickle.load(f) - - assert scan_pickle['plan_type'] == "FlyPlan2D", "Code only for FlyPlan2D." - - # Load hdf5 file - scan_hdf5 = h5py.File(os.path.join(data_dir, scan_str+".h5"), 'r') - - - ## Let's attempt to convert this bad boy - scan_cxi = cdtdata.create_cxi(save_str) - - - ## Add source - cdtdata.add_source(scan_cxi, wavelength=wavelength) - scan_cxi['entry_1/instrument_1/source_1']['name'] = scan_pickle['beamline_id'] - - - ## Add sample - theta = np.radians(theta) - sample_unit_vecs = Rotation.from_rotvec(-theta * np.array([0,1,0])).as_dcm() - orientation = np.hstack((sample_unit_vecs[:,0], sample_unit_vecs[:,1])) - translation = np.zeros(3) - sample_info_dict = { - "name" : "TaTe4", - "orientation" : orientation, - "translation" : translation - } - cdtdata.add_sample_info(scan_cxi, sample_info_dict) - - ## Add other metadata for experiment - metadata['start_time'] = datetime.fromtimestamp(scan_pickle['time']) - cdtdata.add_entry_info(scan_cxi, metadata) - - - ## Add detector - - # Constant detector parameters - detector_pixel_size = 55e-6 # meters - detector_height_px = 515 # px ### WARNING: NEED TO CHECK THIS - detector_width_px = 515 # px - - # Geometry parameters from scan files - distance = scan_pickle['dist_detector'] * 1e-3 # assuming mm, almost sure - gamma = np.radians(scan_pickle['gamma_detector']) - delta = np.radians(scan_pickle['delta_detector']) - Rg = Rotation.from_rotvec(-gamma * np.array([0,1,0])).as_dcm() # cw about y - Rd = Rotation.from_rotvec(-delta * Rg[:,0]).as_dcm() # cw about rotated x - RdRg = np.matmul(Rd, Rg) - - # Define detector basis: row vectors for y and x detector axes - basis = detector_pixel_size * np.array([[0.,-1.,0.],[-1.,0.,0.]]) - basis = np.matmul(RdRg,basis.T).T - - # Define corner posiiton: first find center, then offset it - corner_pos = np.dot(RdRg, distance * np.array([0.,0.,1.])) - if ROI_corner_xy[0] is None: - ROI_corner_xy[0] = 0. - if ROI_corner_xy[1] is None: - ROI_corner_xy[1] = 0. - corner_pos -= basis[0,:] * (detector_width_px/2. - ROI_corner_xy[0]) - corner_pos -= basis[1,:] * (detector_height_px/2. - ROI_corner_xy[1]) - - # Add detector data finally - cdtdata.add_detector(scan_cxi, distance, basis.T, corner=corner_pos) - - - ## Add data - axes = ['translation'] + scan_pickle['axes'] # THIS IS ONLY FOR FLY2D - data = np.copy(scan_hdf5['entry']['instrument']['detector']['data']) - data[data == 0] = 1 # to prevent divide by zero in log error - cdtdata.add_data(scan_cxi, data, axes) - - - ## Add translations - x_bounds = scan_pickle['scan_range'][0] - y_bounds = scan_pickle['scan_range'][1] - xx, yy = np.meshgrid(np.linspace(*x_bounds, scan_pickle['num1']), - np.linspace(*y_bounds, scan_pickle['num2'])) - translations = (1e-6 * - np.stack((xx.ravel(), yy.ravel(), np.zeros_like(xx.ravel())), axis=1)) - cdtdata.add_ptycho_translations(scan_cxi, translations) - - - ## Close hdf5 file - scan_hdf5.close() diff --git a/converters/TITAN_stem.py b/converters/TITAN_stem.py deleted file mode 100644 index 611d6c4..0000000 --- a/converters/TITAN_stem.py +++ /dev/null @@ -1,208 +0,0 @@ -""" -Purpose: Convert file collection from Jim Lebeau's TITAN microscope to .CXI -Author: Abe Levitan -Date: January 2019 -""" - -import numpy as np -import pickle -import h5py -import os -import CDTools -from CDTools.tools import data as cdtdata -from CDTools.datasets import Ptycho2DDataset -from matplotlib import pyplot as plt -from scipy.spatial.transform import Rotation -from datetime import datetime -import xml.etree.ElementTree as ET - - - -def load_raw_image_stack(filename): - # The resulting data is an array of (exposure, image-i, image-j), - # with image0i corresponding to y and image-j corresponding to x - # Note that the real-space scanning is done from the bottom right - # corner, first heading left (in x) then scanning up. - rawdata = np.fromfile(filename,dtype='=1.0", "scipy>=1.0", - "matplotlib>=2.0", + "matplotlib>=2.0", # Matplotlib 2.0 introduces better colormaps and no I'm not sorry "python-dateutil", - "torch>=1.9.0", #1.9.0 implements support for autograd on indexed complex tensors, key to allowing us to use complex tensors in the forward models - "h5py>=2.1", - "pathlib2 ; python_version<'3.4'"], + "torch>=1.9.0", #1.9.0 implements support for autograd on indexed complex tensors, which we need in order to use complex tensors in the forward models + "h5py>=2.1"], extras_require={ 'tests': ["pytest"], - 'docs': ["sphinx","sphinx-argparse","sphinx_rtd_theme"], - ":python_version<'3.4'": ["pathlib2"], + 'docs': ["sphinx","sphinx-argparse","sphinx_rtd_theme"] }, packages=setuptools.find_packages(), classifiers=[