import gradio as gr import yfinance as yf import pandas as pd import numpy as np import plotly.graph_objects as go from plotly.subplots import make_subplots from sklearn.preprocessing import MinMaxScaler from sklearn.linear_model import LinearRegression from sklearn.ensemble import RandomForestRegressor, GradientBoostingRegressor from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score import warnings warnings.filterwarnings("ignore") # ── Theme colours (dark terminal-finance aesthetic) ────────────────────────── BG = "#0d1117" SURFACE = "#161b22" BORDER = "#30363d" GREEN = "#3fb950" RED = "#f85149" BLUE = "#58a6ff" YELLOW = "#d29922" TEXT = "#e6edf3" MUTED = "#8b949e" POPULAR_TICKERS = [ "AAPL", "MSFT", "GOOGL", "AMZN", "TSLA", "META", "NVDA", "JPM", "BRK-B", "V", "NFLX", "DIS", "BABA", "AMD", "INTC", "UBER", "SPOT", "PYPL", "SQ", "SNAP", ] MODELS = { "Linear Regression": LinearRegression(), "Random Forest": RandomForestRegressor(n_estimators=100, random_state=42), "Gradient Boosting": GradientBoostingRegressor(n_estimators=100, random_state=42), } PERIODS = { "6 Months": "6mo", "1 Year": "1y", "2 Years": "2y", "5 Years": "5y", "10 Years": "10y", } INTERVALS = { "Daily": "1d", "Weekly": "1wk", } # ── Feature engineering ─────────────────────────────────────────────────────── def make_features(df: pd.DataFrame) -> pd.DataFrame: df = df.copy() df["MA7"] = df["Close"].rolling(7).mean() df["MA21"] = df["Close"].rolling(21).mean() df["MA50"] = df["Close"].rolling(50).mean() df["EMA12"] = df["Close"].ewm(span=12, adjust=False).mean() df["EMA26"] = df["Close"].ewm(span=26, adjust=False).mean() df["MACD"] = df["EMA12"] - df["EMA26"] df["Vol_MA"] = df["Volume"].rolling(7).mean() df["Return1"] = df["Close"].pct_change(1) df["Return5"] = df["Close"].pct_change(5) df["High_Low"] = df["High"] - df["Low"] df["Close_Open"] = df["Close"] - df["Open"] delta = df["Close"].diff() gain = delta.clip(lower=0).rolling(14).mean() loss = (-delta.clip(upper=0)).rolling(14).mean() rs = gain / loss.replace(0, np.nan) df["RSI"] = 100 - (100 / (1 + rs)) df["Target"] = df["Close"].shift(-1) return df.dropna() # ── Core prediction logic ───────────────────────────────────────────────────── def predict_stock(ticker, period_label, interval_label, model_name, future_days): ticker = ticker.strip().upper() period = PERIODS[period_label] interval = INTERVALS[interval_label] try: raw = yf.download(ticker, period=period, interval=interval, progress=False) if raw.empty: return None, None, f"❌ No data found for **{ticker}**. Check the ticker symbol." if isinstance(raw.columns, pd.MultiIndex): raw.columns = raw.columns.get_level_values(0) raw = raw[["Open","High","Low","Close","Volume"]].dropna() except Exception as e: return None, None, f"❌ Data fetch error: {e}" if len(raw) < 60: return None, None, f"❌ Not enough data ({len(raw)} rows). Try a longer period." df = make_features(raw) feature_cols = [ "MA7","MA21","MA50","EMA12","EMA26","MACD", "Vol_MA","Return1","Return5","High_Low","Close_Open","RSI", "Open","High","Low","Volume", ] X = df[feature_cols].values y = df["Target"].values split = int(len(X) * 0.8) X_train, X_test = X[:split], X[split:] y_train, y_test = y[:split], y[split:] scaler = MinMaxScaler() X_train_s = scaler.fit_transform(X_train) X_test_s = scaler.transform(X_test) model = MODELS[model_name] model.fit(X_train_s, y_train) y_pred = model.predict(X_test_s) mae = mean_absolute_error(y_test, y_pred) rmse = np.sqrt(mean_squared_error(y_test, y_pred)) r2 = r2_score(y_test, y_pred) acc = max(0.0, r2) * 100 # ── Future forecast ─────────────────────────────────────────────────────── future_days = int(future_days) last_features = X[-1].reshape(1, -1) future_prices = [] cur = last_features.copy() for _ in range(future_days): cur_s = scaler.transform(cur) nxt = model.predict(cur_s)[0] future_prices.append(float(nxt)) cur[0, 0] = nxt # crude: update Close proxy last_date = df.index[-1] freq = "B" if interval == "1d" else "W" fut_dates = pd.date_range(last_date, periods=future_days + 1, freq=freq)[1:] # ── Candlestick + prediction chart ─────────────────────────────────────── fig = make_subplots( rows=3, cols=1, shared_xaxes=True, row_heights=[0.55, 0.25, 0.20], vertical_spacing=0.04, subplot_titles=("Price & Prediction", "Volume", "RSI"), ) test_dates = df.index[split:] # Candlestick (all history) fig.add_trace(go.Candlestick( x=raw.index, open=raw["Open"], high=raw["High"], low=raw["Low"], close=raw["Close"], increasing_line_color=GREEN, decreasing_line_color=RED, name="Price", showlegend=False, ), row=1, col=1) # Moving averages for col, color, label in [("MA21", BLUE, "MA 21"), ("MA50", YELLOW, "MA 50")]: fig.add_trace(go.Scatter( x=df.index, y=df[col], line=dict(color=color, width=1.2), name=label, opacity=0.85, ), row=1, col=1) # Test-set predictions fig.add_trace(go.Scatter( x=test_dates, y=y_pred, line=dict(color="#a371f7", width=1.8, dash="dot"), name="Model (test)", opacity=0.9, ), row=1, col=1) # Future forecast fig.add_trace(go.Scatter( x=list(fut_dates), y=future_prices, line=dict(color="#f0883e", width=2), name=f"Forecast ({future_days}d)", mode="lines+markers", marker=dict(size=5), ), row=1, col=1) # Vertical "today" line fig.add_vline(x=str(last_date.date()), line_width=1, line_dash="dash", line_color=MUTED, row=1, col=1) # Volume bars colors_v = [GREEN if c >= o else RED for c, o in zip(raw["Close"], raw["Open"])] fig.add_trace(go.Bar( x=raw.index, y=raw["Volume"], marker_color=colors_v, showlegend=False, name="Volume", ), row=2, col=1) # RSI fig.add_trace(go.Scatter( x=df.index, y=df["RSI"], line=dict(color=BLUE, width=1.5), name="RSI", showlegend=False, ), row=3, col=1) fig.add_hline(y=70, line_dash="dash", line_color=RED, line_width=0.8, row=3, col=1) fig.add_hline(y=30, line_dash="dash", line_color=GREEN, line_width=0.8, row=3, col=1) fig.update_layout( paper_bgcolor=BG, plot_bgcolor=SURFACE, font=dict(family="'JetBrains Mono', monospace", color=TEXT, size=12), legend=dict(bgcolor=SURFACE, bordercolor=BORDER, borderwidth=1, x=0.01, y=0.99, font=dict(size=11)), margin=dict(l=10, r=10, t=40, b=10), xaxis_rangeslider_visible=False, height=700, ) for row in [1, 2, 3]: fig.update_xaxes( gridcolor=BORDER, showgrid=True, zeroline=False, row=row, col=1, ) fig.update_yaxes( gridcolor=BORDER, showgrid=True, zeroline=False, row=row, col=1, ) # ── Metrics card (markdown) ─────────────────────────────────────────────── cur_price = float(raw["Close"].iloc[-1]) fst_fcast = future_prices[0] if future_prices else cur_price lst_fcast = future_prices[-1] if future_prices else cur_price delta_pct = (lst_fcast - cur_price) / cur_price * 100 arrow = "▲" if delta_pct >= 0 else "▼" clr_tag = "🟢" if delta_pct >= 0 else "🔴" rsi_now = float(df["RSI"].iloc[-1]) rsi_sig = ("Overbought ⚠️" if rsi_now > 70 else "Oversold 💡" if rsi_now < 30 else "Neutral ✅") stats_md = f""" ## {ticker} — {model_name} | Metric | Value | |--------|-------| | **Current Price** | `${cur_price:,.2f}` | | **Next-Period Forecast** | `${fst_fcast:,.2f}` | | **{future_days}-Day Forecast** | `${lst_fcast:,.2f}` | | **Expected Change** | `{clr_tag} {arrow} {abs(delta_pct):.2f}%` | --- ### Model Performance (test set) | | | |---|---| | MAE | `${mae:.4f}` | | RMSE | `${rmse:.4f}` | | R² | `{r2:.4f}` | | Accuracy proxy | `{acc:.1f}%` | --- ### Technical Signals | Indicator | Value | Signal | |-----------|-------|--------| | RSI (14) | `{rsi_now:.1f}` | {rsi_sig} | | MACD | `{float(df['MACD'].iloc[-1]):.4f}` | {'Bullish 📈' if float(df['MACD'].iloc[-1]) > 0 else 'Bearish 📉'} | | MA21 vs MA50 | — | {'Golden Cross ✨' if float(df['MA21'].iloc[-1]) > float(df['MA50'].iloc[-1]) else 'Death Cross 💀'} | > ⚠️ **Disclaimer:** This tool is for educational purposes only. Not financial advice. """ return fig, stats_md, "" # ── UI layout ───────────────────────────────────────────────────────────────── css = f""" @import url('https://fonts.googleapis.com/css2?family=JetBrains+Mono:wght@300;400;600&family=Inter:wght@400;600&display=swap'); * {{ box-sizing: border-box; }} body, .gradio-container {{ background: {BG} !important; color: {TEXT} !important; font-family: 'Inter', sans-serif !important; }} /* Header */ .app-header {{ background: {SURFACE}; border-bottom: 1px solid {BORDER}; padding: 28px 32px 20px; margin-bottom: 24px; }} .app-header h1 {{ font-family: 'JetBrains Mono', monospace; font-size: 2rem; font-weight: 600; color: {TEXT}; margin: 0 0 4px; letter-spacing: -0.5px; }} .app-header p {{ color: {MUTED}; font-size: 0.9rem; margin: 0; }} .accent {{ color: {GREEN}; }} /* Cards */ .card {{ background: {SURFACE} !important; border: 1px solid {BORDER} !important; border-radius: 8px !important; padding: 16px !important; }} /* Controls */ label {{ color: {MUTED} !important; font-size: 0.78rem !important; font-weight: 600 !important; text-transform: uppercase !important; letter-spacing: 0.08em !important; margin-bottom: 4px !important; }} input, select, .svelte-1gfkn6j {{ background: {BG} !important; border: 1px solid {BORDER} !important; color: {TEXT} !important; border-radius: 6px !important; font-family: 'JetBrains Mono', monospace !important; }} input:focus {{ border-color: {BLUE} !important; outline: none !important; }} /* Run button */ .run-btn button {{ background: {GREEN} !important; color: #0d1117 !important; font-weight: 700 !important; font-size: 0.95rem !important; border: none !important; border-radius: 6px !important; padding: 12px 0 !important; width: 100% !important; cursor: pointer !important; font-family: 'JetBrains Mono', monospace !important; letter-spacing: 0.05em !important; transition: opacity .15s; }} .run-btn button:hover {{ opacity: 0.85; }} /* Metrics markdown */ .stats-box {{ background: {BG} !important; border: 1px solid {BORDER} !important; border-radius: 8px !important; padding: 20px !important; font-family: 'JetBrains Mono', monospace !important; font-size: 0.82rem !important; }} .stats-box table {{ width: 100%; border-collapse: collapse; }} .stats-box td, .stats-box th {{ padding: 6px 10px; border-bottom: 1px solid {BORDER}; text-align: left; }} .stats-box th {{ color: {MUTED}; font-weight: 600; }} .stats-box code {{ background: {SURFACE}; padding: 2px 6px; border-radius: 4px; color: {BLUE}; }} /* Error box */ .error-box textarea {{ background: transparent !important; color: {RED} !important; border: none !important; font-family: 'JetBrains Mono', monospace !important; font-size: 0.85rem !important; }} /* Ticker pills */ .ticker-pills {{ display: flex; flex-wrap: wrap; gap: 6px; margin-top: 8px; }} .ticker-pills button {{ background: {SURFACE} !important; border: 1px solid {BORDER} !important; color: {TEXT} !important; border-radius: 4px !important; padding: 3px 10px !important; font-size: 0.75rem !important; font-family: 'JetBrains Mono', monospace !important; cursor: pointer !important; transition: border-color .15s; }} .ticker-pills button:hover {{ border-color: {GREEN} !important; color: {GREEN} !important; }} /* Plotly panel */ .plot-panel {{ border: 1px solid {BORDER}; border-radius: 8px; overflow: hidden; }} /* Slider */ input[type=range] {{ accent-color: {BLUE}; }} /* Footer */ .footer {{ text-align: center; color: {MUTED}; font-size: 0.75rem; margin-top: 32px; padding: 16px; border-top: 1px solid {BORDER}; }} """ # ── Build Gradio app ────────────────────────────────────────────────────────── with gr.Blocks(css=css, title="StockSense — ML Stock Predictor") as demo: # Header gr.HTML("""
Machine-learning powered stock analysis & price forecasting · Educational use only