mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-30 13:52:10 +02:00
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:
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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
@@ -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,\
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user