feat: allow engine args through VLLM_ and checks the engine args

This commit is contained in:
velaraptor-runpod
2026-02-24 16:05:18 -06:00
parent b7c6d4f9a2
commit 42443f735e
+94
View File
@@ -1,6 +1,7 @@
import os
import json
import logging
from typing import get_origin, get_args
from torch.cuda import device_count
from vllm import AsyncEngineArgs
from vllm.model_executor.model_loader.tensorizer import TensorizerConfig
@@ -13,6 +14,96 @@ RENAME_ARGS_MAP = {
"MAX_CONTEXT_LEN_TO_CAPTURE": "max_seq_len_to_capture"
}
VLLM_ENV_PREFIX = "VLLM_"
def _resolve_field_type(field_type: type) -> type:
"""Resolve Optional/Union to the concrete type for conversion."""
origin = get_origin(field_type)
args = get_args(field_type) if hasattr(field_type, "__args__") else ()
if origin is not None:
# Optional[X] is Union[X, None]; X | None is UnionType
non_none = [a for a in args if a is not type(None)]
if non_none:
return non_none[0]
return field_type
def _convert_env_value_to_field_type(value: str, field_name: str, field_type: type):
"""Convert env var string to the type expected by AsyncEngineArgs for this field."""
val = value.strip() if isinstance(value, str) else value
if val in ("", "None", "none"):
args = get_args(field_type) if hasattr(field_type, "__args__") else ()
if type(None) in (args or ()):
return None
raise ValueError("empty value not allowed for non-optional field")
effective_type = _resolve_field_type(field_type)
# bool
if effective_type is bool:
return str(val).lower() in ("true", "1", "yes", "on")
# int
if effective_type is int:
return int(val)
# float
if effective_type is float:
return float(val)
# str
if effective_type is str:
return str(val)
# dict, list, or complex (try JSON)
origin = get_origin(effective_type)
if effective_type in (dict, list) or origin in (dict, list):
try:
return json.loads(val)
except json.JSONDecodeError:
return val
# Fallback: try int, float, then str
try:
return int(val)
except ValueError:
pass
try:
return float(val)
except ValueError:
pass
return str(val)
def _get_vllm_env_overrides() -> dict:
"""Collect engine arg overrides from env vars with prefix VLLM_.
Any env var VLLM_<ARG> maps to the engine arg <arg> (lowercase).
E.g. VLLM_MAX_MODEL_LEN=4096 -> max_model_len=4096.
Only keys that exist on AsyncEngineArgs are applied; values are
converted to the field type (int, float, bool, str, json for dict/list).
"""
overrides = {}
valid_fields = AsyncEngineArgs.__dataclass_fields__
for key, value in os.environ.items():
if not key.startswith(VLLM_ENV_PREFIX) or len(key) <= len(VLLM_ENV_PREFIX):
continue
suffix = key[len(VLLM_ENV_PREFIX) :]
arg_name = suffix.lower()
if arg_name not in valid_fields:
continue
field = valid_fields[arg_name]
try:
overrides[arg_name] = _convert_env_value_to_field_type(
value, arg_name, field.type
)
except (ValueError, TypeError, json.JSONDecodeError) as e:
logging.warning(
"Skip VLLM_ env override %s=%r: %s", key, value, e
)
continue
if overrides:
logging.info(
"Applying engine arg overrides from VLLM_ env vars: %s",
list(overrides.keys()),
)
return overrides
DEFAULT_ARGS = {
"disable_log_stats": os.getenv('DISABLE_LOG_STATS', 'False').lower() == 'true',
# disable_log_requests is deprecated, use enable_log_requests instead
@@ -268,6 +359,9 @@ def get_engine_args():
# Rename and match to vllm args
args = match_vllm_args(args)
# Apply any VLLM_* env vars as overrides (e.g. VLLM_MAX_MODEL_LEN=4096 -> max_model_len=4096)
args.update(_get_vllm_env_overrides())
if args.get("load_format") == "bitsandbytes":
args["quantization"] = args["load_format"]