mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-10-01 06:12:17 +02:00
Prevent dataset from plotting/saving outside of Rank 0
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user