mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-30 13:52:10 +02:00
changed polarization and polarized_fancy_ptycho interaction
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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, :, :]
|
||||
|
||||
@@ -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')
|
||||
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user