fix: check max_num_batched_tokenz if max_model_len not set

This commit is contained in:
velaraptor-runpod
2026-02-13 03:29:52 -06:00
parent 767c66c301
commit 4705ba8a7c
+31 -5
View File
@@ -186,6 +186,24 @@ def get_speculative_config():
return None return None
def _resolve_max_model_len(model, trust_remote_code=False, revision=None):
"""Resolve max_model_len from the model's HuggingFace config."""
try:
from transformers import AutoConfig
config = AutoConfig.from_pretrained(
model,
trust_remote_code=trust_remote_code,
revision=revision,
)
for attr in ('max_position_embeddings', 'n_positions', 'max_seq_len', 'seq_length'):
val = getattr(config, attr, None)
if val is not None:
logging.info(f"Resolved max_model_len={val} from model config ({attr})")
return val
except Exception as e:
logging.warning(f"Could not resolve max_model_len from model config: {e}")
return None
limit_mm_env = os.getenv('LIMIT_MM_PER_PROMPT') limit_mm_env = os.getenv('LIMIT_MM_PER_PROMPT')
if limit_mm_env is not None: if limit_mm_env is not None:
DEFAULT_ARGS["limit_mm_per_prompt"] = convert_limit_mm_per_prompt(limit_mm_env) DEFAULT_ARGS["limit_mm_per_prompt"] = convert_limit_mm_per_prompt(limit_mm_env)
@@ -268,11 +286,19 @@ def get_engine_args():
# os.environ["VLLM_ATTENTION_BACKEND"] = "FLASHINFER" # os.environ["VLLM_ATTENTION_BACKEND"] = "FLASHINFER"
# logging.info("Using FLASHINFER for gemma-2 model.") # logging.info("Using FLASHINFER for gemma-2 model.")
# When max_num_batched_tokens is None (env var was 0), set to max_model_len # Set max_num_batched_tokens to max_model_len for unlimited batching.
# to preserve "unlimited" behavior. vLLM defaults None to 2048. # vLLM defaults max_num_batched_tokens to 2048 when None, which is too low.
if args.get("max_num_batched_tokens") is None and args.get("max_model_len") is not None: if args.get("max_num_batched_tokens") is None:
args["max_num_batched_tokens"] = args["max_model_len"] max_model_len = args.get("max_model_len")
logging.info(f"Setting max_num_batched_tokens to max_model_len ({args['max_model_len']}) for unlimited batching.") if max_model_len is None:
max_model_len = _resolve_max_model_len(
args.get("model"),
trust_remote_code=args.get("trust_remote_code", False),
revision=args.get("revision"),
)
if max_model_len is not None:
args["max_num_batched_tokens"] = max_model_len
logging.info(f"Setting max_num_batched_tokens to {max_model_len}")
# VLLM_ATTENTION_BACKEND is deprecated, migrate to attention_backend # VLLM_ATTENTION_BACKEND is deprecated, migrate to attention_backend
if os.getenv('VLLM_ATTENTION_BACKEND'): if os.getenv('VLLM_ATTENTION_BACKEND'):