CoolFace
Apppublic

ariesljm/kronos

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
app.py73 linesDownload Raw Back to root
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))