Spaces:
Sleeping
Sleeping
| 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)))) | |