diff --git a/src/cdtools/datasets/base.py b/src/cdtools/datasets/base.py index 3f8ec8c..cbe7386 100644 --- a/src/cdtools/datasets/base.py +++ b/src/cdtools/datasets/base.py @@ -17,7 +17,7 @@ import torch as t from copy import copy import h5py import pathlib -from cdtools.tools import data as cdtdata +from cdtools.tools import data as cdtdata, multigpu from torch.utils import data as torchdata __all__ = ['CDataset'] @@ -92,6 +92,11 @@ class CDataset(torchdata.Dataset): self.get_as(device='cpu') + # This is a flag related to multi-GPU operation which prevents + # saving/plotting functions from being executed on GPUs outside of + # rank 0 + self.rank = multigpu.get_rank() + def to(self, *args, **kwargs): """Sends the relevant data to the given device and dtype diff --git a/src/cdtools/datasets/ptycho_2d_dataset.py b/src/cdtools/datasets/ptycho_2d_dataset.py index 3825d6d..05de7d7 100644 --- a/src/cdtools/datasets/ptycho_2d_dataset.py +++ b/src/cdtools/datasets/ptycho_2d_dataset.py @@ -154,6 +154,8 @@ class Ptycho2DDataset(CDataset): # Generate a base dataset dataset = CDataset.from_cxi(cxi_file) + + # Mutate the class to this subclass (BasicPtychoDataset) dataset.__class__ = cls @@ -198,7 +200,11 @@ class Ptycho2DDataset(CDataset): cxi_file : str, pathlib.Path, or h5py.File The .cxi file to write to """ - + # FOR MULTI-GPU: Dont run this block of code if it isn't + # called by the rank 0 GPU + if self.rank != 0: + return + # If a bare string is passed if isinstance(cxi_file, str) or isinstance(cxi_file, pathlib.Path): with cdtdata.create_cxi(cxi_file) as f: @@ -230,7 +236,10 @@ class Ptycho2DDataset(CDataset): can display a base-10 log plot of the detector readout at each position. """ - + # FOR MULTI-GPU: Dont run this block of code if it isn't + # called by the rank 0 GPU + if self.rank != 0: + return def get_images(idx): inputs, output = self[idx] @@ -288,6 +297,11 @@ class Ptycho2DDataset(CDataset): level. """ + # FOR MULTI-GPU: Dont run this block of code if it isn't + # called by the rank 0 GPU + if self.rank != 0: + return + mean_pattern, bins, ssnr = analysis.calc_spectral_info(self) cmap_label = f'Log Base 10 of Intensity + {log_offset}' title = 'Scaled mean diffraction pattern'