mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-30 22:02:09 +02:00
Merge branch 'master' into holo-reconstructions
This commit is contained in:
@@ -361,3 +361,89 @@ class Ptycho2DDataset(CDataset):
|
||||
self.background = t.nn.functional.pad(self.background, to_pad)
|
||||
|
||||
|
||||
def downsample(self, factor=2):
|
||||
"""Downsamples all diffraction patterns by the specified factor
|
||||
|
||||
This is an easy way to shrink the amount of data you need to work with
|
||||
if the speckle size is much larger than the detector pixel size.
|
||||
|
||||
The downsampling factor must be an integer. The size of the output
|
||||
patterns are reduced by the specified factor, with each output pixel
|
||||
equal to the sum of a <factor> x <factor> region of pixels in the
|
||||
input pattern. This summation is done by pytorch.functional.avg_pool2d.
|
||||
|
||||
Any mask and background data which is stored with the dataset is
|
||||
downsampled with the data. The background is downsampled using the same
|
||||
method as the data. The mask is expanded so that any output pixel
|
||||
containing a masked pixel will be masked.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
factor : int
|
||||
Default 2, the factor to downsample by
|
||||
|
||||
"""
|
||||
self.patterns = t.nn.functional.avg_pool2d(
|
||||
self.patterns.unsqueeze(0), factor, divisor_override=1)[0]
|
||||
self.mask = t.logical_not(t.nn.functional.max_pool2d(
|
||||
(1-self.mask.to(dtype=t.uint8)).unsqueeze(0).unsqueeze(0),
|
||||
factor
|
||||
)[0,0].to(dtype=t.bool))
|
||||
|
||||
self.detector_geometry['basis'] = \
|
||||
self.detector_geometry['basis'] * factor
|
||||
|
||||
if self.background is not None:
|
||||
self.background = t.nn.functional.avg_pool2d(
|
||||
self.background.unsqueeze(0).unsqueeze(0),
|
||||
factor,
|
||||
divisor_override=1)[0,0]
|
||||
|
||||
|
||||
def crop_translations(self, roi):
|
||||
"""Shrinks the range of translation positions that are analyzed
|
||||
|
||||
This deletes all diffraction patterns associated with x- and
|
||||
y-translations that lie outside of a specified rectangular
|
||||
region of interest. In essence, this operation crops the "relative
|
||||
displacement map" (shown in self.inspect()) down to the region of
|
||||
interest.
|
||||
|
||||
Parameters:
|
||||
----------
|
||||
roi : tuple(float, float, float, float)
|
||||
The translation-x and -y coordinates that define the rectangular
|
||||
region of interest as (in units of meters)
|
||||
(left, right, bottom, top). The definition of these bounds are
|
||||
based on how an image is normally displayed with matplotlib's
|
||||
imshow. The order in which these elements are defined in roi
|
||||
do not matter as long as roi[:2] and roi[2:] correspond with
|
||||
the x and y coordinates, respectively.
|
||||
"""
|
||||
|
||||
# Pull out the bounds of the ROI, ensuring that left < right and
|
||||
# top < bottom
|
||||
x_left, x_right = sorted(roi[:2])
|
||||
y_top, y_bottom = sorted(roi[2:])
|
||||
|
||||
# Create pointers to the x- and y-translation positions in
|
||||
# self.translations
|
||||
x = self.translations[:, 0]
|
||||
y = self.translations[:, 1]
|
||||
|
||||
# Go look for all translation values that lie inside of the roi
|
||||
# and store their indices.
|
||||
inside_roi = (x >= x_left) & (x <= x_right) & (y <= y_bottom) & (y >= y_top)
|
||||
|
||||
# Throw a value error if inside_roi is empty
|
||||
if not t.any(inside_roi):
|
||||
raise ValueError('The roi does not contain any positions from the dataset '
|
||||
'(i.e., patterns and translations will be empty).'
|
||||
' Please redefine the bounds of the roi.')
|
||||
|
||||
# Update patterns and translations
|
||||
self.patterns = self.patterns[inside_roi]
|
||||
self.translations = self.translations[inside_roi]
|
||||
|
||||
if hasattr(self, 'intensities') and self.intensities is not None:
|
||||
self.intensities = self.intensities[inside_roi]
|
||||
@@ -703,11 +703,14 @@ class CDIModel(t.nn.Module):
|
||||
A string with basic info on the latest iteration
|
||||
"""
|
||||
if hasattr(self, 'latest_iteration_time'):
|
||||
return 'Epoch ' + str(len(self.loss_history)) + \
|
||||
' completed in %0.2f s with loss ' %\
|
||||
self.latest_iteration_time + str(self.loss_history[-1])
|
||||
epoch = len(self.loss_history)
|
||||
dt = self.latest_iteration_time
|
||||
loss = self.loss_history[-1]
|
||||
msg = f'Epoch {epoch:3d} completed in {dt:0.2f} s with loss {loss:.5e}'
|
||||
else:
|
||||
return 'No reconstruction iterations performed yet!'
|
||||
msg = 'No reconstruction iterations performed yet!'
|
||||
|
||||
return msg
|
||||
|
||||
# By default, the plot_list is empty
|
||||
plot_list = []
|
||||
|
||||
@@ -79,6 +79,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
lens=False,
|
||||
units='um',
|
||||
dtype=t.float32,
|
||||
obj_view_crop=0,
|
||||
):
|
||||
|
||||
# We need the detector geometry
|
||||
@@ -147,6 +148,19 @@ class Bragg2DPtycho(CDIModel):
|
||||
self.probe = t.nn.Parameter(probe_guess / self.probe_norm)
|
||||
self.obj = t.nn.Parameter(obj_guess)
|
||||
|
||||
# NOTE: I think it makes sense to protect against obj_view_crop
|
||||
# being zero or below, because there is nothing else to show outside
|
||||
# the object array. No reason to throw an error if, e.g., the user
|
||||
# asks for a big padding which goes outside of the actual object array.
|
||||
# Just show the full array.
|
||||
|
||||
if obj_view_crop > 0:
|
||||
self.obj_view_slice = np.s_[obj_view_crop:-obj_view_crop,
|
||||
obj_view_crop:-obj_view_crop]
|
||||
else:
|
||||
self.obj_view_slice = np.s_[:,:]
|
||||
|
||||
|
||||
if probe_support is None:
|
||||
probe_support = t.ones_like(self.probe[0], dtype=t.bool)
|
||||
self.register_buffer('probe_support',
|
||||
@@ -155,8 +169,8 @@ class Bragg2DPtycho(CDIModel):
|
||||
|
||||
if background is None:
|
||||
raise NotImplementedError('Issues with this due to probe fourier padding')
|
||||
background = 1e-6 * t.ones(self.probe[0].shape,
|
||||
dtype=t.float32)
|
||||
shape = [s//oversampling for s in self.probe[0]]
|
||||
background = 1e-6 * t.ones(shape, dtype=t.float32)
|
||||
|
||||
self.background = t.nn.Parameter(background)
|
||||
|
||||
@@ -190,10 +204,11 @@ class Bragg2DPtycho(CDIModel):
|
||||
tools.propagators.generate_high_NA_k_intensity_map(
|
||||
self.obj_basis,
|
||||
self.get_detector_geometry()['basis'] / oversampling,
|
||||
self.background.shape,
|
||||
[oversampling * d for d in self.background.shape],
|
||||
self.get_detector_geometry()['distance'],
|
||||
self.wavelength,dtype=t.float32,
|
||||
lens=lens)
|
||||
|
||||
self.register_buffer('k_map',
|
||||
t.as_tensor(k_map, dtype=dtype))
|
||||
self.register_buffer('intensity_map',
|
||||
@@ -241,6 +256,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
correct_tilt=True,
|
||||
lens=False,
|
||||
obj_padding=200,
|
||||
obj_view_crop=None,
|
||||
units='um',
|
||||
):
|
||||
wavelength = dataset.wavelength
|
||||
@@ -311,7 +327,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
|
||||
|
||||
obj_size, min_translation = tools.initializers.calc_object_setup(
|
||||
det_shape,
|
||||
[s * oversampling for s in det_shape],
|
||||
pix_translations,
|
||||
padding=obj_padding,
|
||||
)
|
||||
@@ -353,9 +369,22 @@ class Bragg2DPtycho(CDIModel):
|
||||
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 = t.exp(1j*(randomize_ang * (t.rand(obj_size)-0.5)))
|
||||
|
||||
|
||||
pfc = (probe_fourier_crop if probe_fourier_crop else [0,0])
|
||||
if obj_view_crop is None:
|
||||
obj_view_crop = min(
|
||||
probe.shape[-2] // 2 + pfc[0],
|
||||
probe.shape[-1] // 2 + pfc[1]
|
||||
)
|
||||
if obj_view_crop < 0:
|
||||
obj_view_crop += min(
|
||||
probe.shape[-2] // 2 + pfc[0],
|
||||
probe.shape[-1] // 2 + pfc[1]
|
||||
)
|
||||
|
||||
obj_view_crop += obj_padding
|
||||
|
||||
det_geo = dataset.detector_geometry
|
||||
|
||||
translation_offsets = 0 * (t.rand((len(dataset),2)) - 0.5)
|
||||
@@ -401,6 +430,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
propagate_probe=propagate_probe,
|
||||
correct_tilt=correct_tilt,
|
||||
lens=lens,
|
||||
obj_view_crop=obj_view_crop,
|
||||
units=units,
|
||||
)
|
||||
|
||||
@@ -473,11 +503,13 @@ class Bragg2DPtycho(CDIModel):
|
||||
|
||||
|
||||
def measurement(self, wavefields):
|
||||
return tools.measurements.quadratic_background(wavefields,
|
||||
self.background,
|
||||
measurement=tools.measurements.incoherent_sum,
|
||||
saturation=self.saturation,
|
||||
oversampling=self.oversampling)
|
||||
return tools.measurements.quadratic_background(
|
||||
wavefields,
|
||||
self.background,
|
||||
measurement=tools.measurements.incoherent_sum,
|
||||
saturation=self.saturation,
|
||||
oversampling=self.oversampling,
|
||||
)
|
||||
|
||||
|
||||
def loss(self, sim_data, real_data, mask=None):
|
||||
@@ -566,21 +598,21 @@ class Bragg2DPtycho(CDIModel):
|
||||
)),
|
||||
('Object Amplitude, Surface Normal View',
|
||||
lambda self, fig: p.plot_amplitude(
|
||||
self.obj,
|
||||
self.obj[self.obj_view_slice],
|
||||
fig=fig,
|
||||
basis=self.obj_basis,
|
||||
units=self.units,
|
||||
)),
|
||||
('Object Phase, Surface Normal View',
|
||||
lambda self, fig: p.plot_phase(
|
||||
self.obj,
|
||||
self.obj[self.obj_view_slice],
|
||||
fig=fig,
|
||||
basis=self.obj_basis,
|
||||
units=self.units,
|
||||
)),
|
||||
('Object Amplitude, Beam View',
|
||||
lambda self, fig: p.plot_amplitude(
|
||||
self.obj,
|
||||
self.obj[self.obj_view_slice],
|
||||
fig=fig,
|
||||
basis=self.obj_basis,
|
||||
view_basis=beam_basis,
|
||||
@@ -588,7 +620,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
)),
|
||||
('Object Phase, Beam View',
|
||||
lambda self, fig: p.plot_phase(
|
||||
self.obj,
|
||||
self.obj[self.obj_view_slice],
|
||||
fig=fig,
|
||||
basis=self.obj_basis,
|
||||
view_basis=beam_basis,
|
||||
@@ -596,7 +628,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
)),
|
||||
('Object Amplitude, Detector View',
|
||||
lambda self, fig: p.plot_amplitude(
|
||||
self.obj,
|
||||
self.obj[self.obj_view_slice],
|
||||
fig=fig,
|
||||
basis=self.obj_basis,
|
||||
view_basis=self.det_basis,
|
||||
@@ -604,7 +636,7 @@ class Bragg2DPtycho(CDIModel):
|
||||
)),
|
||||
('Object Phase, Detector View',
|
||||
lambda self, fig: p.plot_phase(
|
||||
self.obj,
|
||||
self.obj[self.obj_view_slice],
|
||||
fig=fig,
|
||||
basis=self.obj_basis,
|
||||
view_basis=self.det_basis,
|
||||
|
||||
@@ -103,9 +103,16 @@ class FancyPtycho(CDIModel):
|
||||
self.probe = t.nn.Parameter(probe_guess / self.probe_norm)
|
||||
self.obj = t.nn.Parameter(obj_guess)
|
||||
|
||||
|
||||
self.obj_view_slice = np.s_[obj_view_crop:-obj_view_crop,
|
||||
obj_view_crop:-obj_view_crop]
|
||||
# NOTE: I think it makes sense to protect against obj_view_crop
|
||||
# being zero or below, because there is nothing else to show outside
|
||||
# the object array. No reason to throw an error if, e.g., the user
|
||||
# asks for a big padding which goes outside of the actual object array.
|
||||
# Just show the full array.
|
||||
if obj_view_crop > 0:
|
||||
self.obj_view_slice = np.s_[obj_view_crop:-obj_view_crop,
|
||||
obj_view_crop:-obj_view_crop]
|
||||
else:
|
||||
self.obj_view_slice = np.s_[:,:]
|
||||
|
||||
# TODO: perhaps not working anymore for fourier cropped probes
|
||||
if background is None:
|
||||
@@ -498,6 +505,7 @@ class FancyPtycho(CDIModel):
|
||||
shift_probe=True,
|
||||
multiple_modes=True,
|
||||
probe_support=self.probe_support)
|
||||
|
||||
return exit_waves
|
||||
|
||||
|
||||
@@ -515,7 +523,7 @@ class FancyPtycho(CDIModel):
|
||||
self.background,
|
||||
measurement=tools.measurements.incoherent_sum,
|
||||
saturation=self.saturation,
|
||||
oversampling=int(self.oversampling),
|
||||
oversampling=self.oversampling,
|
||||
simulate_finite_pixels=self.simulate_finite_pixels,
|
||||
)
|
||||
|
||||
|
||||
@@ -352,8 +352,7 @@ def synthesize_reconstructions(probes, objects, use_probe=False, obj_slice=None,
|
||||
else:
|
||||
shift = ip.find_shift(synth_obj[obj_slice],obj[obj_slice], resolution=50)
|
||||
|
||||
|
||||
obj = ip.sinc_subpixel_shift(obj,np.array(shift))
|
||||
obj = ip.sinc_subpixel_shift(obj, shift)
|
||||
|
||||
if len(probe.shape) == 3:
|
||||
probe = t.stack([ip.sinc_subpixel_shift(p,tuple(shift))
|
||||
|
||||
@@ -134,18 +134,18 @@ def get_sample_info(cxi_file):
|
||||
metadata[attr] = np.float32(s1[attr][()])
|
||||
|
||||
if 'unit_cell' in s1:
|
||||
metadata['unit_cell'] = np.array(s1['unit_cell']).astype(np.float32)
|
||||
metadata['unit_cell'] = s1['unit_cell'][()].astype(np.float32)
|
||||
|
||||
|
||||
if 'geometry_1/orientation' in s1:
|
||||
orient = np.array(s1['geometry_1/orientation']).astype(np.float32)
|
||||
orient = s1['geometry_1/orientation'][()].astype(np.float32)
|
||||
xvec = orient[:3] / np.linalg.norm(orient[:3])
|
||||
yvec = orient[3:] / np.linalg.norm(orient[3:])
|
||||
metadata['orientation'] = np.array([xvec,yvec,
|
||||
np.cross(xvec,yvec)])
|
||||
|
||||
if 'geometry_1/surface_normal' in s1:
|
||||
snorm = np.array(s1['geometry_1/surface_normal']).astype(np.float32)
|
||||
snorm = s1['geometry_1/surface_normal'][()].astype(np.float32)
|
||||
xvec = np.cross(np.array([0.,1.,0.]), snorm)
|
||||
xvec /= np.linalg.norm(xvec)
|
||||
yvec = np.cross(snorm, xvec)
|
||||
@@ -218,7 +218,7 @@ def get_detector_geometry(cxi_file):
|
||||
d1 = i1['detector_1']
|
||||
|
||||
if 'detector_1/basis_vectors' in i1:
|
||||
basis_vectors = np.array(d1['basis_vectors'])
|
||||
basis_vectors = d1['basis_vectors'][()]
|
||||
if basis_vectors.shape == (2,3):
|
||||
basis_vectors = basis_vectors.T
|
||||
else:
|
||||
@@ -244,11 +244,11 @@ def get_detector_geometry(cxi_file):
|
||||
[-x_pixel_size,0,0]]).transpose()
|
||||
|
||||
try:
|
||||
distance = np.float32(d1['distance'])
|
||||
distance = np.float32(d1['distance'][()])
|
||||
except:
|
||||
distance = None
|
||||
try:
|
||||
corner_position = np.array(d1['corner_position'])
|
||||
corner_position = d1['corner_position'][()]
|
||||
except:
|
||||
corner_position = None
|
||||
|
||||
@@ -292,7 +292,7 @@ def get_mask(cxi_file):
|
||||
|
||||
i1 = cxi_file['entry_1/instrument_1']
|
||||
if 'detector_1/mask' in i1:
|
||||
mask = np.array(i1['detector_1/mask']).astype(np.uint32)
|
||||
mask = i1['detector_1/mask'][()].astype(np.uint32)
|
||||
mask_on = np.equal(mask,np.uint32(0))
|
||||
mask_has_signal = np.equal(mask,np.uint32(0x00001000))
|
||||
return np.logical_or(mask_on,mask_has_signal).astype(bool)
|
||||
@@ -324,7 +324,7 @@ def get_dark(cxi_file):
|
||||
|
||||
i1 = cxi_file['entry_1/instrument_1']
|
||||
if 'detector_1/data_dark' in i1:
|
||||
darks = np.array(i1['detector_1/data_dark'])
|
||||
darks = i1['detector_1/data_dark'][()]
|
||||
dims = tuple(range(len(darks.shape) - 2))
|
||||
darks = np.nanmean(darks,axis=dims)
|
||||
else:
|
||||
@@ -426,7 +426,7 @@ def get_shot_to_shot_info(cxi_file, field_name):
|
||||
else:
|
||||
raise KeyError('Data is not defined within cxi file')
|
||||
|
||||
return np.array(cxi_file[pull_from]).astype(np.float32)
|
||||
return cxi_file[pull_from][()].astype(np.float32)
|
||||
|
||||
|
||||
def get_ptycho_translations(cxi_file):
|
||||
@@ -485,9 +485,9 @@ def add_entry_info(cxi_file, metadata):
|
||||
# included in case the cxi spec becomes more permissive
|
||||
for key, value in metadata.items():
|
||||
if isinstance(value,(str,bytes)):
|
||||
cxi_file['entry_1'][key] = np.string_(value)
|
||||
cxi_file['entry_1'][key] = np.bytes_(value)
|
||||
elif isinstance(value, datetime.datetime):
|
||||
cxi_file['entry_1'][key] = np.string_(value.isoformat())
|
||||
cxi_file['entry_1'][key] = np.bytes_(value.isoformat())
|
||||
elif isinstance(value, numbers.Number):
|
||||
cxi_file['entry_1'][key] = value
|
||||
elif isinstance(value, (np.ndarray,list,tuple)):
|
||||
@@ -525,9 +525,9 @@ def add_sample_info(cxi_file, metadata):
|
||||
if key == 'orientation':
|
||||
continue # this is a special case
|
||||
if isinstance(value,(str,bytes)):
|
||||
s1[key] = np.string_(value)
|
||||
s1[key] = np.bytes_(value)
|
||||
elif isinstance(value, datetime.datetime):
|
||||
s1[key] = np.string_(value.isoformat())
|
||||
s1[key] = np.bytes_(value.isoformat())
|
||||
elif isinstance(value, numbers.Number):
|
||||
s1[key] = value
|
||||
elif isinstance(value, (np.ndarray,list,tuple)):
|
||||
@@ -701,7 +701,7 @@ def add_data(cxi_file, data, axes=None, compression='gzip',
|
||||
axes_str = ':'.join(axes)
|
||||
else:
|
||||
axes_str = str(axes)
|
||||
det1['data'].attrs['axes'] = np.string_(axes_str)
|
||||
det1['data'].attrs['axes'] = np.bytes_(axes_str)
|
||||
|
||||
|
||||
def add_shot_to_shot_info(cxi_file, data, field_name):
|
||||
@@ -851,11 +851,12 @@ def h5_to_nested_dict(h5_file):
|
||||
for key in h5_file.keys():
|
||||
value = h5_file[key]
|
||||
if isinstance(value, h5py.Dataset):
|
||||
arr = np.array(value)
|
||||
arr = value[()]
|
||||
if arr.dtype == object:
|
||||
d[key] = arr.ravel()[0].decode('utf-8')
|
||||
elif arr.ndim == 0:
|
||||
d[key] = arr.ravel()[0]
|
||||
# TODO is this needed with arr = value[()]?
|
||||
d[key] = arr.ravel()[0]
|
||||
else:
|
||||
d[key] = arr
|
||||
|
||||
|
||||
@@ -57,7 +57,7 @@ def exit_wave_geometry(det_basis, det_shape, wavelength, distance, oversampling=
|
||||
|
||||
# This method should work for a general parallelogram-shaped detector
|
||||
det_shape = det_basis * det_shape.to(t.float32)
|
||||
pinv_basis = t.tensor(np.linalg.pinv(det_shape).transpose()).to(t.float32)
|
||||
pinv_basis = t.linalg.pinv(det_shape).transpose(0,1)
|
||||
real_space_basis = pinv_basis * wavelength * distance
|
||||
|
||||
return real_space_basis
|
||||
|
||||
+14
-14
@@ -77,28 +77,28 @@ def ptycho_cxi_1():
|
||||
expected['entry metadata'] = {}
|
||||
e1e = expected['entry metadata']
|
||||
e1e['start_time'] = datetime.datetime.now()
|
||||
e1f['start_time'] = np.string_(e1e['start_time'].isoformat())
|
||||
e1f['start_time'] = np.bytes_(e1e['start_time'].isoformat())
|
||||
e1e['end_time'] = datetime.datetime.now()
|
||||
e1f['end_time'] = np.string_(e1e['end_time'].isoformat())
|
||||
e1f['end_time'] = np.bytes_(e1e['end_time'].isoformat())
|
||||
e1e['experiment_identifier'] = 'Fake Experiment 1'
|
||||
e1f['experiment_identifier'] = np.string_(e1e['experiment_identifier'])
|
||||
e1f['experiment_identifier'] = np.bytes_(e1e['experiment_identifier'])
|
||||
e1e['experiment_description'] = 'A fully defined ptychography experiment to test the data loading'
|
||||
e1f['experiment_description'] = np.string_(e1e['experiment_description'])
|
||||
e1f['experiment_description'] = np.bytes_(e1e['experiment_description'])
|
||||
e1e['program_name'] = 'cdtools'
|
||||
e1f['program_name'] = np.string_(e1e['program_name'])
|
||||
e1f['program_name'] = np.bytes_(e1e['program_name'])
|
||||
e1e['title'] = 'The one experiment we did'
|
||||
e1f['title'] = np.string_(e1e['title'])
|
||||
e1f['title'] = np.bytes_(e1e['title'])
|
||||
|
||||
# Set up the sample info
|
||||
s1f = e1f.create_group('sample_1')
|
||||
expected['sample info'] = {}
|
||||
s1e = expected['sample info']
|
||||
s1e['name'] = 'Fake Sample'
|
||||
s1f['name'] = np.string_(s1e['name'])
|
||||
s1f['name'] = np.bytes_(s1e['name'])
|
||||
s1e['description'] = 'A sample that isn\'t real'
|
||||
s1f['description'] = np.string_(s1e['description'])
|
||||
s1f['description'] = np.bytes_(s1e['description'])
|
||||
s1e['unit_cell_group'] = 'P1'
|
||||
s1f['unit_cell_group'] = np.string_(s1e['unit_cell_group'])
|
||||
s1f['unit_cell_group'] = np.bytes_(s1e['unit_cell_group'])
|
||||
s1e['concentration'] = np.float32(np.random.rand())
|
||||
s1f['concentration'] = s1e['concentration']
|
||||
s1e['mass'] = np.float32(np.random.rand())
|
||||
@@ -151,7 +151,7 @@ def ptycho_cxi_1():
|
||||
d1f.create_dataset('data',data=data)
|
||||
data1f['data'] = h5py.SoftLink('/entry_1/instrument_1/detector_1/data')
|
||||
|
||||
d1f['data'].attrs['axes'] = np.string_('translation:y:x')
|
||||
d1f['data'].attrs['axes'] = np.bytes_('translation:y:x')
|
||||
expected['axes'] = ['translation','y','x']
|
||||
|
||||
g1f = s1f.create_group('geometry_1')
|
||||
@@ -196,7 +196,7 @@ def ptycho_cxi_2():
|
||||
expected['entry metadata'] = {}
|
||||
e1e = expected['entry metadata']
|
||||
e1e['title'] = 'The one experiment we did'
|
||||
e1f['title'] = np.string_(e1e['title'])
|
||||
e1f['title'] = np.bytes_(e1e['title'])
|
||||
|
||||
# Set up the sample info
|
||||
s1f = e1f.create_group('sample_1')
|
||||
@@ -277,9 +277,9 @@ def ptycho_cxi_3():
|
||||
expected['entry metadata'] = {}
|
||||
e1e = expected['entry metadata']
|
||||
e1e['start_time'] = datetime.datetime.now()
|
||||
e1f['start_time'] = np.string_(e1e['start_time'].isoformat())
|
||||
e1f['start_time'] = np.bytes_(e1e['start_time'].isoformat())
|
||||
e1e['end_time'] = datetime.datetime.now()
|
||||
e1f['end_time'] = np.string_(e1e['end_time'].isoformat())
|
||||
e1f['end_time'] = np.bytes_(e1e['end_time'].isoformat())
|
||||
|
||||
# Set up the sample info
|
||||
expected['sample info'] = None
|
||||
@@ -314,7 +314,7 @@ def ptycho_cxi_3():
|
||||
expected['data'] = data
|
||||
data1f.create_dataset('data',data=data)
|
||||
|
||||
data1f['data'].attrs['axes'] = np.string_('translation:y:x')
|
||||
data1f['data'].attrs['axes'] = np.bytes_('translation:y:x')
|
||||
expected['axes'] = ['translation','y','x']
|
||||
|
||||
translations = np.arange(300).reshape((100,3)).astype(np.float32)
|
||||
|
||||
+135
-1
@@ -4,7 +4,9 @@ import numpy as np
|
||||
import torch as t
|
||||
import h5py
|
||||
import datetime
|
||||
|
||||
from copy import deepcopy
|
||||
import pytest
|
||||
import itertools
|
||||
|
||||
#
|
||||
# We start by testing the CDataset base class
|
||||
@@ -276,3 +278,135 @@ def test_Ptycho2DDataset_get_as(ptycho_cxi_1):
|
||||
assert t.allclose(pattern.to(device='cpu'),
|
||||
t.tensor(expected['data'][3,:,:]))
|
||||
|
||||
|
||||
def test_Ptycho2DDataset_downsample(test_ptycho_cxis):
|
||||
for cxi, expected in test_ptycho_cxis:
|
||||
dataset = Ptycho2DDataset.from_cxi(cxi)
|
||||
|
||||
# First we test the case of downsampling by 2 against some explicit
|
||||
# calculations
|
||||
copied_dataset = deepcopy(dataset)
|
||||
copied_dataset.downsample(2)
|
||||
|
||||
# May start failing if the test datasets are changed to include
|
||||
# a dataset with any dimension not even. That's a problem with the
|
||||
# test, not the code. Sorry! -Abe
|
||||
assert t.allclose(
|
||||
copied_dataset.patterns,
|
||||
dataset.patterns[:,::2,::2] +
|
||||
dataset.patterns[:,1::2,::2] +
|
||||
dataset.patterns[:,::2,1::2] +
|
||||
dataset.patterns[:,1::2,1::2]
|
||||
)
|
||||
|
||||
assert t.allclose(
|
||||
copied_dataset.mask,
|
||||
t.logical_and(
|
||||
t.logical_and(dataset.mask[::2,::2],
|
||||
dataset.mask[1::2,::2]),
|
||||
t.logical_and(dataset.mask[::2,1::2],
|
||||
dataset.mask[1::2,1::2]),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
||||
if dataset.background is not None:
|
||||
assert t.allclose(
|
||||
copied_dataset.background,
|
||||
dataset.background[::2,::2] +
|
||||
dataset.background[1::2,::2] +
|
||||
dataset.background[::2,1::2] +
|
||||
dataset.background[1::2,1::2]
|
||||
)
|
||||
|
||||
# And then we just test the shape for a few factors, and check that
|
||||
# it doesn't fail on edge cases (e.g. factor=1)
|
||||
for factor in [1, 2, 3]:
|
||||
copied_dataset = deepcopy(dataset)
|
||||
copied_dataset.downsample(factor=factor)
|
||||
|
||||
expected_pattern_shape = np.concatenate(
|
||||
[[dataset.patterns.shape[0]],
|
||||
np.array(dataset.patterns.shape[-2:]) // factor]
|
||||
)
|
||||
|
||||
assert np.allclose(expected_pattern_shape,
|
||||
np.array(copied_dataset.patterns.shape))
|
||||
|
||||
assert np.allclose(np.array(dataset.mask.shape) // factor,
|
||||
np.array(copied_dataset.mask.shape))
|
||||
|
||||
if dataset.background is not None:
|
||||
assert np.allclose(np.array(dataset.background.shape) // factor,
|
||||
np.array(copied_dataset.background.shape))
|
||||
|
||||
|
||||
def test_Ptycho2DDataset_crop_translations(ptycho_cxi_1):
|
||||
# Grab dataset
|
||||
cxi, expected = ptycho_cxi_1
|
||||
dataset = Ptycho2DDataset.from_cxi(cxi)
|
||||
copied_dataset = deepcopy(dataset)
|
||||
|
||||
# Test 1: Complain when the the bounds of an ROI are correctly defined,
|
||||
# but it does not contain any sample positions inside of it. The
|
||||
# translations in ptycho_cxi_1 ranges from 0m to -300m in both x and y
|
||||
# (it looks like a line scan). We select an ROI at (0, -300) which should
|
||||
# not contain any translation positions.
|
||||
with pytest.raises(ValueError) as excinfo:
|
||||
copied_dataset.crop_translations(roi=(0,-5,-295,-300))
|
||||
assert('(i.e., patterns and translations will be empty)') in str(excinfo.value)
|
||||
|
||||
# Test 2: Draw an ROI that's centered in the middle of the x/y translation range
|
||||
# and make sure that the first and last x/y elements in dataset.translate
|
||||
# contain x_left, x_right, y_top, and y_bottom
|
||||
|
||||
# Make tuples that will store the x and y positions. This will be used for
|
||||
# making permutations of the x/y positions in the ROI later.
|
||||
x_left, y_top = dataset.translations[10,:2]
|
||||
x_right, y_bottom = dataset.translations[-11,:2]
|
||||
|
||||
x_permutations = ((x_left, x_right), (x_right, x_left))
|
||||
y_permutations = ((y_top, y_bottom), (y_bottom, y_top))
|
||||
roi_permutations = tuple((x1, x2, y1, y2) for (x1, x2), (y1, y2) in
|
||||
itertools.product(x_permutations, y_permutations))
|
||||
|
||||
# Get the dataset
|
||||
copied_dataset.crop_translations(roi=roi_permutations[0])
|
||||
|
||||
# Execute the actual test
|
||||
assert (copied_dataset.translations[0,0] in x_permutations[0]) and \
|
||||
(copied_dataset.translations[-1,0] in x_permutations[0])
|
||||
|
||||
assert (copied_dataset.translations[0,1] in y_permutations[0]) and \
|
||||
(copied_dataset.translations[-1,1] in y_permutations[0])
|
||||
|
||||
# Test 3: Check if the shape of dataset.patterns and dataset.translate is correct
|
||||
# (designed to be 20 fewer rows here)
|
||||
# In the future, this should include a check for dataset.intensities once
|
||||
# an appropriate cxi file is set up for conftest.
|
||||
expected_patterns_shape = np.concatenate([[dataset.patterns.shape[0] - 20],
|
||||
dataset.patterns.shape[-2:]])
|
||||
|
||||
expected_translations_shape = np.concatenate([[dataset.translations.shape[0] - 20],
|
||||
dataset.translations.shape[-1:]])
|
||||
|
||||
assert np.allclose(np.array(copied_dataset.patterns.shape), expected_patterns_shape)
|
||||
|
||||
assert np.allclose(np.array(copied_dataset.translations.shape), expected_translations_shape)
|
||||
|
||||
# Test 4: Make sure that we always get the same result no matter what order we
|
||||
# define the left/right and bottom/top values in roi, provided that roi[:2]
|
||||
# and roi[2:] correspond with the x and y coordinates, respectively.
|
||||
|
||||
# Check each permutation
|
||||
for roi in roi_permutations:
|
||||
# Copy the dataset again; dataset will be modified each time crop_translation
|
||||
# is successfully executed.
|
||||
copied_dataset = deepcopy(dataset)
|
||||
copied_dataset.crop_translations(roi=roi)
|
||||
|
||||
# Check if the contents of dataset.patterns and dataset.translate is correct
|
||||
assert t.allclose(copied_dataset.patterns, dataset.patterns[10:-10,:])
|
||||
|
||||
assert t.allclose(copied_dataset.translations, dataset.translations[10:-10,:])
|
||||
|
||||
+15
-13
@@ -182,11 +182,11 @@ def test_add_detector(tmp_path):
|
||||
# Check this directly since we want to make sure it saved
|
||||
# the pixel sizes
|
||||
d1 = f['entry_1/instrument_1/detector_1']
|
||||
read_basis = np.array(d1['basis_vectors'])
|
||||
read_x_pix = np.float32(d1['x_pixel_size'])
|
||||
read_y_pix = np.float32(d1['y_pixel_size'])
|
||||
read_distance = np.float32(d1['distance'])
|
||||
read_corner = np.array(d1['corner_position'])
|
||||
read_basis = d1['basis_vectors'][()]
|
||||
read_x_pix = d1['x_pixel_size'][()]
|
||||
read_y_pix = d1['y_pixel_size'][()]
|
||||
read_distance = d1['distance'][()]
|
||||
read_corner = d1['corner_position'][()]
|
||||
|
||||
assert np.isclose(distance, read_distance)
|
||||
assert np.allclose(basis, read_basis)
|
||||
@@ -232,8 +232,8 @@ def test_add_data(tmp_path):
|
||||
with h5py.File(tmp_path / 'test_add_data.cxi','r') as f:
|
||||
# Check this directly since we want to make sure it saved
|
||||
# it in all the places it should have
|
||||
read_data_1 = np.array(f['entry_1/data_1/data'])
|
||||
read_data_2 = np.array(f['entry_1/instrument_1/detector_1/data'])
|
||||
read_data_1 = f['entry_1/data_1/data'][()]
|
||||
read_data_2 = f['entry_1/instrument_1/detector_1/data'][()]
|
||||
read_axes = str(f['entry_1/instrument_1/detector_1/data'].attrs['axes'].decode())
|
||||
|
||||
assert np.allclose(fake_data, read_data_1)
|
||||
@@ -261,9 +261,10 @@ def test_add_shot_to_shot_info(tmp_path):
|
||||
with h5py.File(tmp_path / 'test_add_shot_to_shot_info.cxi') as f:
|
||||
# Check this directly since we want to make sure it saved
|
||||
# it in all the places it should have
|
||||
read_analyzer_1 = np.array(f['entry_1/data_1/analyzer_angle'])
|
||||
read_analyzer_2 = np.array(f['entry_1/instrument_1/detector_1/analyzer_angle'])
|
||||
read_analyzer_3 = np.array(f['entry_1/sample_1/geometry_1/analyzer_angle'])
|
||||
read_analyzer_1 = f['entry_1/data_1/analyzer_angle'][()]
|
||||
read_analyzer_2 = \
|
||||
f['entry_1/instrument_1/detector_1/analyzer_angle'][()]
|
||||
read_analyzer_3 = f['entry_1/sample_1/geometry_1/analyzer_angle'][()]
|
||||
|
||||
assert np.allclose(analyzer, read_analyzer_1)
|
||||
assert np.allclose(analyzer, read_analyzer_2)
|
||||
@@ -280,9 +281,10 @@ def test_add_ptycho_translations(tmp_path):
|
||||
with h5py.File(tmp_path / 'test_add_ptycho_translations.cxi','r') as f:
|
||||
# Check this directly since we want to make sure it saved
|
||||
# it in all the places it should have
|
||||
read_translations_1 = np.array(f['entry_1/data_1/translation'])
|
||||
read_translations_2 = np.array(f['entry_1/instrument_1/detector_1/translation'])
|
||||
read_translations_3 = np.array(f['entry_1/sample_1/geometry_1/translation'])
|
||||
read_translations_1 = f['entry_1/data_1/translation'][()]
|
||||
read_translations_2 = \
|
||||
f['entry_1/instrument_1/detector_1/translation'][()]
|
||||
read_translations_3 = f['entry_1/sample_1/geometry_1/translation'][()]
|
||||
|
||||
assert np.allclose(-translations, read_translations_1)
|
||||
assert np.allclose(-translations, read_translations_2)
|
||||
|
||||
@@ -6,7 +6,7 @@ import numpy as np
|
||||
def test_intensity():
|
||||
wavefields = t.rand((5,10,10)) + 1j * t.rand((5,10,10))
|
||||
epsilon=1e-6
|
||||
np_result = np.abs(t.as_tensor(wavefields))**2 + epsilon
|
||||
np_result = np.abs(wavefields.numpy())**2 + epsilon
|
||||
assert t.allclose(measurements.intensity(wavefields,epsilon=epsilon),
|
||||
t.as_tensor(np_result))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user