mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 21:12:42 +02:00
360 lines
15 KiB
Python
360 lines
15 KiB
Python
import torch as t
|
|
from CDTools.models import CDIModel
|
|
from CDTools import tools
|
|
from CDTools.tools import plotting as p
|
|
from CDTools.tools.interactions import RPI_interaction
|
|
from CDTools.tools import initializers
|
|
from scipy.ndimage.morphology import binary_dilation
|
|
import numpy as np
|
|
from copy import copy
|
|
|
|
__all__ = ['MultimodeRPI']
|
|
|
|
|
|
|
|
__all__ = ['RPI']
|
|
|
|
class MultimodeRPI(CDIModel):
|
|
|
|
@property
|
|
def obj(self):
|
|
return t.complex(self.obj_real, self.obj_imag)
|
|
|
|
@property
|
|
def weights(self):
|
|
ws = t.complex(self.weights_real, self.weights_imag)
|
|
return ws / 10# / self.obj_real.size().numel()
|
|
|
|
def __init__(self, wavelength, detector_geometry, probe_basis,
|
|
probe, obj_guess, detector_slice=None,
|
|
background=None, mask=None, saturation=None,
|
|
obj_support=None, oversampling=1, weight_matrix=False):
|
|
|
|
super(MultimodeRPI, 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.probe_basis = t.tensor(probe_basis)
|
|
|
|
scale_factor = t.tensor([probe.shape[-1]/obj_guess.shape[-1],
|
|
probe.shape[-2]/obj_guess.shape[-2]])
|
|
self.obj_basis = self.probe_basis * scale_factor
|
|
self.detector_slice = detector_slice
|
|
|
|
# Maybe something to include in a bit
|
|
# self.surface_normal = t.tensor(surface_normal)
|
|
|
|
self.saturation = saturation
|
|
|
|
if mask is None:
|
|
self.mask = mask
|
|
else:
|
|
self.mask = t.tensor(mask, dtype=t.bool)
|
|
|
|
|
|
self.probe = t.tensor(probe, dtype=t.complex64)
|
|
|
|
obj_guess = t.tensor(obj_guess, dtype=t.complex64)
|
|
|
|
self.obj_real = t.nn.Parameter(obj_guess.real)
|
|
self.obj_imag = t.nn.Parameter(obj_guess.imag)
|
|
|
|
self.weights_real = t.nn.Parameter(t.eye(probe.shape[0])* 10)# * self.obj_real.size().numel())
|
|
self.weights_imag = t.nn.Parameter(t.zeros(probe.shape[0]))
|
|
|
|
if not weight_matrix:
|
|
self.weights_real.requires_grad=False
|
|
self.weights_imag.requires_grad=False
|
|
|
|
# Wait for LBFGS to be updated for complex-valued parameters
|
|
# self.obj = t.nn.Parameter(obj_guess.to(t.float32))
|
|
|
|
if background is None:
|
|
if detector_slice is not None:
|
|
background = 1e-6 * t.ones(
|
|
self.probe[0][self.detector_slice].shape,
|
|
dtype=t.float32)
|
|
else:
|
|
background = 1e-6 * t.ones(self.probe[0].shape,
|
|
dtype=t.float32)
|
|
|
|
self.background = t.tensor(background, dtype=t.float32)
|
|
|
|
if obj_support is not None:
|
|
self.obj_support = obj_support
|
|
self.obj.data = self.obj * obj_support[None, ...]
|
|
else:
|
|
self.obj_support = t.ones_like(self.obj[0, ...])
|
|
|
|
self.oversampling = oversampling
|
|
|
|
|
|
@classmethod
|
|
def from_dataset(cls, dataset, probe, obj_size=None, background=None, mask=None, padding=0, n_modes=1, saturation=None, scattering_mode=None, oversampling=1, auto_center=False, initialization='random', opt_for_fft=False, weight_matrix=False, probe_threshold=0):
|
|
raise NotImplementedError()
|
|
|
|
wavelength = dataset.wavelength
|
|
det_basis = dataset.detector_geometry['basis']
|
|
det_shape = dataset[0][1].shape
|
|
distance = dataset.detector_geometry['distance']
|
|
|
|
# always do this on the cpu
|
|
get_as_args = dataset.get_as_args
|
|
dataset.get_as(device='cpu')
|
|
# We only need the patterns here, not the inputs associated with them.
|
|
_, patterns = dataset[:]
|
|
dataset.get_as(*get_as_args[0],**get_as_args[1])
|
|
|
|
# Set to none to avoid issues with things outside the detector
|
|
if auto_center:
|
|
center = tools.image_processing.centroid(t.sum(patterns,dim=0))
|
|
else:
|
|
center = None
|
|
|
|
# Then, generate the probe geometry from the dataset
|
|
ewg = tools.initializers.exit_wave_geometry
|
|
probe_basis, probe_shape, det_slice = ewg(det_basis,
|
|
det_shape,
|
|
wavelength,
|
|
distance,
|
|
center=center,
|
|
padding=padding,
|
|
opt_for_fft=opt_for_fft,
|
|
oversampling=oversampling)
|
|
|
|
if not isinstance(probe,t.Tensor):
|
|
probe = t.as_tensor(probe)
|
|
|
|
# Potentially need all of this orientation stuff later
|
|
|
|
#if hasattr(dataset, 'sample_info') and \
|
|
# dataset.sample_info is not None and \
|
|
# 'orientation' in dataset.sample_info:
|
|
# surface_normal = dataset.sample_info['orientation'][2]
|
|
#else:
|
|
# surface_normal = np.array([0.,0.,1.])
|
|
|
|
# If this information is supplied when the function is called,
|
|
# then we override the information in the .cxi file
|
|
#if scattering_mode in {'t', 'transmission'}:
|
|
# surface_normal = np.array([0.,0.,1.])
|
|
#elif scattering_mode in {'r', 'reflection'}:
|
|
# outgoing_dir = np.cross(det_basis[:,0], det_basis[:,1])
|
|
# outgoing_dir /= np.linalg.norm(outgoing_dir)
|
|
# surface_normal = outgoing_dir + np.array([0.,0.,1.])
|
|
# surface_normal /= np.linalg.norm(surface_normal)
|
|
|
|
|
|
if background is None and hasattr(dataset, 'background') \
|
|
and dataset.background is not None:
|
|
background = t.sqrt(dataset.background)
|
|
elif background is not None:
|
|
background = t.sqrt(t.Tensor(background).to(dtype=t.float32))
|
|
|
|
det_geo = dataset.detector_geometry
|
|
|
|
# If no mask is given, but one exists in the dataset, load it.
|
|
if mask is None and hasattr(dataset, 'mask') \
|
|
and dataset.mask is not None:
|
|
mask = dataset.mask.to(t.bool)
|
|
|
|
# Now we initialize the object
|
|
if obj_size is None:
|
|
# This is a standard size for a well-matched probe and detector
|
|
obj_size = (np.array(probe_shape) // 2).astype(int)
|
|
|
|
if initialization.lower().strip() == 'random':
|
|
# I think something to do with the fact that the object is defined
|
|
# on a coarser grid needs to be accounted for here that is not
|
|
# accounted for yet
|
|
scale = t.sum(patterns[0]) / t.sum(t.abs(probe)**2)
|
|
obj_guess = scale * t.exp(2j * np.pi * t.rand([n_modes,]+obj_size))
|
|
elif initialization.lower().strip() == 'spectral':
|
|
if background is not None:
|
|
obj_guess = initializers.RPI_spectral_init(
|
|
patterns[0], probe, obj_size, mask=mask,
|
|
background=background**2, n_modes=n_modes)
|
|
else:
|
|
obj_guess = initializers.RPI_spectral_init(
|
|
patterns[0], probe, obj_size, mask=mask,
|
|
n_modes=n_modes)
|
|
|
|
else:
|
|
raise KeyError('Initialization "' + str(initialization) + \
|
|
'" invalid - use "spectral" or "random"')
|
|
|
|
probe_intensity = t.sqrt(t.sum(t.abs(probe)**2,axis=0))
|
|
probe_fft = tools.propagators.far_field(probe_intensity)
|
|
pad0l = (probe.shape[-2] - obj_size[-2])//2
|
|
pad0r = probe.shape[-2] - obj_size[-2] - pad0l
|
|
pad1l = (probe.shape[-1] - obj_size[-1])//2
|
|
pad1r = probe.shape[-1] - obj_size[-1] - pad1l
|
|
probe_lr_fft = probe_fft[pad0l:-pad0r,pad1l:-pad1r]
|
|
probe_lr = t.abs(tools.propagators.inverse_far_field(probe_lr_fft))
|
|
|
|
obj_support = probe_lr > t.max(probe_lr) * probe_threshold
|
|
obj_support = t.as_tensor(binary_dilation(obj_support))
|
|
|
|
return cls(wavelength, det_geo, probe_basis,
|
|
probe, obj_guess, detector_slice=det_slice,
|
|
background=background, mask=mask, saturation=saturation,
|
|
obj_support=obj_support, oversampling=oversampling,
|
|
weight_matrix=weight_matrix)
|
|
|
|
|
|
def random_init(self, pattern):
|
|
scale = t.sum(pattern) / t.sum(t.abs(self.probe)**2)
|
|
self.obj.data = scale * t.exp(
|
|
2j * np.pi * t.rand(self.obj.shape)).to(
|
|
dtype=self.obj.dtype, device=self.obj.device)
|
|
|
|
def spectral_init(self, pattern):
|
|
if self.background is not None:
|
|
self.obj.data = initializers.RPI_spectral_init(
|
|
pattern, self.probe, self.obj.shape[-3:-1], mask=self.mask,
|
|
background=self.background**2, n_modes=self.obj.shape[0]).to(
|
|
dtype=self.obj.dtype, device=self.obj.device)
|
|
else:
|
|
self.obj.data = initializers.RPI_spectral_init(
|
|
pattern, self.probe, self.obj.shape[-3:-1], mask=self.mask,
|
|
n_modes=self.obj.shape[0]).to(
|
|
dtype=self.obj.dtype, device=self.obj.device)
|
|
|
|
# Needs work
|
|
def interaction(self, index, *args):
|
|
# including *args allows this to work with all sorts of datasets
|
|
# that might include other information in with the index in their
|
|
# "input" parameters (such as translations for a ptychography dataset).
|
|
# This makes it seamless to use such a dataset even though those
|
|
# extra arguments will not be used.
|
|
|
|
|
|
all_exit_waves = []
|
|
|
|
# Mix the probes with the weight matrix
|
|
prs = t.sum(self.weights[..., None, None] * self.probe, axis=-3)
|
|
|
|
for i in range(self.probe.shape[0]):
|
|
pr = prs[i]
|
|
# Here we have a 3D probe (one single mode)
|
|
# and a 4D object (multiple modes mixing incoherently)
|
|
exit_waves = RPI_interaction(pr,
|
|
self.obj_support * self.obj[i])
|
|
all_exit_waves.append(exit_waves.unsqueeze(0))
|
|
|
|
# This creates a bunch of modes generated from all possible combos
|
|
# of the probe and object modes all strung out along the first index
|
|
|
|
output = t.cat(all_exit_waves)
|
|
|
|
# If we have multiple indexes input, we unsqueeze and repeat the stack
|
|
# of wavefields enough times to simulate each requested index. This
|
|
# seems silly, but it enables (for example) one to do a reconstruction
|
|
# from a set of diffraction patterns that are all known to be from the
|
|
# same object.
|
|
try:
|
|
# will fail if index has no length, for example when index
|
|
# is just an int. In this case, we just do nothing instead
|
|
output = output.unsqueeze(0).repeat(1,len(index),1,1,1)
|
|
except TypeError:
|
|
pass
|
|
return output
|
|
|
|
|
|
def forward_propagator(self, wavefields):
|
|
return tools.propagators.far_field(wavefields)
|
|
|
|
|
|
def backward_propagator(self, wavefields):
|
|
return tools.propagators.inverse_far_field(wavefields)
|
|
|
|
|
|
def measurement(self, wavefields):
|
|
# Here I'm taking advantage of an undocumented feature in the
|
|
# incoherent_sum measurement function where it will work with
|
|
# a 4D wavefield array as well as a 5D array.
|
|
return tools.measurements.quadratic_background(wavefields,
|
|
self.background,
|
|
detector_slice=self.detector_slice,
|
|
measurement=tools.measurements.incoherent_sum,
|
|
saturation=self.saturation,
|
|
oversampling=self.oversampling)
|
|
|
|
def loss(self, sim_data, real_data, mask=None):
|
|
return tools.losses.amplitude_mse(real_data, sim_data, mask=mask)
|
|
#return tools.losses.poisson_nll(real_data, sim_data, mask=mask)
|
|
|
|
def regularizer(self, factors):
|
|
return factors[0] * t.sum(t.abs(self.obj[0,:,:])**2) \
|
|
+ factors[1] * t.sum(t.abs(self.obj[1:,:,:])**2)
|
|
|
|
def to(self, *args, **kwargs):
|
|
super(MultimodeRPI, 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.probe = self.probe.to(*args,**kwargs)
|
|
self.probe_basis = self.probe_basis.to(*args,**kwargs)
|
|
self.obj_basis = self.obj_basis.to(*args,**kwargs)
|
|
self.obj_support = self.obj_support.to(*args,**kwargs)
|
|
self.background = self.background.to(*args, **kwargs)
|
|
|
|
# Maybe include in a bit
|
|
#self.surface_normal = self.surface_normal.to(*args, **kwargs)
|
|
|
|
def sim_to_dataset(self, args_list):
|
|
raise NotImplementedError('No sim to dataset yet, sorry!')
|
|
|
|
plot_list = [
|
|
('Root Sum Squared Amplitude of all Probes',
|
|
lambda self, fig: p.plot_amplitude(
|
|
np.sqrt(np.sum((t.abs(t.sum(self.weights[..., None, None].detach() * self.probe, axis=-3))**2).cpu().numpy(),axis=0)),
|
|
fig=fig, basis=self.probe_basis)),
|
|
('Object Amplitudes',
|
|
lambda self, fig: p.plot_amplitude(self.obj, fig=fig,
|
|
basis=self.obj_basis)),
|
|
('Object Phases',
|
|
lambda self, fig: p.plot_phase(self.obj, fig=fig,
|
|
basis=self.obj_basis))
|
|
]
|
|
|
|
|
|
def save_results(self, dataset=None, full_obj=False):
|
|
# dataset is set as a kwarg here because it isn't needed, but the
|
|
# common pattern is to pass a dataset. This makes it okay if one
|
|
# continues to use that standard pattern
|
|
probe_basis = self.probe_basis.detach().cpu().numpy()
|
|
obj_basis = self.obj_basis.detach().cpu().numpy()
|
|
probe = self.probe.detach().cpu().numpy()
|
|
# Provide the option to save out the subdominant objects or
|
|
# just the dominant one
|
|
if full_obj:
|
|
obj = self.obj.detach().cpu().numpy()
|
|
else:
|
|
obj = self.obj[0].detach().cpu().numpy()
|
|
background = self.background.detach().cpu().numpy()**2
|
|
|
|
return {'probe_basis': probe_basis, 'obj_basis': obj_basis,
|
|
'probe': probe,'obj': obj,
|
|
'background': background}
|
|
|