changed polarization and polarized_fancy_ptycho interaction

This commit is contained in:
Anastasiia Kutakh
2021-07-29 02:59:37 -04:00
parent a4d70bc375
commit 6924ce58fe
4 changed files with 63 additions and 85 deletions
+8 -2
View File
@@ -123,8 +123,11 @@ class FancyPtycho(CDIModel):
@classmethod
def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um'):
def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um', polarized=False):
# if polarized=True, this function takes care of the dataset iterator by taking into account the polarizer and analyzer components
# however, it drops them afterwards and treats the dataset as if it's not polarized
wavelength = dataset.wavelength
det_basis = dataset.detector_geometry['basis']
det_shape = dataset[0][1].shape
@@ -133,7 +136,10 @@ class FancyPtycho(CDIModel):
# always do this on the cpu
get_as_args = dataset.get_as_args
dataset.get_as(device='cpu')
(indices, translations), patterns = dataset[:]
if not polarized:
(indices, translations), patterns = dataset[:]
else:
(indices, translations, polarizer, analyzer), patterns = dataset[:]
dataset.get_as(*get_as_args[0],**get_as_args[1])
# Set to none to avoid issues with things outside the detector
+24 -63
View File
@@ -9,6 +9,7 @@ from datetime import datetime
import numpy as np
from scipy import linalg as sla
from copy import copy
from CDTools.tools import polarization
__all__ = ['PolarizedFancyPtycho']
@@ -49,60 +50,28 @@ class PolarizedFancyPtycho(FancyPtycho):
self.analyzer_offsets = t.nn.Parameter(t.tensor(analyzer_offsets).to(dtype=t.float32)) / analyzer_scale
@classmethod
def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um'):
def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um', left_polarized=True):
super(PolarizedFancyPtycho, cls).from_dataset(dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um')
model = FancyPtycho.from_dataset(dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um', polarized=True)
# Mutate the class to its subclass
model.__class__ = cls
# always do this on the cpu
get_as_args = dataset.get_as_args
dataset.get_as(device='cpu')
(indices, translations), patterns = dataset[:]
dataset.get_as(*get_as_args[0],**get_as_args[1])
# Set to none to avoid issues with things outside the detector
model.probe.data = model.probe.data.unsqueeze(-3)
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)
scalar_probe_shape = probe_shape.clone()
probe_shape = t.stack((probe_shape[:-2], t.tensor([2,]), probe_shape([-2:])))
obj_size, min_translation = tools.initializers.calc_object_setup(scalar_probe_shape, pix_translations, padding=200)
obj_size = t.cat((t.tensor([2, 2]), obj_size))
tensor vs tensor.data
# Finally, initialize the probe and object using this information
if probe_size is None:
model.probe.data = tools.initializers.SHARP_style_probe(dataset, scalar_probe_shape, det_slice, propagation_distance=propagation_distance, oversampling=oversampling, polarized=True)
if left_polarized:
x = 1j
else:
probe = tools.initializers.gaussian_probe(dataset, probe_basis, scalar_probe_shape, probe_size, propagation_distance=propagation_distance, polarized=True)
x = -1j
# model.probe.data = t.stack((model.probe.data.to(dtype=t.cfloat), x * model.probe.data.to(dtype=t.cfloat)), dim=-3)
# obj = t.stack((model.obj.data, model.obj.data), dim=-3)
# model.obj.data = t.stack((obj, obj), dim=-4)
print('probe guess shape:', model.probe.shape)
print('object guess shape:', model.obj.shape)
# tensor vs tensor.data
return model
return cls(wavelength, det_geo, probe_basis, probe, obj,
detector_slice=det_slice,
surface_normal=surface_normal,
min_translation=min_translation,
translation_offsets = translation_offsets,
weights=Ws, mask=mask, background=background,
translation_scale=translation_scale,
saturation=saturation,
probe_support=probe_support,
obj_support=obj_support,
oversampling=oversampling,
loss=loss,units=units)
def interaction(self, index, translations, polarizer, analyzer):
@@ -118,7 +87,7 @@ class PolarizedFancyPtycho(FancyPtycho):
# This restricts the basis probes to stay within the probe support
basis_prs = self.probe * self.probe_support[...,:,:] # This makes no sense
basis_prs = self.probe * self.probe_support[...,:,:] # This makes no sense
# self.probe is an Nx2xXxY stach of probes
# Now we construct the probes for each shot from the basis probes
@@ -129,34 +98,26 @@ class PolarizedFancyPtycho(FancyPtycho):
prs = Ws[...,None,None,None,None] * basis_prs
else:
raise NotImplementedError('Unstable Modes not Implemented for polarized light')
# If a frame-by-frame weight matrix is defined
# This takes the dot product of all the weight matrices with
# the probes. The output has dimensions of translation, then
# coherent mode index, then x,y, and then complex index
# Maybe this can be done with a matmul now?
prs = t.sum(Ws[...,None,None] * basis_prs, axis=-3)
# Now we actually do the interaction, using the sinc subpixel
# translation model as per usual
# I DON'T KNOW WHAT PROBE NORM IS (AS WELL AS OBJ SUPP AND PROBE SUPP)
pol_probes = polarization.apply_linear_polarizer(prs, polarizer)
pol_probes = polarization.apply_polarizer(polarizer, pol_probes)
exit_waves = self.probe_norm * tools.interactions.ptycho_2D_sinc(
prs, self.obj_support * self.obj,pix_trans,
shift_probe=True, multiple_modes=True, polarized=True)
analyzed_exit_waves = polarization.apply_polarizer(analyzer, exit_waves)
analyzed_exit_waves = polarization.apply_linear_polarizer(exit_waves, analyzer)
#exit_waves = self.probe_norm * tools.interactions.ptycho_2D_round(
# prs, self.obj_support * self.obj,pix_trans,
# multiple_modes=True)
return exit_waves
return analyzed_exit_waves
def vectorial_wavefields(wavefields, func, *args. **kwargs):
def vectorial_wavefields(wavefields, func, *args, **kwargs):
wavefields_x = wavefields[..., 0, :, :, :]
wavefields_y = wavefields[..., 1, :, :, :]
out_x = func(wavefields_x. *args, **kwargs)
out_x = func(wavefields_x, *args, **kwargs)
out_y = func(wavefields_y, *args, **kwargs)
out = t.stack((out_x, out_y), dim=-4)
return out[..., None, :, :]
+1 -5
View File
@@ -438,8 +438,6 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multi
tr[1]:tr[1]+probe.shape[-1]]
for tr in integer_translations])
if polarized:
polarized_probes = t.stack([polarization.apply_linear_polarizer(probe, polarizer[idx]) for idx in range(polarizer.shape)])
# Nx(P)x2x1xMxL tensor
selections = t.stack([obj[:, :,tr[0]:tr[0]+probe.shape[-2],
tr[1]:tr[1]+probe.shape[-1]]
for tr in integer_translations])
@@ -491,9 +489,7 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multi
shift_probe = shifted_probe.transpose(-1, -3).transpose(-2, -4)
selections = selections.transpose(-1, -3).transpose(-2, -4)
output = t.matmul(selections, shift_probe).transpose(-1, -3).transpose(-2, -4)
output = t.stack([polarization.apply_linear_polarizer(output[idx, ...], analyzer[idx]) for idx in range(analyzer.shape)])
else:
raise NotImplementedError('Object shift not yet implemented')
+30 -15
View File
@@ -11,7 +11,7 @@ __all__ = ['apply_linear_polarizer',
'apply_circular_polarizer',
'apply_jones_matrix']
def apply_linear_polarizer(probe, polar_angle):
def apply_linear_polarizer(probe, polarizer):
"""
Applies a linear polarizer to the probe
@@ -19,8 +19,9 @@ def apply_linear_polarizer(probe, polar_angle):
----------
probe: t.Tensor
A (...)x2x1xMxL tensor representing the probe, MxL - the size of the probe
polar_angle: float
The angle between the fast-axis of the linear polarizer and the horizontal axis
polarizer: t.Tensor
A 1D tensor representing the polarizer angles for each of the patterns (or a single tensor of shape (1))
Returns:
--------
@@ -28,18 +29,21 @@ def apply_linear_polarizer(probe, polar_angle):
(...)x2x1xMxL
"""
probe = probe.to(dtype=t.cfloat)
theta = math.radians(polar_angle)
polarizer = t.tensor([[(cos(theta)) ** 2, sin(2 * theta) / 2], [sin(2 * theta) / 2, sin(theta) ** 2]]).to(dtype=t.cfloat)
if len(polarizer) == 1:
theta = math.radians(polarizer)
polarizer = t.tensor([[(cos(theta)) ** 2, sin(2 * theta) / 2], [sin(2 * theta) / 2, sin(theta) ** 2]]).to(dtype=t.cfloat)
else:
pol_cos = lambda idx: cos(math.radians(polarizer[idx]))
pol_sin = lambda idx: sin(math.radians(polarizer[idx]))
jones_matrices = t.stack(([t.tensor([[(pol_cos(idx)) ** 2, pol_sin(idx) * pol_cos(idx)], [pol_sin(idx) * pol_cos(idx), (pol_sin(idx)) ** 2]]).to(dtype=t.cfloat) for idx in range(len(polarizer))]))
# I haven't figured out how to multiply tensors using tensordot yet,
# so we'll be temporarily using matmul on the previously tranposed vector
# (since it returns the matrix multiplication product over the last two dimensions
# (since it returns the matrix multiplication product over the last two dimensions)
#Swap the dimensions for the prober to be (...)xMxLx2x1 to perform matmul on it
probe = probe.transpose(-1, -3).transpose(-2, -4)
polarized_probe = t.matmul(polarizer, probe)
# Transpose it back
return polarized_probe.transpose(-1, -3).transpose(-2, -4)
return apply_jones_matrix(probe, jones_matrices)
def apply_jones_matrix(probe, jones_matrix):
"""
@@ -48,17 +52,23 @@ def apply_jones_matrix(probe, jones_matrix):
Parameters:
----------
probe: t.Tensor
A (...)x2xMxL tensor representing the probe
A (...N)x2xMxL tensor representing the probe
jones_matrix: t.tensor
(...)x2x2
(N)x2x2
Returns:
--------
linearly polarized probe: t.Tensor
(...)x2xMxL
(...N)x2xMxL
"""
return t.tensordot(jones_matrix,probe,dims=[[-1,],[-3]])
jones_matrix = jones_matrix[..., None, None, :, :]
# make it (N)x1x1x2x2
probe = probe[..., None, :, :]
probe = probe.transpose(-1, -3).transpose(-2, -4)
# (...N)xMxLx2x1
output = t.matmul(jones_matrix, probe).transpose(-2, -4).transpose(-1, -3)
# (...N)x2x1xMxL
return output.squeeze(-3)
def apply_phase_retardance(probe, phase_shift):
"""
@@ -164,3 +174,8 @@ def apply_half_wave_plate(probe, fast_axis_angle):
# print(apply_phase_retardance(probe, 29).shape)
# print(apply_half_wave_plate(probe, 29).shape)
# print(apply_quarter_wave_plate(probe, 29).shape)
# d = t.cat(([a for i in range(3)]))
# print(d.shape)