diff --git a/CDTools/datasets/polarized_ptycho_2d_dataset.py b/CDTools/datasets/polarized_ptycho_2d_dataset.py index b1e22a7..01d21d2 100644 --- a/CDTools/datasets/polarized_ptycho_2d_dataset.py +++ b/CDTools/datasets/polarized_ptycho_2d_dataset.py @@ -62,18 +62,10 @@ class PolarizedPtycho2DDataset(Ptycho2DDataset): *args, **kwargs) - # self.polarizer = t.tensor(polarizer_angles, dtype=t.float32) - # self.analyzer = t.tensor(analyzer_angles, dtype=t.float32) + self.polarizer = t.tensor(polarizer_angles, dtype=t.float32) + self.analyzer = t.tensor(analyzer_angles, dtype=t.float32) + - - polarizer = [] - analyzer = [] - for k in range(t.tensor(translations).shape[0]): - polarizer.append((k//3)%3 * 45) - analyzer.append((k%3 * 45)) - self.polarizer = t.tensor(polarizer) - self.analyzer = t.tensor(analyzer) - def _load(self, index): """ Internal function to load data @@ -157,17 +149,8 @@ class PolarizedPtycho2DDataset(Ptycho2DDataset): dataset.__class__ = cls # Now, we save out the polarizer and analyzer states - # polarizer = cdtdata.get_shot_to_shot_info(cxi_file, 'polarizer_angle') - # analyzer = cdtdata.get_shot_to_shot_info(cxi_file, 'analyzer_angle') - - polarizer = [] - analyzer = [] - for k in range(dataset.translations.shape[0]): - polarizer.append((k//3)%3) - analyzer.append((k%3)) - dataset.polarizer = t.tensor(polarizer) - dataset.analyzer = t.tensor(analyzer) - + polarizer = cdtdata.get_shot_to_shot_info(cxi_file, 'polarizer_angle') + analyzer = cdtdata.get_shot_to_shot_info(cxi_file, 'analyzer_angle') dataset.analyzer = t.tensor(analyzer, dtype=t.float32) dataset.polarizer = t.tensor(polarizer, dtype=t.float32) diff --git a/tests/test_datasets.py b/tests/test_datasets.py index 7a825f0..02c9e99 100644 --- a/tests/test_datasets.py +++ b/tests/test_datasets.py @@ -308,6 +308,8 @@ def test_PolarizedPtycho2DDataset_init(): detector_geometry=detector_geometry, mask=mask) + print(analyzer.dtype) + print(dataset.analyzer.dtype) assert t.all(t.eq(dataset.mask,t.BoolTensor(mask))) assert dataset.entry_info == entry_info assert dataset.sample_info == sample_info diff --git a/tests/tools/test_initializers.py b/tests/tools/test_initializers.py index e340315..4b81475 100644 --- a/tests/tools/test_initializers.py +++ b/tests/tools/test_initializers.py @@ -145,9 +145,10 @@ def test_gaussian_probe(ptycho_cxi_1): normalization += np.sum(im.cpu().numpy()) normalization /= len(dataset) - normalization_1 = normalization / np.sum(np.abs(np_probe)**2) + normalization_1 = np.sqrt(normalization / np.sum(np.abs(np_probe)**2)) probe = initializers.gaussian_probe(dataset, basis, shape, sigma).numpy() + assert np.allclose(probe, normalization_1*np_probe) # And then a propagated probe @@ -159,7 +160,7 @@ def test_gaussian_probe(ptycho_cxi_1): Rz = z * (1 + (zr / z)**2) np_probe = np.exp(-Rs**2 / wz**2) * np.exp(-1j * k * Rs**2 / (2 * Rz)) - normalization_2 = normalization / np.sum(np.abs(np_probe)**2) + normalization_2 = np.sqrt(normalization / np.sum(np.abs(np_probe)**2)) probe = initializers.gaussian_probe(dataset, basis, shape, sigma, propagation_distance=z).numpy()