misc updates

This commit is contained in:
Abe Levitan
2020-01-25 19:01:03 -05:00
parent e71c266fe4
commit a1cca0dcea
6 changed files with 21 additions and 32 deletions
+1 -1
View File
@@ -177,7 +177,7 @@ class Ptycho2DDataset(CDataset):
# If a bare string is passed
if isinstance(cxi_file, str) or isinstance(cxi_file, pathlib.Path):
with h5py.File(cxi_file,'w') as f:
with cdtdata.create_cxi(cxi_file) as f:
return self.to_cxi(f)
super(Ptycho2DDataset,self).to_cxi(cxi_file)
+7 -4
View File
@@ -96,7 +96,7 @@ class FancyPtycho(CDIModel):
@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):
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):
wavelength = dataset.wavelength
det_basis = dataset.detector_geometry['basis']
@@ -110,8 +110,11 @@ class FancyPtycho(CDIModel):
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))
if auto_center:
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,
@@ -142,7 +145,7 @@ class FancyPtycho(CDIModel):
surface_normal = outgoing_dir + np.array([0.,0.,1.])
surface_normal /= np.linalg.norm(outgoing_dir)
# 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)
+4
View File
@@ -232,7 +232,11 @@ def gaussian_probe(dataset, basis, shape, sigma, propagation_distance=0):
# First, we want to generate the parameters (sigma and curvature) for the
# propagated gaussian. Ignore the purely z-dependent phases
wavelength = dataset.wavelength
if propagation_distance is None:
propagation_distance = 0
z = propagation_distance # for shorthand
sigma = np.array(sigma)
k = 2 * np.pi / wavelength
zr = k * sigma**2
+8 -26
View File
@@ -43,20 +43,11 @@ def intensity(wavefield, detector_slice=None, epsilon=1e-7, saturation=None, ove
sim_patterns : torch.Tensor
A real MxN array storing the wavefield's intensities
"""
output = cmath.cabssq(wavefield) + epsilon
output = cmath.cabssq(wavefield)
# Now we apply oversampling
if oversampling != 1:
dim = output.dim()
if dim == 2:
output = output[None,None,:,:]
if dim == 3:
output = output[None,:,:,:]
output = avg_pool2d(output, 2, 2)
if dim == 2:
output = output[0,0,:,:]
if dim == 3:
output = output[0,:,:,:]
output = avg_pool2d(output, oversampling)
# Then we grab the detector slice
if detector_slice is not None:
@@ -67,9 +58,9 @@ def intensity(wavefield, detector_slice=None, epsilon=1e-7, saturation=None, ove
# And now saturation
if saturation is None:
return output
return output + epsilon
else:
return t.clamp(output,0,saturation)
return t.clamp(output + epsilon,0,saturation)
def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=None, oversampling=1):
@@ -102,20 +93,11 @@ def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=Non
"""
# This syntax just adds an axis to the slice to preserve the J direction
output = t.sum(cmath.cabssq(wavefields),dim=0) + epsilon
output = t.sum(cmath.cabssq(wavefields),dim=0)
# Now we apply oversampling
if oversampling != 1:
dim = output.dim()
if dim == 2:
output = output[None,None,:,:]
if dim == 3:
output = output[None,:,:,:]
output = avg_pool2d(output, 2, 2)
if dim == 2:
output = output[0,0,:,:]
if dim == 3:
output = output[0,:,:,:]
output = avg_pool2d(output, oversampling)
# Then we grab the detector slice
if detector_slice is not None:
@@ -125,9 +107,9 @@ def incoherent_sum(wavefields, detector_slice=None, epsilon=1e-7, saturation=Non
output = output[(np.s_[:],) + detector_slice]
if saturation is None:
return output
return output + epsilon
else:
return t.clamp(output,0,saturation)
return t.clamp(output + epsilon,0,saturation)
def quadratic_background(wavefield, background, detector_slice=None, measurement=intensity, epsilon=1e-7, saturation=None, oversampling=1):
Binary file not shown.
+1 -1
View File
@@ -28,7 +28,7 @@ dataset.get_as(device='cuda')
# Run the reconstruction
for i, loss in enumerate(model.Adam_optimize(250, dataset,batch_size=5)):
for i, loss in enumerate(model.Adam_optimize(100, dataset,batch_size=5)):
print(i,loss)
model.inspect(dataset)