Jenjo79/kronos-forecast
0
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>▲ Kronos Forecast API</h1>206 <p>Endpoints:</p>207 <ul>208 <li><code>POST /api/predict</code> — run a Kronos forecast (~30s)</li>209 <li><code>POST /api/spot</code> — just current price (fast, no model)</li>210 <li><code>GET /health</code> — liveness check</li>211 </ul>212 <p><a href="/docs">Interactive API docs →</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 