File size: 1,888 Bytes
1fa7db9
 
 
0636548
567b333
1fa7db9
 
567b333
 
0636548
1fa7db9
567b333
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1fa7db9
 
567b333
 
 
 
 
 
 
 
a30ddbc
567b333
 
 
 
 
 
 
 
 
 
 
 
a30ddbc
1fa7db9
 
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
import matplotlib.pyplot as plt
import numpy as np
import math
from scipy.fft import fft2, fftshift, ifftshift
from zern_generator import Zernike


def denoise(screen: np.ndarray):
    return np.where(np.abs(screen) < 1e-4, 0, screen)


class Fraunhofer:
    def __init__(self, zernike: Zernike):
        self.focal_distance = 100 * 10 ** -3
        self.distance = 100 * 10 ** -3
        self.radius = 5 * 10 ** -3
        self.span = (-(10 ** -2), 10 ** -2)
        self.wavelength = 500 * 10 ** -9
        self.k = 2*math.pi/self.wavelength
        
        self.zernike = zernike
    
    def _denoise(screen: np.ndarray):
        return np.where(np.abs(screen) < 1e-4, 0, screen)
    
    def get_focal_point(self,input_screen):
        phase_mul = ((np.exp(1j*self.k*self.distance) * np.exp(1j*self.k*(self.zernike.xv**2 + self.zernike.yv**2)/(2*self.distance)))/(1j*self.wavelength*self.distance))
        U_fraun = fft2(ifftshift(input_screen), overwrite_x=True)
        U_fraun = fftshift(U_fraun * phase_mul)


        I_fraun = np.abs(U_fraun)**2
        return I_fraun / np.max(I_fraun)
    
    
    def generate_diff(self, defocus_amount, *zernikes):
        zernikes = list(zernikes)
        if len(zernikes) < 5:
            zernikes += [0] * (5-len(zernikes))
        

 
        wf = self.zernike.generate_zern_screen([1], self.radius)
        

        f_plus = self.zernike.generate_zern_screen(self.zernike.add_defocus(defocus_amount,zernikes), norm_radius=self.radius)
        f_minus = self.zernike.generate_zern_screen(self.zernike.add_defocus(-defocus_amount,zernikes), norm_radius=self.radius)

        focal_image_plus = denoise(self.get_focal_point(wf*np.exp(f_plus * 1j)))
        focal_image_minus = denoise(self.get_focal_point(wf*np.exp(f_minus * 1j)))

        return focal_image_plus, focal_image_minus, focal_image_plus-focal_image_minus