Add a tool to split ptycho datasets and a context manager to save .mat files, and slightly adjust the save format

This commit is contained in:
Abe Levitan
2024-02-11 10:29:07 -05:00
parent e48970f7a9
commit 44a6fac37a
10 changed files with 153 additions and 27 deletions
Binary file not shown.
+10 -13
View File
@@ -11,21 +11,18 @@ dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(filename)
model = cdtools.models.FancyPtycho.from_dataset(dataset, n_modes=2)
# Let's do this reconstruction on the GPU, shall we?
model.to(device='cuda')
dataset.get_as(device='cuda')
#model.to(device='cuda')
#dataset.get_as(device='cuda')
# Now, we run a short reconstruction from the dataset
for loss in model.Adam_optimize(10, dataset, batch_size=50, schedule=True):
# And we liveplot the updates to the model as they happen
print(model.report())
model.inspect(dataset)
with model.save_on_exit('example_reconstructions/gold_balls.mat', dataset):
# Now, we run a short reconstruction from the dataset
for loss in model.Adam_optimize(10, dataset, batch_size=50):
# And we liveplot the updates to the model as they happen
print(model.report())
model.inspect(dataset)
# This orthogonalizes the incoherent probe modes
model.tidy_probes()
# And we save out the results as a .mat file
io.savemat('example_reconstructions/gold_balls.mat',
model.save_results(dataset))
# This orthogonalizes the incoherent probe modes
model.tidy_probes()
# Finally, we plot the results
model.inspect(dataset)
+36
View File
@@ -0,0 +1,36 @@
import cdtools
from matplotlib import pyplot as plt
from scipy import io
# First, we load an example dataset from a .cxi file
filename = 'example_data/AuBalls_700ms_30nmStep_3_6SS_filter.cxi'
dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(filename)
datasets = dataset.split()
for idx, dataset in enumerate(datasets):
print(f'Working on half {idx}')
# Next, we create a ptychography model from the dataset
# Note that we explicitly ask for two incoherent probe modes
model = cdtools.models.FancyPtycho.from_dataset(dataset, n_modes=2)
# Let's do this reconstruction on the GPU, shall we?
#model.to(device='cuda')
#dataset.get_as(device='cuda')
with model.save_on_exit(f'example_reconstructions/gold_balls_half{idx}.mat',
dataset):
# Now, we run a short reconstruction from the dataset
for loss in model.Adam_optimize(10, dataset, batch_size=50):
# And we liveplot the updates to the model as they happen
print(model.report())
model.inspect(dataset)
# This orthogonalizes the incoherent probe modes
model.tidy_probes()
# Finally, we plot the results
model.inspect(dataset)
model.compare(dataset)
plt.show()
+31
View File
@@ -4,8 +4,10 @@ from copy import copy
import h5py
import pathlib
from cdtools.datasets import CDataset
from cdtools.datasets.random_selection import random_selection
from cdtools.tools import data as cdtdata
from cdtools.tools import plotting
from copy import deepcopy
__all__ = ['Ptycho2DDataset']
@@ -253,3 +255,32 @@ class Ptycho2DDataset(CDataset):
return plotting.plot_nanomap_with_images(self.translations.detach().cpu(), get_images, values=nanomap_values, nanomap_units=units, image_title='Diffraction Pattern', image_colorbar_title=cbar_title)
def split(self):
"""Splits a dataset into two pseudorandomly selected sub-datasets
"""
# the selection is only 5,000 items long, so we repeat it to be long
# enough for the dataset
repeated_random_selection = (random_selection
* int(np.ceil(len(self) / len(random_selection))))
repeated_random_selection = np.array(repeated_random_selection)
# Here, I use a fixed random selection for reproducibility
cut_random_selection =repeated_random_selection.astype(bool)[:len(self)]
dataset_1 = deepcopy(self)
dataset_1.translations = self.translations[cut_random_selection]
dataset_1.patterns = self.patterns[cut_random_selection]
if hasattr(self, 'intensities') and self.intensities is not None:
dataset_1.intensities = self.intensities[cut_random_selection]
dataset_2 = deepcopy(self)
dataset_2.translations = self.translations[~cut_random_selection]
dataset_2.patterns = self.patterns[~cut_random_selection]
if hasattr(self, 'intensities') and self.intensities is not None:
dataset_2.intensities = self.intensities[~cut_random_selection]
return dataset_1, dataset_2
File diff suppressed because one or more lines are too long
+46 -1
View File
@@ -37,6 +37,8 @@ import numpy as np
import threading
import queue
import time
from scipy import io
from contextlib import contextmanager
from .complex_lbfgs import MyLBFGS
__all__ = ['CDIModel']
@@ -169,7 +171,50 @@ class CDIModel(t.nn.Module):
results : dict
A dictionary containing all the parameters and buffers of the model, i.e. the result of self.state_dict(), converted to numpy.
"""
return {k: v.cpu().numpy() for k, v in self.state_dict().items()}
state_dict = {k: v.cpu().numpy() for k, v in self.state_dict().items()}
return {
'state_dict': state_dict,
'loss_train': np.array(self.loss_train),
}
def save_to_mat(self, filename, *args):
"""Saves the results to a .mat file
Parameters
----------
filename : str
The filename to save under
*args
Accepts any additional args that model.save_results needs, for this model
"""
return io.savemat(filename, self.save_results(*args))
@contextmanager
def save_on_exit(self, filename, *args, exception_filename=None):
"""Saves the results of the model when the context is exited
If you wrap the main body of your code in this context manager,
it will either save the results to a .mat file upon completion,
or when any exception is raised during execution.
Parameters
----------
filename : str
The filename to save under, upon completion
*args
Accepts any additional args that model.save_results needs, for this model
exception_filename : str
Optional, a separate filename to use if an exception is raised during execution. Default is equal to filename
"""
try:
yield
self.save_to_mat(filename, *args)
except Exception as e:
if exception_filename is None:
exception_filename = filename
self.save_to_mat(exception_filename, *args)
raise e
def AD_optimize(self, iterations, data_loader, optimizer,\
+24 -8
View File
@@ -564,17 +564,33 @@ class Bragg2DPtycho(CDIModel):
def save_results(self, dataset):
# This will save out everything needed to recreate the object
# in the same state, but it's not the best formatted. For example,
# "background" stores the square root of the background, etc.
base_results = super().save_results()
# We also save out the main results in a more readable format
basis = self.probe_basis.detach().cpu().numpy()
translations = self.corrected_translations(dataset).detach().cpu().numpy()
translations=self.corrected_translations(dataset).detach().cpu().numpy()
original_translations = dataset.translations.detach().cpu().numpy()
probe = self.probe.detach().cpu().numpy()
probe = probe * self.probe_norm.detach().cpu().numpy()
obj = self.obj.detach().cpu().numpy()
background = self.background.detach().cpu().numpy()**2
weights = self.weights.detach().cpu().numpy()
losses = np.array(self.loss_train)
return {'basis':basis, 'translation':translations,
'probe':probe,'obj':obj,
'background':background,
'weights':weights,
'losses':losses}
oversampling = self.oversampling
wavelength = self.wavelength.cpu().numpy()
results = {
'basis': basis,
'translations': translations,
'original_translations': original_translations,
'probe': probe,
'obj': obj,
'background': background,
'oversampling': oversampling,
'weights': weights,
'wavelength': wavelength,
}
return {**base_results, **results}
+5 -5
View File
@@ -724,9 +724,9 @@ class FancyPtycho(CDIModel):
# This will save out everything needed to recreate the object
# in the same state, but it's not the best formatted. For example,
# "background" stores the square root of the background, etc.
state_dict = super().save_results()
base_results = super().save_results()
# So, we also save out the main results in a more readable format
# We also save out the main results in a more readable format
basis = self.probe_basis.detach().cpu().numpy()
translations=self.corrected_translations(dataset).detach().cpu().numpy()
original_translations = dataset.translations.detach().cpu().numpy()
@@ -738,7 +738,7 @@ class FancyPtycho(CDIModel):
oversampling = self.oversampling
wavelength = self.wavelength.cpu().numpy()
return {
results = {
'basis': basis,
'translations': translations,
'original_translations': original_translations,
@@ -748,6 +748,6 @@ class FancyPtycho(CDIModel):
'oversampling': oversampling,
'weights': weights,
'wavelength': wavelength,
'state_dict': state_dict,
}
return {**base_results, **results}