Spaces:
Sleeping
Sleeping
| 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" | |
| 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 | |
| async def health() -> dict: | |
| if _model is None: | |
| raise HTTPException(status_code=503, detail="Model not loaded") | |
| return {"status": "ok", "model": "loaded"} | |
| 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()) | |