This commit is contained in:
Anastasiia Kutakh
2021-08-24 01:26:11 -04:00
parent aadd720f1e
commit d0bca7633b
2 changed files with 177 additions and 31 deletions
+147 -6
View File
@@ -73,10 +73,6 @@ class PolarizedFancyPtycho(FancyPtycho):
else:
x = -1j
# if probe_size is None:
# probe = tools.initializers.SHARP_style_probe(dataset, probe_shape, det_slice, propagation_distance=propagation_distance, oversampling=oversampling)
# else:
# probe = tools.initializers.gaussian_probe(dataset, probe_basis, probe_shape, probe_size, propagation_distance=propagation_distance)
probe = model.probe.detach()
probe = t.cat((probe, probe * x), dim=-3)
probe_max = t.max(t.abs(probe))
@@ -104,7 +100,153 @@ class PolarizedFancyPtycho(FancyPtycho):
polarizers = [tools.polarization.generate_linear_polarizer(i * 45) for i in range(3)]
# WHAT IS INDEX?
@classmethod
def from_dataset2(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, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um'):
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 include the *extras to make this work even with datasets, like
# polarization dependent datasets, that might toss out extra inputs
(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
if auto_center:
center = tools.image_processing.centroid(t.sum(patterns, dim=0))
else:
center = None
if left_polarized:
x = 1j
else:
x = -1j
# 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)
probe_shape = t.stack((2, probe_shape), dim=-3)
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)
# Next generate the object geometry from the probe geometry and
# the translations
pix_translations = tools.interactions.translations_to_pixel(probe_basis, translations, surface_normal=surface_normal)
obj_size, min_translation = tools.initializers.calc_object_setup(probe_shape, pix_translations, padding=200)
if hasattr(dataset, 'background') and dataset.background is not None:
background = t.sqrt(dataset.background)
else:
background = None
# Finally, initialize the probe and object using this information
if probe_size is None:
probe = tools.initializers.SHARP_style_probe(dataset, probe_shape, det_slice, propagation_distance=propagation_distance, oversampling=oversampling)
else:
probe = tools.initializers.gaussian_probe(dataset, probe_basis, probe_shape, probe_size, propagation_distance=propagation_distance)
# Now we initialize all the subdominant probe modes
probe_max = t.max(t.abs(probe))
probe_stack = [0.01 * probe_max * t.rand(probe.shape, dtype=probe.dtype) for i in range(n_modes - 1)]
probe = t.stack([probe, ] + probe_stack)
# probe = t.stack([tools.propagators.far_field(probe),] + probe_stack)
probe_x, probe_y = probe, probe * x
probe = t.stact((probe_x, probe_y), dim=-3)
a = t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5))
b = t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5))
c = t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5))
d = t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5))
ab = t.stack((a, b), dim=-3)
cd = t.stack((c, d), dim=-3)
obj = t.stack((ab, cd), dim=-4)
det_geo = dataset.detector_geometry
translation_offsets = 0 * (t.rand((len(dataset), 2)) - 0.5)
if dm_rank is not None and dm_rank != 0:
if dm_rank > n_modes:
raise KeyError('Density matrix rank cannot be greater than the number of modes. Use dm_rank = -1 to use a full rank matrix.')
elif dm_rank == -1:
# dm_rank == -1 is defined to mean full-rank
dm_rank = n_modes
Ws = t.zeros(len(dataset), dm_rank, n_modes, dtype=t.complex64)
# Start with as close to the identity matrix as possible,
# cutting of when we hit the specified maximum rank
for i in range(0, dm_rank):
Ws[:, i, i] = 1
else:
# dm_rank == None or dm_rank = 0 triggers a special case where
# a standard incoherent multi-mode model is used. This is the
# default, because it is so common.
# In this case, we define a set of weights which only has one index
Ws = t.ones(len(dataset))
if hasattr(dataset, 'mask') and dataset.mask is not None:
mask = dataset.mask.to(t.bool)
else:
mask = None
if probe_support_radius is not None:
probe_support = t.zeros(probe[0].shape, dtype=t.bool)
xs, ys = np.mgrid[:probe.shape[-2], :probe.shape[-1]]
xs = xs - np.mean(xs)
ys = ys - np.mean(ys)
Rs = np.sqrt(xs**2 + ys**2)
probe_support[Rs < probe_support_radius] = 1
probe = probe * probe_support[None, :, :]
else:
probe_support = None
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,
oversampling=oversampling,
loss=loss, units=units)
def interaction(self, index, translations, polarizer, analyzer, test=False):
# Step 1 is to convert the translations for each position into a
@@ -137,7 +279,6 @@ class PolarizedFancyPtycho(FancyPtycho):
shift_probe=True, multiple_modes=True, polarized=True)
analyzed_exit_waves = polarization.apply_linear_polarizer(exit_waves, analyzer)
# print('POLARIZED FANCY PTYCHO INTERACTION OBJ')
return analyzed_exit_waves
+30 -25
View File
@@ -14,7 +14,7 @@ __all__ = ['apply_linear_polarizer',
'apply_circular_polarizer',
'apply_jones_matrix',
'generate_linear_polarizer',
'generate_phase_retarder']
'generate_birefringent_obj']
# Abe - split these into two functions
@@ -36,10 +36,10 @@ def generate_linear_polarizer(pol_angle):
cd = t.stack((c, d), dim=-1)
jones_matrices = t.stack((ab, cd), dim=-2)
if single_angle:
return jones_matrices[0].to(dtype=t.cfloat)
return jones_matrices[0].to(dtype=t.cfloat)
else:
return jones_matrices.to(dtype=t.cfloat)
def apply_linear_polarizer(probe, polarizer, multiple_modes=True, transpose=True):
"""
@@ -56,7 +56,7 @@ def apply_linear_polarizer(probe, polarizer, multiple_modes=True, transpose=True
Returns:
--------
linearly polarized probe: t.Tensor
(N)(P)x2x1xMxL
(N)(P)x2x1xMxL
"""
jones_matrices = generate_linear_polarizer(polarizer)
return apply_jones_matrix(probe, jones_matrices, transpose=transpose, multiple_modes=multiple_modes)
@@ -75,19 +75,19 @@ def apply_jones_matrix(probe, jones_matrix, transpose=True, multiple_modes=True)
probe: t.Tensor
A (N)(P)x2xMxL tensor representing the probe
jones_matrix: t.tensor
(N)x2x2x(M)x(L)
(N)x2x2x(M)x(L)
Returns:
--------
a probe with the jones matrix applied: t.Tensor
(N)(P)x2xMxL
(N)(P)x2xMxL
"""
if transpose:
if jones_matrix.dim() < 4:
jones_matrix = jones_matrix[..., None, None]
if multiple_modes:
jones_matrix = jones_matrix.unsqueeze(-5)
jones_matrix = jones_matrix.unsqueeze(-5)
probe = probe[..., None, :, :]
# if jones matrices do not differ from pattern to pattern
if probe.dim() > jones_matrix.dim():
@@ -96,19 +96,19 @@ def apply_jones_matrix(probe, jones_matrix, transpose=True, multiple_modes=True)
elif jones_matrix.dim() > probe.dim():
probe = probe.unsqueeze(0)
# print('apply jonesmatrix: probe', probe.shape, 'matrix:', jones_matrix)
jones_matrix = jones_matrix.transpose(-1, -3).transpose(-2, -4)
jones_matrix = jones_matrix.transpose(-1, -3).transpose(-2, -4)
probe = probe.transpose(-1, -3).transpose(-2, -4)
output = t.matmul(jones_matrix, probe).transpose(-2, -4).transpose(-1, -3).squeeze(-3)
else:
raise NotImplementedError
return output
def apply_phase_retardance(probe, phase_shift, multiple_modes=True):
"""
Shifts the y-component of the field wrt the x-component by a given phase shift
Shifts the y-component of the field wrt the x-component by a given phase shift
Parameters:
----------
@@ -120,7 +120,7 @@ def apply_phase_retardance(probe, phase_shift, multiple_modes=True):
Returns:
--------
probe: t.Tensor
(...)x2x1xMxL
(...)x2x1xMxL
"""
theta = t.as_tensor(phase_shift, dtype=t.float32)
theta = t.deg2rad(theta)
@@ -140,11 +140,11 @@ def apply_circular_polarizer(probe, left_polarized=True, multiple_modes=True):
A (...)x2xMxL tensor representing the probe
left_polarizd: bool
True for the left-polarization, False for the right
Returns:
--------
circularly polarized probe: t.Tensor
(...)x2xMxL
(...)x2xMxL
"""
probe = probe.to(dtype=t.cfloat)
if left_polarized:
@@ -166,7 +166,7 @@ def apply_quarter_wave_plate(probe, fast_axis_angle, multiple_modes=True):
Returns:
--------
polarized probe: t.Tensor
(...)x2x1xMxL
(...)x2x1xMxL
"""
probe = probe.to(dtype=t.cfloat)
theta = math.radians(fast_axis_angle)
@@ -174,7 +174,7 @@ def apply_quarter_wave_plate(probe, fast_axis_angle, multiple_modes=True):
jones_matrix = exponent* t.tensor([[(cos(theta))**2 + 1j * (sin(theta))**2, (1 - 1j) * sin(theta) * cos(theta)], [(1 - 1j) * sin(theta) * cos(theta), (sin(theta))**2 + 1j * (cos(theta))**2]]).to(dtype=t.cfloat)
out = apply_jones_matrix(probe, jones_matrix, multiple_modes=multiple_modes)
return out
return out
def apply_half_wave_plate(probe, fast_axis_angle, multiple_modes=True):
"""
@@ -188,7 +188,7 @@ def apply_half_wave_plate(probe, fast_axis_angle, multiple_modes=True):
Returns:
--------
polarized probe: t.Tensor
(...)x2x1xMxL
(...)x2x1xMxL
"""
probe = probe.to(dtype=t.cfloat)
theta = math.radians(fast_axis_angle)
@@ -196,19 +196,24 @@ def apply_half_wave_plate(probe, fast_axis_angle, multiple_modes=True):
jones_matrix = exponent * t.tensor([[(cos(theta))**2 - (sin(theta))**2, 2 * sin(theta) * cos(theta)], [2 * sin(theta) * cos(theta), (sin(theta))**2 - (cos(theta))**2]]).to(dtype=t.cfloat)
out = apply_jones_matrix(probe, jones_matrix, multiple_modes=multiple_modes)
return out
def generate_phase_retarder(fast_axis=0, phase=0):
phase = t.as_tensor(phase).to(dtype=t.float32)
phase = t.deg2rad(phase)
def coord_rot(angle):
return out
def generate_birefringent_obj(fast_axis=90, phase_ret=10, atten_fast=1, atten_ret=1, global_phase=0):
def to_rad(angle):
angle = t.as_tensor(angle, dtype=t.float32)
angle = t.deg2rad(angle)
return angle
fast_axis = to_rad(fast_axis)
phase_ret = to_rad(phase_ret)
global_phase = to_rad(global_phase)
def coord_rot(angle):
a = t.stack((t.cos(angle), t.sin(angle)), dim=-1)
b = t.stack((-t.sin(angle), t.cos(angle)), dim=-1)
return t.stack((a, b), dim=-2).to(dtype=t.cfloat)
r1 = coord_rot(-fast_axis)
r2 = coord_rot(fast_axis)
p = t.as_tensor([[1, 0], [0, t.exp(phase*1j)]], dtype=t.cfloat)
return t.matmul(r1, t.matmul(p, r2))
p = t.exp(global_phase * 1j) * t.as_tensor([[atten_fast, 0], [0, atten_ret * t.exp(phase_ret*1j)]], dtype=t.cfloat)
return t.matmul(r1, t.matmul(p, r2))