mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 21:12:42 +02:00
Start making changes to enable multi-GPU, beginning with simple_ptycho
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user