Team Ai
Apppublic

AmitPandit175/Gradient_Descent_Visualization_App

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
5_Visualizer.py227 linesDownload Raw Back to pages
1import streamlit as st
2import numpy as np
3import plotly.graph_objects as go
4from sympy import symbols, diff, lambdify, sin, cos, tan, exp, log
5import sympy as sp
6
7def gradient_descent(formula, start, learning_rate, iterations=50, threshold=1e10, grad_clip_value=1e5):
8    x_sym = symbols('x')
9    try:
10        formula_sym = eval(formula, {"x": x_sym, "sin": sin, "cos": cos, "tan": tan, "exp": exp, "log": log, "np": np})
11        derivative = diff(formula_sym, x_sym)
12        grad_func = lambdify(x_sym, derivative, modules="numpy")
13    except Exception as e:
14        st.error(f"Error in formula evaluation: {e}")
15        return []
16
17    x = start
18    trajectory = [x]
19
20    for _ in range(iterations):
21        try:
22            grad = grad_func(x)
23            grad = np.clip(grad, -grad_clip_value, grad_clip_value)
24            x = x - learning_rate * grad
25            if abs(x) > threshold:
26                st.error("Gradient descent diverged due to large values.")
27                break
28            trajectory.append(x)
29        except Exception as e:
30            st.error(f"Error in gradient calculation: {e}")
31            break
32
33    return trajectory
34
35st.set_page_config(layout="wide")
36st.title("Gradient Descent Visualizer")
37
38# Session state initialization
39if "iterations" not in st.session_state:
40    st.session_state.iterations = 0
41if "formula" not in st.session_state:
42    st.session_state.formula = "x**2"
43if "start_point" not in st.session_state:
44    st.session_state.start_point = 5.0
45if "learning_rate" not in st.session_state:
46    st.session_state.learning_rate = 0.25
47if "show_all_iterations" not in st.session_state:
48    st.session_state.show_all_iterations = False
49if "tolerance" not in st.session_state:
50    st.session_state.tolerance = 0.01
51
52def get_ml_formula(algorithm):
53    formulas = {
54        "Linear Regression": "(x - 2)**2",
55        "Logistic Regression": "log(1 + exp(-x))"
56    }
57    return formulas.get(algorithm, "x**2")
58
59with st.sidebar:
60    st.write("## Inputs")
61    formula_input = st.text_input("Function", value=st.session_state.formula)
62    if formula_input != st.session_state.formula:
63        st.session_state.formula = formula_input
64        st.session_state.iterations = 0
65        st.session_state.trajectory = []
66
67    # Function buttons
68    col1, col2, col3 = st.columns(3)
69    with col1:
70        if st.button("x^2"):
71            st.session_state.formula = "x**2"
72            st.session_state.iterations = 0
73            st.session_state.trajectory = []
74    with col2:
75        if st.button("sin(x)"):
76            st.session_state.formula = "sin(x)"
77            st.session_state.iterations = 0
78            st.session_state.trajectory = []
79    with col3:
80        if st.button("cos(x)"):
81            st.session_state.formula = "cos(x)"
82            st.session_state.iterations = 0
83            st.session_state.trajectory = []
84
85    col4, col5, col6 = st.columns(3)
86    with col4:
87        if st.button("tan(x)"):
88            st.session_state.formula = "tan(x)"
89            st.session_state.iterations = 0
90            st.session_state.trajectory = []
91    with col5:
92        if st.button("exp(x)"):
93            st.session_state.formula = "exp(x)"
94            st.session_state.iterations = 0
95            st.session_state.trajectory = []
96    with col6:
97        if st.button("log(x)"):
98            st.session_state.formula = "log(x)"
99            st.session_state.iterations = 0
100            st.session_state.trajectory = []
101
102    # ML function buttons
103    st.write("## ML Optimized Equations")
104    col7, col8 = st.columns(2)
105    with col7:
106        if st.button("Linear Regression"):
107            st.session_state.formula = get_ml_formula("Linear Regression")
108            st.session_state.iterations = 0
109            st.session_state.trajectory = []
110    with col8:
111        if st.button("Logistic Regression"):
112            st.session_state.formula = get_ml_formula("Logistic Regression")
113            st.session_state.iterations = 0
114            st.session_state.trajectory = []
115
116    # User inputs
117    start_point_input = st.number_input("Starting Point", value=st.session_state.start_point)
118    if start_point_input != st.session_state.start_point:
119        st.session_state.start_point = start_point_input
120        st.session_state.iterations = 0
121        st.session_state.trajectory = []
122
123    learning_rate_input = st.number_input("Learning Rate", value=st.session_state.learning_rate, min_value=0.01, step=0.01)
124    if learning_rate_input != st.session_state.learning_rate:
125        st.session_state.learning_rate = learning_rate_input
126        st.session_state.iterations = 0
127        st.session_state.trajectory = []
128
129    tolerance_input = st.number_input("Tolerance", value=st.session_state.tolerance, min_value=0.001, step=0.001)
130    if tolerance_input != st.session_state.tolerance:
131        st.session_state.tolerance = tolerance_input
132        st.session_state.iterations = 0
133        st.session_state.trajectory = []
134
135    # Action buttons
136    if st.button("Next Iteration", key="next_iteration_button"):
137        st.session_state.iterations += 1
138
139    if st.button("Show All Iterations", key="toggle_iterations_button"):
140        st.session_state.show_all_iterations = not st.session_state.show_all_iterations
141
142    zoom_factor = st.slider("Zoom Level", min_value=1, max_value=30, value=10, step=1)
143
144# Function evaluation
145trajectory = []
146try:
147    x_sym = symbols('x')
148    formula = st.session_state.formula.replace('np.maximum', 'Max').replace('np.piecewise', 'Piecewise')
149    formula = formula.replace('Max', 'np.maximum')
150    formula_sym = eval(formula, {"x": x_sym, "sin": sin, "cos": cos, "tan": tan, "exp": exp, "log": log, "np": np})
151    y_func = lambdify(x_sym, formula_sym, 'numpy')
152
153    def safe_y_func(x):
154        if st.session_state.formula == "log(x)":
155            x = np.clip(x, 1e-10, None)
156        return y_func(x)
157
158    trajectory = gradient_descent(st.session_state.formula, st.session_state.start_point, st.session_state.learning_rate, iterations=st.session_state.iterations)
159except Exception as e:
160    st.error(f"Error in formula evaluation: {e}")
161
162# Minima check
163if trajectory:
164    current_x = trajectory[-1]
165    previous_x = trajectory[-2] if len(trajectory) > 1 else None
166    difference = abs(current_x - previous_x) if previous_x is not None else None
167    if difference is not None and difference <= st.session_state.tolerance:
168        st.success(f"Found minima at position: {current_x:.4f}")
169
170# Plot
171x_start = min(trajectory) - zoom_factor if trajectory else -10
172x_end = max(trajectory) + zoom_factor if trajectory else 10
173x = np.linspace(x_start, x_end, 500)
174try:
175    y = safe_y_func(x)
176except Exception as e:
177    st.error(f"Error in generating graph: {e}")
178    y = None
179
180if y is not None:
181    fig = go.Figure()
182    fig.add_trace(go.Scatter(x=x, y=y, mode='lines', name=f"y = {st.session_state.formula}"))
183
184    if trajectory:
185        fig.add_trace(go.Scatter(x=trajectory[:-1], y=[safe_y_func(t) for t in trajectory[:-1]],
186                                 mode='markers', name="Previous Points", marker=dict(color='yellow')))
187        fig.add_trace(go.Scatter(x=[trajectory[-1]], y=[safe_y_func(trajectory[-1])],
188                                 mode='markers', name="Current Point", marker=dict(color='red', size=10)))
189
190        grad_at_current = (safe_y_func(current_x + 1e-5) - safe_y_func(current_x)) / 1e-5
191        tangent_y = grad_at_current * (x - current_x) + safe_y_func(current_x)
192        fig.add_trace(go.Scatter(x=x, y=tangent_y, mode='lines', name="Tangent Line", line=dict(color='orange')))
193
194    fig.update_layout(title=f"Graph of {st.session_state.formula}",
195                      xaxis_title="x", yaxis_title="y",
196                      template="plotly_white",
197                      width=1000, height=600)
198
199    st.plotly_chart(fig)
200
201# Current iteration table
202if trajectory:
203    current_data = {
204        "Iteration": [st.session_state.iterations],
205        "Current X": [f"{current_x:.2f}"],
206        "Previous X": [f"{previous_x:.2f}" if previous_x is not None else "N/A"],
207        "Difference": [f"{difference:.2f}" if difference is not None else "N/A"]
208    }
209    st.write("### Current Iteration Details")
210    st.table(current_data)
211
212# All iterations table
213if trajectory and st.session_state.show_all_iterations:
214    all_iterations_data = []
215    for i, val in enumerate(trajectory[1:], start=1):
216        prev_val = trajectory[i - 1]
217        diff_val = val - prev_val
218        all_iterations_data.append({
219            "Iteration": i,
220            "Current X": f"{val:.2f}",
221            "Previous X": f"{prev_val:.2f}",
222            "Difference": f"{diff_val:.2f}"
223        })
224    st.write("### All Iterations")
225    st.table(all_iterations_data)
226
227