x
This commit is contained in:
@@ -15,8 +15,7 @@
|
|||||||
├── app
|
├── app
|
||||||
│ ├── config.py
|
│ ├── config.py
|
||||||
│ ├── model_catalog.py
|
│ ├── model_catalog.py
|
||||||
│ ├── start_openai.py
|
│ └── start_openai.py
|
||||||
│ └── schemas.py
|
|
||||||
├── .dockerignore
|
├── .dockerignore
|
||||||
├── config.json
|
├── config.json
|
||||||
├── docker-compose.yml
|
├── docker-compose.yml
|
||||||
|
|||||||
@@ -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,
|
|
||||||
)
|
|
||||||
-50
@@ -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)
|
|
||||||
@@ -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
|
|
||||||
@@ -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()
|
|
||||||
Reference in New Issue
Block a user