diff --git a/CDTools/models/__init__.py b/CDTools/models/__init__.py index 72034f8..3bf347c 100644 --- a/CDTools/models/__init__.py +++ b/CDTools/models/__init__.py @@ -95,6 +95,10 @@ class CDIModel(t.nn.Module): raise NotImplementedError() + def inspect(self): + raise NotImplementedError() + + def AD_optimize(self, iterations, data_loader, optimizer, scheduler=None): for it in range(iterations): diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py index 87664c0..1bc3c39 100644 --- a/CDTools/models/fancy_ptycho.py +++ b/CDTools/models/fancy_ptycho.py @@ -4,6 +4,8 @@ import torch as t from CDTools.models import CDIModel from CDTools import tools from CDTools.tools import cmath +from CDTools.tools import plotting as p +from matplotlib import pyplot as plt import numpy as np from copy import copy @@ -225,3 +227,47 @@ class FancyPtycho(CDIModel): pass + def corrected_translations(self,dataset): + translations = dataset.translations.to(dtype=self.probe.dtype,device=self.probe.device) + t_offset = tools.interactions.pixel_to_translations(self.probe_basis,self.translation_offsets*self.translation_scale) + return translations + t_offset + + + def inspect(self, dataset=None): + p.plot_amplitude(self.probe[0], basis=self.probe_basis) + plt.title('Dominant Probe Amplitude') + p.plot_phase(self.probe[0], basis=self.probe_basis) + plt.title('Dominant Probe Phase') + + if len(self.probe) >=2: + p.plot_amplitude(self.probe[1], basis=self.probe_basis) + plt.title('Subdominant Probe Amplitude') + p.plot_phase(self.probe[1], basis=self.probe_basis) + plt.title('Subdominant Probe Phase') + + p.plot_amplitude(self.obj, basis=self.probe_basis) + plt.title('Object Amplitude') + p.plot_phase(self.obj, basis=self.probe_basis) + plt.title('Object Phase') + + if dataset is not None: + p.plot_translations(self.corrected_translations(dataset)) + + plt.figure() + plt.imshow(self.background.detach().cpu().numpy()**2) + plt.title('Background') + + + def save_results(self, dataset): + basis = self.probe_basis.detach().cpu().numpy() + translations = self.corrected_translations(dataset).detach().cpu().numpy() + probe = cmath.torch_to_complex(self.probe.detach().cpu()) + probe = probe * self.probe_norm.detach().cpu().numpy() + obj = cmath.torch_to_complex(self.obj.detach().cpu()) + background = self.background.detach().cpu().numpy() + weights = self.weights.detach().cpu().numpy() + + return {'basis':basis, 'translation':translations, + 'probe':probe,'obj':obj, + 'background':background, + 'weights':weights} diff --git a/CDTools/models/simple_ptycho.py b/CDTools/models/simple_ptycho.py index dc9d27f..de6c5a5 100644 --- a/CDTools/models/simple_ptycho.py +++ b/CDTools/models/simple_ptycho.py @@ -3,8 +3,10 @@ from __future__ import division, print_function, absolute_import import torch as t from CDTools.models import CDIModel from CDTools import tools +from CDTools.tools import plotting as p from copy import copy from torch.utils import data as torchdata +from matplotlib import pyplot as plt class SimplePtycho(CDIModel): @@ -132,9 +134,27 @@ class SimplePtycho(CDIModel): def sim_to_dataset(self, args_list): - pass + raise NotImplementedError() + def inspect(self): + p.plot_amplitude(self.probe, basis=self.probe_basis) + plt.title('Probe Amplitude') + p.plot_phase(self.probe, basis=self.probe_basis) + plt.title('Probe Phase') + p.plot_amplitude(self.obj, basis=self.probe_basis) + plt.title('Object Amplitude') + p.plot_phase(self.obj, basis=self.probe_basis) + plt.title('Object Phase') + + + def save_results(self): + probe = tools.cmath.torch_to_complex(self.probe.detach().cpu()) + probe = probe * self.probe_norm.detach().cpu().numpy() + obj = tools.cmath.torch_to_complex(self.obj.detach().cpu()) + 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: @@ -148,6 +168,8 @@ class SimplePtycho(CDIModel): if self.mask is not None: mask = self.mask[...,None] + else: + mask=None def probe_update(exit_wave, exit_wave_corrected, probe, object, translation): return probe+tools.cmath.cmult(beta*(tools.cmath.cconj(object)/t.max(tools.cmath.cabssq(object)))[translation[0]:translation[0]+probe_shape[0],translation[1]:translation[1]+probe_shape[1]], \ diff --git a/CDTools/scripts/synthesize.py b/CDTools/scripts/synthesize.py index 518807e..15ada9b 100644 --- a/CDTools/scripts/synthesize.py +++ b/CDTools/scripts/synthesize.py @@ -28,35 +28,30 @@ if __name__ == '__main__': synth_probe, synth_obj, aligned_objs = synthesize_reconstructions( dataset['probe'], dataset['obj'], args.use_probe) - print(np.max(np.abs(synth_obj))) - print(np.max(np.abs(aligned_objs[0]))) freqs, prtf = calc_consistency_prtf(synth_obj, aligned_objs, dataset['basis']) - - plotting.plot_phase(dataset['probe'][0][0],basis=1e6*dataset['basis']) - plotting.plot_amplitude(dataset['probe'][0][0],basis=1e6*dataset['basis']) - plotting.plot_colorized(dataset['probe'][0][0],basis=1e6*dataset['basis']) - + plotting.plot_phase(dataset['probe'][0][0]),basis=dataset['basis']) + plotting.plot_amplitude(dataset['probe'][0][0]),basis=dataset['basis']) + plotting.plot_colorized(dataset['probe'][0][0]),basis=dataset['basis']) try: - plotting.plot_phase(synth_probe[1],basis=1e6*dataset['basis']) - plotting.plot_amplitude(synth_probe[1],basis=1e6*dataset['basis']) - plotting.plot_colorized(synth_probe[1],basis=1e6*dataset['basis']) + plotting.plot_phase(synth_probe[1],basis=dataset['basis']) + plotting.plot_amplitude(synth_probe[1],basis=dataset['basis']) + plotting.plot_colorized(synth_probe[1],basis=dataset['basis']) except: pass - - - plotting.plot_amplitude(synth_obj,basis=1e6*dataset['basis']) - plotting.plot_colorized(synth_obj,basis=1e6*dataset['basis']) - plotting.plot_phase(synth_obj,basis=1e6*dataset['basis']) + + plotting.plot_amplitude(synth_obj,basis=dataset['basis']) + plotting.plot_colorized(synth_obj,basis=dataset['basis']) + plotting.plot_phase(synth_obj,basis=dataset['basis']) plt.figure() try: real_translations = dataset['basis'].dot(dataset['translation'][0].transpose()) real_translations -= np.min(real_translations,axis=1)[:,None] - plt.plot(real_translations[0]*1e6,real_translations[1]*1e6,'k.') - plt.plot(real_translations[0]*1e6,real_translations[1]*1e6,'b-',linewidth=0.5) + real_translations = real_translations.transpose() + plotting.plot_translations(real_translations) plt.figure() except: pass diff --git a/CDTools/tools/plotting.py b/CDTools/tools/plotting.py index 3d2212a..3d6297d 100644 --- a/CDTools/tools/plotting.py +++ b/CDTools/tools/plotting.py @@ -8,7 +8,7 @@ from matplotlib.colors import hsv_to_rgb __all__ = ['colorize','plot_1D','plot_amplitude','plot_phase', - 'plot_colorized'] + 'plot_colorized', 'plot_translations','get_units_factor'] def colorize(z): @@ -38,6 +38,34 @@ def colorize(z): return hsv_to_rgb(np.dstack((h,s,v))) +def get_units_factor(units): + """Gets the multiplicative factor associated with a length unit + + Args: + units (str) : The abbreviation for the unit type + + Returns: + (float) : The factor meters / (unit) + """ + + u = units.lower() + if u=='m': + factor=1 + if u=='cm': + factor=1e2 + if u=='mm': + factor=1e3 + if u=='um': + factor=1e6 + if u=='nm': + factor=1e9 + if u=='a': + factor=1e10 + if u=='pm': + factor=1e12 + return factor + + def plot_1D(arr, fig = None, **kwargs): """Simple 1D plotter @@ -51,60 +79,108 @@ def plot_1D(arr, fig = None, **kwargs): if fig is None: fig = plt.figure() ax = fig.add_subplot(111, **kwargs) + else: + plt.figure(fig.number) + plt.scatter(np.arange(arr.shape[-1]), arr) -def plot_amplitude(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), **kwargs): +def plot_amplitude(im, fig = None, basis=None, units='um', **kwargs): """ Plots the amplitude of a complex Tensor or numpy array with dimensions NxMx2. Args: im (t.Tensor) : An image with dimensions NxMx2. fig (matplotlib.figure.Figure) : A matplotlib figure to use to plot. If None, a new figure is created with an Axes subplot at 111. - basis (numpy array) : The probe basis, used to put the axis labels in real space units. - Should have dimensions 3x2 + basis (numpy array) : Optional, the 3x2 probe basis, used to put the axis labels in real space units. + units (str) : The units to convert the basis to **kwargs: Can be used to set any keyword arguments for the matplotlib.axes.Axes class (see https://matplotlib.org/api/axes_api.html#the-axes-class) """ if fig is None: fig = plt.figure() ax = fig.add_subplot(111, **kwargs) - basis_norm = np.linalg.norm(basis, axis = 0) + else: + plt.figure(fig.number) + if isinstance(im, t.Tensor): absolute = cmath.cabs(im).detach().cpu().numpy() else: absolute = np.absolute(im) - plt.imshow(absolute, cmap = 'viridis', extent = [0, absolute.shape[-1]*basis_norm[1], 0, absolute.shape[-2]*basis_norm[0]]) + + #Plot in a basis if it exists, otherwise dont + if basis is not None: + if isinstance(basis,t.Tensor): + basis = basis.detach().cpu().numpy() + basis_norm = np.linalg.norm(basis, axis = 0) + basis_norm = basis_norm * get_units_factor(units) + + extent = [0, absolute.shape[-1]*basis_norm[1], 0, absolute.shape[-2]*basis_norm[0]] + else: + extent=None + + plt.imshow(absolute, cmap = 'viridis', extent = extent) plt.colorbar() + if basis is not None: + plt.xlabel('X (' + units + ')') + plt.ylabel('Y (' + units + ')') + else: + plt.xlabel('j (pixels)') + plt.ylabel('i (pixels)') + return fig -def plot_phase(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), **kwargs): +def plot_phase(im, fig=None, basis=None, units='um', **kwargs): """ Plots the phase of a complex Tensor or numpy array with dimensions NxMx2. Args: im (t.Tensor) : An image with dimensions NxMx2. fig (matplotlib.figure.Figure) : A matplotlib figure to use to plot. If None, a new figure is created with an Axes subplot at 111. - basis (numpy array) : The probe basis, used to put the axis labels in real space units. - Should have dimensions 3x2 + basis (numpy array) : Optional, the 3x2 probe basis, used to put the axis labels in real space units. **kwargs: Can be used to set any keyword arguments for the matplotlib.axes.Axes class (see https://matplotlib.org/api/axes_api.html#the-axes-class) """ if fig is None: fig = plt.figure() ax = fig.add_subplot(111, **kwargs) + else: + plt.figure(fig.number) + # If the user has matplotlib >=3.0, use the preferred colormap if isinstance(im, t.Tensor): phase = cmath.cphase(im).detach().cpu().numpy() else: phase = np.angle(im) - basis_norm = np.linalg.norm(basis, axis = 0) - try: plt.imshow(phase, cmap = 'twilight', extent = [0, phase.shape[-1]*basis_norm[1], 0, phase.shape[-2]*basis_norm[0]]) - except: plt.imshow(phase, cmap = 'hsv', extent = [0, phase.shape[-1]*basis_norm[1], 0, phase.shape[-2]*basis_norm[0]]) + + if basis is not None: + if isinstance(basis,t.Tensor): + basis = basis.detach().cpu().numpy() + basis_norm = np.linalg.norm(basis, axis = 0) + basis_norm = basis_norm * get_units_factor(units) + + extent = [0, phase.shape[-1]*basis_norm[1], 0, phase.shape[-2]*basis_norm[0]] + else: + extent=None + + try: + plt.imshow(phase, cmap = 'twilight', extent=extent) + except: + plt.imshow(phase, cmap = 'hsv', extent=extent) + plt.colorbar() + + if basis is not None: + plt.xlabel('X (' + units + ')') + plt.ylabel('Y (' + units + ')') + else: + plt.xlabel('j (pixels)') + plt.ylabel('i (pixels)') + + return fig -def plot_colorized(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), **kwargs): +def plot_colorized(im, fig=None, basis=None, units='um', **kwargs): """ Plots the colorized version of a complex Tensor or numpy array with dimensions NxMx2. The darkness corresponds to the intensity of the image, and the color corresponds to the phase. @@ -113,17 +189,69 @@ def plot_colorized(im, fig = None, basis = np.array([[0,-1], [-1,0], [0,0]]), * im (t.Tensor) : An image with dimensions NxMx2. fig (matplotlib.figure.Figure) : A matplotlib figure to use to plot. If None, a new figure is created with an Axes subplot at 111. - basis (numpy array) : The probe basis, used to put the axis labels in real space units. - Should have dimensions 3x2 + basis (numpy array) : Optional, the 3x2 probe basis, used to put the axis labels in real space units. **kwargs: Can be used to set any keyword arguments for the matplotlib.axes.Axes class (see https://matplotlib.org/api/axes_api.html#the-axes-class) """ if fig is None: fig = plt.figure() ax = fig.add_subplot(111, **kwargs) + else: + plt.figure(fig.number) + if isinstance(im, t.Tensor): im = cmath.torch_to_complex(im.detach().cpu()) - basis_norm = np.linalg.norm(basis, axis = 0) + + if basis is not None: + if isinstance(basis,t.Tensor): + basis = basis.detach().cpu().numpy() + basis_norm = np.linalg.norm(basis, axis = 0) + basis_norm = basis_norm * get_units_factor(units) + + extent = [0, im.shape[-1]*basis_norm[1], 0, im.shape[-2]*basis_norm[0]] + else: + extent=None + colorized = colorize(im) - plt.imshow(colorized, extent = [0, im.shape[-1]*basis_norm[1], 0, im.shape[-2]*basis_norm[0]]) + plt.imshow(colorized, extent=extent) + + if basis is not None: + plt.xlabel('X (' + units + ')') + plt.ylabel('Y (' + units + ')') + else: + plt.xlabel('j (pixels)') + plt.ylabel('i (pixels)') + return fig + + + +def plot_translations(translations, fig=None, units='um'): + """Plots a set of probe translations in a nicely formatted way + + Args: + translations: An Nx2 or Nx3 set of translations in real space + fig : Optional, a figure to plot into + units : Default is um, units to report in (assuming input in m) + + Returns: + + """ + + factor = get_units_factor(units) + + if fig is None: + fig = plt.figure() + ax = fig.add_subplot(111) + else: + plt.figure(fig.number) + + if isinstance(translations, t.Tensor): + translations = translations.detach().cpu().numpy() + + translations = translations * factor + plt.plot(translations[:,0], translations[:,1],'k.') + plt.plot(translations[:,0], translations[:,1],'b-', linewidth=0.5) + plt.xlabel('X (' + units + ')') + plt.ylabel('Y (' + units + ')') + diff --git a/examples/gold_ball_ptycho.py b/examples/gold_ball_ptycho.py index 911e7c2..1e3cee4 100644 --- a/examples/gold_ball_ptycho.py +++ b/examples/gold_ball_ptycho.py @@ -7,20 +7,22 @@ from CDTools.tools import interactions import h5py import numpy as np from matplotlib import pyplot as plt +import torch as t +import pickle filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' #filename = '/media/Data Bank/CSX_3_19/Processed_CXIs/115195_p.cxi' with h5py.File(filename,'r') as f: dataset = CDTools.datasets.Ptycho_2D_Dataset.from_cxi(f) - darks = np.array(f['entry_1/instrument_1/detector_1/data_dark']) + #darks = np.array(f['entry_1/instrument_1/detector_1/data_dark']) -old_patterns = dataset.patterns.clone() -dataset.patterns -= t.tensor(np.nanmean(darks,axis=0)) -dataset.patterns = t.clamp(dataset.patterns,min=0) +#old_patterns = dataset.patterns.clone() +#dataset.patterns -= t.tensor(np.nanmean(darks,axis=0)) +#dataset.patterns = t.clamp(dataset.patterns,min=0) model = CDTools.models.FancyPtycho.from_dataset(dataset,n_modes=3,randomize_ang=0.1*np.pi) -dataset.patterns = old_patterns +#dataset.patterns = old_patterns # default is CPU with 32-bit floats model.to(device='cuda') @@ -38,13 +40,9 @@ for i, loss in enumerate(model.Adam_optimize(50, dataset, batch_size=100, lr=0.0 print(i,loss) -# Show some figures of merit -plot_amplitude(model.probe[0], basis=model.probe_basis.cpu()*1e6) -plot_phase(model.probe[0], basis=model.probe_basis.cpu()*1e6) -plot_amplitude(model.obj, basis=model.probe_basis.cpu()*1e6) -plot_phase(model.obj, basis=model.probe_basis.cpu()*1e6) -translations = (interactions.translations_to_pixel(model.probe_basis.cpu(), dataset.translations) + model.translation_offsets.detach().cpu()).numpy() -plt.figure() -plt.plot(translations[:,1],translations[:,0],'k-',linewidth=0.5) -plt.plot(translations[:,1],translations[:,0],'b.') +#with open('test_results.pickle', 'wb') as f: +# pickle.dump(model.save_results(dataset),f) + +model.inspect(dataset) plt.show() +exit() diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index f5a0f19..5a97b4a 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -6,12 +6,14 @@ from CDTools.tools.plotting import * import h5py import torch as t import numpy as np +from matplotlib import pyplot as plt +import pickle -filename = '../../Downloads/114429_p.cxi' +#filename = '../../Downloads/114429_p.cxi' #filename = '../../../Projects/CSX_3_19/cxis/processed/114429_p.cxi' #filename = '../../../Projects/CSX_3_19/cxis/processed/115145_p.cxi' -#filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' +filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' with h5py.File(filename,'r') as f: @@ -20,21 +22,18 @@ with h5py.File(filename,'r') as f: model = CDTools.models.SimplePtycho.from_dataset(dataset) - # Uncomment these to use on the CPU # default is CPU with 32-bit floats -#model.to(device='cuda') +model.to(device='cuda') #dataset.to(device='cuda') -#dataset.get_as(device='cuda') +dataset.get_as(device='cuda') -for loss in model.ePIE(1, dataset): +for loss in model.Adam_optimize(1, dataset): print(loss) -from matplotlib import pyplot as plt +#with open('test_results.pickle', 'wb') as f: +# pickle.dump(model.save_results(),f) -plot_amplitude(model.probe) -plot_phase(model.probe) -plot_amplitude(model.obj) -plot_phase(model.obj) +model.inspect() plt.show() diff --git a/tests/tools/test_analysis.py b/tests/tools/test_analysis.py index cdb085e..926e452 100644 --- a/tests/tools/test_analysis.py +++ b/tests/tools/test_analysis.py @@ -143,10 +143,11 @@ def test_synthesize_reconstructions(): obj = np.copy(obj) s_probe, s_obj, obj_stack = analysis.synthesize_reconstructions(probes,objects) - assert np.max(s_probe - probe) < 1e-5 - assert np.max(s_obj - obj) < 1e-5 + assert np.max(s_probe - probe) < 2e-5 + assert np.max(s_obj - obj) < 2e-5 for t_obj in obj_stack: - assert np.max(t_obj - obj) < 1e-5 + assert np.max(t_obj - obj) < 5e-5 + def test_calc_consistency_prtf():