From deb4cc2c404b89dba9c2cc5057d38dcfede21c14 Mon Sep 17 00:00:00 2001 From: Yu Yao Date: Sat, 15 Aug 2026 04:24:59 -0700 Subject: [PATCH] [misc] fix: Cap server requests by token budget Signed-off-by: Yu Yao --- .../bridge/inference/text_generation.py | 8 +++++-- .../scripts/test_text_generation.py | 24 +++++++++++++++++-- 2 files changed, 28 insertions(+), 4 deletions(-) diff --git a/src/megatron/bridge/inference/text_generation.py b/src/megatron/bridge/inference/text_generation.py index 6d791f8ed8..bbff86bf2f 100644 --- a/src/megatron/bridge/inference/text_generation.py +++ b/src/megatron/bridge/inference/text_generation.py @@ -320,8 +320,8 @@ def build_inference_config( one source of truth. Pure: never mutates caller state. ``max_requests`` resolves to ``max_batch_size`` if set, else ``num_prompts``. When both are - ``None`` (e.g. a server that should auto-size to the KV-cache memory buffer), ``max_requests`` - is left as ``None`` for the engine to size. + ``None``, an explicit token budget caps server capacity; otherwise ``max_requests`` is left as + ``None`` for the engine to size from the KV-cache memory buffer. """ effective_block_size = block_size_tokens if getattr(getattr(model, "config", None), "cache_mla_latents", False) and block_size_tokens != 64: @@ -345,6 +345,10 @@ def build_inference_config( max_requests = rounded max_tokens_limit = max_tokens or DynamicInferenceContext.DEFAULT_MAX_TOKENS + if max_requests is None and max_tokens is not None: + max_requests = max_tokens_limit // tp * tp + if max_requests == 0: + raise ValueError(f"--max_tokens ({max_tokens_limit}) must be at least --tp ({tp}).") if max_requests is not None and max_requests > max_tokens_limit: if max_batch_size is not None: raise ValueError(f"--max_batch_size ({max_batch_size}) cannot exceed --max_tokens ({max_tokens_limit}).") diff --git a/tests/unit_tests/scripts/test_text_generation.py b/tests/unit_tests/scripts/test_text_generation.py index 5cca9605ee..a635281a00 100644 --- a/tests/unit_tests/scripts/test_text_generation.py +++ b/tests/unit_tests/scripts/test_text_generation.py @@ -186,8 +186,28 @@ def test_build_inference_config_rounds_max_requests_up_to_tp(text_generation): assert config.kwargs["materialize_only_last_token_logits"] is True -def test_build_inference_config_auto_sizes_max_requests_when_unset(text_generation): - """Server path: max_batch_size and num_prompts both None -> max_requests left as None.""" +def test_build_inference_config_caps_auto_sized_server_requests_at_token_budget(text_generation): + model = types.SimpleNamespace(position_embedding_type="rope", max_sequence_length=8192) + + config = text_generation.build_inference_config( + model=model, + max_sequence_length=4096, + max_batch_size=None, + num_prompts=None, + tp=2, + block_size_tokens=256, + kv_cache_buffer_size_gb=20.0, + max_tokens=128, + return_log_probs=False, + enable_chunked_prefill=False, + ) + + assert config.kwargs["max_requests"] is not None + assert config.kwargs["max_requests"] <= 128 + assert config.kwargs["max_requests"] % 2 == 0 + + +def test_build_inference_config_preserves_kv_auto_sizing_without_token_limit(text_generation): model = types.SimpleNamespace(position_embedding_type="rope", max_sequence_length=8192) config = text_generation.build_inference_config(