x
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user