[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:
@@ -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
|
||||
|
||||
@@ -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=(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user