Compare commits

..
4 Commits
Author SHA1 Message Date
chrisvelaandGitHub d69cc021e8 Merge pull request #268 from runpod-workers/fix/spec-config-0-to-none
Release / release (push) Waiting to run
fix: spec config env vars should be none if zero
2026-02-18 15:51:51 -06:00
velaraptor-runpod 61faa8f137 fix: spec config env vars should be none if zero 2026-02-18 15:41:19 -06:00
chrisvelaandGitHub 1606cff557 Merge pull request #265 from runpod-workers/fix/zero-max-model-num_batches
Release / release (push) Waiting to run
fix: check for zero param and set to None
2026-02-13 15:26:06 -06:00
velaraptor-runpod e705c9494b fix: check for zero param and set to None 2026-02-13 15:23:54 -06:00
+18 -6
View File
@@ -122,9 +122,14 @@ def get_speculative_config():
# Option 2: Build config from individual environment variables
spec_method = os.getenv('SPECULATIVE_METHOD')
spec_model = os.getenv('SPECULATIVE_MODEL')
num_spec_tokens = os.getenv('NUM_SPECULATIVE_TOKENS')
ngram_max = os.getenv('NGRAM_PROMPT_LOOKUP_MAX')
ngram_min = os.getenv('NGRAM_PROMPT_LOOKUP_MIN')
_num_spec_tokens = os.getenv('NUM_SPECULATIVE_TOKENS')
_ngram_max = os.getenv('NGRAM_PROMPT_LOOKUP_MAX')
_ngram_min = os.getenv('NGRAM_PROMPT_LOOKUP_MIN')
# Convert numeric vars to int so '0' (hub.json default) is treated as unset
num_spec_tokens = (int(_num_spec_tokens) or None) if _num_spec_tokens else None
ngram_max = (int(_ngram_max) or None) if _ngram_max else None
ngram_min = (int(_ngram_min) or None) if _ngram_min else None
if not any([spec_method, spec_model, ngram_max]):
return None
@@ -150,11 +155,11 @@ def get_speculative_config():
if spec_model:
config['model'] = spec_model
if num_spec_tokens:
config['num_speculative_tokens'] = int(num_spec_tokens)
config['num_speculative_tokens'] = num_spec_tokens
if ngram_max:
config['prompt_lookup_max'] = int(ngram_max)
config['prompt_lookup_max'] = ngram_max
if ngram_min:
config['prompt_lookup_min'] = int(ngram_min)
config['prompt_lookup_min'] = ngram_min
draft_tp = os.getenv('SPECULATIVE_DRAFT_TENSOR_PARALLEL_SIZE')
if draft_tp:
@@ -288,6 +293,13 @@ def get_engine_args():
# Set max_num_batched_tokens to max_model_len for unlimited batching.
# vLLM defaults max_num_batched_tokens to 2048 when None, which is too low.
if args.get("max_model_len") == 0:
args["max_model_len"] = None
if args.get("max_num_batched_tokens") == 0:
args["max_num_batched_tokens"] = None
if args.get("max_num_batched_tokens") is None:
max_model_len = args.get("max_model_len")
if max_model_len is None: