Merge branch 'master' into holo-reconstructions

This commit is contained in:
gnzng
2025-02-08 17:25:27 -08:00
11 changed files with 338 additions and 73 deletions
+86
View File
@@ -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]
+7 -4
View File
@@ -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 = []
+49 -17
View File
@@ -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,
+12 -4
View File
@@ -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,
)
+1 -2
View File
@@ -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))
+17 -16
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+1 -1
View File
@@ -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))