x
This commit is contained in:
@@ -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