ariesljm/kronos
0
1from fastapi import FastAPI, HTTPException2from pydantic import BaseModel3import pandas as pd4import torch5import sys6import os7from datetime import timedelta8 9# 引入 Kronos10sys.path.append(os.path.join(os.path.dirname(__file__), "Kronos"))11from model import Kronos, KronosTokenizer, KronosPredictor12 13app = FastAPI()14 15# 全局加载模型 (启动时加载一次)16device = "cuda" if torch.cuda.is_available() else "cpu"17print(f"Loading model on {device}...")18 19tokenizer = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-base")20# 如果 HF 免费版 CPU 跑不动 Base,可以改回 small21model = Kronos.from_pretrained("NeoQuasar/Kronos-base") 22model = model.to(device)23model.eval()24 25predictor = KronosPredictor(model, tokenizer, device=device, max_context=512)26print("Model loaded!")27 28class CandleData(BaseModel):29 # 接收的数据格式30 pair: str31 data: list[dict] # 包含 open, high, low, close, volume, date 的列表32 33@app.post("/predict")34def predict(payload: CandleData):35 try:36 # 1. 重建 DataFrame37 df = pd.DataFrame(payload.data)38 if 'date' in df.columns:39 df['timestamps'] = pd.to_datetime(df['date'])40 else:41 # 如果传来的是时间戳42 df['timestamps'] = pd.to_datetime(df['date'])43 44 input_cols = ["open", "high", "low", "close", "volume"]45 46 # 2. 准备时间戳47 x_timestamp = df["timestamps"]48 pred_steps = 349 last_time = x_timestamp.iloc[-1]50 y_timestamp = pd.Series([51 last_time + timedelta(minutes=5 * (i + 1)) 52 for i in range(pred_steps)53 ])54 55 # 3. 推理56 with torch.no_grad():57 forecast = predictor.predict(58 df=df[input_cols],59 x_timestamp=x_timestamp,60 y_timestamp=y_timestamp,61 pred_len=pred_steps,62 T=0.8,63 top_p=0.9,64 sample_count=10 # CPU如果太慢,可以把这里改为 3 或 565 )66 67 # 4. 返回结果68 raw_pred = forecast["close"].mean()69 return {"prediction": float(raw_pred)}70 71 except Exception as e:72 print(f"Error: {e}")73 raise HTTPException(status_code=500, detail=str(e))