NeuroOptics / zern_generator.py
Niccko's picture
add show indexes btn
a30ddbc
Raw
History Blame Contribute Delete
4.92 kB
import math
import matplotlib.pyplot as plt
import numpy as np
from zernike import RZern
import utils
SAMPLE_SIZE = 256
span = (-1, 1)
class Zernike:
def __init__(self):
self.set_image_params()
self.term = [
"PISTON",
"V-TILT",
"H-TILT",
"O-PR-ASTIGMATISM",
"DEFOCUS",
"V-PR-ASTIGMATISM",
"V-TREFOIL",
"V-PR-COMA",
"H-PR-COMA",
"O-TREFOIL",
"O-QUADRAFOIL",
"O-SC-ASTIGMATISM",
"P-SPHERICAL",
"V-SC-ASTIGMATISM",
"V-QUADRAFOIL",
# "V-SC-COMA",
# "H-SC-COMA",
# "V-SC-TREFOIL",
# "O-SC-TREFOIL",
# "O-PENTAFOIL",
# "V-PENTAFOIL"
]
def set_image_params(self, span = (-(10 ** -2), 10 ** -2), npix = 2**8):
self.span = span
self.npix = npix
ddx = np.linspace(*self.span, self.npix)
ddy = np.linspace(*self.span, self.npix)
self.xv, self.yv = np.meshgrid(ddx, ddy)
def generate_zern_screen(self, zerns, norm_radius = 1):
zerns = self._osa_to_noll_list(zerns)
num_zerns = len(zerns)
mask = utils.circle_mask(self.xv, self.yv, norm_radius)
degree = math.ceil((-3 + math.sqrt(8 * num_zerns + 1)) / 2)
zerns = list(zerns) + [0] * int((((degree + 1) * (degree + 2) / 2) - len(zerns)))
cart = RZern(degree)
cart.make_cart_grid(self.xv/norm_radius, self.yv/norm_radius)
screen = np.array(cart.eval_grid(np.array(zerns), matrix=True))
screen[mask == 0] = 0
screen[np.isnan(screen)] = 0
return screen
def generate_zern_wavefront_fig(self, *zerns):
Phi = self.generate_zern_screen(zerns)
# Phi[np.isnan(Phi)] = 0
fig_2d, ax_2d = plt.subplots()
fig_3d, ax_3d = plt.subplots(subplot_kw={"projection": "3d"})
a = ax_2d.imshow(Phi, extent=(*span, *span), cmap='gray')
ax_3d.plot_surface(self.xv, self.yv, Phi)
fig_2d.colorbar(a)
plt.show()
return fig_2d, fig_3d
def add_defocus(self, defocus_value, zernikes):
zernikes = list(zernikes)
if len(zernikes) < 5:
zernikes += [0] * (5-len(zernikes))
zernikes[4] += defocus_value
return zernikes
def generate_batch(self, resolution):
#term = ["PISTON","V-TILT","H-TILT","DEFOCUS","O-PR-ASTIGMATISM","V-PR-ASTIGMATISM","V-PR-COMA","H-PR-COMA","V-TREFOIL","O-TREFOIL","P-SPHERICAL","O-SC-ASTIGMATISM","V-SC-ASTIGMATISM","V-QUADRAFOIL","O-QUADRAFOIL","V-SC-COMA","H-SC-COMA","V-SC-TREFOIL","O-SC-TREFOIL","O-PENTAFOIL","V-PENTAFOIL"]
fig = plt.figure()
fig.set_figwidth(30)
fig.set_figheight(20)
zern_index = 1
for i in range(5):
for j in range(6):
if j > i:
continue
zern = [0] * (zern_index - 1) + [1]
ax = plt.subplot2grid((5, 6), (i, j))
Phi = self.generate_zern_screen(zern, norm_radius=self.span[1]/2)
ax.set_title(f" \n Zernike {zern_index-1}{self.term[zern_index-1]} ")
ax.imshow(Phi, extent=(*span, *span), cmap='gray')
zern_index += 1
plt.show()
return fig
def _osa_to_mn(self, idx):
n = math.ceil((-3 + math.sqrt((9 + 8 * idx))) / 2)
m = 2 * idx - n * (n + 2)
return m, n
def _mn_to_noll(self, m, n):
first = (n * (n + 1)) // 2
second = abs(m)
third = 0
# I wanted to simplify the logic but it seems to work like a colander
if m > 0 and n % 4 in [0, 1]:
third = 0
elif m < 0 and n % 4 in [2, 3]:
third = 0
elif m >= 0 and n % 4 in [2, 3]:
third = 1
elif m <= 0 and n % 4 in [0, 1]:
third = 1
j = first + second + third
return j
def _osa_to_noll(self, idx):
return self._mn_to_noll(*self._osa_to_mn(idx))
def _osa_to_noll_list(self, zernikes):
indexed_zerns = sorted(map(lambda x: (self._osa_to_noll(x[0]), x[1]), enumerate(zernikes)), key=lambda x: x[0])
used_indeces = [x for x, _ in indexed_zerns]
mn, mx = min(used_indeces), max(used_indeces)
missing_indeces = set(range(mn, mx+1)) - set(used_indeces)
return [x[1] for x in list(sorted(indexed_zerns + [(x, 0) for x in missing_indeces]))]
term = ["PISTON","V-TILT","H-TILT","DEFOCUS","O-PR-ASTIGMATISM","V-PR-ASTIGMATISM","V-PR-COMA","H-PR-COMA","V-TREFOIL","O-TREFOIL","P-SPHERICAL","O-SC-ASTIGMATISM","V-SC-ASTIGMATISM","V-QUADRAFOIL","O-QUADRAFOIL","V-SC-COMA","H-SC-COMA","V-SC-TREFOIL","O-SC-TREFOIL","O-PENTAFOIL","V-PENTAFOIL"]
print(Zernike()._osa_to_noll_list(list(range(0,15))))