from fastapi import FastAPI, HTTPException import chess from contextlib import asynccontextmanager import sys from src.model import ChessPolicyModel, PolicyModelInference from src.tokenizer import tokenizer import torch from pydantic import BaseModel ml = {} @asynccontextmanager async def lifespan(app: FastAPI): tokenizer = torch.load("./model/tokenizer.pt", weights_only=False, map_location=torch.device('cpu')) model = ChessPolicyModel(vocab_size=tokenizer.language_size) model.load_state_dict( torch.load("./model/policy_model.pt", weights_only=False, map_location=torch.device('cpu')) ) ml["inference"] = PolicyModelInference(model, tokenizer, device="cpu") yield ml.clear() app = FastAPI(lifespan=lifespan) class InferenceRequest(BaseModel): moves: list[str] @app.post("/inference") def model_inference(req: InferenceRequest): board = chess.Board() for move in req.moves: try: board.push_uci(move) except ValueError as e: raise HTTPException(status_code=400, detail=f"Incorrect move {move}: {e}") try: return {"move" : ml["inference"](board)} except ValueError as e: raise HTTPException(status_code=500, detail=f"Model Failed to evaluate: {e}")