Team Ai
Apppublic

helboukkouri/interactive-plot

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
app.py180 linesDownload Raw Back to root
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