Start making changes to enable multi-GPU, beginning with simple_ptycho

This commit is contained in:
Abe Levitan
2022-10-16 12:01:55 -07:00
parent 978c9d3206
commit 1a34aaef60
3 changed files with 49 additions and 49 deletions
+4
View File
@@ -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)
+36 -5
View File
@@ -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,
+9 -44
View File
@@ -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) <https://www.osapublishing.org/optica/abstract.cfm?uri=optica-4-7-736>`_.
Optional parameters are: