import logging from contextlib import asynccontextmanager import torch from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from peft import PeftModel from pydantic import BaseModel from transformers import AutoModelForCausalLM, AutoTokenizer logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) _model = None _tokenizer = None BASE_MODEL_ID = "microsoft/Phi-3-mini-4k-instruct" ADAPTER_ID = "rohan1324/phi3-mini-finance-qlora" @asynccontextmanager async def lifespan(app: FastAPI): global _model, _tokenizer logger.info("Loading tokenizer...") _tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL_ID, trust_remote_code=True) logger.info("Loading base model in float16 on CPU...") base = AutoModelForCausalLM.from_pretrained( BASE_MODEL_ID, torch_dtype=torch.float16, device_map="cpu", trust_remote_code=True, low_cpu_mem_usage=True, ) logger.info("Applying LoRA adapter...") _model = PeftModel.from_pretrained(base, ADAPTER_ID) _model.eval() logger.info("Model ready.") yield logger.info("Shutting down.") app = FastAPI(title="Finance Hub Model Server", lifespan=lifespan) app.add_middleware( CORSMiddleware, allow_origins=[ "https://*.render.com", "https://*.vercel.app", "http://localhost:3000", "http://localhost:8000", ], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) class GenerateRequest(BaseModel): prompt: str max_tokens: int = 512 temperature: float = 0.7 class GenerateResponse(BaseModel): generated_text: str @app.get("/health") async def health() -> dict: if _model is None: raise HTTPException(status_code=503, detail="Model not loaded") return {"status": "ok", "model": "loaded"} @app.post("/generate", response_model=GenerateResponse) async def generate(request: GenerateRequest) -> GenerateResponse: if _model is None: raise HTTPException(status_code=503, detail="Model not loaded") if not request.prompt.strip(): raise HTTPException(status_code=400, detail="Prompt cannot be empty") inputs = _tokenizer(request.prompt, return_tensors="pt") input_len = inputs.input_ids.shape[1] with torch.no_grad(): outputs = _model.generate( **inputs, max_new_tokens=request.max_tokens, temperature=request.temperature, do_sample=request.temperature > 0, pad_token_id=_tokenizer.eos_token_id, use_cache=False ) new_tokens = outputs[0][input_len:] generated = _tokenizer.decode(new_tokens, skip_special_tokens=True) return GenerateResponse(generated_text=generated.strip())