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()