x
This commit is contained in:
+38
-26
@@ -2,44 +2,56 @@ from functools import lru_cache
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.model_catalog import resolve_model_profile
|
||||
from app.model_catalog import load_catalog, resolve_model_profile, resolve_runtime_settings
|
||||
|
||||
|
||||
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")
|
||||
class Settings(BaseModel):
|
||||
config_file: str = "config.json"
|
||||
model_key: Optional[str] = None
|
||||
selected_model: Optional[str] = None
|
||||
model_name: str = Field(default="", alias="MODEL_NAME")
|
||||
model_name: str = ""
|
||||
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")
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8000
|
||||
openai_host: str = "0.0.0.0"
|
||||
openai_port: int = 8001
|
||||
max_model_len: int = 8192
|
||||
gpu_memory_utilization: float = 0.92
|
||||
tensor_parallel_size: int = 2
|
||||
max_num_seqs: int = 64
|
||||
max_tokens: int = 4096
|
||||
dtype: str = "bfloat16"
|
||||
enforce_eager: bool = False
|
||||
trust_remote_code: bool = False
|
||||
tool_call_parser: Optional[str] = None
|
||||
enable_auto_tool_choice: bool = False
|
||||
revision: Optional[str] = None
|
||||
api_key: Optional[str] = None
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_settings() -> Settings:
|
||||
settings = Settings()
|
||||
catalog = load_catalog("config.json")
|
||||
runtime = resolve_runtime_settings(catalog)
|
||||
settings = Settings(
|
||||
config_file="config.json",
|
||||
model_key=runtime["model_key"],
|
||||
host=runtime["host"],
|
||||
port=runtime["port"],
|
||||
openai_host=runtime["openai_host"],
|
||||
openai_port=runtime["openai_port"],
|
||||
api_key=runtime["api_key"],
|
||||
tensor_parallel_size=runtime["tensor_parallel_size"],
|
||||
dtype=runtime["dtype"],
|
||||
revision=runtime["revision"],
|
||||
)
|
||||
_, updates, env_vars = resolve_model_profile(
|
||||
catalog_path=settings.config_file,
|
||||
content=catalog,
|
||||
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)
|
||||
return settings.model_copy(update=updates | runtime)
|
||||
|
||||
+28
-5
@@ -29,14 +29,37 @@ def _to_float(value: Any, default: float) -> float:
|
||||
return default
|
||||
|
||||
|
||||
def resolve_model_profile(
|
||||
catalog_path: str, requested_model: str | None, requested_tp: int
|
||||
) -> tuple[str, dict[str, Any], dict[str, str]]:
|
||||
def load_catalog(catalog_path: str = "config.json") -> dict[str, Any]:
|
||||
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"}
|
||||
return content
|
||||
|
||||
|
||||
def resolve_runtime_settings(content: dict[str, Any]) -> dict[str, Any]:
|
||||
services = content.get("services", {})
|
||||
api_service = dict(services.get("api", {}))
|
||||
openai_service = dict(services.get("openai", {}))
|
||||
models = dict(content.get("models", {}))
|
||||
return {
|
||||
"host": str(api_service.get("host", "0.0.0.0")),
|
||||
"port": _to_int(api_service.get("port"), 8000),
|
||||
"openai_host": str(openai_service.get("host", "0.0.0.0")),
|
||||
"openai_port": _to_int(openai_service.get("port"), 8001),
|
||||
"api_key": str(content.get("api_key", "")).strip() or None,
|
||||
"tensor_parallel_size": _to_int(content.get("tensor_parallel_size"), 2),
|
||||
"dtype": str(content.get("dtype", "bfloat16")),
|
||||
"revision": str(content.get("revision", "")).strip() or None,
|
||||
"model_key": str(models.get("selected", "")).strip() or None,
|
||||
}
|
||||
|
||||
|
||||
def resolve_model_profile(
|
||||
content: dict[str, Any], requested_model: str | None, requested_tp: int
|
||||
) -> tuple[str, dict[str, Any], dict[str, str]]:
|
||||
models = dict(content.get("models", {}))
|
||||
profiles = dict(models.get("profiles", {}))
|
||||
default_model = models.get("default")
|
||||
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")
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
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()
|
||||
+12
-12
@@ -2,25 +2,25 @@ import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from app.model_catalog import resolve_model_profile
|
||||
from app.model_catalog import load_catalog, resolve_model_profile, resolve_runtime_settings
|
||||
|
||||
|
||||
def build_command() -> list[str]:
|
||||
config_file = os.getenv("MODEL_CONFIG_FILE", "config.json")
|
||||
model_key = os.getenv("MODEL_KEY")
|
||||
requested_tp = int(os.getenv("TENSOR_PARALLEL_SIZE", "2"))
|
||||
config_file = "config.json"
|
||||
catalog = load_catalog(config_file)
|
||||
runtime = resolve_runtime_settings(catalog)
|
||||
_, updates, env_vars = resolve_model_profile(
|
||||
catalog_path=config_file,
|
||||
requested_model=model_key,
|
||||
requested_tp=requested_tp,
|
||||
content=catalog,
|
||||
requested_model=runtime["model_key"],
|
||||
requested_tp=runtime["tensor_parallel_size"],
|
||||
)
|
||||
for key, value in env_vars.items():
|
||||
os.environ[key] = value
|
||||
host = os.getenv("OPENAI_HOST", "0.0.0.0")
|
||||
port = os.getenv("OPENAI_PORT", "8001")
|
||||
dtype = os.getenv("DTYPE", "bfloat16")
|
||||
revision = os.getenv("REVISION", "").strip()
|
||||
api_key = os.getenv("API_KEY", "").strip()
|
||||
host = str(runtime["openai_host"])
|
||||
port = str(runtime["openai_port"])
|
||||
dtype = str(runtime["dtype"])
|
||||
revision = runtime["revision"] or ""
|
||||
api_key = runtime["api_key"] or ""
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
|
||||
Reference in New Issue
Block a user