[Fix] Restore data_parallel_rank alias on native /generate (dp-aware gateway routing is silently dropped) (#33565)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Sam Shleifer
2026-08-08 14:59:16 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent c9444deef4
commit d238e36b24
2 changed files with 41 additions and 0 deletions
+17
View File
@@ -263,6 +263,11 @@ class GenerateReqInput:
# For DP routing — external router assigns a specific DP worker
routed_dp_rank: Optional[int] = None
# Deprecated alias for `routed_dp_rank`, still accepted because
# sgl-model-gateway's dp-aware mode injects this spelling into every
# request it forwards (DPAwareWorker::prepare_request), and the OpenAI
# entrypoints and Engine.generate() accept it as well.
data_parallel_rank: Optional[int] = None
# For PD disagg — hint telling decode which prefill DP worker has the KV cache
disagg_prefill_dp_rank: Optional[int] = None
# Routing key for routing-key schedule policy
@@ -365,6 +370,18 @@ class GenerateReqInput:
ValueError: If inputs are not properly specified (e.g., none or all of
text, input_ids, input_embeds are provided)
"""
if self.data_parallel_rank is not None:
import warnings
warnings.warn(
"'data_parallel_rank' is deprecated, use 'routed_dp_rank' instead.",
DeprecationWarning,
stacklevel=2,
)
if self.routed_dp_rank is None:
self.routed_dp_rank = self.data_parallel_rank
self.data_parallel_rank = None
self._validate_inputs()
self._determine_batch_size()
if self.session_id is not None and self.session_params is not None:
@@ -703,6 +703,30 @@ class TestGenerateReqInputNormalization(CustomTestCase):
)
req.normalize_batch_and_arguments()
def test_data_parallel_rank_alias_maps_to_routed_dp_rank(self):
req = GenerateReqInput(text="Hello", sampling_params={}, data_parallel_rank=2)
req.normalize_batch_and_arguments()
self.assertEqual(req.routed_dp_rank, 2)
self.assertIsNone(req.data_parallel_rank)
def test_data_parallel_rank_alias_does_not_override_routed_dp_rank(self):
req = GenerateReqInput(
text="Hello", sampling_params={}, data_parallel_rank=2, routed_dp_rank=1
)
req.normalize_batch_and_arguments()
self.assertEqual(req.routed_dp_rank, 1)
def test_data_parallel_rank_alias_propagates_to_batch_items(self):
req = GenerateReqInput(
text=["Hello", "World"],
sampling_params=[{}, {}],
rid=["id1", "id2"],
data_parallel_rank=3,
)
req.normalize_batch_and_arguments()
self.assertEqual(req[0].routed_dp_rank, 3)
self.assertEqual(req[1].routed_dp_rank, 3)
class TestEmbeddingReqInputGetItem(CustomTestCase):
"""Test EmbeddingReqInput.__getitem__."""