[PD-Disagg] Fully support external DP dispatch w/ PD-disaggregation mode. (#19268)

Co-authored-by: Ratish P <114130421+ratish1@users.noreply.github.com>
This commit is contained in:
Liangsheng Yin
2026-02-24 19:58:01 -08:00
committed by GitHub
parent 241ee90164
commit 539f772f54
18 changed files with 253 additions and 62 deletions

View File

@@ -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

View File

@@ -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=(

View File

@@ -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,

View File

@@ -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,

View File

@@ -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):

View File

@@ -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,

View File

@@ -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,

View File

@@ -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

View File

@@ -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,
)

View File

@@ -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

View File

@@ -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),

View File

@@ -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.

View File

@@ -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,

View File

@@ -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,
)
)

View File

@@ -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]