diff --git a/CDTools/models/fancy_ptycho.py b/CDTools/models/fancy_ptycho.py index 89a3750..3ef8a51 100644 --- a/CDTools/models/fancy_ptycho.py +++ b/CDTools/models/fancy_ptycho.py @@ -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 diff --git a/CDTools/models/polarized_fancy_ptycho.py b/CDTools/models/polarized_fancy_ptycho.py index cd0027b..e64b374 100644 --- a/CDTools/models/polarized_fancy_ptycho.py +++ b/CDTools/models/polarized_fancy_ptycho.py @@ -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, :, :] diff --git a/CDTools/tools/interactions/interactions.py b/CDTools/tools/interactions/interactions.py index 2be19a1..009c441 100644 --- a/CDTools/tools/interactions/interactions.py +++ b/CDTools/tools/interactions/interactions.py @@ -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') diff --git a/CDTools/tools/polarization/polarization.py b/CDTools/tools/polarization/polarization.py index c3ad200..bcf63df 100644 --- a/CDTools/tools/polarization/polarization.py +++ b/CDTools/tools/polarization/polarization.py @@ -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) \ No newline at end of file