myspacev1 / app.py
omkar2323's picture
Update app.py
4dab8dd verified
Raw
History Blame Contribute Delete
1.46 kB
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"}