Merge branch 'polarization' of github.mit.edu:Scattering/CDTools into polarization

This commit is contained in:
Abe Levitan
2021-08-25 17:06:34 -04:00
6 changed files with 522 additions and 281 deletions
+52 -52
View File
@@ -17,14 +17,14 @@ __all__ = ['translations_to_pixel', 'pixel_to_translations',
def translations_to_pixel(basis, translations, surface_normal=t.Tensor([0.,0.,1.])):
"""Takes real space translations and outputs them in pixel space
This works for any 2D ptychography geometry. It takes in
A set of translations in (x,y) space and outputs the same translations
in internal pixel units perpendicular to the detector.
in internal pixel units perpendicular to the detector.
It uses information on the wavefield basis and, if defined, the
sample normal, to perform the conversion.
The assumed geometry is incoming radiation with a wavevector parallel
to the +z axis, [0,0,1]. The default sample orientation has a surface
normal parallel to this direction
@@ -33,7 +33,7 @@ def translations_to_pixel(basis, translations, surface_normal=t.Tensor([0.,0.,1.
----------
basis : torch.Tensor
The real space basis the wavefields are defined in
translations : torch.Tensor
translations : torch.Tensor
A Jx3 stack of real-space translations, or a single translation
surface_normal : torch.Tensor
Optional, the sample's surface normal
@@ -68,18 +68,18 @@ def translations_to_pixel(basis, translations, surface_normal=t.Tensor([0.,0.,1.
return pixel_translations[0]
else:
return pixel_translations
def pixel_to_translations(basis, pixel_translations, surface_normal=t.Tensor([0,0,1])):
"""Takes pixel-space translations and outputs them in real space
This works for any 2D ptychography geometry. It takes in
A set of internal pixel unit translations in (i,j) space and
outputs the same translations real (x,y) space
It uses information on the wavefield basis and, if defined, the
sample normal, to perform the conversion.
The assumed geometry is incoming radiation with a wavevector parallel
to the +z axis, [0,0,1]. The default sample orientation has a surface
normal parallel to this direction. Because of this, the z direction
@@ -96,7 +96,7 @@ def pixel_to_translations(basis, pixel_translations, surface_normal=t.Tensor([0,
Returns
-------
real_translations : torch.Tensor
real_translations : torch.Tensor
A Jx3 stack of real-space translations, or a single translation
"""
projection_1 = t.Tensor([[1,0,0],
@@ -129,7 +129,7 @@ def pixel_to_translations(basis, pixel_translations, surface_normal=t.Tensor([0,
def project_translations_to_sample(sample_basis, translations):
"""Takes real space translations and outputs them in pixels in a sample basis
This projection function is designed for the Bragg2DPtycho class. More
broadly, it works to take a set of translations in the lab frame and
convert each one into two values. First, an (i,j) value in pixels
@@ -145,7 +145,7 @@ def project_translations_to_sample(sample_basis, translations):
relative amount the probe needs to be propagated to reach any given
location), a positive motion along the z-axis of the probe forming optics
will lead to a negative propagation distance.
The assumed geometry is incoming radiation with a wavevector parallel
to the +z axis, [0,0,1].
@@ -153,7 +153,7 @@ def project_translations_to_sample(sample_basis, translations):
----------
sample_basis : torch.Tensor
The real space basis the wavefields are defined in
translations : torch.Tensor
translations : torch.Tensor
A Jx3 stack of real-space translations, or a single translation
Returns
@@ -171,7 +171,7 @@ def project_translations_to_sample(sample_basis, translations):
# Then we calculate a matrix which can do the projection
propagation_dir = t.Tensor(np.array([0,0,1])).to(
device=surface_normal.device,
dtype=surface_normal.dtype)
@@ -179,13 +179,13 @@ def project_translations_to_sample(sample_basis, translations):
I = t.eye(3).to(
device=surface_normal.device,
dtype=surface_normal.dtype)
# Here we're setting up a matrix-vector equation mat*answer=input
# At some point ger will need to be replaced by outer, but for now
# outer many places still don't have new enough versions of torch.
mat = t.cat((I - t.ger(propagation_dir,propagation_dir),
surface_normal.unsqueeze(0)))
# And we invert the matrix to do the projection
projector = t.pinverse(mat)[:,:3].to(device=translations.device,
dtype=translations.dtype)
@@ -200,7 +200,7 @@ def project_translations_to_sample(sample_basis, translations):
device=translations.device,
dtype=translations.dtype)
sample_projection = t.mm(basis_vectors_inv, projector).t()
prop_projection = t.mm(propagation_dir_inv, projector).t()
@@ -218,9 +218,9 @@ def project_translations_to_sample(sample_basis, translations):
return pixel_translations[0], propagations[0]
else:
return pixel_translations, propagations
def ptycho_2D_round(probe, obj, translations, multiple_modes=False, upsample_obj=False):
"""Returns a stack of exit waves without accounting for subpixel shifts
@@ -229,15 +229,15 @@ def ptycho_2D_round(probe, obj, translations, multiple_modes=False, upsample_obj
dimension as the translation index and the final dimensions
corresponding to the detector. The exit waves are calculated by
shifting the probe by the rounded value of the translation
If multiple_modes is set to False, any additional dimensions in the
ptycho_2D_round function will be assumed to correspond to the translation
index. If multiple_modes is set to true, the (-4th) dimension of the probe
will always be assumed to be defining a set of (P) incoherently mixing
modes to be broadcast all translation indices. If any additional dimensions
closer to the start exist, they will be assumed to be translation indices
Parameters
----------
probe : torch.Tensor
@@ -251,7 +251,7 @@ def ptycho_2D_round(probe, obj, translations, multiple_modes=False, upsample_obj
Returns
-------
exit_waves : torch.Tensor
exit_waves : torch.Tensor
An (N)x(P)xMxL tensor of the calculated exit waves
"""
@@ -260,9 +260,9 @@ def ptycho_2D_round(probe, obj, translations, multiple_modes=False, upsample_obj
translations = translations[None,:]
single_translation = True
integer_translations = t.round(translations).to(dtype=t.int32)
if upsample_obj:
selections = t.stack([obj[tr[0]:tr[0]+probe.shape[-2]//2,
tr[1]:tr[1]+probe.shape[-1]//2]
@@ -292,7 +292,7 @@ def ptycho_2D_round(probe, obj, translations, multiple_modes=False, upsample_obj
def ptycho_2D_linear(probe, obj, translations, shift_probe=True):
"""Returns a stack of exit waves accounting for subpixel shifts
This function returns a collection of exit waves, with the first
dimension as the translation index and the final dimensions
corresponding to the detector. The exit waves are calculated by
@@ -322,7 +322,7 @@ def ptycho_2D_linear(probe, obj, translations, shift_probe=True):
if translations.dim() == 1:
translations = translations[None,:]
single_translation = True
# Separate the translations into a part that chooses the window
# And a part that defines the windowing function
integer_translations = t.floor(translations)
@@ -342,15 +342,15 @@ def ptycho_2D_linear(probe, obj, translations, shift_probe=True):
sel01 = t.cat((probe[:,-1:],probe[:,:-1]),dim=1)
sel10 = t.cat((probe[-1:,:],probe[:-1,:]),dim=0)
sel11 = t.cat((sel01[-1:,:],sel01[:-1,:]),dim=0)
selection = sel00 * (1-sp[0])*(1-sp[1]) + \
sel10 * sp[0]*(1-sp[1]) + \
sel01 * (1-sp[0])*sp[1] + \
sel11 * sp[0]*sp[1]
obj_slice = obj[tr[0]:tr[0]+probe.shape[0],
tr[1]:tr[1]+probe.shape[1]]
exit_waves.append(selection * obj_slice)
else:
for tr, sp in zip(integer_translations,
@@ -359,16 +359,16 @@ def ptycho_2D_linear(probe, obj, translations, shift_probe=True):
# Here we subpixel shift the object by (-i,-j) after
# slicing out the correct translation of the probe
#
sel00 = obj[tr[0]:tr[0]+probe.shape[0],
tr[1]:tr[1]+probe.shape[1]]
sel01 = obj[tr[0]:tr[0]+probe.shape[0],
tr[1]+1:tr[1]+1+probe.shape[1]]
sel10 = obj[tr[0]+1:tr[0]+1+probe.shape[0],
tr[1]:tr[1]+probe.shape[1]]
sel11 = obj[tr[0]+1:tr[0]+1+probe.shape[0],
tr[1]+1:tr[1]+1+probe.shape[1]]
@@ -387,7 +387,7 @@ def ptycho_2D_linear(probe, obj, translations, shift_probe=True):
def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multiple_modes=True, polarized=False, polarizer=None, analyzer=None):
"""Returns a stack of exit waves accounting for subpixel shifts
This function returns a collection of exit waves, with the first
dimension as the translation index and the final dimensions
corresponding to the detector. The exit waves are calculated by
@@ -427,7 +427,7 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multi
if translations.dim() == 1:
translations = translations[None,:]
single_translation = True
# Separate the translations into a part that chooses the window
# And a part that defines the windowing function
integer_translations = t.floor(translations)
@@ -480,7 +480,7 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multi
else:
raise NotImplementedError('Object shift not yet implemented')
print('ptyvho 2d sinc', output.shape)
if single_translation:
return output[0]
else:
@@ -489,7 +489,7 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multi
def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, padding=10):
"""Returns a stack of exit waves accounting for subpixel shifts
This function returns a collection of exit waves, with the first
dimension as the translation index and the final dimensions
corresponding to the detector. The exit waves are calculated by
@@ -505,7 +505,7 @@ def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, pad
on the input wavefield, and the first two indexes index differences from
that pixel. It is easier to interpret the resulting matrix though if the
latter two indices index locations in the output plane. NOTE: I believe
this change has now been made
this change has now been made
Parameters
----------
@@ -529,17 +529,17 @@ def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, pad
if translations.dim() == 1:
translations = translations[None,:]
single_translation = True
# Separate the translations into a part that chooses the window
# And a part that defines the windowing function
integer_translations = t.floor(translations)
subpixel_translations = translations - integer_translations
integer_translations = integer_translations.to(dtype=t.int32)
exit_waves = []
B = s_matrix.shape[0]//2
if shift_probe:
i = t.arange(probe.shape[-2]) - probe.shape[-2]//2
j = t.arange(probe.shape[-1]) - probe.shape[-1]//2
@@ -548,14 +548,14 @@ def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, pad
J = 2 * np.pi * J.to(t.float32) / probe.shape[-1]
I = I.to(dtype=probe.dtype,device=probe.device)
J = J.to(dtype=probe.dtype,device=probe.device)
for tr, sp in zip(integer_translations,
subpixel_translations):
fft_probe = t.fft.fftshift(t.fft.fft2(probe), dim=(-1,-2))
shifted_fft_probe = fft_probe * t.exp(1j*(-sp[0]*I - sp[1]*J))
shifted_probe = t.fft.ifft2(t.fft.ifftshift(shifted_fft_probe,
dim=(-1,-2)))
s_matrix_slice = s_matrix[:,:,tr[0]:tr[0]+probe.shape[-2]+2*B,
tr[1]:tr[1]+probe.shape[-1]+2*B]
@@ -564,14 +564,14 @@ def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, pad
device=s_matrix_slice.device,
dtype=s_matrix_slice.dtype)
for i in range(s_matrix.shape[0]):
for j in range(s_matrix.shape[1]):
output [i:i+probe.shape[-2],j:j+probe.shape[-1]] += \
shifted_probe * s_matrix_slice[i,j,i:i+probe.shape[-2],j:j+probe.shape[-1]]
exit_waves.append(output)
else:
raise NotImplementedError('Object shift not yet implemented')
@@ -579,11 +579,11 @@ def ptycho_2D_sinc_s_matrix(probe, s_matrix, translations, shift_probe=True, pad
return exit_waves[0]
else:
return t.stack(exit_waves)
def RPI_interaction(probe, obj):
"""Returns an exit wave from a high-res probe and a low-res obj
In this interaction, the probe and object arrays are assumed to cover
the same physical region of space, but with the probe array sampling that
region of space more finely. Thus, to do the interaction, the object
@@ -593,7 +593,7 @@ def RPI_interaction(probe, obj):
method and is not commonly used elsewhere.
This also works with object functions that have an extra first dimension
for an incoherently mixing model.
for an incoherently mixing model.
Parameters
@@ -610,7 +610,7 @@ def RPI_interaction(probe, obj):
"""
# TODO: The upsampling only works for arrays of even dimension!
# The far-field propagator is just a 2D FFT but with an fftshift
fftobj = propagators.far_field(obj)
# We calculate the padding that we need to do the upsampling
@@ -618,7 +618,7 @@ def RPI_interaction(probe, obj):
pad0r = probe.shape[-2] - obj.shape[-2] - pad0l
pad1l = (probe.shape[-1] - obj.shape[-1])//2
pad1r = probe.shape[-1] - obj.shape[-1] - pad1l
if obj.dim() == 2:
fftobj = t.nn.functional.pad(fftobj, (pad1l, pad1r, pad0l, pad0r))
elif obj.dim() == 3:
@@ -626,7 +626,7 @@ def RPI_interaction(probe, obj):
fftobj, (pad1l, pad1r, pad0l, pad0r, 0,0))
else:
raise NotImplementedError('RPI interaction with obj of dimension higher than 4 (including complex dimension) is not supported.')
# Again, just an inverse FFT but with an fftshift
upsampled_obj = propagators.inverse_far_field(fftobj)
+45 -41
View File
@@ -17,7 +17,11 @@ from matplotlib import ticker, patheffects
__all__ = ['colorize', 'plot_amplitude', 'plot_phase',
'plot_colorized', 'plot_translations', 'get_units_factor',
'plot_nanomap', 'plot_real', 'plot_imag',
'plot_nanomap_with_images']
'plot_nanomap_with_images',
'polarized_plot_component_amplitudes',
'polarized_plot_phase_ret',
'polarized_plot_global_phases',
'polarized_plot_ellipses']
def colorize(z):
@@ -84,7 +88,7 @@ def get_units_factor(units):
def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m', cmap='viridis', cmap_label=None, **kwargs):
"""Plots an image with a colorbar and on an appropriate spatial grid
If a figure is given explicitly, it will clear that existing figure and
plot over it. Otherwise, it will generate a new figure.
@@ -94,7 +98,7 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
Finally, if a function is passed to the plot_func argument, this function
will be called on each slice of data before it is plotted. This is used
internally to enable the plot_real, plot_image, plot_phase, etc. functions.
Parameters
----------
@@ -120,7 +124,7 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
used_fig : matplotlib.figure.Figure
The figure object that was actually plotted to.
"""
# convert to numpy
if isinstance(im, t.Tensor):
# If final dimension is 2, assume it is a complex array. If not,
@@ -142,7 +146,7 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
title = plt.gca().get_title()
fig.clear()
# If im only has two dimensions, this reshape will add a leading
# dimension, and update will be called on index 0. If it has 3 or more
# dimensions, then all the leading dimensions will be compressed into
@@ -151,9 +155,9 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
reshaped_im = im.reshape(-1,s[-2],s[-1])
num_images = reshaped_im.shape[0]
fig.plot_idx = idx % num_images
to_plot = plot_func(reshaped_im[fig.plot_idx])
#Plot in a basis if it exists, otherwise dont
if basis is not None:
if isinstance(basis,t.Tensor):
@@ -181,7 +185,7 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
plt.xlabel('j (pixels)')
plt.ylabel('i (pixels)')
plt.title(title)
if len(im.shape) >= 3:
@@ -194,14 +198,14 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
result = make_plot(0)
update = make_plot
def on_action(event):
if not hasattr(event, 'button'):
event.button = None
if not hasattr(event, 'key'):
event.key = None
if event.key == 'up' or event.button == 'up':
update(fig.plot_idx - 1)
elif event.key == 'down' or event.button == 'down':
@@ -217,9 +221,9 @@ def plot_image(im, plot_func=lambda x: x, fig=None, basis=None, units='$\\mu$m',
fig.my_callbacks = []
fig.my_callbacks.append(fig.canvas.mpl_connect('key_press_event',on_action))
fig.my_callbacks.append(fig.canvas.mpl_connect('scroll_event',on_action))
return result
def plot_real(im, fig = None, basis=None, units='$\\mu$m', cmap='viridis', cmap_label='Real Part (a.u.)', **kwargs):
"""Plots the real part of a complex array with dimensions NxM
@@ -256,7 +260,7 @@ def plot_real(im, fig = None, basis=None, units='$\\mu$m', cmap='viridis', cmap_
return plot_image(im, plot_func=plot_func, fig=fig, basis=basis,
units=units, cmap=cmap, cmap_label=cmap_label,
**kwargs)
def plot_imag(im, fig = None, basis=None, units='$\\mu$m', cmap='viridis', cmap_label='Imaginary Part (a.u.)', **kwargs):
@@ -304,7 +308,7 @@ def plot_amplitude(im, fig = None, basis=None, units='$\\mu$m', cmap='viridis',
If a basis is explicitly passed, the image will be plotted in real-space
coordinates.
Parameters
----------
im : array
@@ -520,12 +524,12 @@ def plot_nanomap(translations, values, fig=None, units='$\\mu$m', convention='pr
def plot_nanomap_with_images(translations, get_image_func, values=None, mask=None, basis=None, fig=None, nanomap_units='$\\mu$m', image_units='$\\mu$m', convention='probe', image_title='Image', image_colorbar_title='Image Amplitude', nanomap_colorbar_title='Integrated Intensity', cmap='viridis', **kwargs):
"""Plots a nanomap, with an image or stack of images for each point
In many situations, ptychography data or the output of ptychography
reconstructions is formatted as a set of images associated with various
points in real space. This function is designed to allow for browsing
through this kind of data, by making it possible to visualize a
"""
# This should pull heavily from the dataset.inspect function
@@ -558,11 +562,11 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
s0 = bbox.width * bbox.height / translations.shape[0] * 72**2 #72 is points per inch
s0 /= 4 # A rough value to make the size work out
s = np.ones(translations.shape[0]) * s0
s[idx] *= 4
return s
def update_colorbar(im):
#
# This solves the problem of the colorbar being changed
@@ -571,29 +575,29 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
if hasattr(im, 'norecurse') and im.norecurse:
im.norecurse=False
return
im.norecurse=True
# This is needed to update the colorbar
# only change limits if array contains multiple values
if np.min(im.get_array()) != np.max(im.get_array()):
im.set_clim(vmin=np.min(im.get_array()),
vmax=np.max(im.get_array()))
#
# The meatiest part of this program, here we just go through and
# set up the plot how we want it
#
# First we set up the left-hand plot, which shows an overview map
axes[0].set_title('Relative Displacement Map')
translations = translations.detach().cpu().numpy()
if convention.lower() != 'probe':
translations = translations * -1
s = calculate_sizes(0)
nanomap_units_factor = get_units_factor(nanomap_units)
nanomap = axes[0].scatter(nanomap_units_factor * translations[:,0],
nanomap_units_factor * translations[:,1],
@@ -614,7 +618,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
# where the colorbar should have been to avoid stretching the
# nanomap plot, while still not showing the (now useless) colorbar.
cb1.remove()
# Now we set up the second plot, which shows the individual
# diffraction patterns
axes[1].set_title(image_title)
@@ -634,7 +638,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
# This fails if the basis is not rectangular
basis_norm = np.linalg.norm(np_basis, axis = 0)
basis_norm = basis_norm * get_units_factor(image_units)
extent = [0, example_im.shape[-1]*basis_norm[1], 0,
example_im.shape[-2]*basis_norm[0]]
else:
@@ -655,9 +659,9 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
axes[1].text_box.set_path_effects(
[patheffects.Stroke(linewidth=2, foreground='black'),
patheffects.Normal()])
meas = axes[1].imshow(im, extent=extent, cmap=cmap)
cb2 = plt.colorbar(meas, ax=axes[1], orientation='horizontal',
format='%.2e',
ticks=ticker.LinearLocator(numticks=5),
@@ -665,7 +669,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
cb2.ax.tick_params(labelrotation=20)
cb2.ax.set_title(image_colorbar_title, size="medium", pad=5)
cb2.ax.callbacks.connect('xlim_changed', lambda ax: update_colorbar(meas))
# This function handles all the updating, except for moving the
# slider value. This is done because the slider widget is
# ultimately responsible for triggering an update, so all other
@@ -675,7 +679,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
# We have to explicitly make it an integer because the slider will
# output floats (even if they are still integer-valued)
idx = int(idx)
# Get the new data for this index
im = get_image_func(idx)
if len(im.shape) >= 3:
@@ -686,22 +690,22 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
axes[1].image_idx = im_idx
axes[1].text_box.set_text(str(im_idx))
im = im.reshape(-1,im.shape[-2],im.shape[-1])[im_idx]
# Now we resize the nanomap to show the new selection
axes[0].collections[0].set_sizes(calculate_sizes(idx))
# And we update the data in the image as well
ax_im = axes[1].images[-1]
ax_im.set_data(im)
update_colorbar(ax_im)
#
# Now we define the functions to handle various kinds of events
# that can be thrown our way
#
# We start by creating the slider here, so it can be used
# by the update hooks.
slider = Slider(axslider, 'Image #', 0, translations.shape[0]-1, valstep=1, valfmt="%d")
@@ -715,7 +719,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
# while the mouse is within the image display
im = im.reshape(-1,im.shape[-2],im.shape[-1])
im_idx = axes[1].image_idx
if event.key == 'up' or event.button == 'up' \
or event.key == 'left':
im_idx = (im_idx - 1) % im.shape[0]
@@ -733,7 +737,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
event.button = None
if not hasattr(event, 'key'):
event.key = None
if event.key == 'up' or event.button == 'up' or event.key == 'left':
idx = slider.val - 1
elif event.key == 'down' or event.button == 'down' or event.key == 'right':
@@ -742,7 +746,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
# This prevents errors from being thrown on irrelevant key
# or mouse input
return
# Handle the wraparound and trigger the update
idx = int(idx) % translations.shape[0]
slider.set_val(idx)
@@ -753,16 +757,16 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non
# for example, scroll events that happen over the nanomap
if event.mouseevent.button == 1:
slider.set_val(event.ind[0])
# Here we connect the various update functions
cid1 = fig.canvas.mpl_connect('pick_event',on_pick)
cid2 = fig.canvas.mpl_connect('key_press_event',on_action)
cid3 = fig.canvas.mpl_connect('scroll_event',on_action)
# It's so dumb that matplotlib doesn't automatically track this for you
fig.nanomap_cids = [cid1,cid2,cid3]
fig.nanomap_cids = [cid1,cid2,cid3]
slider.on_changed(update)
# Throw an extra update into the mix just to get rid of any things
# (like the nanomap dot sizes) that otherwise would change on the
# first update
+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))