From 1a34aaef60f9dadef93ca250f36cf9a450c7afe4 Mon Sep 17 00:00:00 2001 From: Abe Levitan Date: Sun, 16 Oct 2022 12:01:55 -0700 Subject: [PATCH] Start making changes to enable multi-GPU, beginning with simple_ptycho --- examples/simple_ptycho.py | 4 +++ src/cdtools/models/base.py | 41 +++++++++++++++++++--- src/cdtools/models/simple_ptycho.py | 53 +++++------------------------ 3 files changed, 49 insertions(+), 49 deletions(-) diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index 5917f6b..e270277 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -8,10 +8,14 @@ dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(filename) # Next, we create a ptychography model from the dataset model = cdtools.models.SimplePtycho.from_dataset(dataset) +model.to(device='cuda:0') +dataset.get_as(device='cuda:0') + # Now, we run a short reconstruction from the dataset! for loss in model.Adam_optimize(20, dataset): print(model.report()) +print(model.save_results().keys()) # Finally, we plot the results model.inspect(dataset) model.compare(dataset) diff --git a/src/cdtools/models/base.py b/src/cdtools/models/base.py index d7f1203..2ffd774 100644 --- a/src/cdtools/models/base.py +++ b/src/cdtools/models/base.py @@ -94,16 +94,47 @@ class CDIModel(t.nn.Module): def loss(self, sim_data, real_data): raise NotImplementedError() + + def store_detector_geometry(self, detector_geometry): + if 'distance' in detector_geometry: + self.register_buffer('det_distance', + t.as_tensor(detector_geometry['distance'])) + if 'basis' in detector_geometry: + self.register_buffer('det_basis', + t.as_tensor(detector_geometry['basis'])) + if 'corner' in detector_geometry: + self.register_buffer('det_corner', + t.as_tensor(detector_geometry['corner'])) - def to(self, *args, **kwargs): - super(CDIModel,self).to(*args,**kwargs) - - + def get_detector_geometry(self): + detector_geometry = {} + if hasattr(self, 'det_distance'): + detector_geometry['distance'] = self.det_distance + if hasattr(self, 'det_basis'): + detector_geometry['basis'] = self.det_basis + if hasattr(self, 'det_corner'): + detector_geometry['corner'] = self.det_corner + return detector_geometry + def simulate_to_dataset(self, args_list): raise NotImplementedError() def save_results(self): - raise NotImplementedError() + """A convenience function to get the state dict as numpy arrays + + This function exists for two reasons, even though it is just a thin + wrapper on top of t.module.state_dict(). First, because the model + parameters for Automatic Differentiation ptychography and + related CDI methods *are* the results, it's nice to explicitly + recognize the role of extracting the state_dict as saving the + results of the reconstruction + + Second, because display, further processing, long-term storage, + etc. are often done with dictionaries of numpy files, it's useful + to have a convenience function which does that conversion + automatically. + """ + return {k: v.cpu().numpy() for k, v in self.state_dict().items()} def AD_optimize(self, iterations, data_loader, optimizer,\ scheduler=None, regularization_factor=None, thread=True, diff --git a/src/cdtools/models/simple_ptycho.py b/src/cdtools/models/simple_ptycho.py index 9dbf679..8090f5a 100644 --- a/src/cdtools/models/simple_ptycho.py +++ b/src/cdtools/models/simple_ptycho.py @@ -22,34 +22,27 @@ class SimplePtycho(CDIModel): surface_normal=np.array([0.,0.,1.]), mask=None): super(SimplePtycho,self).__init__() - self.wavelength = t.tensor(wavelength) - self.detector_geometry = copy(detector_geometry) - det_geo = self.detector_geometry - if hasattr(det_geo, 'distance'): - det_geo['distance'] = t.tensor(det_geo['distance']) - if hasattr(det_geo, 'basis'): - det_geo['basis'] = t.tensor(det_geo['basis']) - if hasattr(det_geo, 'corner'): - det_geo['corner'] = t.tensor(det_geo['corner']) + self.register_buffer('wavelength', t.as_tensor(wavelength)) + self.store_detector_geometry(detector_geometry) - self.min_translation = t.tensor(min_translation) - - self.probe_basis = t.tensor(probe_basis) + self.register_buffer('min_translation', t.as_tensor(min_translation)) + self.register_buffer('probe_basis', t.as_tensor(probe_basis)) self.detector_slice = copy(detector_slice) - self.surface_normal = t.tensor(surface_normal) + self.register_buffer('surface_normal', t.as_tensor(surface_normal)) + if mask is None: - self.mask = None + self.register_buffer('mask', None) else: - self.mask = t.tensor(mask, dtype=t.bool) + self.register_buffer('mask', t.as_tensor(mask, dtype=t.bool)) probe_guess = t.tensor(probe_guess, dtype=t.complex64) obj_guess = t.tensor(obj_guess, dtype=t.complex64) # We rescale the probe here so it learns at the same rate as the # object - self.probe_norm = t.max(t.abs(probe_guess)) + self.register_buffer('probe_norm', t.max(t.abs(probe_guess))) self.probe = t.nn.Parameter(probe_guess / self.probe_norm) self.obj = t.nn.Parameter(obj_guess) @@ -133,26 +126,6 @@ class SimplePtycho(CDIModel): return tools.losses.amplitude_mse(real_data, sim_data, mask=mask) - def to(self, *args, **kwargs): - super(SimplePtycho, self).to(*args, **kwargs) - self.wavelength = self.wavelength.to(*args,**kwargs) - # move the detector geometry too - det_geo = self.detector_geometry - if hasattr(det_geo, 'distance'): - det_geo['distance'] = det_geo['distance'].to(*args,**kwargs) - if hasattr(det_geo, 'basis'): - det_geo['basis'] = det_geo['basis'].to(*args,**kwargs) - if hasattr(det_geo, 'corner'): - det_geo['corner'] = det_geo['corner'].to(*args,**kwargs) - - if self.mask is not None: - self.mask = self.mask.to(*args, **kwargs) - - self.min_translation = self.min_translation.to(*args,**kwargs) - self.probe_basis = self.probe_basis.to(*args,**kwargs) - self.probe_norm = self.probe_norm.to(*args,**kwargs) - self.surface_normal = self.surface_normal.to(*args, **kwargs) - def sim_to_dataset(self, args_list): # In the future, potentially add more control # over what metadata is saved (names, etc.) @@ -203,14 +176,6 @@ class SimplePtycho(CDIModel): ] - - def save_results(self): - probe = self.probe.detach().cpu().numpy() - probe = probe * self.probe_norm.detach().cpu().numpy() - obj = self.obj.detach().cpu().numpy() - return {'probe':probe,'obj':obj} - - def ePIE(self, iterations, dataset, beta = 1.0): """Runs an ePIE reconstruction as described in `Maiden et al. (2017) `_. Optional parameters are: