from fastapi import FastAPI, HTTPException from pydantic import BaseModel from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline from contextlib import asynccontextmanager import os # Model configuration - use environment variable if available (for Spaces) repo_id = os.getenv("MODEL_ID", "omkar2323/my-llm-model") # Global variable for the pipeline pipe = None # Lifespan context manager to load model at startup @asynccontextmanager async def lifespan(app: FastAPI): # Load model on startup global pipe tokenizer = AutoTokenizer.from_pretrained(repo_id) model = AutoModelForCausalLM.from_pretrained(repo_id) pipe = pipeline("text-generation", model=model, tokenizer=tokenizer) yield # Clean up resources on shutdown del pipe # Create FastAPI app app = FastAPI(lifespan=lifespan) # Define request and response models class GenerateRequest(BaseModel): prompt: str max_new_tokens: int = 50 class GenerateResponse(BaseModel): generated_text: str @app.post("/generate", response_model=GenerateResponse) async def generate_text(request: GenerateRequest): try: result = pipe(request.prompt, max_new_tokens=request.max_new_tokens) return GenerateResponse(generated_text=result[0]["generated_text"]) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @app.get("/") async def root(): return {"message": "Text generation API is running"}