File size: 4,918 Bytes
5d61225
 
 
 
 
567b333
5d61225
 
 
 
 
567b333
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a30ddbc
567b333
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c7180e6
567b333
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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))))