Prevent dataset from plotting/saving outside of Rank 0

This commit is contained in:
yoshikisd
2025-11-08 06:46:23 +00:00
parent 8d8da151d7
commit b0604d4595
2 changed files with 22 additions and 3 deletions
+6 -1
View File
@@ -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
+16 -2
View File
@@ -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'