AmitPandit175/Gradient_Descent_Visualization_App
0
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 