mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-27 04:32:09 +02:00
Enhance split method in Ptycho2DDataset to allow non-random dataset splitting
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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])
|
||||
|
||||
Reference in New Issue
Block a user