| 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__) |
|
|
| |
| 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): |
| |
| logger.info("Starting up...") |
| download_model() |
| init_engine(MODEL_PATH, n_ctx=4096, n_threads=4) |
| logger.info("Ready!") |
| yield |
| |
| 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: |
| |
| 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} |