CoolFace
Apppublic

Jenjo79/kronos-forecast

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
app.py292 linesDownload Raw Back to root
1"""2Kronos Forecast API — FastAPI version (HF Space backend).3 4Endpoints:5  GET  /            — landing page6  GET  /health      — liveness check7  POST /api/predict — full Kronos forecast for a ticker8  POST /api/spot    — just the current/recent price for a ticker (cheap, no model)9"""10 11import os12import sys13import json14from datetime import datetime15 16from fastapi import FastAPI17from fastapi.middleware.cors import CORSMiddleware18from fastapi.responses import HTMLResponse, JSONResponse19from pydantic import BaseModel20 21import numpy as np22import pandas as pd23import yfinance as yf24import torch  # noqa: F40125 26sys.path.insert(0, os.path.join(os.path.dirname(__file__), "Kronos"))27from model import Kronos, KronosTokenizer, KronosPredictor  # noqa: E40228 29# ---------------------------------------------------------------------------30# Model loading at startup31# ---------------------------------------------------------------------------32print("Loading Kronos tokenizer + model (this takes ~30s on first run)...")33TOKENIZER = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-2k")34MODEL = Kronos.from_pretrained("NeoQuasar/Kronos-mini")35PREDICTOR = KronosPredictor(MODEL, TOKENIZER, device="cpu", max_context=512)36print("Model ready.")37 38INTERVAL_MAP = {39    "1h": {"yf_interval": "1h", "yf_period": "60d", "freq": "1H"},40    "1d": {"yf_interval": "1d", "yf_period": "2y", "freq": "1D"},41}42 43 44def fetch_ohlcv(ticker: str, interval: str = "1h") -> pd.DataFrame:45    cfg = INTERVAL_MAP[interval]46    df = yf.download(47        ticker,48        period=cfg["yf_period"],49        interval=cfg["yf_interval"],50        auto_adjust=False,51        progress=False,52    )53    if df.empty:54        raise ValueError(f"No data for ticker {ticker!r} at interval {interval}")55 56    if isinstance(df.columns, pd.MultiIndex):57        df.columns = df.columns.get_level_values(0)58 59    df = df.rename(60        columns={61            "Open": "open",62            "High": "high",63            "Low": "low",64            "Close": "close",65            "Volume": "volume",66        }67    )68    df = df[["open", "high", "low", "close", "volume"]].dropna()69    df["amount"] = df["close"] * df["volume"]70    df.index.name = "timestamp"71    return df72 73 74def run_forecast(75    ticker: str = "SPY",76    interval: str = "1d",77    lookback: int = 200,78    horizon: int = 5,79    n_samples: int = 30,80    temperature: float = 1.0,81    top_p: float = 0.9,82):83    df = fetch_ohlcv(ticker, interval=interval)84    df = df.tail(lookback).reset_index()85    df = df.rename(columns={df.columns[0]: "timestamp"})86 87    x_df = df[["open", "high", "low", "close", "volume", "amount"]]88    x_ts = pd.to_datetime(df["timestamp"])89 90    freq = INTERVAL_MAP[interval]["freq"]91    last_ts = x_ts.iloc[-1]92    y_ts = pd.Series(pd.date_range(start=last_ts, periods=horizon + 1, freq=freq)[1:])93 94    preds = []95    for _ in range(n_samples):96        out = PREDICTOR.predict(97            df=x_df,98            x_timestamp=x_ts,99            y_timestamp=y_ts,100            pred_len=horizon,101            T=temperature,102            top_p=top_p,103            sample_count=1,104            verbose=False,105        )106        preds.append(out["close"].values)107 108    preds = np.stack(preds, axis=0)109    mean = preds.mean(axis=0)110    low = np.percentile(preds, 10, axis=0)111    high = np.percentile(preds, 90, axis=0)112 113    last_close = float(x_df["close"].iloc[-1])114    terminal = preds[:, -1]115    bullish_prob = float((terminal > last_close).mean())116 117    recent_returns = np.diff(np.log(x_df["close"].values[-horizon:]))118    recent_vol = float(np.std(recent_returns)) if len(recent_returns) > 1 else 0.0119    pred_returns = np.diff(np.log(preds), axis=1)120    pred_vols = np.std(pred_returns, axis=1)121    vol_expansion_prob = (122        float((pred_vols > recent_vol).mean()) if recent_vol > 0 else 0.5123    )124 125    expected_change_pct = float((mean[-1] - last_close) / last_close * 100.0)126 127    history = [128        {129            "t": ts.isoformat(),130            "open": float(o),131            "high": float(h),132            "low": float(l),133            "close": float(c),134        }135        for ts, o, h, l, c in zip(136            x_ts, x_df["open"], x_df["high"], x_df["low"], x_df["close"]137        )138    ]139    forecast_mean = [140        {"t": ts.isoformat(), "close": float(v)} for ts, v in zip(y_ts, mean)141    ]142    forecast_low = [143        {"t": ts.isoformat(), "close": float(v)} for ts, v in zip(y_ts, low)144    ]145    forecast_high = [146        {"t": ts.isoformat(), "close": float(v)} for ts, v in zip(y_ts, high)147    ]148 149    return {150        "ticker": ticker,151        "interval": interval,152        "generated_at": datetime.utcnow().isoformat() + "Z",153        "last_close": last_close,154        "history": history,155        "forecast_mean": forecast_mean,156        "forecast_low": forecast_low,157        "forecast_high": forecast_high,158        "metrics": {159            "bullish_prob": bullish_prob,160            "vol_expansion_prob": vol_expansion_prob,161            "expected_change_pct": expected_change_pct,162            "n_samples": n_samples,163            "horizon": horizon,164            "lookback": lookback,165        },166    }167 168 169# ---------------------------------------------------------------------------170# FastAPI app171# ---------------------------------------------------------------------------172app = FastAPI(title="Kronos Forecast API")173 174app.add_middleware(175    CORSMiddleware,176    allow_origins=["*"],177    allow_credentials=False,178    allow_methods=["*"],179    allow_headers=["*"],180)181 182 183class PredictRequest(BaseModel):184    # {"data": [ticker, interval, lookback, horizon, n_samples]}185    data: list186 187 188class SpotRequest(BaseModel):189    # {"data": [ticker, interval]}  — interval optional, defaults to "1d"190    data: list191 192 193@app.get("/", response_class=HTMLResponse)194def root():195    return """196    <html><head><title>Kronos Forecast API</title>197    <style>198      body { font-family: ui-monospace, monospace; background: #0e0e0c; color: #e8e8e6;199             padding: 40px; max-width: 720px; margin: auto; line-height: 1.6; }200      h1 { color: #ff6b00; }201      code { background: #1a1a17; padding: 2px 6px; color: #ff9b50; }202      pre { background: #16161300; border: 1px solid #26261f; padding: 16px; overflow-x: auto; }203      a { color: #ff6b00; }204    </style></head><body>205      <h1>&#9650; Kronos Forecast API</h1>206      <p>Endpoints:</p>207      <ul>208        <li><code>POST /api/predict</code> &mdash; run a Kronos forecast (~30s)</li>209        <li><code>POST /api/spot</code> &mdash; just current price (fast, no model)</li>210        <li><code>GET /health</code> &mdash; liveness check</li>211      </ul>212      <p><a href="/docs">Interactive API docs &rarr;</a></p>213    </body></html>214    """215 216 217@app.get("/health")218def health():219    return {"status": "ok", "model": "Kronos-mini", "device": "cpu"}220 221 222@app.post("/api/predict")223def api_predict(req: PredictRequest):224    """Returns Gradio-compatible envelope so the existing frontend works."""225    try:226        if len(req.data) < 5:227            raise ValueError(228                "Expected 5 args: [ticker, interval, lookback, horizon, n_samples]"229            )230        ticker, interval, lookback, horizon, n_samples = req.data[:5]231        result = run_forecast(232            ticker=str(ticker).upper().strip(),233            interval=str(interval),234            lookback=int(lookback),235            horizon=int(horizon),236            n_samples=int(n_samples),237        )238        return {"data": [json.dumps(result)]}239    except Exception as e:240        return JSONResponse(241            status_code=500,242            content={"data": [json.dumps({"error": str(e)})]},243        )244 245 246@app.post("/api/spot")247def api_spot(req: SpotRequest):248    """249    Cheap endpoint: returns the most recent close + a few recent bars for a250    ticker. Used by the History view to compare predictions against actuals251    without burning a full Kronos forecast.252    """253    try:254        if len(req.data) < 1:255            raise ValueError("Expected at least 1 arg: [ticker, interval?]")256        ticker = str(req.data[0]).upper().strip()257        interval = str(req.data[1]) if len(req.data) > 1 else "1d"258        df = fetch_ohlcv(ticker, interval=interval)259        # Return the last 30 bars so the client can compare predicted vs actual260        df = df.tail(30).reset_index()261        df = df.rename(columns={df.columns[0]: "timestamp"})262        bars = [263            {264                "t": pd.Timestamp(row["timestamp"]).isoformat(),265                "open": float(row["open"]),266                "high": float(row["high"]),267                "low": float(row["low"]),268                "close": float(row["close"]),269            }270            for _, row in df.iterrows()271        ]272        result = {273            "ticker": ticker,274            "interval": interval,275            "fetched_at": datetime.utcnow().isoformat() + "Z",276            "last_close": float(df["close"].iloc[-1]),277            "last_t": pd.Timestamp(df["timestamp"].iloc[-1]).isoformat(),278            "bars": bars,279        }280        return {"data": [json.dumps(result)]}281    except Exception as e:282        return JSONResponse(283            status_code=500,284            content={"data": [json.dumps({"error": str(e)})]},285        )286 287 288if __name__ == "__main__":289    import uvicorn290 291    uvicorn.run(app, host="0.0.0.0", port=7860)292