This commit is contained in:
2026-03-29 05:06:19 +08:00
parent 67eb27f2d2
commit 617665c6ab
12 changed files with 452 additions and 1 deletions
+45
View File
@@ -0,0 +1,45 @@
from functools import lru_cache
import os
from typing import Optional
from pydantic import Field
from pydantic_settings import BaseSettings, SettingsConfigDict
from app.model_catalog import resolve_model_profile
class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8")
config_file: str = Field(default="config.json", alias="MODEL_CONFIG_FILE")
model_key: Optional[str] = Field(default=None, alias="MODEL_KEY")
selected_model: Optional[str] = None
model_name: str = Field(default="", alias="MODEL_NAME")
served_model_name: Optional[str] = None
host: str = Field(default="0.0.0.0", alias="HOST")
port: int = Field(default=8000, alias="PORT")
max_model_len: int = Field(default=8192, alias="MAX_MODEL_LEN")
gpu_memory_utilization: float = Field(default=0.92, alias="GPU_MEMORY_UTILIZATION")
tensor_parallel_size: int = Field(default=2, alias="TENSOR_PARALLEL_SIZE")
max_num_seqs: int = Field(default=64, alias="MAX_NUM_SEQS")
max_tokens: int = Field(default=4096, alias="MAX_TOKENS")
dtype: str = Field(default="bfloat16", alias="DTYPE")
enforce_eager: bool = Field(default=False, alias="ENFORCE_EAGER")
trust_remote_code: bool = Field(default=False, alias="TRUST_REMOTE_CODE")
tool_call_parser: Optional[str] = Field(default=None, alias="TOOL_CALL_PARSER")
enable_auto_tool_choice: bool = Field(default=False, alias="ENABLE_AUTO_TOOL_CHOICE")
revision: Optional[str] = Field(default=None, alias="REVISION")
api_key: Optional[str] = Field(default=None, alias="API_KEY")
@lru_cache(maxsize=1)
def get_settings() -> Settings:
settings = Settings()
_, updates, env_vars = resolve_model_profile(
catalog_path=settings.config_file,
requested_model=settings.model_key,
requested_tp=settings.tensor_parallel_size,
)
for key, value in env_vars.items():
os.environ[key] = value
return settings.model_copy(update=updates)
+46
View File
@@ -0,0 +1,46 @@
from threading import Lock
from vllm import LLM, SamplingParams
from app.config import Settings
from app.schemas import GenerateRequest, GenerateResponse
class InferenceEngine:
def __init__(self, settings: Settings) -> None:
self.settings = settings
self.lock = Lock()
self.model = LLM(
model=settings.model_name,
tensor_parallel_size=settings.tensor_parallel_size,
max_model_len=settings.max_model_len,
gpu_memory_utilization=settings.gpu_memory_utilization,
max_num_seqs=settings.max_num_seqs,
dtype=settings.dtype,
enforce_eager=settings.enforce_eager,
trust_remote_code=settings.trust_remote_code,
revision=settings.revision,
)
def generate(self, req: GenerateRequest) -> GenerateResponse:
sampling_params = SamplingParams(
temperature=req.temperature,
top_p=req.top_p,
max_tokens=req.max_tokens,
repetition_penalty=req.repetition_penalty,
stop=req.stop,
)
with self.lock:
outputs = self.model.generate([req.prompt], sampling_params, use_tqdm=False)
output = outputs[0]
completion = output.outputs[0].text
usage_prompt = len(output.prompt_token_ids)
usage_completion = len(output.outputs[0].token_ids)
return GenerateResponse(
text=completion,
prompt=req.prompt,
model=self.settings.model_name,
usage_prompt_tokens=usage_prompt,
usage_completion_tokens=usage_completion,
usage_total_tokens=usage_prompt + usage_completion,
)
+43
View File
@@ -0,0 +1,43 @@
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
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.model_name)
@app.post("/v1/generate", response_model=GenerateResponse, dependencies=[Depends(verify_api_key)])
def generate(req: GenerateRequest) -> GenerateResponse:
if engine is None:
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Engine not ready")
return engine.generate(req)
+66
View File
@@ -0,0 +1,66 @@
import json
from pathlib import Path
from typing import Any
def _to_bool(value: Any, default: bool) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in {"true", "1", "yes", "y"}:
return True
if normalized in {"false", "0", "no", "n"}:
return False
return default
def _to_int(value: Any, default: int) -> int:
try:
return int(value)
except (TypeError, ValueError):
return default
def _to_float(value: Any, default: float) -> float:
try:
return float(value)
except (TypeError, ValueError):
return default
def resolve_model_profile(
catalog_path: str, requested_model: str | None, requested_tp: int
) -> tuple[str, dict[str, Any], dict[str, str]]:
content = json.loads(Path(catalog_path).read_text(encoding="utf-8"))
if not isinstance(content, dict):
raise ValueError("config.json must be a JSON object")
default_model = content.get("default_model")
profiles = {k: v for k, v in content.items() if k != "default_model"}
model_key = requested_model or default_model
if not model_key or model_key not in profiles:
raise ValueError(f"model profile '{model_key}' not found in config.json")
profile = profiles[model_key]
if not isinstance(profile, dict):
raise ValueError(f"model profile '{model_key}' must be a JSON object")
valid_tp_raw = profile.get("valid_tp", [])
valid_tp = [_to_int(item, 0) for item in valid_tp_raw if _to_int(item, 0) > 0]
resolved_tp = requested_tp
if valid_tp and resolved_tp not in valid_tp:
resolved_tp = valid_tp[0]
updates = {
"selected_model": model_key,
"model_name": profile.get("hf_model_id", model_key),
"served_model_name": profile.get("served_model_name", model_key),
"max_model_len": _to_int(profile.get("ctx"), 8192),
"max_num_seqs": _to_int(profile.get("max_num_seqs"), 64),
"max_tokens": _to_int(profile.get("max_tokens"), 4096),
"gpu_memory_utilization": _to_float(profile.get("gpu_util"), 0.92),
"trust_remote_code": _to_bool(profile.get("trust_remote"), False),
"enforce_eager": _to_bool(profile.get("enforce_eager"), False),
"tensor_parallel_size": resolved_tp,
"tool_call_parser": profile.get("tool_call_parser"),
"enable_auto_tool_choice": _to_bool(profile.get("enable_auto_tool_choice"), False),
}
env_vars = {str(k): str(v) for k, v in dict(profile.get("env", {})).items()}
return model_key, updates, env_vars
+26
View File
@@ -0,0 +1,26 @@
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
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