diff --git a/src/cdtools/datasets/ptycho_2d_dataset.py b/src/cdtools/datasets/ptycho_2d_dataset.py index e270714..bdf4e60 100644 --- a/src/cdtools/datasets/ptycho_2d_dataset.py +++ b/src/cdtools/datasets/ptycho_2d_dataset.py @@ -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 x 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] \ No newline at end of file diff --git a/src/cdtools/models/base.py b/src/cdtools/models/base.py index d2de6f7..b31061d 100644 --- a/src/cdtools/models/base.py +++ b/src/cdtools/models/base.py @@ -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 = [] diff --git a/src/cdtools/models/bragg_2d_ptycho.py b/src/cdtools/models/bragg_2d_ptycho.py index 473c35f..bc10ab4 100644 --- a/src/cdtools/models/bragg_2d_ptycho.py +++ b/src/cdtools/models/bragg_2d_ptycho.py @@ -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, diff --git a/src/cdtools/models/fancy_ptycho.py b/src/cdtools/models/fancy_ptycho.py index ce5c9c4..46e7db6 100644 --- a/src/cdtools/models/fancy_ptycho.py +++ b/src/cdtools/models/fancy_ptycho.py @@ -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, ) diff --git a/src/cdtools/tools/analysis/analysis.py b/src/cdtools/tools/analysis/analysis.py index d6f2f78..0fd913e 100644 --- a/src/cdtools/tools/analysis/analysis.py +++ b/src/cdtools/tools/analysis/analysis.py @@ -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)) diff --git a/src/cdtools/tools/data/data.py b/src/cdtools/tools/data/data.py index 19e7467..556aa38 100644 --- a/src/cdtools/tools/data/data.py +++ b/src/cdtools/tools/data/data.py @@ -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 diff --git a/src/cdtools/tools/initializers/initializers.py b/src/cdtools/tools/initializers/initializers.py index c0385ee..7e727af 100644 --- a/src/cdtools/tools/initializers/initializers.py +++ b/src/cdtools/tools/initializers/initializers.py @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py index 9eceb0c..ca3a2c7 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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) diff --git a/tests/test_datasets.py b/tests/test_datasets.py index d1b5f82..9e72fb7 100644 --- a/tests/test_datasets.py +++ b/tests/test_datasets.py @@ -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,:]) diff --git a/tests/tools/test_data.py b/tests/tools/test_data.py index 03998dd..bf048e5 100644 --- a/tests/tools/test_data.py +++ b/tests/tools/test_data.py @@ -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) diff --git a/tests/tools/test_measurements.py b/tests/tools/test_measurements.py index 78eceb0..bfa06de 100644 --- a/tests/tools/test_measurements.py +++ b/tests/tools/test_measurements.py @@ -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))