From 3a71d742d2e447c863c17d68ac356c05d5b4599a Mon Sep 17 00:00:00 2001 From: gnzng Date: Tue, 29 Jul 2025 14:28:07 -0700 Subject: [PATCH] Enhance split method in Ptycho2DDataset to allow non-random dataset splitting --- src/cdtools/datasets/ptycho_2d_dataset.py | 84 ++++++++++++++++------- tests/test_datasets.py | 35 ++++++++++ 2 files changed, 95 insertions(+), 24 deletions(-) diff --git a/src/cdtools/datasets/ptycho_2d_dataset.py b/src/cdtools/datasets/ptycho_2d_dataset.py index 3825d6d..1b6b58a 100644 --- a/src/cdtools/datasets/ptycho_2d_dataset.py +++ b/src/cdtools/datasets/ptycho_2d_dataset.py @@ -296,35 +296,71 @@ class Ptycho2DDataset(CDataset): cmap_label=cmap_label, title=title, ) - - - def split(self): - """Splits a dataset into two pseudorandomly selected sub-datasets + + def split(self, select_randomly: bool = True): + """ + Splits a dataset into two pseudorandomly selected sub-datasets + + select_randomly : bool + If True, the dataset is split into two disjoint datasets + using a pseudorandom selection. If False, the dataset is split + into two halves. """ - # the selection is only 5,000 items long, so we repeat it to be long - # enough for the dataset - repeated_random_selection = (random_selection - * int(np.ceil(len(self) / len(random_selection)))) + if select_randomly is True: + # the selection is only 5,000 items long, so we repeat it to be long + # enough for the dataset + repeated_random_selection = (random_selection * int(np.ceil(len(self) / len(random_selection)))) - repeated_random_selection = np.array(repeated_random_selection) - # Here, I use a fixed random selection for reproducibility - cut_random_selection =repeated_random_selection.astype(bool)[:len(self)] - - dataset_1 = deepcopy(self) - dataset_1.translations = self.translations[cut_random_selection] - dataset_1.patterns = self.patterns[cut_random_selection] - if hasattr(self, 'intensities') and self.intensities is not None: - dataset_1.intensities = self.intensities[cut_random_selection] - - dataset_2 = deepcopy(self) - dataset_2.translations = self.translations[~cut_random_selection] - dataset_2.patterns = self.patterns[~cut_random_selection] - if hasattr(self, 'intensities') and self.intensities is not None: - dataset_2.intensities = self.intensities[~cut_random_selection] + repeated_random_selection = np.array(repeated_random_selection) + # Here, I use a fixed random selection for reproducibility + cut_random_selection = repeated_random_selection.astype(bool)[:len(self)] - return dataset_1, dataset_2 + dataset_1 = deepcopy(self) + dataset_1.translations = self.translations[cut_random_selection] + dataset_1.patterns = self.patterns[cut_random_selection] + if hasattr(self, 'intensities') and self.intensities is not None: + dataset_1.intensities = self.intensities[cut_random_selection] + dataset_2 = deepcopy(self) + dataset_2.translations = self.translations[~cut_random_selection] + dataset_2.patterns = self.patterns[~cut_random_selection] + if hasattr(self, 'intensities') and self.intensities is not None: + dataset_2.intensities = self.intensities[~cut_random_selection] + + return dataset_1, dataset_2 + + elif select_randomly is False: + # If we are not randomly selecting, we just split the dataset in half + # by taking every second item of the dataset + + if len(self.translations) < 2: + raise ValueError( + 'The dataset is too small to split. It must contain at least 2 items.' + ) + if len(self.translations) % 2 != 0: + warnings.warn( + 'The dataset has an odd number of items, so the first half will be one item larger than the second half.' + ) + + dataset_1 = deepcopy(self) + dataset_1.translations = self.translations[::2] + dataset_1.patterns = self.patterns[::2] + if hasattr(self, 'intensities') and self.intensities is not None: + dataset_1.intensities = self.intensities[::2] + + dataset_2 = deepcopy(self) + dataset_2.translations = self.translations[1::2] + dataset_2.patterns = self.patterns[1::2] + if hasattr(self, 'intensities') and self.intensities is not None: + dataset_2.intensities = self.intensities[1::2] + + return dataset_1, dataset_2 + else: + raise ValueError( + 'select_randomly must be True or a list of booleans, not ' + f'{type(select_randomly)}' + ) def pad(self, to_pad, value=0, mask=True): """Pads all the diffraction patterns by a speficied amount diff --git a/tests/test_datasets.py b/tests/test_datasets.py index 4d17f0b..ac9c7bc 100644 --- a/tests/test_datasets.py +++ b/tests/test_datasets.py @@ -510,3 +510,38 @@ def test_Ptycho2DDataset_crop_translations(ptycho_cxi_1): assert t.allclose(copied_dataset.patterns, dataset.patterns[10:-10, :]) assert t.allclose(copied_dataset.translations, dataset.translations[10:-10, :]) + + +def test_Ptycho2DDataset_split(ptycho_cxi_1): + # Grab dataset + cxi, expected = ptycho_cxi_1 + dataset = Ptycho2DDataset.from_cxi(cxi) + + # Test: Split the dataset into two datasets, select randomly, default True + dataset_1, dataset_2 = dataset.split() + + # This is using a fixed random selection, so we can check the beginning of this. + # from random_selection.py: random_selection = [0, 0, 1, 0, 0, 1, 0, ...] + assert len(dataset_1) + len(dataset_2) == len(dataset) + + assert t.allclose(dataset_1.patterns[0], dataset.patterns[2]) + assert t.allclose(dataset_1.patterns[1], dataset.patterns[5]) + assert t.allclose(dataset_2.patterns[0], dataset.patterns[0]) + assert t.allclose(dataset_2.patterns[1], dataset.patterns[1]) + assert t.allclose(dataset_2.patterns[2], dataset.patterns[3]) + + # Test the translations + assert t.allclose(dataset_1.translations[0], dataset.translations[2]) + assert t.allclose(dataset_1.translations[1], dataset.translations[5]) + assert t.allclose(dataset_2.translations[0], dataset.translations[0]) + assert t.allclose(dataset_2.translations[1], dataset.translations[1]) + assert t.allclose(dataset_2.translations[2], dataset.translations[3]) + + # Test: Split the dataset into two datasets + dataset_1, dataset_2 = dataset.split(select_randomly=False) + + assert len(dataset_1) + len(dataset_2) == len(dataset) + assert t.allclose(dataset_1.patterns, dataset.patterns[::2]) + assert t.allclose(dataset_2.patterns, dataset.patterns[1::2]) + assert t.allclose(dataset_1.translations, dataset.translations[::2]) + assert t.allclose(dataset_2.translations, dataset.translations[1::2])