Adding catch-all plotting tools and saving tools to models

This commit is contained in:
Abe Levitan
2019-04-25 15:46:36 -04:00
parent 1104f338d1
commit d0a56c2d99
8 changed files with 256 additions and 63 deletions
+4
View File
@@ -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):
+46
View File
@@ -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}
+23 -1
View File
@@ -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) <https://www.osapublishing.org/optica/abstract.cfm?uri=optica-4-7-736>`_.
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]], \
+12 -17
View File
@@ -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
+145 -17
View File
@@ -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 + ')')
+12 -14
View File
@@ -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()
+10 -11
View File
@@ -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()
+4 -3
View File
@@ -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():