mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 21:12:42 +02:00
620 lines
24 KiB
Python
620 lines
24 KiB
Python
"""Contains functions to sensibly initialize reconstructions
|
|
|
|
The functions in this module both do the geometric calculations needed to
|
|
initialize the reconstrucions, and the heuristic calculations for
|
|
geierating sensible initializations for the probe guess.
|
|
"""
|
|
|
|
import numpy as np
|
|
import torch as t
|
|
from CDTools.tools.propagators import *
|
|
from CDTools.tools.analysis import orthogonalize_probes
|
|
from CDTools.tools import image_processing
|
|
from scipy.fftpack import next_fast_len
|
|
from scipy.sparse import linalg as spla
|
|
from torch.nn.functional import pad
|
|
import numpy as np
|
|
from functools import *
|
|
|
|
__all__ = ['exit_wave_geometry', 'calc_object_setup', 'gaussian',
|
|
'gaussian_probe', 'SHARP_style_probe', 'STEM_style_probe',
|
|
'RPI_spectral_init',
|
|
'generate_subdominant_modes']
|
|
|
|
def exit_wave_geometry(det_basis, det_shape, wavelength, distance, center=None, opt_for_fft=True, padding=0, oversampling=1):
|
|
"""Returns an exit wave basis and shape, as well as a detector slice for the given detector geometry
|
|
|
|
It takes in the parameters for a given detector - the basis defining
|
|
the pixel pitch and the shape, as well as the wavelength and propagation
|
|
distance. Optionally, it accepts a defined "center", the zero-frequency
|
|
pixel's location. It then will automatically define a larger detector
|
|
if necessary, define the exit wave basis associated with a far-field
|
|
diffraction experiment, and return that basis, shape, and detector slice
|
|
|
|
Parameters
|
|
----------
|
|
det_basis : torch.Tensor
|
|
The detector basis, as defined elsewhere
|
|
det_shape : torch.Size
|
|
The (i,j) shape of the detector
|
|
wavelength : float)
|
|
The wavelength of light for the experiment, in m
|
|
distance : float
|
|
The sample-detector distance, in m
|
|
center : torch.Tensor
|
|
If defined, the location of the zero frequency pixel
|
|
opt_for_fft : bool
|
|
Default is true, whether to increase detector size to improve fft performance
|
|
padding : int
|
|
Default is 0, the size of an extra border of nonphysical pixels around the detector
|
|
oversampling : int
|
|
Default is 1, the amount to multiply the exit wave shape by.
|
|
|
|
Returns
|
|
-------
|
|
basis : torch.Tensor
|
|
The exit wave basis
|
|
shape : torch.Tensor
|
|
The exit wave's shape
|
|
slice : slice
|
|
The slice corresponding to the physical detector
|
|
"""
|
|
|
|
det_shape = t.as_tensor(tuple(det_shape), dtype=t.int32)
|
|
det_basis = t.as_tensor(det_basis)
|
|
# First, set the center if it's not already specified
|
|
# This definition matches the center pixel of an fftshifted array
|
|
if center is None:
|
|
# this is center//2, but in pytorch 1.9.0 it throws a warning
|
|
# if you just do that
|
|
|
|
center = t.div(det_shape,2,rounding_mode='floor')# // 2
|
|
else:
|
|
center = t.as_tensor(center, dtype=t.int32)
|
|
|
|
# Then, calculate the required detector size from the centering
|
|
# This is a bit opaque but was worth doing accurately
|
|
min_left = center * 2
|
|
min_right = (det_shape - center) * 2 - 1
|
|
full_shape = t.max(min_left,min_right).to(t.int32) + 2 * padding
|
|
|
|
# In some edge cases this shape can be smaller than the detector shape
|
|
full_shape = t.max(full_shape, det_shape)
|
|
|
|
|
|
if opt_for_fft:
|
|
full_shape = t.as_tensor([next_fast_len(dim) for dim in full_shape],
|
|
dtype=t.int32)
|
|
|
|
# Then, generate a slice that pops the actual detector from the full
|
|
# detector shape
|
|
# this is full_center//2, but in pytorch 1.9.0 it throws a warning
|
|
# if you just do that
|
|
full_center = t.div(full_shape,2, rounding_mode='floor')# // 2
|
|
det_slice = np.s_[int(full_center[0]-center[0]):
|
|
int(full_center[0]-center[0]+det_shape[0]),
|
|
int(full_center[1]-center[1]):
|
|
int(full_center[1]-center[1]+det_shape[1])]
|
|
|
|
|
|
# Finally, generate the basis for the exit wave in real space
|
|
|
|
# This method should work for a general parallelogram
|
|
# shaped detector
|
|
det_shape = det_basis * full_shape.to(t.float32)
|
|
pinv_basis = t.tensor(np.linalg.pinv(det_shape).transpose()).to(t.float32)
|
|
real_space_basis = pinv_basis * wavelength * distance
|
|
|
|
# This is definitely correct, but less simple. Included here
|
|
# So future me can check that both versions are consistent.
|
|
#oop_dir = np.cross(det_basis[:,0],det_basis[:,1])
|
|
#oop_dir /= np.linalg.norm(oop_dir)
|
|
#full_basis = np.array([np.array(det_basis[:,0]),np.array(det_basis[:,1]),oop_dir]).transpose()
|
|
#inv_basis = t.tensor(np.linalg.inv(full_basis)[:2,:].transpose()).to(t.float32)
|
|
#real_space_basis = inv_basis*wavelength * distance / \
|
|
# full_shape.to(t.float32)
|
|
|
|
# Finally, convert the shape back to a torch.Size
|
|
full_shape = t.Size([dim * oversampling for dim in full_shape])
|
|
|
|
|
|
return real_space_basis, full_shape, det_slice
|
|
|
|
|
|
def calc_object_setup(probe_shape, translations, padding=0):
|
|
"""Returns an object shape and minimum pixel translation
|
|
|
|
Based on the given pixel-space translations, it will calculate the
|
|
required size for an object array and calculate the pixel translation
|
|
that corresponds to a shift by (0,0) of the probe.
|
|
|
|
Optionally a small extra border can be defined via the padding
|
|
attribute. If this is done, the calculated pixel translation will
|
|
correspond to (padding,padding)
|
|
|
|
Parameters
|
|
----------
|
|
probe_shape : torch.Size
|
|
The size of the probe array
|
|
translations : torch.Tensor
|
|
A Jx2 stack of pixel-valued (i,j) translations
|
|
padding : int
|
|
Optional, the size of an extra border to include
|
|
|
|
Returns
|
|
-------
|
|
obj_shape : torch.Size
|
|
The minimum required size of the object array
|
|
min_translation : torch.Tensor
|
|
The minimum pixel-valued translation
|
|
"""
|
|
|
|
# First we look at the translations to find the minimum translation
|
|
# and the range of translations
|
|
min_translation = t.min(translations, dim=0)[0]
|
|
translation_range = t.max(translations, dim=0)[0] - min_translation
|
|
|
|
# Calculate the required shape
|
|
translation_range = t.ceil(translation_range).numpy().astype(np.int32)
|
|
shape = translation_range + np.array(probe_shape) + 2 * padding
|
|
shape = t.Size(shape)
|
|
|
|
# And the minimum translation
|
|
min_translation = min_translation - padding
|
|
|
|
return shape, min_translation
|
|
|
|
|
|
|
|
def gaussian(shape, sigma, amplitude=1, center = None, curvature=[0,0]):
|
|
"""Returns an array with a centered Gaussian
|
|
|
|
Takes in the shape and standard deviation of a gaussian
|
|
and returns a complex torch tensor (trailing dimension is 2) with
|
|
values corresponding to a two-dimensional gaussian function
|
|
|
|
Note that [0, 0] is taken to be at the upper left corner of the array.
|
|
Default is centered at ((shape[0]-1)/2, (shape[1]-1)/2)) because x and y are zero-indexed.
|
|
By default, the phase is uniformly 0, however a curvature can be
|
|
specified to simulate a probe that has been propagated a known distance
|
|
from it's focal point. The curvature is implemented by adding a quadratic
|
|
phase phi = exp(i*curvature/2 r^2) to the Gaussian
|
|
|
|
Parameters
|
|
----------
|
|
shape : array
|
|
A 1x2 array specifying the dimensions of the output array in the form (i shape, j shape)
|
|
sigma : array
|
|
A 1x2 array specifying the i- and j- standard deviation of the gaussian in the form (i stdev, j stdev)
|
|
amplitude : float
|
|
Default 1, the amplitude the gaussian to simulate
|
|
center : array
|
|
Optional, a 1x2 array specifying the location of the center of the gaussian (i center, j center)
|
|
curvature : array
|
|
Optional, a complex part to add to the gaussian coefficient
|
|
|
|
Returns
|
|
-------
|
|
torch.Tensor
|
|
The complex-style tensor storing the Gaussian
|
|
"""
|
|
if center is None:
|
|
center = ((shape[0]-1)/2, (shape[1]-1)/2)
|
|
|
|
i, j = np.mgrid[:shape[0], :shape[1]]
|
|
isq = (i - center[0])**2
|
|
jsq = (j - center[1])**2
|
|
result = np.exp((1j*curvature[0] / 2 - 1 / (2 * sigma[0]**2)) * isq + \
|
|
(1j*curvature[1] / 2 - 1 / (2 * sigma[1]**2)) * jsq)
|
|
return t.as_tensor(amplitude*result,dtype=t.complex64)
|
|
|
|
|
|
|
|
def gaussian_probe(dataset, basis, shape, sigma, propagation_distance=0):
|
|
"""Initializes a gaussian probe based on experimental parameters
|
|
|
|
This function generates a gaussian probe initialization which has a
|
|
total fluence matching the order of magnitude of the intensity in
|
|
the observed dataset, provided the object function is of order 1.
|
|
|
|
The initialization is done using parameters defined in physical units,
|
|
such as sigma (in meters) and the propagation distance (in meters).
|
|
The internal conversion to pixel space is done with a provided probe
|
|
basis and probe shape.
|
|
|
|
TODO: Should be updated to accept a mask
|
|
|
|
Sigma can be provided either as a scalar for a uniform beam, or as
|
|
an iterable of length 2 with [sigma_i, sigma_j] being the components
|
|
of sigma in the directions parallel to the i and j basis vectors of
|
|
the probe basis
|
|
|
|
Parameters
|
|
----------
|
|
dataset : Ptycho_2D_Dataset
|
|
The dataset whose intensity we want to match
|
|
basis : array
|
|
The real space basis for exit waves in our experiment
|
|
shape : array
|
|
The shape of the simulated real space arrays
|
|
sigma : array
|
|
The standard deviation of the probe at it's focus
|
|
propagation_distance : float
|
|
Default 0, a distance to propagate the gaussian from it's focus
|
|
|
|
Returns
|
|
-------
|
|
torch.Tensor
|
|
The complex-style tensor storing the Gaussian
|
|
"""
|
|
# First, we want to generate the parameters (sigma and curvature) for the
|
|
# propagated gaussian. Ignore the purely z-dependent phases
|
|
wavelength = dataset.wavelength
|
|
if propagation_distance is None:
|
|
propagation_distance = 0
|
|
|
|
z = propagation_distance # for shorthand
|
|
|
|
sigma = np.array(sigma)
|
|
k = 2 * np.pi / wavelength
|
|
zr = k * sigma**2
|
|
sigmaz = sigma * np.sqrt(1 + (z / zr)**2)
|
|
curvature = -k * z / (z**2 + zr**2)
|
|
|
|
# The conversion must then be done to pixel space
|
|
sigma_pix = sigmaz / np.array([np.linalg.norm(basis[:,0]),
|
|
np.linalg.norm(basis[:,1])])
|
|
curvature_pix = curvature * np.array([np.linalg.norm(basis[:,0]),
|
|
np.linalg.norm(basis[:,1])])**2
|
|
|
|
# Then we can generate the gaussian array
|
|
probe = gaussian(shape, sigma=sigma_pix, curvature=curvature_pix)
|
|
|
|
# Finally, we should calculate the average pattern intensity from the
|
|
# dataset and normalize the gaussian probe. This should be done by
|
|
avg_intensities = [t.sum(dataset[idx][1]) for idx in range(len(dataset))]
|
|
avg_intensity = t.mean(t.tensor(avg_intensities))
|
|
probe_intensity = t.sum(t.abs(probe)**2)
|
|
return avg_intensity / probe_intensity * probe
|
|
|
|
|
|
def SHARP_style_probe(dataset, shape, det_slice, propagation_distance=None, oversampling=1):
|
|
"""Generates a SHARP style probe guess from a dataset
|
|
|
|
What we call the "SHARP" style probe guess is to take a mean of all
|
|
the diffraction patterns and use that as an initial guess of the
|
|
Fourier space distribution of the probe. We set all the phases to
|
|
zero, which would for many simple beams (like a zone plate) generate
|
|
a first guess of the probe that is very close to the focal spot of
|
|
the probe beam.
|
|
|
|
If the probe is simulated in higher resolution than the detector,
|
|
a common occurence, these undefined pixels are set to zero for the
|
|
purposes of defining the guess
|
|
|
|
We make a small tweak to this procedure to lower the central pixel of
|
|
the probe generated this way, which can often overwhelm the rest of the
|
|
probe if there is significant noise on the detector
|
|
|
|
Parameters
|
|
----------
|
|
dataset : Ptycho_2D_Dataset
|
|
The dataset to work from
|
|
shape : torch.Size
|
|
The size of the probe array to simulate
|
|
det_slice : slice
|
|
A slice or tuple of slices corresponding to the detector region in Fourier space
|
|
propagation_distance : float
|
|
Default is no propagation, an amount to propagate the guessed probe from it's focal point
|
|
oversampling : int
|
|
Default 1, the width of the region of pixels in the wavefield to bin into a single detector pixel
|
|
|
|
Returns
|
|
-------
|
|
torch.Tensor
|
|
The complex-style tensor storing the probe guess
|
|
"""
|
|
|
|
# NOTE: I don't love the way np and torch are mixed here, I think this
|
|
# function deserves some love.
|
|
|
|
# to use the mask or not?
|
|
intensities = np.zeros([dim // oversampling for dim in shape])
|
|
for params, im in dataset:
|
|
if hasattr(dataset,'mask') and dataset.mask is not None:
|
|
intensities[det_slice] += dataset.mask.cpu().numpy() * im.cpu().numpy()
|
|
else:
|
|
intensities[det_slice] += im.cpu().numpy()
|
|
intensities /= len(dataset)
|
|
|
|
# Subtract off a known background if it's stored
|
|
if hasattr(dataset, 'background') and dataset.background is not None:
|
|
intensities[det_slice] = np.clip(intensities[det_slice] - dataset.background.cpu().numpy(), a_min=0,a_max=None)
|
|
|
|
probe_fft = t.tensor(np.sqrt(intensities)).to(dtype=t.complex64)
|
|
|
|
probe_guess = inverse_far_field(probe_fft).numpy()
|
|
# Now we remove the central pixel
|
|
center = np.array(probe_guess.shape) // 2
|
|
|
|
# I'm always unsure whether to use this modification:
|
|
|
|
probe_guess[center[0], center[1]]=np.mean([
|
|
probe_guess[center[0]-1, center[1]],
|
|
probe_guess[center[0]+1, center[1]],
|
|
probe_guess[center[0], center[1]-1],
|
|
probe_guess[center[0], center[1]+1]])
|
|
|
|
probe_guess = t.as_tensor(probe_guess, dtype=t.complex64)
|
|
|
|
if propagation_distance is not None:
|
|
# First generate the propagation array
|
|
|
|
probe_shape = t.as_tensor(tuple(probe_guess.shape))
|
|
|
|
# Start by recalculating the probe basis from the given information
|
|
det_basis = t.as_tensor(dataset.detector_geometry['basis'])
|
|
basis_dirs = det_basis / t.norm(det_basis, dim=0)
|
|
distance = dataset.detector_geometry['distance']
|
|
probe_basis = basis_dirs * dataset.wavelength * distance / \
|
|
(probe_shape * t.norm(det_basis,dim=0))
|
|
|
|
# Then package everything as it's needed
|
|
probe_spacing = t.norm(probe_basis,dim=0).numpy()
|
|
probe_shape = probe_shape.numpy().astype(np.int32)
|
|
|
|
# And generate the propagator
|
|
AS_prop = generate_angular_spectrum_propagator(probe_shape, probe_spacing, dataset.wavelength, propagation_distance)
|
|
|
|
probe_guess = near_field(probe_guess,AS_prop)
|
|
|
|
# Finally, place this probe in a full-sized array if there is oversampling
|
|
final_probe = t.zeros(shape,dtype=t.complex64)
|
|
left = shape[0]//2 - probe_guess.shape[0] // 2
|
|
top = shape[1]//2 - probe_guess.shape[1] // 2
|
|
final_probe[left:left+probe_guess.shape[0],
|
|
top:top+probe_guess.shape[1]] = probe_guess
|
|
|
|
return final_probe
|
|
|
|
|
|
def STEM_style_probe(dataset, shape, det_slice, convergence_semiangle, propagation_distance=None, oversampling=1):
|
|
"""Generates a STEM style probe guess from a dataset
|
|
|
|
What we call the "STEM" style probe guess is a probe generated by
|
|
a uniform aperture in Fourier space, with no optical aberrations.
|
|
This is the kind of probe than an ideal, aberration-corrected STEM
|
|
with perfect coherence would produce, so it is a good starting
|
|
guess for STEM datasets.
|
|
|
|
We set the initial intensity of the probe with relation to the
|
|
diffraction patterns so that a typical object will have an intensity
|
|
of around 1. We also set the initial phase ramp of the probe such that
|
|
the undiffracted probe has a centroid on the detector which matches the
|
|
centroid of the diffraction patterns
|
|
|
|
Parameters
|
|
----------
|
|
dataset : Ptycho_2D_Dataset
|
|
The dataset to work from
|
|
shape : torch.Size
|
|
The size of the probe array to simulate
|
|
det_slice : slice
|
|
A slice or tuple of slices corresponding to the detector region in Fourier space
|
|
convergence_angle : float
|
|
The convergence angle of the probe, in mrad.
|
|
propagation_distance : float
|
|
Default is no propagation, an amount to propagate the guessed probe from it's focal point
|
|
oversampling : int
|
|
Default 1, the width of the region of pixels in the wavefield to bin into a single detector pixel
|
|
|
|
Returns
|
|
-------
|
|
torch.Tensor
|
|
The complex-style tensor storing the probe guess
|
|
"""
|
|
|
|
# The basis in the dataset should describe the basis of the probe's
|
|
# Fourier transform, even if the probe is simulated on larger stage than
|
|
# the detector. The only issue is if the probe is oversampled in
|
|
# Fourier space (simulated on a larger stage in real space). That factor
|
|
# is defined by oversampling.
|
|
|
|
probe_basis = dataset.detector_geometry['basis'] / oversampling
|
|
|
|
mean_im = t.mean(dataset.patterns,dim=0)
|
|
center = image_processing.centroid(mean_im)
|
|
|
|
Is = t.arange(shape[0], dtype=t.float32) - center[0]
|
|
Js = t.arange(shape[1], dtype=t.float32) - center[1]
|
|
Is,Js = t.meshgrid(Is,Js)
|
|
|
|
Rs = t.tensordot(probe_basis,t.stack([Is,Js]),dims=1)
|
|
forward = t.Tensor([0,0,1])
|
|
Rs = (forward * dataset.detector_geometry['distance'])[:,None,None] + Rs
|
|
dirs = Rs / t.linalg.norm(Rs,dim=(0,))
|
|
angles = t.acos(t.tensordot(forward,dirs,dims=1)) * 1000 #in mrad
|
|
|
|
probe_fft = t.zeros(shape,dtype=t.complex64)
|
|
probe_fft[angles<convergence_semiangle] = 1
|
|
|
|
probe_mean = t.mean(t.abs(probe_fft)**2)
|
|
diff_mean = t.mean(mean_im)
|
|
probe_fft = probe_fft * t.sqrt(diff_mean / probe_mean)
|
|
|
|
probe_guess = inverse_far_field(probe_fft)
|
|
|
|
if propagation_distance is not None:
|
|
# It's probably worth checking if this is correct when oversampling
|
|
# is not 1
|
|
|
|
probe_shape = t.as_tensor(tuple(probe_guess.shape))
|
|
|
|
# Start by recalculating the probe basis from the given information
|
|
det_basis = t.as_tensor(dataset.detector_geometry['basis'])
|
|
basis_dirs = det_basis / t.norm(det_basis, dim=0)
|
|
distance = dataset.detector_geometry['distance']
|
|
probe_basis = basis_dirs * dataset.wavelength * distance / \
|
|
(probe_shape * t.norm(det_basis,dim=0))
|
|
|
|
# Then package everything as it's needed
|
|
probe_spacing = t.norm(probe_basis,dim=0).numpy()
|
|
probe_shape = probe_shape.numpy().astype(np.int32)
|
|
|
|
# And generate the propagator
|
|
AS_prop = generate_angular_spectrum_propagator(probe_shape, probe_spacing, dataset.wavelength, propagation_distance)
|
|
|
|
probe_guess = near_field(probe_guess,AS_prop)
|
|
|
|
# Finally, place this probe in a full-sized array if there is oversampling
|
|
final_probe = t.zeros(shape,dtype=t.complex64)
|
|
left = shape[0]//2 - probe_guess.shape[0] // 2
|
|
top = shape[1]//2 - probe_guess.shape[1] // 2
|
|
final_probe[left:left+probe_guess.shape[0],
|
|
top:top+probe_guess.shape[1]] = probe_guess
|
|
|
|
return final_probe
|
|
|
|
|
|
def RPI_spectral_init(pattern, probe, obj_shape, n_modes=1, mask=None, background=None):
|
|
|
|
# First, check if the probe is a single mode or many modes.
|
|
# If the probe is many modes, orthogonalize it and use the top mode
|
|
# for initialization
|
|
if probe.dim() == 4:
|
|
probe = orthogonalize_probes(probe)[0]
|
|
|
|
pad0l = (probe.shape[-2] - obj_shape[0])//2
|
|
pad0r = probe.shape[-2] - obj_shape[0] - pad0l
|
|
pad1l = (probe.shape[-1] - obj_shape[1])//2
|
|
pad1r = probe.shape[-1] - obj_shape[1] - pad1l
|
|
|
|
def a_dagger(im):
|
|
im = t.tensor(im.reshape(obj_shape)).to(dtype=t.complex64)
|
|
im = inverse_far_field(pad(far_field(im), (pad1l,pad1r,pad0l,pad0r)))
|
|
exit_wave = probe * im
|
|
return far_field(exit_wave).numpy().ravel()
|
|
|
|
def a(measured):
|
|
measured = t.tensor(measured.reshape(pattern.shape[0],pattern.shape[1])).to(dtype=t.complex64)
|
|
im = inverse_far_field(measured)
|
|
multiplied = t.conj(probe) * im
|
|
backplane = far_field(multiplied)
|
|
clipped = backplane[pad0l:pad0l+obj_shape[0],
|
|
pad1l:pad1l+obj_shape[1]]
|
|
return (inverse_far_field(clipped)).numpy().ravel()
|
|
|
|
patsize = pattern.shape[0]*pattern.shape[1]
|
|
imsize = obj_shape[0]*obj_shape[1]
|
|
probesize = probe.shape[0]*probe.shape[1]
|
|
A_dagger = spla.LinearOperator((patsize, imsize),matvec=a_dagger)
|
|
A = spla.LinearOperator((imsize,patsize),matvec=a)
|
|
|
|
# Correct the pattern for the background and mask
|
|
np_pattern = pattern.numpy()
|
|
if background is not None:
|
|
np_pattern = np_pattern - background.numpy()
|
|
if mask is not None:
|
|
np_pattern = np_pattern * mask.numpy()
|
|
np_pattern = np.abs(np_pattern.ravel())
|
|
|
|
realspace_intensities = A * A_dagger * np.ones(imsize)
|
|
vec = A_dagger * realspace_intensities
|
|
def y(measured):
|
|
#return measured * np_pattern
|
|
# This normalizes the intensities to account for the fact
|
|
# that some pixels draw from way more spots on the detector than
|
|
# others
|
|
return measured * np_pattern / np.abs(vec)
|
|
|
|
Y = spla.LinearOperator((patsize, patsize),matvec=y)
|
|
eigval, z0 = spla.eigs(A * Y * A_dagger, k=n_modes, which='LM')
|
|
z0 = z0.transpose().reshape(n_modes,obj_shape[0], obj_shape[1])
|
|
|
|
# Now we set the overall scale and relative weights of the guess
|
|
scale_factor = np.sqrt(np.sum(np_pattern) /
|
|
t.sum(t.abs(probe)**2).numpy())
|
|
relative_weights = eigval / np.sum(eigval**2)
|
|
|
|
z0 = z0 * (scale_factor * relative_weights[:,None,None])
|
|
# Now we have to normalize the modes by their eigenvalues
|
|
|
|
return t.as_tensor(z0, dtype=t.complex64)
|
|
|
|
|
|
def generate_subdominant_modes(dominant_mode, n_modes, circular=True):
|
|
"""Generates guesses of subdominant modes based on spatial derivatives
|
|
|
|
The idea here is that vibration is an extremely common cause of
|
|
spatial incoherence. This typically presents as subdominant modes that
|
|
look like the derivative of the dominant mode along various axes.
|
|
|
|
Therefore, it is a good starting guess if the guess of the dominant
|
|
mode is reasonable. The output probes will all be normalized to have
|
|
the same intensity as the input dominant probe
|
|
|
|
Parameters
|
|
----------
|
|
dominant_mode : array
|
|
The complex-valued dominant mode to work from
|
|
n_modes : int
|
|
The number of additional modes to create
|
|
circular : bool
|
|
Default True, whether to use circular modes (x+iy) or linear (x,y)
|
|
Returns
|
|
-------
|
|
torch.Tensor
|
|
The complex-style tensor storing the probe guesses
|
|
"""
|
|
|
|
# This generates a list of tuples, where each tuple (a,b) corresponds
|
|
# a term (kx^a)*(ky^b), or (kx+iky)^a*(kx-iky)^b for circular modes
|
|
def make_orders(n_orders):
|
|
if n_orders >1024:
|
|
raise KeyError('Are you sure you want this many orders?')
|
|
i = 1
|
|
n = 0
|
|
total = 0
|
|
while True:
|
|
if total >= n_orders:
|
|
break
|
|
|
|
yield (n,i-n)
|
|
total += 1
|
|
if n<i:
|
|
n+=1
|
|
else:
|
|
n = 0
|
|
i += 1
|
|
|
|
# Gotta check and convert to pytorch if it's numpy
|
|
|
|
# First step is to take the FFT of the probe
|
|
dominant_fft = far_field(dominant_mode)
|
|
|
|
shape = dominant_mode.shape
|
|
center = ((shape[-2]-1)//2, (shape[-1]-1)//2)
|
|
|
|
i, j = np.mgrid[:shape[-2], :shape[-1]]
|
|
i = t.tensor(i - center[0]).to(dtype=dominant_fft.dtype,
|
|
device=dominant_fft.device)
|
|
j = t.tensor(j - center[1]).to(dtype=dominant_fft.dtype,
|
|
device=dominant_fft.device)
|
|
|
|
if circular:
|
|
a = i + 1j * j
|
|
b = i - 1j * j
|
|
else:
|
|
a = i
|
|
b = j
|
|
|
|
probe_norm = t.sum(t.abs(dominant_mode)**2)
|
|
# Then we need to multiply that FFT by various phase masks and IFFT
|
|
probes = []
|
|
for a_order,b_order in make_orders(n_modes):
|
|
mask = reduce(t.mul,[a]*a_order+[b]*b_order)
|
|
new_probe = inverse_far_field(mask * dominant_fft)
|
|
probes.append(new_probe * (probe_norm / t.sum(t.abs(new_probe))))
|
|
|
|
return t.stack(probes)
|