Team Ai
Apppublic

AmitPandit175/Gradient_Descent_Visualization_App

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
2_Function_requirements.py166 linesDownload Raw Back to pages
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