fomext's picture
Upload 5 files
2ba36a2 verified
Raw
History Blame
7.29 kB
"""
OpenAI-compatible FastAPI wrapper for Qwen3-14B (GGUF / llama-cpp-python)
Endpoints: GET /v1/models, POST /v1/chat/completions
Supports streaming (SSE) and non-streaming responses.
"""
import os
import time
import uuid
import json
import asyncio
import logging
from typing import AsyncIterator, List, Optional
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import StreamingResponse, JSONResponse
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
from llama_cpp import Llama
# ---------------------------------------------------------------------------
# Logging
# ---------------------------------------------------------------------------
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Config (override via environment variables)
# ---------------------------------------------------------------------------
MODEL_PATH = os.environ.get("MODEL_PATH", "/models/qwen3-14b-q4_k_m.gguf")
MODEL_ID = os.environ.get("MODEL_ID", "qwen3-14b")
N_CTX = int(os.environ.get("N_CTX", "4096"))
N_THREADS = int(os.environ.get("N_THREADS", str(os.cpu_count() or 4)))
N_BATCH = int(os.environ.get("N_BATCH", "512"))
VERBOSE = os.environ.get("VERBOSE", "false").lower() == "true"
# ---------------------------------------------------------------------------
# Load model at startup
# ---------------------------------------------------------------------------
logger.info(f"Loading model from {MODEL_PATH} — this may take a few minutes on CPU …")
llm = Llama(
model_path=MODEL_PATH,
n_ctx=N_CTX,
n_threads=N_THREADS,
n_batch=N_BATCH,
n_gpu_layers=0, # CPU only
verbose=VERBOSE,
chat_format="chatml", # Qwen3 uses ChatML
)
logger.info("Model loaded ✓")
# ---------------------------------------------------------------------------
# FastAPI app
# ---------------------------------------------------------------------------
app = FastAPI(title="Qwen3-14B OpenAI-compatible API", version="1.0.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
# ---------------------------------------------------------------------------
# Pydantic schemas (OpenAI-compatible subset)
# ---------------------------------------------------------------------------
class Message(BaseModel):
role: str
content: str
class ChatCompletionRequest(BaseModel):
model: str = MODEL_ID
messages: List[Message]
max_tokens: Optional[int] = Field(default=1024, ge=1, le=8192)
temperature: Optional[float] = Field(default=0.7, ge=0.0, le=2.0)
top_p: Optional[float] = Field(default=0.9, ge=0.0, le=1.0)
stream: Optional[bool] = False
stop: Optional[List[str]] = None
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_chunk(delta_content: str, finish_reason: Optional[str], request_id: str) -> str:
chunk = {
"id": request_id,
"object": "chat.completion.chunk",
"created": int(time.time()),
"model": MODEL_ID,
"choices": [
{
"index": 0,
"delta": {"content": delta_content} if delta_content else {},
"finish_reason": finish_reason,
}
],
}
return f"data: {json.dumps(chunk)}\n\n"
async def _stream_response(request: ChatCompletionRequest, request_id: str) -> AsyncIterator[str]:
"""Run llama-cpp in a thread pool and yield SSE chunks."""
messages = [{"role": m.role, "content": m.content} for m in request.messages]
loop = asyncio.get_event_loop()
def _run():
return llm.create_chat_completion(
messages=messages,
max_tokens=request.max_tokens,
temperature=request.temperature,
top_p=request.top_p,
stop=request.stop or [],
stream=True,
)
# llama-cpp streaming returns a generator; run initial call in thread pool
gen = await loop.run_in_executor(None, _run)
# Yield first role delta
yield _make_chunk("", None, request_id)
for chunk in gen:
choice = chunk["choices"][0]
delta = choice.get("delta", {})
content = delta.get("content", "")
finish = choice.get("finish_reason")
if content:
yield _make_chunk(content, None, request_id)
if finish:
yield _make_chunk("", finish, request_id)
break
yield "data: [DONE]\n\n"
# ---------------------------------------------------------------------------
# Routes
# ---------------------------------------------------------------------------
@app.get("/")
async def root():
return {"status": "ok", "model": MODEL_ID}
@app.get("/v1/models")
async def list_models():
return {
"object": "list",
"data": [
{
"id": MODEL_ID,
"object": "model",
"created": 1700000000,
"owned_by": "local",
}
],
}
@app.post("/v1/chat/completions")
async def chat_completions(request: ChatCompletionRequest):
messages = [{"role": m.role, "content": m.content} for m in request.messages]
if request.stream:
request_id = f"chatcmpl-{uuid.uuid4().hex}"
return StreamingResponse(
_stream_response(request, request_id),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)
# Non-streaming
loop = asyncio.get_event_loop()
def _run():
return llm.create_chat_completion(
messages=messages,
max_tokens=request.max_tokens,
temperature=request.temperature,
top_p=request.top_p,
stop=request.stop or [],
stream=False,
)
result = await loop.run_in_executor(None, _run)
choice = result["choices"][0]
usage = result.get("usage", {})
return {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": int(time.time()),
"model": MODEL_ID,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": choice["message"]["content"],
},
"finish_reason": choice.get("finish_reason", "stop"),
}
],
"usage": {
"prompt_tokens": usage.get("prompt_tokens", 0),
"completion_tokens": usage.get("completion_tokens", 0),
"total_tokens": usage.get("total_tokens", 0),
},
}
# ---------------------------------------------------------------------------
# Health check (useful for Docker HEALTHCHECK)
# ---------------------------------------------------------------------------
@app.get("/health")
async def health():
return {"status": "healthy", "model": MODEL_ID}