import math import matplotlib.pyplot as plt import numpy as np import cupy as cp 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): if isinstance(zerns, cp.ndarray): zerns = zerns.get() 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 = cart.eval_grid(np.array(zerns), matrix=True) screen = np.array(screen) 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 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))))