NeuroOpticsTest / app.py
Niccko's picture
Add download button
d1d7e79
Raw
History Blame Contribute Delete
5.45 kB
import os
import gradio as gr
import matplotlib.pyplot as plt
#from fraunhofer_test import generate_diff
from zern_generator import Zernike
from fraunhofer_test import Fraunhofer
from PIL import Image
import numpy as np
import shutil
css = """
.container {
# height: 50vh;
# width: 100%;
# overflow-x: auto !important;
# overflow-y: auto !important;
# scrollbar-width: thin !important;
}
"""
zerns = {}
def update_zerns(key, value):
global zerns
zerns[key] = value
def get_zern_from_file(file):
with open(file) as f:
coeffs = [float(l) for l in f]
return tuple(coeffs)
zernike = Zernike()
fraunhofer = Fraunhofer(zernike)
curr_diff = None
theme = gr.themes.Default(primary_hue=gr.themes.colors.red, secondary_hue=gr.themes.colors.pink)
with gr.Blocks(theme=theme, css=css) as demo:
with gr.Row():
zern_sliders = []
with gr.Column(scale=1):
with gr.Tab(label="Генератор"):
defocus_amt = gr.Number(
label="Величина дефокуса",
value=2,
interactive=True
)
with gr.Row():
num_zernikes = gr.Number(
label="Количество коэффициентов Цернике",
value=6,
interactive=True
)
show_zernikes_btn = gr.Button("Показать")
sliders_container = gr.Column(elem_classes=["container"])
@gr.render(inputs=[num_zernikes], triggers=[show_zernikes_btn.click])
def render_count(count):
global zerns
zerns = {}
sliders_container.children.clear()
count = min(count, 100)
for i in range(count):
term_name = f" - {zernike.term[i]}" if i < len(zernike.term) else ""
sld = gr.Slider(
key=i,
label=f"Zernike {i+1}{term_name}",
minimum=-100,
interactive=True,
step=0.01,
value=0
)
sld.change(update_zerns, inputs=[gr.Number(value=i, visible=False), sld])
sliders_container.add_child(sld)
return gr.update()
with gr.Tab(label="Настройки"):
resolution_num = gr.Number(value=256, label="Разрешение")
with gr.Column(scale=5):
with gr.Tab(label="Экран"):
with gr.Row():
plot_2d = gr.Plot()
plot_3d = gr.Plot()
btn = gr.Button("Рассчитать")
with gr.Tab(label="Тестовые экраны"):
plot_batch = gr.Plot(label="Тестовые экраны (ANSI scheme)")
btn_batch = gr.Button("Рассчитать тестовые экраны")
with gr.Tab(label="Разница у фокуса"):
plot_diff = gr.Plot(label="Разница в фокусе")
btn_diff = gr.Button("Рассчитать разницу в фокусе")
file_download = gr.File(label="Скачать")
def on_button_click(resolution):
return zernike.generate_zern_wavefront_fig(*[zerns[key] for key in sorted(zerns)])
def on_settings_click(resolution):
print(resolution)
zernike.set_image_params(npix=resolution)
def on_diff_button_click(defocus_amount):
shutil.rmtree("temp", ignore_errors=True)
os.makedirs("temp")
global curr_diff
pl, mn, diff = fraunhofer.generate_diff(defocus_amount, *[zerns[key] for key in sorted(zerns)])
curr_diff = diff
fig, axs = plt.subplots(1, 3, figsize = (20, 10))
axs[0].imshow(pl, cmap='grey')
axs[1].imshow(mn, cmap='grey')
axs[2].imshow(diff, cmap='grey')
axs[0].axis('off')
axs[0].set_title(f"+{defocus_amount}")
axs[1].axis('off')
axs[1].set_title(f"-{defocus_amount}")
axs[2].axis('off')
axs[2].set_title(f"Difference")
a = np.min(curr_diff)
b = np.max(curr_diff)
img_arr = ((curr_diff - a) / (b - a) * 255).astype(np.uint8)
img = Image.fromarray(np.stack([img_arr] * 3, axis=-1), mode='RGB')
z = [str(zerns[key]) for key in sorted(zerns)]
filename = f'temp/{"_".join(z[1:])}.png'
img.save(filename)
return fig, filename
btn.click(on_button_click, inputs=[resolution_num], outputs=[plot_2d, plot_3d])
btn_batch.click(zernike.generate_batch, inputs=[resolution_num], outputs=[plot_batch])
resolution_num.change(on_settings_click, inputs=[resolution_num])
btn_diff.click(on_diff_button_click, inputs=[defocus_amt], outputs=[plot_diff, file_download])
demo.launch()