| from fastapi import FastAPI, HTTPException |
| from pydantic import BaseModel |
| from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline |
| from contextlib import asynccontextmanager |
| import os |
|
|
| |
| repo_id = os.getenv("MODEL_ID", "omkar2323/my-llm-model") |
|
|
| |
| pipe = None |
|
|
| |
| @asynccontextmanager |
| async def lifespan(app: FastAPI): |
| |
| global pipe |
| tokenizer = AutoTokenizer.from_pretrained(repo_id) |
| model = AutoModelForCausalLM.from_pretrained(repo_id) |
| pipe = pipeline("text-generation", model=model, tokenizer=tokenizer) |
| yield |
| |
| del pipe |
|
|
| |
| app = FastAPI(lifespan=lifespan) |
|
|
| |
| 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"} |
|
|
|
|