from fastapi import FastAPI, HTTPException from pydantic import BaseModel import uvicorn app = FastAPI(title="TimesFM 2.5 API Service") try: import torch # ⚠️ 2.5 版本的最新导入路径,旧版已被 Google 废弃 from timesfm.timesfm_2p5.timesfm_2p5_torch import TimesFM_2p5_200M_torch from timesfm import ForecastConfig # 针对免费 CPU 环境的优化设置 torch.set_float32_matmul_precision("high") # 加载 2.5 版本的轻量化权重 (200M参数) tfm = TimesFM_2p5_200M_torch.from_pretrained("google/timesfm-2.5-200m-pytorch") # 2.5 版本的全新配置方法 (Compile) tfm.compile( ForecastConfig( max_context=512, # 最大支持的历史上下文长度 max_horizon=128, # 最大支持的预测步数 normalize_inputs=True, # 自动对数据进行归一化处理 ) ) print("TimesFM 2.5 模型已成功加载并编译!") except Exception as e: print(f"模型加载失败: {e}") tfm = None class ForecastRequest(BaseModel): history: list[float] @app.post("/predict") async def predict(data: ForecastRequest): if tfm is None: raise HTTPException(status_code=500, detail="模型未成功加载,请查看后台运行日志") if len(data.history) < 8: raise HTTPException(status_code=400, detail="历史数据至少需要 8 位") # 截取最后的 32 位数据作为预测基准 input_series = data.history[-32:] try: # 2.5 版本的全新预测调用格式 point_forecast, _ = tfm.forecast( horizon=7, # 预测未来 7 个时间步 inputs=[input_series] # 新版直接接收 Python List ) # 返回预测结果 return {"status": "success", "forecast": point_forecast[0].tolist()} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @app.get("/") async def root(): return {"status": "healthy", "message": "TimesFM 2.5 极速版运行中!"} if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=7860)