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