financehub-api / app.py
rohan1324's picture
update generate call
47f4333 verified
Raw
History Blame Contribute Delete
2.8 kB
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())