diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 37d50cf59..148392922 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -355,15 +355,15 @@ class DecodePreallocQueue: req.retraction_mb_id = None self.retracted_queue.append(req) else: - dp_rank = self._resolve_dp_rank(req) - if dp_rank is None: + prefill_dp_rank = self._resolve_prefill_dp_rank(req) + if prefill_dp_rank is None: self.pending_reqs.append(req) return - self._create_receiver_and_enqueue(req, dp_rank) + self._create_receiver_and_enqueue(req, prefill_dp_rank) - def _resolve_dp_rank(self, req: Req) -> Optional[int]: - if req.data_parallel_rank is not None: - return req.data_parallel_rank + def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]: + if req.disagg_prefill_dp_rank is not None: + return req.disagg_prefill_dp_rank if _is_fake_transfer(req, self.scheduler.server_args): return 0 @@ -379,7 +379,7 @@ class DecodePreallocQueue: return None - def _create_receiver_and_enqueue(self, req: Req, dp_rank: int) -> None: + def _create_receiver_and_enqueue(self, req: Req, prefill_dp_rank: int) -> None: backend = ( TransferBackend.FAKE if _is_fake_transfer(req, self.scheduler.server_args) @@ -391,7 +391,7 @@ class DecodePreallocQueue: mgr=self.kv_manager, bootstrap_addr=f"{req.bootstrap_host}:{req.bootstrap_port}", bootstrap_room=req.bootstrap_room, - prefill_dp_rank=dp_rank, + prefill_dp_rank=prefill_dp_rank, ) self.queue.append( @@ -493,16 +493,16 @@ class DecodePreallocQueue: raise ValueError(f"Unexpected poll case: {poll}") def _resolve_pending_reqs(self) -> None: - """Batch-resolve dp_ranks for pending requests and create receivers.""" + """Batch-resolve prefill_dp_ranks for pending requests and create receivers.""" if not self.pending_reqs: return bootstrap_addr = f"{self.pending_reqs[0].bootstrap_host}:{self.pending_reqs[0].bootstrap_port}" # If a request is following the bootstrap room, - # we need get the prefill info before resolving the dp_rank, + # we need get the prefill info before resolving the prefill_dp_ranks # which is a conflict with the lazy resolve logic in CommonKVReceiver, - # so we need to ensure the parallel info before resolving the dp_rank + # so we need to ensure the parallel info before resolving it. if not self.kv_manager.ensure_parallel_info(bootstrap_addr): return @@ -510,9 +510,9 @@ class DecodePreallocQueue: need_query = [] for req in self.pending_reqs: # NOTE: we need resolve it again because we may ensure the parallel info here - dp_rank = self._resolve_dp_rank(req) - if dp_rank is not None: - resolved.append((req, dp_rank)) + prefill_dp_rank = self._resolve_prefill_dp_rank(req) + if prefill_dp_rank is not None: + resolved.append((req, prefill_dp_rank)) else: need_query.append(req) @@ -534,8 +534,8 @@ class DecodePreallocQueue: else: self.pending_reqs = [] - for req, dp_rank in resolved: - self._create_receiver_and_enqueue(req, dp_rank) + for req, prefill_dp_rank in resolved: + self._create_receiver_and_enqueue(req, prefill_dp_rank) def pop_preallocated( self, rids_to_check: Optional[List[str]] = None diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index a82a52b03..b3d7cfbeb 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -341,7 +341,7 @@ class MMReceiverHTTP(MMReceiverBase): skip_mm_pool=True, ) - def create_req(self, recv_req): + def create_req(self, recv_req: TokenizedGenerateReqInput): req = Req( recv_req.rid, recv_req.input_text, @@ -362,7 +362,8 @@ class MMReceiverHTTP(MMReceiverBase): bootstrap_port=recv_req.bootstrap_port, bootstrap_room=recv_req.bootstrap_room, disagg_mode=self.scheduler.disaggregation_mode, - data_parallel_rank=recv_req.data_parallel_rank, + routed_dp_rank=recv_req.routed_dp_rank, + disagg_prefill_dp_rank=recv_req.disagg_prefill_dp_rank, vocab_size=self.scheduler.model_config.vocab_size, priority=recv_req.priority, metrics_collector=( diff --git a/python/sglang/srt/entrypoints/EngineBase.py b/python/sglang/srt/entrypoints/EngineBase.py index 4f4bd85e7..7bcac278d 100644 --- a/python/sglang/srt/entrypoints/EngineBase.py +++ b/python/sglang/srt/entrypoints/EngineBase.py @@ -28,6 +28,8 @@ class EngineBase(ABC): bootstrap_host: Optional[Union[List[str], str]] = None, bootstrap_port: Optional[Union[List[int], int]] = None, bootstrap_room: Optional[Union[List[int], int]] = None, + routed_dp_rank: Optional[int] = None, + disagg_prefill_dp_rank: Optional[int] = None, data_parallel_rank: Optional[int] = None, rid: Optional[Union[List[str], str]] = None, priority: Optional[int] = None, diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index da8d56508..b03888c3e 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -202,6 +202,35 @@ class Engine(EngineBase): self.loop = asyncio.new_event_loop() asyncio.set_event_loop(self.loop) + def _resolve_routed_dp_rank( + self, + routed_dp_rank: Optional[int], + data_parallel_rank: Optional[int], + ) -> Optional[int]: + if data_parallel_rank is not None: + import warnings + + warnings.warn( + "'data_parallel_rank' is deprecated, use 'routed_dp_rank' instead.", + DeprecationWarning, + stacklevel=3, + ) + if routed_dp_rank is None: + routed_dp_rank = data_parallel_rank + + if self.server_args.enable_dp_attention: + if routed_dp_rank is None: + logger.debug("routed_dp_rank not provided, using default dispatch") + elif routed_dp_rank < 0: + raise ValueError("routed_dp_rank must be non-negative") + elif routed_dp_rank >= self.server_args.dp_size: + raise ValueError( + f"routed_dp_rank must be less than dp_size: {self.server_args.dp_size}" + ) + + logger.debug(f"routed_dp_rank: {routed_dp_rank}") + return routed_dp_rank + def generate( self, # The input prompt. It can be a single prompt or a batch of prompts. @@ -232,6 +261,9 @@ class Engine(EngineBase): bootstrap_host: Optional[Union[List[str], str]] = None, bootstrap_port: Optional[Union[List[int], int]] = None, bootstrap_room: Optional[Union[List[int], int]] = None, + routed_dp_rank: Optional[int] = None, + disagg_prefill_dp_rank: Optional[int] = None, + # Deprecated: use routed_dp_rank instead data_parallel_rank: Optional[int] = None, external_trace_header: Optional[Dict] = None, rid: Optional[Union[List[str], str]] = None, @@ -242,15 +274,9 @@ class Engine(EngineBase): The arguments of this function is the same as `sglang/srt/managers/io_struct.py::GenerateReqInput`. Please refer to `GenerateReqInput` for the documentation. """ - if self.server_args.enable_dp_attention: - if data_parallel_rank is None: - logger.debug("data_parallel_rank not provided, using default dispatch") - elif data_parallel_rank < 0: - raise ValueError("data_parallel_rank must be non-negative") - elif data_parallel_rank >= self.server_args.dp_size: - raise ValueError( - f"data_parallel_rank must be less than dp_size: {self.server_args.dp_size}" - ) + routed_dp_rank = self._resolve_routed_dp_rank( + routed_dp_rank, data_parallel_rank + ) obj = GenerateReqInput( text=prompt, @@ -271,7 +297,8 @@ class Engine(EngineBase): bootstrap_host=bootstrap_host, bootstrap_port=bootstrap_port, bootstrap_room=bootstrap_room, - data_parallel_rank=data_parallel_rank, + routed_dp_rank=routed_dp_rank, + disagg_prefill_dp_rank=disagg_prefill_dp_rank, external_trace_header=external_trace_header, rid=rid, session_params=session_params, @@ -324,6 +351,9 @@ class Engine(EngineBase): bootstrap_host: Optional[Union[List[str], str]] = None, bootstrap_port: Optional[Union[List[int], int]] = None, bootstrap_room: Optional[Union[List[int], int]] = None, + routed_dp_rank: Optional[int] = None, + disagg_prefill_dp_rank: Optional[int] = None, + # Deprecated: use routed_dp_rank instead data_parallel_rank: Optional[int] = None, external_trace_header: Optional[Dict] = None, rid: Optional[Union[List[str], str]] = None, @@ -334,18 +364,10 @@ class Engine(EngineBase): The arguments of this function is the same as `sglang/srt/managers/io_struct.py::GenerateReqInput`. Please refer to `GenerateReqInput` for the documentation. """ + routed_dp_rank = self._resolve_routed_dp_rank( + routed_dp_rank, data_parallel_rank + ) - if self.server_args.enable_dp_attention: - if data_parallel_rank is None: - logger.debug("data_parallel_rank not provided, using default dispatch") - elif data_parallel_rank < 0: - raise ValueError("data_parallel_rank must be non-negative") - elif data_parallel_rank >= self.server_args.dp_size: - raise ValueError( - f"data_parallel_rank must be in range [0, {self.server_args.dp_size-1}]" - ) - - logger.debug(f"data_parallel_rank: {data_parallel_rank}") obj = GenerateReqInput( text=prompt, input_ids=input_ids, @@ -365,7 +387,8 @@ class Engine(EngineBase): bootstrap_host=bootstrap_host, bootstrap_port=bootstrap_port, bootstrap_room=bootstrap_room, - data_parallel_rank=data_parallel_rank, + routed_dp_rank=routed_dp_rank, + disagg_prefill_dp_rank=disagg_prefill_dp_rank, external_trace_header=external_trace_header, rid=rid, session_params=session_params, diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 99bd64a06..1550c707a 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -233,6 +233,20 @@ class BatchResponse(BaseModel): metadata: Optional[dict] = None +def _migrate_deprecated_dp_rank(values: dict) -> dict: + if isinstance(values, dict) and values.get("data_parallel_rank") is not None: + import warnings + + warnings.warn( + "'data_parallel_rank' is deprecated, use 'routed_dp_rank' instead.", + DeprecationWarning, + stacklevel=2, + ) + if values.get("routed_dp_rank") is None: + values["routed_dp_rank"] = values["data_parallel_rank"] + return values + + class CompletionRequest(BaseModel): # Ordered by official OpenAI API documentation # https://platform.openai.com/docs/api-reference/completions/create @@ -285,7 +299,11 @@ class CompletionRequest(BaseModel): bootstrap_port: Optional[Union[List[Optional[int]], int]] = None bootstrap_room: Optional[Union[List[int], int]] = None - # For data parallel rank routing + # For DP routing — external router assigns a specific DP worker + routed_dp_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 + # Deprecated: use routed_dp_rank instead data_parallel_rank: Optional[int] = None # For request id @@ -300,6 +318,11 @@ class CompletionRequest(BaseModel): # For custom metric labels custom_labels: Optional[Dict[str, str]] = None + @model_validator(mode="before") + @classmethod + def _handle_deprecated_dp_rank(cls, values): + return _migrate_deprecated_dp_rank(values) + @field_validator("max_tokens") @classmethod def validate_max_tokens_positive(cls, v): @@ -614,7 +637,11 @@ class ChatCompletionRequest(BaseModel): bootstrap_port: Optional[Union[List[Optional[int]], int]] = None bootstrap_room: Optional[Union[List[int], int]] = None - # For data parallel rank routing + # For DP routing — external router assigns a specific DP worker + routed_dp_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 + # Deprecated: use routed_dp_rank instead data_parallel_rank: Optional[int] = None # OpenAI/SGLang default sampling parameters @@ -626,6 +653,11 @@ class ChatCompletionRequest(BaseModel): "repetition_penalty": 1.0, } + @model_validator(mode="before") + @classmethod + def _handle_deprecated_dp_rank(cls, values): + return _migrate_deprecated_dp_rank(values) + @model_validator(mode="before") @classmethod def set_tool_choice_default(cls, values): diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 7b94b5407..fde200f82 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -296,7 +296,8 @@ class OpenAIServingChat(OpenAIServingBase): bootstrap_host=request.bootstrap_host, bootstrap_port=request.bootstrap_port, bootstrap_room=request.bootstrap_room, - data_parallel_rank=request.data_parallel_rank, + routed_dp_rank=request.routed_dp_rank, + disagg_prefill_dp_rank=request.disagg_prefill_dp_rank, return_hidden_states=request.return_hidden_states, return_routed_experts=request.return_routed_experts, rid=request.rid, diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index f071027f0..fd61942d7 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -111,7 +111,8 @@ class OpenAIServingCompletion(OpenAIServingBase): bootstrap_host=request.bootstrap_host, bootstrap_port=request.bootstrap_port, bootstrap_room=request.bootstrap_room, - data_parallel_rank=request.data_parallel_rank, + routed_dp_rank=request.routed_dp_rank, + disagg_prefill_dp_rank=request.disagg_prefill_dp_rank, return_hidden_states=request.return_hidden_states, return_routed_experts=request.return_routed_experts, rid=request.rid, diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 62477a173..ae5b38ec7 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -494,9 +494,9 @@ class DataParallelController: self.max_req_input_len = scheduler_info[0]["max_req_input_len"] def maybe_external_dp_rank_routing(self, req: Req): - if req.data_parallel_rank is not None: - logger.debug(f"Direct routing to DP rank {req.data_parallel_rank}") - self.workers[req.data_parallel_rank].send_pyobj(req) + if req.routed_dp_rank is not None: + logger.debug(f"Direct routing to DP rank {req.routed_dp_rank}") + self.workers[req.routed_dp_rank].send_pyobj(req) return True return False diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 3ebfac7a0..2c6db4b2f 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -400,6 +400,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): retraction_counts=recv_obj.retraction_counts, token_steps=recv_obj.token_steps, load=recv_obj.load, + dp_ranks=recv_obj.dp_ranks, time_stats=recv_obj.time_stats, ) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 393f1cf06..1252f9a81 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -187,7 +187,11 @@ class GenerateReqInput(BaseReq): # Require reasoning for the request (hybrid reasoning model only) require_reasoning: bool = False - # For data parallel rank routing + # For DP routing — external router assigns a specific DP worker + routed_dp_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 + # Deprecated: use routed_dp_rank instead data_parallel_rank: Optional[int] = None # For background responses (OpenAI responses API) @@ -251,6 +255,18 @@ class GenerateReqInput(BaseReq): 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() self._handle_parallel_sampling() @@ -624,9 +640,8 @@ class GenerateReqInput(BaseReq): decode_tp_size=( self.decode_tp_size[i] if self.decode_tp_size is not None else None ), - data_parallel_rank=( - self.data_parallel_rank if self.data_parallel_rank is not None else None - ), + routed_dp_rank=self.routed_dp_rank, + disagg_prefill_dp_rank=self.disagg_prefill_dp_rank, conversation_id=self.conversation_id, priority=self.priority, extra_key=self.extra_key, @@ -693,8 +708,10 @@ class TokenizedGenerateReqInput(BaseReq): # Require reasoning for the request (hybrid reasoning model only) require_reasoning: bool = False - # For data parallel rank routing - data_parallel_rank: Optional[int] = None + # For DP routing + routed_dp_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 # Priority for the request priority: Optional[int] = None @@ -897,8 +914,8 @@ class TokenizedEmbeddingReqInput(BaseReq): token_type_ids: List[int] # Dummy sampling params for compatibility sampling_params: SamplingParams - # For data parallel rank routing - data_parallel_rank: Optional[int] = None + # For DP routing + routed_dp_rank: Optional[int] = None # Priority for the request priority: Optional[int] = None # The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings. @@ -984,6 +1001,8 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): customized_info: Optional[Dict[str, List[Any]]] = None # Detailed breakdown of cached tokens by source (device/host/storage) cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None + # DP rank of the scheduler that processed each request + dp_ranks: Optional[List[int]] = None # For observability time_stats: Optional[List[SchedulerReqTimeStats]] = None @@ -1076,6 +1095,8 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): customized_info: Optional[Dict[str, List[Any]]] = None # Detailed breakdown of cached tokens by source (device/host/storage) cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None + # DP rank of the scheduler that processed each request + dp_ranks: Optional[List[int]] = None # For observability time_stats: Optional[List[SchedulerReqTimeStats]] = None diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 9e7679dc8..f2586497d 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -273,6 +273,7 @@ def _handle_output_by_index(output, i): customized_info=_extract_field_by_index( output, "customized_info", i, check_length=False ), + dp_ranks=_extract_field_by_index(output, "dp_ranks", i, check_length=False), placeholder_tokens_idx=None, placeholder_tokens_val=None, retraction_counts=_extract_field_by_index(output, "retraction_counts", i), diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 3663a0a2d..5576cd577 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -509,7 +509,8 @@ class Req(ReqDllmMixin): bootstrap_port: Optional[int] = None, bootstrap_room: Optional[int] = None, disagg_mode: Optional[DisaggregationMode] = None, - data_parallel_rank: Optional[int] = None, + routed_dp_rank: Optional[int] = None, + disagg_prefill_dp_rank: Optional[int] = None, vocab_size: Optional[int] = None, priority: Optional[int] = None, metrics_collector: Optional[SchedulerMetricsCollector] = None, @@ -770,8 +771,8 @@ class Req(ReqDllmMixin): self.bootstrap_room: Optional[int] = bootstrap_room self.disagg_kv_sender: Optional[BaseKVSender] = None - # For data parallel rank routing - self.data_parallel_rank: Optional[int] = data_parallel_rank + self.routed_dp_rank: Optional[int] = routed_dp_rank + self.disagg_prefill_dp_rank: Optional[int] = disagg_prefill_dp_rank # the start index of the sent kv cache # We want to send it chunk by chunk for chunked prefill. diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 19dfb0823..bc5a35951 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1505,7 +1505,8 @@ class Scheduler( bootstrap_port=recv_req.bootstrap_port, bootstrap_room=recv_req.bootstrap_room, disagg_mode=self.disaggregation_mode, - data_parallel_rank=recv_req.data_parallel_rank, + routed_dp_rank=recv_req.routed_dp_rank, + disagg_prefill_dp_rank=recv_req.disagg_prefill_dp_rank, vocab_size=self.model_config.vocab_size, priority=recv_req.priority, metrics_collector=( @@ -1811,6 +1812,7 @@ class Scheduler( recv_req.input_ids, recv_req.sampling_params, token_type_ids=recv_req.token_type_ids, + routed_dp_rank=recv_req.routed_dp_rank, priority=recv_req.priority, dimensions=recv_req.dimensions, lora_id=recv_req.lora_id, diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index ee0409895..23f61659b 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1110,6 +1110,8 @@ class SchedulerOutputProcessorMixin: ): req.log_time_stats() + dp_ranks = [self.dp_rank] * len(rids) if rids else None + # Send to detokenizer if reqs or is_idle_batch: if self.model_config.is_multimodal_gen: @@ -1154,6 +1156,7 @@ class SchedulerOutputProcessorMixin: placeholder_tokens_val=None, retraction_counts=retraction_counts, load=load, + dp_ranks=dp_ranks, ) ) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index db13ddce1..aa7a5d95a 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -929,7 +929,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi require_reasoning=obj.require_reasoning, return_hidden_states=obj.return_hidden_states, return_routed_experts=obj.return_routed_experts, - data_parallel_rank=obj.data_parallel_rank, + routed_dp_rank=obj.routed_dp_rank, + disagg_prefill_dp_rank=obj.disagg_prefill_dp_rank, priority=obj.priority, extra_key=obj.extra_key, routing_key=obj.routing_key, @@ -1518,6 +1519,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi if getattr(recv_obj, "customized_info", None): for k, v in recv_obj.customized_info.items(): meta_info[k] = v[i] + if getattr(recv_obj, "dp_ranks", None): + meta_info["dp_rank"] = recv_obj.dp_ranks[i] if isinstance(recv_obj, BatchStrOutput): state.text += recv_obj.output_strs[i] diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py b/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py index 39e809358..297f9e954 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py @@ -7,6 +7,7 @@ import ipaddress import logging import random import urllib +import warnings from http import HTTPStatus from itertools import chain from typing import Optional @@ -69,6 +70,10 @@ class MiniLoadBalancer: ) self.enable_trace = False + self.test_external_dp_routing = router_args.test_external_dp_routing + self.prefill_dp_size = None + self.decode_dp_size = None + def _validate_router_args(self, router_args: RouterArgs): logger.warning( "\x1b[33mMiniLB is only for debugging purposes, it only supports random policy!\033[0m" @@ -95,6 +100,32 @@ class MiniLoadBalancer: trace_set_thread_info("Mini lb") uvicorn.run(app, host=self.host, port=self.port) + async def _ensure_dp_sizes(self): + if self.prefill_dp_size is not None: + return + async with aiohttp.ClientSession() as session: + async with session.get(f"{self.prefill_urls[0]}/server_info") as resp: + info = await resp.json() + self.prefill_dp_size = len(info.get("internal_states", [1])) + async with session.get(f"{self.decode_urls[0]}/server_info") as resp: + info = await resp.json() + self.decode_dp_size = len(info.get("internal_states", [1])) + logger.info( + f"[MiniLB] DP sizes: prefill={self.prefill_dp_size}, decode={self.decode_dp_size}" + ) + + def _fork_dp_requests(self, request): + p_rank = random.randint(0, self.prefill_dp_size - 1) + d_rank = random.randint(0, self.decode_dp_size - 1) + + prefill_req = request.copy() + decode_req = request.copy() + prefill_req["routed_dp_rank"] = p_rank + decode_req["routed_dp_rank"] = d_rank + decode_req["disagg_prefill_dp_rank"] = p_rank + + return prefill_req, decode_req, d_rank + def select_pair(self): assert len(self.prefill_urls) > 0, "No prefill servers available" assert len(self.decode_urls) > 0, "No decode servers available" @@ -111,6 +142,16 @@ class MiniLoadBalancer: ) -> ORJSONResponse: assert endpoint[0] != "/", f"Endpoint should not start with '/': {endpoint}" + expected_decode_dp_rank = None + if self.test_external_dp_routing: + await self._ensure_dp_sizes() + prefill_req, decode_req, expected_decode_dp_rank = self._fork_dp_requests( + modified_request + ) + else: + prefill_req = modified_request + decode_req = modified_request + async with aiohttp.ClientSession( timeout=aiohttp.ClientTimeout( total=self.timeout @@ -130,12 +171,12 @@ class MiniLoadBalancer: tasks = [ session.post( f"{prefill_server}/{endpoint}", - json=modified_request, + json=prefill_req, headers=headers, ), session.post( f"{decode_server}/{endpoint}", - json=modified_request, + json=decode_req, headers=headers, ), ] @@ -169,6 +210,16 @@ class MiniLoadBalancer: ) trace_req_finish(bootstrap_room) + if expected_decode_dp_rank is not None: + actual = ret_json.get("meta_info", {}).get("dp_rank") + if actual != expected_decode_dp_rank: + return ORJSONResponse( + content={ + "error": f"DP rank mismatch: expected {expected_decode_dp_rank}, got {actual}" + }, + status_code=500, + ) + return ORJSONResponse( content=ret_json, status_code=decode_response.status, @@ -177,6 +228,10 @@ class MiniLoadBalancer: async def generate_stream( self, modified_request, prefill_server, decode_server, endpoint="generate" ): + + if self.test_external_dp_routing: + warnings.warn("--test-external-dp-routing is not supported with streaming") + assert endpoint[0] != "/", f"Endpoint should not start with '/': {endpoint}" async def stream_results(): diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py b/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py index 8f03ae92e..fecb64cf4 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/router_args.py @@ -18,6 +18,7 @@ class RouterArgs: # PD-specific configuration mini_lb: bool = False + test_external_dp_routing: bool = False pd_disaggregation: bool = False # Enable PD disaggregated mode prefill_urls: List[tuple] = dataclasses.field( default_factory=list @@ -360,6 +361,11 @@ class RouterArgs: action="store_true", help="Enable MiniLB", ) + pd_group.add_argument( + f"--{prefix}test-external-dp-routing", + action="store_true", + help="(MiniLB only) Randomly assign routed_dp_rank / disagg_prefill_dp_rank per request and verify the response dp_rank matches.", + ) pd_group.add_argument( f"--{prefix}pd-disaggregation", action="store_true", diff --git a/test/registered/distributed/test_disaggregation_dp_attention.py b/test/registered/distributed/test_disaggregation_dp_attention.py index a723111e0..fed913b88 100644 --- a/test/registered/distributed/test_disaggregation_dp_attention.py +++ b/test/registered/distributed/test_disaggregation_dp_attention.py @@ -126,5 +126,43 @@ class TestDisaggregationDPAttentionRoundRobin(TestDisaggregationDPAttention): self.assertEqual(result["completed"], 1000) +@unittest.skip( + "Skip this test until new testing logic in mini-lb has been updated in docker image." +) +class TestDisaggregationDPAttentionExternalRouting(TestDisaggregationDPAttention): + """Test external DP rank assignment via mini-lb --test-external-dp-routing. + + NOTE: In PD disaggregation the response comes from the decode server, + so meta_info["dp_rank"] reflects the decode-side DP rank. Prefill DP + rank correctness is verified implicitly — if the wrong prefill DP + worker were used, KV transfer would fail and the request would error. + The mini-lb internally verifies meta_info["dp_rank"] matches the + assigned decode dp_rank; a mismatch returns HTTP 500. + """ + + @classmethod + def launch_lb(cls): + from sglang.test.test_utils import popen_with_error_check + + lb_command = [ + "python3", + "-m", + "sglang_router.launch_router", + "--pd-disaggregation", + "--mini-lb", + "--test-external-dp-routing", + "--prefill", + cls.prefill_url, + "--decode", + cls.decode_url, + "--host", + cls.base_host, + "--port", + cls.lb_port, + ] + cls.process_lb = popen_with_error_check(lb_command) + cls.wait_server_ready(cls.lb_url + "/health", process=cls.process_lb) + + if __name__ == "__main__": unittest.main()