NeuroOpticsTest / generate_dataset.py
Niccko's picture
Add download button
d1d7e79
Raw
History Blame Contribute Delete
2.79 kB
from zern_generator import Zernike
from PIL import Image
from fraunhofer_test import Fraunhofer, FraunhoferGPU
import cupy as cp
import matplotlib.pyplot as plt
import itertools
from time import time
from pprint import pprint
import os
from multiprocessing import Pool, cpu_count
import numpy as np
zernike = Zernike()
zernike.set_image_params(npix=2 ** 8)
fraunhofer = FraunhoferGPU(zernike)
def generate_image(z, group_depth):
arr = fraunhofer.generate_diff(1, 0, *z)
a = cp.min(arr[2])
b = cp.max(arr[2])
img_arr = ((arr[2] - a) / (b - a) * 255).astype(cp.uint8)
img = Image.fromarray(np.stack([img_arr] * 3, axis=-1), mode='RGB')
coeffs = '_'.join([f'{x:.2f}' for x in z[group_depth:]])
group = "/".join([str(x) for x in z[:group_depth]])
path = f'./screens_tst/{group}'
if not os.path.exists(path):
try:
os.makedirs(path)
except:
print("Directory already exists. Skipping.")
filename = f'{path}/{coeffs}_.png'
img.save(filename)
def generate_dataset(range, step, kf_cnt):
zs = get_zernike_coeffs(range, step, kf_cnt)
num_workers = 24
if num_workers is None:
num_workers = cpu_count()
with Pool(num_workers) as pool:
results = []
for result in pool.imap_unordered(generate_image, zs):
results.append(result)
for result in results:
result.get()
def generate_dataset_gpu(zs):
c = 0
window_size = 50
times = []
start_time = time()
total_images = len(zs)
num_workers = 24
if num_workers is None:
num_workers = cpu_count()
for z in zs:
generate_image(z)
def get_zernike_coeffs(range, step, kf_cnt):
points = tuple(np.linspace(*range, int((range[1] - range[0]) / step + 1)))
points_as_tuples = tuple(float(x) for x in points)
return np.array(list(set(itertools.product(points_as_tuples, repeat=kf_cnt))))
if __name__ == '__main__':
# rng = (-5, 5)
# step = 0.5
# points = tuple(np.linspace(*rng, int((rng[1] - rng[0]) / step + 1)))
# points_as_tuples = tuple(float(x) for x in points)
zernikes = get_zernike_coeffs((-5,5), 0.25, 2)
print("Total images: ", len(zernikes))
total_gen_time, total_save_time = 0, 0
start = time()
for i, z in enumerate(zernikes):
generate_image(z, 1)
if i % 100 == 0:
print(f"{i}/{len(zernikes)}")
print("Average generation time: ", (time()-start)/len(zernikes))
# generate_dataset(zernikes)
# last = 1
# for i in range(1,6):
# p = get_zernike_coeffs([-5,5], 0.25, i)
# print(f"Количество коэфф-ов: {i} -> Количество комбинаций: {len(p)}")
# last = len(p)