AmitPandit175/Gradient_Descent_Visualization_App
0
1import streamlit as st
2import plotly.graph_objects as go
3import numpy as np
4from scipy.optimize import fsolve
5
6st.set_page_config(page_title="Function Requirements", layout="wide")
7st.title("๐ง Function Requirements for Gradient Descent")
8
9with st.expander("๐ Introduction", expanded=True):
10 st.markdown("""
11 Gradient Descent requires a function to meet **two key conditions**:
12
13 1. ๐งฎ **Differentiability** โ The function must have a well-defined derivative at all points.
14 2. ๐ **Convexity** โ Ideally, the function should be convex to guarantee a single global minimum.
15 """)
16
17# ------------------------ Differentiability Section ------------------------
18st.markdown("---")
19st.header("1๏ธโฃ Differentiability")
20
21with st.expander("๐ Explanation"):
22 st.markdown("""
23 A function is **differentiable** if it has a defined slope (derivative) at every point.
24 Gradient Descent depends on the slope to move toward the minimum value.
25
26 - **Continuity**: No breaks or jumps.
27 - **Differentiability**: Smooth curves with well-defined tangents.
28
29 ### โ๏ธ Why it matters:
30 If a function isn't differentiable at a point, Gradient Descent might behave **unpredictably**.
31
32 ---
33 ### โ
Example: $f(x) = x^2$
34 - It's smooth and differentiable everywhere.
35 - Derivative: $f'(x) = 2x$
36 """)
37
38# Plot f(x) = x^2
39x = np.linspace(-10, 10, 500)
40y1 = x**2
41y2 = 2 * x
42
43fig1 = go.Figure()
44fig1.add_trace(go.Scatter(x=x, y=y1, mode='lines', name='f(x) = xยฒ', line=dict(color='blue')))
45fig1.update_layout(title="๐ท Plot of f(x) = xยฒ", xaxis_title="x", yaxis_title="y")
46
47fig2 = go.Figure()
48fig2.add_trace(go.Scatter(x=x, y=y2, mode='lines', name="f'(x) = 2x", line=dict(color='red')))
49fig2.update_layout(title="๐บ Derivative: f'(x) = 2x", xaxis_title="x", yaxis_title="y")
50
51st.plotly_chart(fig1, use_container_width=True)
52st.plotly_chart(fig2, use_container_width=True)
53
54with st.expander("โ ๏ธ Non-Differentiable Example: $f(x) = |x|$"):
55 st.markdown("""
56 - For $x>0$, derivative = 1
57 - For $x<0$, derivative = -1
58 - At $x=0$: **undefined**
59
60 โ This creates a **sharp corner**, making Gradient Descent unreliable there.
61 """)
62
63fig3 = go.Figure()
64fig3.add_trace(go.Scatter(x=x, y=np.abs(x), mode='lines', name="f(x) = |x|", line=dict(color='green')))
65fig3.update_layout(title="๐ฉ Plot of f(x) = |x| (Not Differentiable at x = 0)", xaxis_title="x", yaxis_title="y")
66st.plotly_chart(fig3, use_container_width=True)
67
68# Gradient Explanation
69with st.expander("๐ Gradient in Higher Dimensions"):
70 st.markdown("""
71 > In higher dimensions, **gradient** is a vector of partial derivatives:
72
73 $$
74 \\nabla f(p) = \\left[ \\frac{\\partial f}{\\partial x_1}, \\frac{\\partial f}{\\partial x_2}, \\dots, \\frac{\\partial f}{\\partial x_n} \\right]
75 $$
76
77 - It **points toward the steepest ascent**.
78 - Gradient Descent moves in the opposite direction to **minimize** the function.
79 """)
80
81# ------------------------ Convexity Section ------------------------
82st.markdown("---")
83st.header("2๏ธโฃ Convexity")
84
85with st.expander("๐ Convexity Explained"):
86 st.markdown("""
87 A function is **convex** if a line between any two points on its curve lies **above or on the curve**.
88
89 ### ๐งฎ Convexity Condition:
90
91 $$
92 f(\\lambda x_1 + (1 - \\lambda)x_2) \\leq \\lambda f(x_1) + (1 - \\lambda)f(x_2)
93 $$
94
95 Convex functions ensure:
96 - **Global minimum** is always reachable.
97 - Gradient Descent won't get stuck in local minima.
98
99 ---
100 ### โ
Convex Function: $f(x) = x^2$
101 """)
102
103# Convex Function Plot
104x = np.linspace(-3, 3, 500)
105fig_convex = go.Figure()
106fig_convex.add_trace(go.Scatter(x=x, y=x**2, mode='lines', name='f(x) = xยฒ', line=dict(color='blue')))
107fig_convex.add_trace(go.Scatter(x=x, y=2*x, mode='lines', name="f'(x) = 2x", line=dict(color='orange')))
108fig_convex.update_layout(title="๐ Convex Function and Its Derivative", xaxis_title="x", yaxis_title="y")
109st.plotly_chart(fig_convex, use_container_width=True)
110
111# Non-Convex Function Plot
112x_non_convex = np.linspace(-10, 10, 500)
113y_non_convex = np.sin(x_non_convex) + 0.1 * x_non_convex**2
114y_prime_non_convex = np.cos(x_non_convex) + 0.2 * x_non_convex
115
116fig_non_convex = go.Figure()
117fig_non_convex.add_trace(go.Scatter(x=x_non_convex, y=y_non_convex, mode='lines', name='f(x) = sin(x) + 0.1xยฒ', line=dict(color='blue')))
118fig_non_convex.add_trace(go.Scatter(x=x_non_convex, y=y_prime_non_convex, mode='lines', name="f'(x)", line=dict(color='red')))
119fig_non_convex.update_layout(title="๐ซ Non-Convex Function and Its Derivative", xaxis_title="x", yaxis_title="y")
120st.plotly_chart(fig_non_convex, use_container_width=True)
121
122# ------------------------ Saddle Points Section ------------------------
123st.markdown("---")
124st.header("3๏ธโฃ Saddle Points & Semi-Convexity")
125
126with st.expander("๐ What are Saddle Points?"):
127 st.markdown("""
128 A **saddle point** is where the gradient is zero, but the point is **neither a minimum nor a maximum**.
129
130 - **Gradient = 0**, but descent is **ambiguous**.
131 - Causes **stalling** in high-dimensional Gradient Descent.
132
133 ---
134 ### ๐ธ Example: $f(x) = x^4 - 4x^2 + x$
135 """)
136
137x = np.linspace(-3, 3, 500)
138y_semi_convex = x**4 - 4*x**2 + x
139
140fig_semi_convex = go.Figure()
141fig_semi_convex.add_trace(go.Scatter(x=x, y=y_semi_convex, mode='lines', name='f(x) = xโด - 4xยฒ + x', line=dict(color='blue')))
142
143def derivative(x):
144 return 4*x**3 - 8*x + 1
145
146critical_points = fsolve(derivative, [-2, 0, 2])
147critical_values = np.interp(critical_points, x, y_semi_convex)
148
149fig_semi_convex.add_trace(go.Scatter(x=critical_points, y=critical_values, mode='markers', name='Critical Points', marker=dict(color='yellow', size=10)))
150fig_semi_convex.add_trace(go.Scatter(x=[0], y=[0], mode='markers+text', name='Saddle Point', marker=dict(color='red', size=12), textposition="top center"))
151
152fig_semi_convex.update_layout(title="๐งฉ Semi-Convex Function with Saddle Points", xaxis_title="x", yaxis_title="y")
153st.plotly_chart(fig_semi_convex, use_container_width=True)
154
155# ------------------------ Summary ------------------------
156with st.expander("๐ Summary"):
157 st.markdown("""
158 โ
For **Gradient Descent** to perform effectively:
159
160 - Function must be **continuous and differentiable**.
161 - **Convexity** ensures convergence to a global minimum.
162 - **Saddle points** and **non-differentiable points** can cause issues.
163
164 Use visualizations and derivatives to **understand and debug** optimization functions!
165 """)
166 