File size: 1,458 Bytes
4dab8dd 3070fb1 4dab8dd 3070fb1 4dab8dd 63ddfb5 4dab8dd 63ddfb5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 | 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"}
|