helboukkouri/interactive-plot
0
1import gradio as gr2import numpy as np3import sympy as sp4import seaborn as sns5from matplotlib import pyplot as plt6 7sns.set_style(style="darkgrid")8sns.set_context(context="notebook", font_scale=0.7)9 10MAX_NOISE = 2011DEFAULT_NOISE = 612SLIDE_NOISE_STEP = 213 14MAX_POINTS = 10015DEFAULT_POINTS = 2016SLIDE_POINTS_STEP = 517 18def generate_equation(process_params):19 process_params = process_params.astype(float).values.tolist()20 21 # Define symbols22 x = sp.symbols('x')23 coefficients = sp.symbols('a b c d e')24 25 # Create the polynomial expression26 polynomial_expression = None27 for i, coef in enumerate(reversed(coefficients)):28 polynomial_expression = polynomial_expression + coef * x**i if polynomial_expression else coef * x**i29 30 # Parameter mapping31 parameters = {coef: value for coef, value in zip(coefficients, process_params[0])}32 33 # Substitute parameter values into the expression34 polynomial_with_values = polynomial_expression.subs(parameters)35 latex_representation = sp.latex(polynomial_with_values)36 return fr"Underlying process $${latex_representation}$$"37 38 39def true_process(x, process_params):40 """The true process we want to model."""41 process_params = process_params.astype(float).values.tolist()42 return (43 process_params[0][0] * (x ** 4)44 + process_params[0][1] * (x ** 3)45 + process_params[0][2] * (x ** 2)46 + process_params[0][3] * x47 + process_params[0][4]48 )49 50 51def generate_data(num_points, noise_level, process_params):52 53 # x is the list of input values54 input_values = np.linspace(-5, 2, num_points)55 input_values_dense = np.linspace(-5, 2, MAX_POINTS)56 57 # y = f(x) is the underlying process we want to model58 y = [true_process(x, process_params) for x in input_values]59 y_dense = [true_process(x, process_params) for x in input_values_dense]60 61 # however, we can only observe a noisy version of f(x)62 noise = np.random.normal(0, noise_level, len(input_values))63 y_noisy = y + noise64 65 return input_values, input_values_dense, y, y_dense, y_noisy66 67 68def make_plot(69 num_points, noise_level, process_params,70 show_true_process, show_original_points, show_added_noise, show_noisy_points,71 ):72 73 x, x_dense, y, y_dense, y_noisy = generate_data(num_points, noise_level, process_params)74 75 fig = plt.figure(dpi=300)76 if show_true_process:77 plt.plot(78 x_dense, y_dense, "-", color="#363A4F",79 label="True Process",80 lw=1.5,81 )82 if show_added_noise:83 plt.vlines(84 x, y, y_noisy, color="#556D9A",85 linestyles="dashed",86 alpha=0.75,87 lw=1,88 label="Added Noise",89 )90 if show_original_points:91 plt.plot(92 x, y, "-o", color="none",93 ms=6,94 markerfacecolor="white",95 markeredgecolor="#556D9A",96 markeredgewidth=1.2,97 label="Original Points",98 )99 if show_noisy_points:100 plt.plot(101 x, y_noisy, "-o", color="none",102 ms=6.5,103 markerfacecolor="#556D9A",104 markeredgecolor="none",105 markeredgewidth=1.5,106 alpha=1,107 label="Noisy Points",108 )109 110 plt.xlabel("\nx")111 plt.ylabel("y") 112 plt.legend(fontsize=7.5)113 plt.tight_layout()114 plt.show()115 return fig116 117# Force main column to be 100 pixels wide, knowing that the parent is a flex container with column direction 118css = """119.gradio-container {120 width: min(1000px, 50%)!important;121 min-width: 800px;122}123.main-plot {124}125"""126with gr.Blocks(css=css) as demo:127 with gr.Row():128 with gr.Column():129 with gr.Row():130 process_params = gr.DataFrame(131 value=[[0.5, 2, -0.5, -2, 1]],132 label="Underlying Process Coefficients",133 type="pandas",134 column_widths=("2", "1", "1", "1", "1w"),135 headers=["x ** 4", "x ** 3", "x ** 2", "x", "1"],136 interactive=True137 )138 equation = gr.Markdown()139 140 with gr.Row():141 with gr.Column():142 num_points = gr.Slider(143 minimum=5,144 maximum=MAX_POINTS,145 value=DEFAULT_POINTS,146 step=SLIDE_POINTS_STEP,147 label="Number of Points"148 )149 with gr.Column():150 noise_level = gr.Slider(151 minimum=0,152 maximum=MAX_NOISE,153 value=DEFAULT_NOISE,154 step=SLIDE_NOISE_STEP,155 label="Noise Level"156 )157 158 show_params = []159 with gr.Row():160 with gr.Column():161 show_params.append(gr.Checkbox(label="Show Underlying Process", value=True))162 show_params.append(gr.Checkbox(label="Show Original Points", value=True))163 with gr.Column():164 show_params.append(gr.Checkbox(label="Show Added Noise", value=True))165 show_params.append(gr.Checkbox(label="Show Noisy Points", value=True))166 167 scatter_plot = gr.Plot(elem_classes=["main-plot"])168 169 num_points.change(fn=make_plot, inputs=[num_points, noise_level, process_params, *show_params], outputs=scatter_plot)170 noise_level.change(fn=make_plot, inputs=[num_points, noise_level, process_params, *show_params], outputs=scatter_plot)171 process_params.change(fn=make_plot, inputs=[num_points, noise_level, process_params, *show_params], outputs=scatter_plot)172 process_params.change(fn=generate_equation, inputs=[process_params], outputs=equation)173 for component in show_params:174 component.change(fn=make_plot, inputs=[num_points, noise_level, process_params, *show_params], outputs=scatter_plot)175 demo.load(fn=make_plot, inputs=[num_points, noise_level, process_params, *show_params], outputs=scatter_plot)176 demo.load(fn=generate_equation, inputs=[process_params], outputs=equation)177 178if __name__ == "__main__":179 demo.launch()180 