Fix dataset loading so it doesn't use so much dang memory

This commit is contained in:
Abe Levitan
2021-08-04 15:38:55 -04:00
parent 9f1a5bcae9
commit 5afff73b53
3 changed files with 34 additions and 11 deletions
+28 -6
View File
@@ -28,10 +28,12 @@ class Ptycho2DDataset(CDataset):
def __init__(self, translations, patterns, axes=None, *args, **kwargs):
"""The __init__ function allows construction from python objects.
The detector_geometry dictionary is defined to have the
entries defined by the outputs of data.get_detector_geometry.
Note that the created dataset object will not copy the data in the
patterns parameter in order to avoid doubling the memory requiement
for large datasets.
Parameters
----------
@@ -62,7 +64,11 @@ class Ptycho2DDataset(CDataset):
super(Ptycho2DDataset,self).__init__(*args, **kwargs)
self.axes = copy(axes)
self.translations = t.tensor(translations, dtype=t.float32)
self.patterns = t.tensor(patterns, dtype=t.float32)
self.patterns = t.as_tensor(patterns, dtype=t.float32)
if self.patterns.dtype == t.float64:
raise NotImplementedError('64-bit floats are not supported and precision will not be retained in reconstructions! Please explicitly convert your data to 32-bit or submit a pull request')
if self.mask is None:
self.mask = t.ones(self.patterns.shape[-2:]).to(dtype=t.bool)
self.mask.masked_fill_(t.isnan(t.sum(self.patterns,dim=(0,))),0)
@@ -147,10 +153,13 @@ class Ptycho2DDataset(CDataset):
# Load the data that is only relevant for this class
patterns, axes = cdtdata.get_data(cxi_file)
translations = cdtdata.get_ptycho_translations(cxi_file)
# And now re-do the stuff from __init__
dataset.translations = t.Tensor(translations).clone()
dataset.patterns = t.Tensor(patterns).clone()
dataset.translations = t.tensor(translations, dtype=t.float32)
dataset.patterns = t.as_tensor(patterns)
if dataset.patterns.dtype == t.float64:
raise NotImplementedError('64-bit floats are not supported and precision will not be retained in reconstructions! Please explicitly convert your data to 32-bit or submit a pull request')
dataset.axes = axes
if dataset.mask is None:
dataset.mask = t.ones(dataset.patterns.shape[-2:]).to(dtype=t.bool)
@@ -210,8 +219,21 @@ class Ptycho2DDataset(CDataset):
return meas_data * mask
translations = self.translations.detach().cpu().numpy()
nanomap_values = (self.mask.to(t.float32) * self.patterns).sum(dim=(1,2)).detach().cpu().numpy()
# This takes about twice as long as it would to just do it all at
# once, but it avoids creating another self.patterns-sized array
# as an intermediate step. This can be super important because
# self.patterns can be more than half the available memory
nanomap_values = np.ones(self.translations.shape[0])
chunk_size = 10
for i in range(0, self.translations.shape[0], chunk_size):
nanomap_values[i:i+chunk_size] = \
t.sum(self.mask * self.patterns[i:i+chunk_size],dim=(1,2))
# This is the faster but more memory-intensive version
# nanomap_values = (self.mask * self.patterns).sum(dim=(1,2)).detach().cpu().numpy()
if logarithmic:
cbar_title='Log Base 10 of Diffraction Intensity'
else:
+4 -3
View File
@@ -317,7 +317,6 @@ def get_dark(cxi_file):
return darks
def get_data(cxi_file, cut_zeroes = True):
"""Returns an array with the full stack of detector data defined in the cxi file object
@@ -350,10 +349,12 @@ def get_data(cxi_file, cut_zeroes = True):
pull_from = 'entry_1/instrument_1/detector_1/data'
else:
raise KeyError('Data is not defined within cxi file')
data = np.array(cxi_file[pull_from]).astype(np.float32)
data = cxi_file[pull_from][:]
# Use maximum in-place to avoid allocating any more memory than is needed
if cut_zeroes:
data[data < 0] = 0
np.maximum(data,0,data)
if 'axes' in cxi_file[pull_from].attrs:
try:
+2 -2
View File
@@ -217,5 +217,5 @@ def apply_half_wave_plate(probe, fast_axis_angle):
# print(exitw[:, 0, :, :])
# print('y', exitw[:, 1, :, :])
a = t.ones(2, 4)
print(t.sum(a, dim=1).shape)
#a = t.ones(2, 4)
#print(t.sum(a, dim=1).shape)