[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
|
# For DP routing — external router assigns a specific DP worker
|
||||||
routed_dp_rank: Optional[int] = None
|
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
|
# For PD disagg — hint telling decode which prefill DP worker has the KV cache
|
||||||
disagg_prefill_dp_rank: Optional[int] = None
|
disagg_prefill_dp_rank: Optional[int] = None
|
||||||
# Routing key for routing-key schedule policy
|
# 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
|
ValueError: If inputs are not properly specified (e.g., none or all of
|
||||||
text, input_ids, input_embeds are provided)
|
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._validate_inputs()
|
||||||
self._determine_batch_size()
|
self._determine_batch_size()
|
||||||
if self.session_id is not None and self.session_params is not None:
|
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()
|
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):
|
class TestEmbeddingReqInputGetItem(CustomTestCase):
|
||||||
"""Test EmbeddingReqInput.__getitem__."""
|
"""Test EmbeddingReqInput.__getitem__."""
|
||||||
|
|||||||
Reference in New Issue
Block a user