Spaces:
Sleeping
Sleeping
| 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(""" | |
| <div class="app-header"> | |
| <h1>π Stock<span class="accent">Sense</span></h1> | |
| <p>Machine-learning powered stock analysis & price forecasting Β· Educational use only</p> | |
| </div> | |
| """) | |
| with gr.Row(): | |
| # ββ Left panel: controls βββββββββββββββββββββββββββββββββββββββββββββ | |
| with gr.Column(scale=1, elem_classes="card"): | |
| gr.HTML("<div style='font-size:0.78rem;color:#8b949e;font-weight:600;text-transform:uppercase;letter-spacing:.08em;margin-bottom:8px'>Quick Pick</div>") | |
| ticker_pills_html = "".join( | |
| f'<button onclick="document.querySelector(\'#ticker_input input\').value=\'{t}\';' | |
| f'document.querySelector(\'#ticker_input input\').dispatchEvent(new Event(\'input\'))">{t}</button>' | |
| for t in POPULAR_TICKERS | |
| ) | |
| gr.HTML(f'<div class="ticker-pills">{ticker_pills_html}</div>') | |
| ticker_input = gr.Textbox( | |
| label="Ticker Symbol", | |
| placeholder="e.g. AAPL, TSLA, MSFT β¦", | |
| value="AAPL", | |
| elem_id="ticker_input", | |
| ) | |
| with gr.Row(): | |
| period_input = gr.Dropdown( | |
| label="History Period", | |
| choices=list(PERIODS.keys()), | |
| value="2 Years", | |
| ) | |
| interval_input = gr.Dropdown( | |
| label="Interval", | |
| choices=list(INTERVALS.keys()), | |
| value="Daily", | |
| ) | |
| model_input = gr.Dropdown( | |
| label="ML Model", | |
| choices=list(MODELS.keys()), | |
| value="Random Forest", | |
| ) | |
| future_input = gr.Slider( | |
| label="Forecast Horizon (days)", | |
| minimum=1, maximum=90, step=1, value=30, | |
| ) | |
| run_btn = gr.Button("βΆ Run Prediction", elem_classes="run-btn") | |
| error_box = gr.Textbox( | |
| visible=True, interactive=False, | |
| show_label=False, elem_classes="error-box", | |
| ) | |
| # Stats card below button | |
| stats_md = gr.Markdown( | |
| value="*Run a prediction to see metrics here.*", | |
| elem_classes="stats-box", | |
| ) | |
| # ββ Right panel: chart βββββββββββββββββββββββββββββββββββββββββββββββ | |
| with gr.Column(scale=3): | |
| chart = gr.Plot(elem_classes="plot-panel", label="") | |
| # Footer | |
| gr.HTML(f""" | |
| <div class="footer"> | |
| StockSense Β· Built with Gradio & scikit-learn Β· | |
| Data via Yahoo Finance Β· | |
| <span style="color:{RED}">Not financial advice</span> | |
| </div> | |
| """) | |
| # ββ Wire up βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def run(ticker, period, interval, model, future): | |
| fig, stats, err = predict_stock(ticker, period, interval, model, future) | |
| return ( | |
| fig if fig else go.Figure(), | |
| stats if stats else "", | |
| err, | |
| ) | |
| run_btn.click( | |
| fn=run, | |
| inputs=[ticker_input, period_input, interval_input, model_input, future_input], | |
| outputs=[chart, stats_md, error_box], | |
| ) | |
| # Auto-run on load | |
| demo.load( | |
| fn=run, | |
| inputs=[ticker_input, period_input, interval_input, model_input, future_input], | |
| outputs=[chart, stats_md, error_box], | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() | |