[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:
co-authored by
Claude Fable 5
parent
c9444deef4
commit
d238e36b24
@@ -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__."""
|
||||
|
||||
Reference in New Issue
Block a user