File size: 5,041 Bytes
74c1022
 
 
 
87627fc
74c1022
ff6d288
74c1022
 
 
 
 
ff6d288
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
87627fc
ff6d288
 
 
 
 
 
87627fc
 
 
ff6d288
 
 
 
87627fc
 
 
 
 
 
ff6d288
 
 
87627fc
ff6d288
 
 
d1d7e79
 
87627fc
d1d7e79
 
 
ff6d288
 
 
 
 
 
 
 
 
 
 
 
 
 
87627fc
ff6d288
87627fc
ff6d288
 
87627fc
ff6d288
 
 
 
87627fc
ff6d288
 
 
 
 
 
 
 
 
 
 
87627fc
 
ff6d288
 
 
 
 
 
 
 
 
 
 
 
 
 
80623b8
ff6d288
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d1d7e79
 
 
 
ff6d288
 
87627fc
d1d7e79
 
 
ff6d288
87627fc
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
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))))