abcd / api /main.py
Karan6933's picture
Upload 5 files
a17c086 verified
Raw
History Blame Contribute Delete
2.79 kB
import os
import asyncio
from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from typing import List, Optional
import logging
from engine import init_engine, get_engine
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# Configuration
MODEL_PATH = os.getenv("MODEL_PATH", "model/model.gguf")
MODEL_URL = os.getenv("MODEL_URL", "https://huggingface.co/prithivMLmods/Nanbeige4.1-3B-f32-GGUF/resolve/main/Nanbeige4.1-3B.Q8_0.gguf")
class GenerateRequest(BaseModel):
prompt: str
max_tokens: int = 256
temperature: float = 0.7
stream: bool = True
class BatchRequest(BaseModel):
prompts: List[str]
max_tokens: int = 256
temperature: float = 0.7
def download_model():
"""Download model if not exists"""
if not os.path.exists(MODEL_PATH):
os.makedirs(os.path.dirname(MODEL_PATH), exist_ok=True)
logger.info(f"Downloading model from {MODEL_URL}")
import urllib.request
urllib.request.urlretrieve(MODEL_URL, MODEL_PATH)
logger.info("Model downloaded")
@asynccontextmanager
async def lifespan(app: FastAPI):
# Startup
logger.info("Starting up...")
download_model()
init_engine(MODEL_PATH, n_ctx=4096, n_threads=4)
logger.info("Ready!")
yield
# Shutdown
logger.info("Shutting down...")
app = FastAPI(title="Nanbeige LLM API", lifespan=lifespan)
@app.post("/generate")
async def generate(req: GenerateRequest):
"""Single prompt generation with streaming"""
engine = get_engine()
if req.stream:
async def stream_generator():
async for token in engine.generate_stream(
req.prompt,
max_tokens=req.max_tokens,
temperature=req.temperature
):
yield token
return StreamingResponse(
stream_generator(),
media_type="text/plain"
)
else:
# Non-streaming: collect all tokens
chunks = []
async for token in engine.generate_stream(
req.prompt,
max_tokens=req.max_tokens,
temperature=req.temperature
):
chunks.append(token)
return {"text": "".join(chunks)}
@app.post("/generate_batch")
async def generate_batch(req: BatchRequest):
"""Batch generation (multiple prompts)"""
engine = get_engine()
results = await engine.generate_batch(
req.prompts,
max_tokens=req.max_tokens,
temperature=req.temperature
)
return {"results": results}
@app.get("/health")
async def health():
return {"status": "ok", "model_loaded": get_engine()._model is not None}