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"}