diff --git a/README.md b/README.md index f450e1d..c1b3d75 100644 --- a/README.md +++ b/README.md @@ -15,8 +15,7 @@ ├── app │ ├── config.py │ ├── model_catalog.py -│ ├── start_openai.py -│ └── schemas.py +│ └── start_openai.py ├── .dockerignore ├── config.json ├── docker-compose.yml diff --git a/app/engine.py b/app/engine.py deleted file mode 100644 index 7c3e82f..0000000 --- a/app/engine.py +++ /dev/null @@ -1,48 +0,0 @@ -import httpx - -from app.config import Settings -from app.schemas import GenerateRequest, GenerateResponse - - -class InferenceEngine: - def __init__(self, settings: Settings) -> None: - self.settings = settings - self.client = httpx.Client(timeout=300.0) - - def close(self) -> None: - self.client.close() - - def generate(self, req: GenerateRequest) -> GenerateResponse: - headers = {"Content-Type": "application/json"} - if self.settings.api_key: - headers["Authorization"] = f"Bearer {self.settings.api_key}" - payload = { - "model": self.settings.served_model_name or self.settings.model_name, - "messages": [{"role": "user", "content": req.prompt}], - "max_tokens": req.max_tokens, - "temperature": req.temperature, - "top_p": req.top_p, - } - if req.stop: - payload["stop"] = req.stop - if req.enable_thinking is not None: - payload["chat_template_kwargs"] = {"enable_thinking": req.enable_thinking} - response = self.client.post( - f"{self.settings.vllm_openai_internal_url}/chat/completions", - headers=headers, - json=payload, - ) - response.raise_for_status() - body = response.json() - completion = body["choices"][0]["message"]["content"] - usage = body.get("usage", {}) - usage_prompt = int(usage.get("prompt_tokens", 0)) - usage_completion = int(usage.get("completion_tokens", 0)) - return GenerateResponse( - text=completion, - prompt=req.prompt, - model=self.settings.served_model_name or self.settings.model_name, - usage_prompt_tokens=usage_prompt, - usage_completion_tokens=usage_completion, - usage_total_tokens=usage_prompt + usage_completion, - ) diff --git a/app/main.py b/app/main.py deleted file mode 100644 index 160b1e9..0000000 --- a/app/main.py +++ /dev/null @@ -1,50 +0,0 @@ -from contextlib import asynccontextmanager - -from fastapi import Depends, FastAPI, Header, HTTPException, status - -from app.config import Settings, get_settings -from app.engine import InferenceEngine -from app.schemas import GenerateRequest, GenerateResponse, HealthResponse - -engine: InferenceEngine | None = None - - -def verify_api_key( - settings: Settings = Depends(get_settings), x_api_key: str | None = Header(default=None) -) -> None: - if settings.api_key and x_api_key != settings.api_key: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail="Invalid API key", - ) - - -@asynccontextmanager -async def lifespan(_: FastAPI): - global engine - settings = get_settings() - engine = InferenceEngine(settings) - yield - if engine is not None: - engine.close() - engine = None - - -app = FastAPI(title="ROCm vLLM Inference API", version="1.0.0", lifespan=lifespan) - - -@app.get("/health", response_model=HealthResponse) -def health(settings: Settings = Depends(get_settings)) -> HealthResponse: - return HealthResponse(status="ok", model=settings.served_model_name or settings.model_name) - - -@app.post("/v1/generate", response_model=GenerateResponse, dependencies=[Depends(verify_api_key)]) -def generate(req: GenerateRequest, settings: Settings = Depends(get_settings)) -> GenerateResponse: - if engine is None: - raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Engine not ready") - if req.max_tokens > settings.max_tokens: - raise HTTPException( - status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, - detail=f"max_tokens must be <= {settings.max_tokens}", - ) - return engine.generate(req) diff --git a/app/schemas.py b/app/schemas.py deleted file mode 100644 index 026e23a..0000000 --- a/app/schemas.py +++ /dev/null @@ -1,27 +0,0 @@ -from typing import List, Optional - -from pydantic import BaseModel, Field - - -class GenerateRequest(BaseModel): - prompt: str - max_tokens: int = Field(default=256, ge=1, le=4096) - temperature: float = Field(default=0.7, ge=0.0, le=2.0) - top_p: float = Field(default=0.95, gt=0.0, le=1.0) - repetition_penalty: float = Field(default=1.0, ge=0.5, le=2.0) - stop: Optional[List[str]] = None - enable_thinking: Optional[bool] = None - - -class GenerateResponse(BaseModel): - text: str - prompt: str - model: str - usage_prompt_tokens: int - usage_completion_tokens: int - usage_total_tokens: int - - -class HealthResponse(BaseModel): - status: str - model: str diff --git a/app/start_api.py b/app/start_api.py deleted file mode 100644 index 400e4c0..0000000 --- a/app/start_api.py +++ /dev/null @@ -1,23 +0,0 @@ -import subprocess -import sys - -from app.config import get_settings - - -def main() -> None: - settings = get_settings() - command = [ - sys.executable, - "-m", - "uvicorn", - "app.main:app", - "--host", - settings.host, - "--port", - str(settings.port), - ] - raise SystemExit(subprocess.call(command)) - - -if __name__ == "__main__": - main()