import gradio as gr import os os.environ['NUMBA_ENABLE_CUDASIM'] = '1' os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' # Suppress TensorFlow warnings unless fatal # !! For now, because it is a private repo !! from huggingface_hub import login, hf_hub_download # login(token=os.environ.get("HF_AUTH_TOKEN")) # from huggingface_hub import hf_hub_download from tensorflow.keras.saving import load_model import numpy as np from numba import cuda from scipy.io import readsav import matplotlib.pyplot as plt import time from filters import GaussianFilter from vec_metrics import get_pred_vis, mae_of_vectors # ============================================= LOAD MODEL =========================================== # Load your trained model model_path = hf_hub_download(repo_id="mervess/FCD-Solar", filename="fcd.keras") model = load_model(model_path, custom_objects={'GaussianFilter': GaussianFilter}, compile=False) print("#) Model loaded successfully.") # Load the F matrix f_matrix = np.load("data/fourier_matrix.npy") print("#) STIX fourier matrix loaded successfully.") # ============================================= PREDICT ============================================== # Define a function for prediction def predict_image(vis): norm_vis = np.empty(vis.shape) alphas = np.empty(vis.shape[0], dtype=np.float32) threadsperblock = 32 blockspergrid = (alphas.size + (threadsperblock - 1)) // threadsperblock calc_alphas_cpu[blockspergrid, threadsperblock](vis, alphas, norm_vis) predicted_img = model.predict(norm_vis, batch_size=128, verbose=False) predicted_img = predicted_img * alphas[:, np.newaxis, np.newaxis, np.newaxis] output_image = np.squeeze(predicted_img) return output_image @cuda.jit def calc_alphas_cpu(vis_arrays, alphas, norm_vis): pos = cuda.grid(1) if pos < alphas.size: nvis = np.sqrt( np.square(vis_arrays[pos, :24]) + np.square(vis_arrays[pos, 24:])) alphas[pos] = np.max(nvis) * 0.5 for i in range(vis_arrays.shape[1]): norm_vis[pos][i] = vis_arrays[pos][i] / alphas[pos] # ============================================ READ FILES ============================================ def read_file( file_name, file_type='.npy' ): if file_type == '.npy': vis_array = np.load( file_name ) elif file_type == '.sav': data = readsav( file_name, python_dict=True, verbose=False ) vis_structure = data['vis'] vis_components = vis_structure['OBSVIS'] vis_array = np.hstack( ( np.real(vis_components), np.imag(vis_components) ) ) return vis_array # ========================================== EXAMPLE INPUTS ========================================== # Provide sample data # ===== 1 ===== vis_array = read_file('example_data/20211216_vis.sav', file_type='.sav') # ===== 2 ===== vis_array_2 = read_file('example_data/20221231_vis.sav', file_type='.sav') # ===== 3 ===== vis_array_3 = read_file('example_data/20220828_vis.sav', file_type='.sav') sample_vectors = { "20211216T233548-20211216T234036_10-15keV": vis_array, "20221231T123044-20221231T123252_4-10keV": vis_array_2, "20220828T155202-20220828T160950_25-50keV": vis_array_3, } print("#) Sample data loaded successfully.") # ============================================ VISUALIZE ============================================= def plot_map( map_data, img_name='', color_map='turbo', is_minorticks_on=True, show_colorbar=True, save_plot=False, plot_title="plot.pdf" ): fig, ax = plt.subplots(figsize=(3, 3)) cax = ax.imshow(map_data, cmap=color_map) ax.set_title(img_name) ax.axis("off") if show_colorbar: cbar = fig.colorbar(cax, ax=ax, orientation='horizontal', fraction=0.046, pad=0.07, extend='both') cbar.ax.tick_params(labelsize='large') # Increase colorbar font size # Set the colorbar ticks to beginning, middle, and end, optimizing the formatting process min_val = map_data.min() max_val = map_data.max() mid_val = (min_val + max_val) / 2 cbar.set_ticks([min_val, mid_val, max_val]) cbar.set_ticklabels([f'{min_val:.4f}', f'{mid_val:.4f}', f'{max_val:.4f}']) if is_minorticks_on: cbar.minorticks_on() plt.tight_layout() if save_plot: f_format = os.path.splitext(plot_title)[-1][1:] if show_colorbar: plt.savefig(plot_title, format=f_format, bbox_inches='tight') else: plt.savefig(plot_title, format=f_format, bbox_inches='tight', pad_inches=0) plt.show() def to_html( msg, color='red' ): return f"
{msg}
" # =============================================== DEMO =============================================== # Gradio interface def demo( sample_choice, input_vector, uploaded_file ): """ Function to demonstrate the FCD model. Parameters: ----------- sample_choice: str Sample choice from the dropdown list. input_vector: str Input vector from the textbox. uploaded_file: file Uploaded file. Returns: -------- numpy array (image) Reconstructed image. str (-) Message to display. / (+) Prediction time. str Chi2 value. """ vis = None # ==== SAMPLE CHOICE ==== if sample_choice != "" and \ sample_choice is not None and \ sample_choice in sample_vectors: vis = sample_vectors[sample_choice] # ==== INPUT VECTOR ==== elif input_vector != "" and input_vector is not None and \ isinstance(input_vector, str): try: vis = np.array( [ float(i) for i in input_vector.split(",") ], dtype=np.float32 ) if vis.shape != (48,): return None, None, None, to_html("Please provide an input vector of 48 real numbers.") except: return None, None, None, to_html("Please provide an input consisting of real numbers.") # ==== UPLOADED FILE ==== elif uploaded_file is not None: try: # Get the file name file_name = uploaded_file.name # Split the file name to get the extension _, file_extension = os.path.splitext(file_name) vis = read_file( file_name, file_type=file_extension ) except Exception as e: return None, None, None, to_html(f"Error processing the uploaded file: {e}.") # ==== PREDICT ==== if vis is None: return None, None, None, to_html("Please provide a sample, input vector, or upload a file.") else: start_time = time.time() image = predict_image(vis.reshape(1, -1)) prediction_time = time.time() - start_time # Calculate MAE between predicted visibility and input visibility pred_vis = get_pred_vis( image, f_matrix ) mae = mae_of_vectors( vis, pred_vis ) plot_map(image, save_plot=True, plot_title="fcd_image.png") return "fcd_image.png", \ f"{prediction_time:.2f} seconds", \ f"{mae:.3f}", \ to_html("Image has been reconstructed.", color='green') # =============================================== MAIN =============================================== # Build Gradio Interface interface = gr.Interface( fn=demo, inputs=[ gr.Dropdown( choices=[""] + list(sample_vectors.keys()), label="Select a sample vector (i.e., visibility, Fourier components):", value="", # Placeholder ), gr.Textbox( label="Enter your vector (comma-separated 48 real numbers):", placeholder="e.g., 0.1, 0.2, 0.3", ), gr.File( label="Upload your file (either a .sav or .npy file):", type="filepath", # Accept file uploads ), ], outputs=[ gr.Image(type="filepath", label="Reconstructed Image"), gr.Textbox(label="Reconstruction Time"), # Line for prediction time gr.Textbox(label="MAE Metric"), # Line for MAE metric gr.HTML(label="Message"), # Line for messages ], title="FCD Demo", description="Select a sample, enter vector (visibility) values, or upload a file to see the reconstructed image.", ) # Launch the app interface.launch()