[PD] Fix cross-rank queue divergence by gating metadata readiness before all-reduce (#26394)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -37,12 +37,12 @@ from sglang.srt.disaggregation.base import KVPoll
|
||||
from sglang.srt.disaggregation.base.conn import StateType
|
||||
from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver
|
||||
from sglang.srt.disaggregation.utils import (
|
||||
FAKE_BOOTSTRAP_HOST,
|
||||
DisaggregationMode,
|
||||
KVClassType,
|
||||
MetadataBuffers,
|
||||
ReqToMetadataIdxAllocator,
|
||||
TransferBackend,
|
||||
_is_fake_transfer,
|
||||
get_kv_class,
|
||||
is_mla_backend,
|
||||
poll_and_all_reduce,
|
||||
@@ -82,18 +82,10 @@ logger = logging.getLogger(__name__)
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
CLIP_MAX_NEW_TOKEN = envs.SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION.get()
|
||||
|
||||
|
||||
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool:
|
||||
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||
req.bootstrap_host is None
|
||||
and server_args.disaggregation_transfer_backend == "fake"
|
||||
)
|
||||
|
||||
|
||||
def _bootstrap_addr(req: Req) -> str:
|
||||
# FIXME: make a property of a req
|
||||
return NetworkAddress(req.bootstrap_host, req.bootstrap_port).to_host_port_str()
|
||||
@@ -1378,12 +1370,7 @@ class DecodeTransferQueue:
|
||||
):
|
||||
self.staging_handler.register_decode_req(dr.req.bootstrap_room, dr)
|
||||
|
||||
def _commit_transfer_to_req(self, decode_req: DecodeRequest) -> bool:
|
||||
"""
|
||||
Returns:
|
||||
True if the request should be removed from the queue (success or corruption)
|
||||
False if metadata not ready yet (keep in queue for next poll)
|
||||
"""
|
||||
def _commit_transfer_to_req(self, decode_req: DecodeRequest):
|
||||
idx = decode_req.metadata_buffer_index
|
||||
(
|
||||
output_id,
|
||||
@@ -1409,11 +1396,25 @@ class DecodeTransferQueue:
|
||||
if _is_fake_transfer(decode_req.req, self.scheduler.server_args):
|
||||
pass
|
||||
elif actual_room == 0:
|
||||
# Case 1: Metadata not ready yet (actual_room == 0)
|
||||
# Keep request in queue and wait for next poll
|
||||
return False
|
||||
# Should never happen: _poll_with_metadata_gate already confirmed
|
||||
# readiness on all TP ranks. Abort deterministically to avoid
|
||||
# cross-rank queue divergence.
|
||||
logger.error(
|
||||
f"Metadata unexpectedly not ready after readiness gate: "
|
||||
f"request {decode_req.req.rid}, bootstrap_room={expected_room}, "
|
||||
f"metadata_buffer_index={idx}"
|
||||
)
|
||||
prepare_abort(
|
||||
decode_req.req,
|
||||
"Metadata unexpectedly not ready after readiness gate "
|
||||
"(bootstrap_room=0)",
|
||||
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
decode_req.kv_receiver.clear()
|
||||
decode_req.kv_receiver = None
|
||||
return
|
||||
elif actual_room != expected_room:
|
||||
# Case 2: Real corruption detected (mismatch)
|
||||
# Real corruption detected (mismatch)
|
||||
# Abort the request and remove from the queue
|
||||
error_msg = (
|
||||
f"Context corruption detected: Request {decode_req.req.rid} "
|
||||
@@ -1430,9 +1431,9 @@ class DecodeTransferQueue:
|
||||
)
|
||||
decode_req.kv_receiver.clear()
|
||||
decode_req.kv_receiver = None
|
||||
return True
|
||||
return
|
||||
|
||||
# Case 3: Success - commit the transfer
|
||||
# Success - commit the transfer
|
||||
decode_req.req.output_ids.append(output_id[0].item())
|
||||
decode_req.req.cached_tokens = cached_tokens[0].item()
|
||||
decode_req.req.cached_tokens_device = cached_tokens[1].item()
|
||||
@@ -1464,11 +1465,24 @@ class DecodeTransferQueue:
|
||||
decode_req.kv_receiver.clear()
|
||||
decode_req.kv_receiver = None
|
||||
decode_req.req.time_stats.set_wait_queue_entry_time()
|
||||
return True
|
||||
return
|
||||
|
||||
def _poll_with_metadata_gate(self) -> List[int]:
|
||||
return poll_and_all_reduce(
|
||||
[dr.kv_receiver for dr in self.queue],
|
||||
self.gloo_group,
|
||||
decode_reqs=self.queue,
|
||||
metadata_buffers=self.metadata_buffers,
|
||||
server_args=self.scheduler.server_args,
|
||||
)
|
||||
|
||||
def _poll_with_staging(self) -> list:
|
||||
return poll_and_all_reduce_with_staging(
|
||||
self.queue, self.staging_handler, self.gloo_group
|
||||
self.queue,
|
||||
self.staging_handler,
|
||||
self.gloo_group,
|
||||
metadata_buffers=self.metadata_buffers,
|
||||
server_args=self.scheduler.server_args,
|
||||
)
|
||||
|
||||
def _init_staging_handler(self, kv_manager):
|
||||
@@ -1489,9 +1503,7 @@ class DecodeTransferQueue:
|
||||
if self.enable_staging:
|
||||
polls = self._poll_with_staging()
|
||||
else:
|
||||
polls = poll_and_all_reduce(
|
||||
[dr.kv_receiver for dr in self.queue], self.gloo_group
|
||||
)
|
||||
polls = self._poll_with_metadata_gate()
|
||||
|
||||
transferred_reqs = []
|
||||
indices_to_remove = set()
|
||||
@@ -1524,26 +1536,23 @@ class DecodeTransferQueue:
|
||||
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
|
||||
continue
|
||||
elif poll == KVPoll.Success:
|
||||
should_remove = self._commit_transfer_to_req(decode_req)
|
||||
if should_remove:
|
||||
indices_to_remove.add(i)
|
||||
# Check if request was aborted due to corruption
|
||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||
self.scheduler.output_streamer.stream_output(
|
||||
[decode_req.req],
|
||||
decode_req.req.return_logprob,
|
||||
self._commit_transfer_to_req(decode_req)
|
||||
indices_to_remove.add(i)
|
||||
# Check if request was aborted due to corruption
|
||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||
self.scheduler.output_streamer.stream_output(
|
||||
[decode_req.req],
|
||||
decode_req.req.return_logprob,
|
||||
)
|
||||
if self.scheduler.enable_hisparse:
|
||||
self.scheduler.hisparse_coordinator.request_finished(
|
||||
decode_req.req
|
||||
)
|
||||
if self.scheduler.enable_hisparse:
|
||||
self.scheduler.hisparse_coordinator.request_finished(
|
||||
decode_req.req
|
||||
)
|
||||
release_kv_cache(
|
||||
decode_req.req, self.tree_cache, is_insert=False
|
||||
)
|
||||
if self.scheduler.metrics_reporter.enable_metrics:
|
||||
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
|
||||
else:
|
||||
transferred_reqs.append(decode_req.req)
|
||||
release_kv_cache(decode_req.req, self.tree_cache, is_insert=False)
|
||||
if self.scheduler.metrics_reporter.enable_metrics:
|
||||
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
|
||||
else:
|
||||
transferred_reqs.append(decode_req.req)
|
||||
elif poll in [
|
||||
KVPoll.Bootstrapping,
|
||||
KVPoll.WaitingForInput,
|
||||
|
||||
@@ -11,6 +11,7 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.disaggregation.base import KVPoll
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.utils import is_npu
|
||||
|
||||
@@ -23,6 +24,7 @@ if TYPE_CHECKING:
|
||||
CommonKVSender,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
#########################
|
||||
# Constants & Enums
|
||||
@@ -52,17 +54,54 @@ class DisaggregationMode(Enum):
|
||||
FAILURE_PROB = float(os.getenv("DISAGGREGATION_TEST_FAILURE_PROB", 0))
|
||||
|
||||
|
||||
def poll_and_all_reduce(pollers, gloo_group: dist.ProcessGroup):
|
||||
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool:
|
||||
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||
req.bootstrap_host is None
|
||||
and server_args.disaggregation_transfer_backend == "fake"
|
||||
)
|
||||
|
||||
|
||||
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> None:
|
||||
"""Downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||
|
||||
Mutates `polls` in-place. Called before all-reduce so that MIN across TP
|
||||
ranks naturally prevents any rank from committing before all ranks are ready.
|
||||
"""
|
||||
for i, poll_val in enumerate(polls):
|
||||
if poll_val == int(KVPoll.Success):
|
||||
decode_req = decode_reqs[i]
|
||||
if _is_fake_transfer(decode_req.req, server_args):
|
||||
continue
|
||||
actual_room = metadata_buffers.bootstrap_room[
|
||||
decode_req.metadata_buffer_index, 0
|
||||
].item()
|
||||
if actual_room == 0:
|
||||
polls[i] = int(KVPoll.Transferring)
|
||||
|
||||
|
||||
def poll_and_all_reduce(
|
||||
pollers,
|
||||
gloo_group: dist.ProcessGroup,
|
||||
decode_reqs=None,
|
||||
metadata_buffers: Optional[MetadataBuffers] = None,
|
||||
server_args: Optional[ServerArgs] = None,
|
||||
):
|
||||
# at a certain prob, the poll is failed to simulate failure
|
||||
if FAILURE_PROB > 0:
|
||||
from sglang.srt.disaggregation.base import KVPoll
|
||||
|
||||
polls = [
|
||||
int(KVPoll.Failed) if random.random() < FAILURE_PROB else int(poller.poll())
|
||||
for poller in pollers
|
||||
]
|
||||
else:
|
||||
polls = [int(poller.poll()) for poller in pollers]
|
||||
|
||||
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||
if (
|
||||
decode_reqs is not None
|
||||
and metadata_buffers is not None
|
||||
and server_args is not None
|
||||
):
|
||||
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
|
||||
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
|
||||
dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=gloo_group)
|
||||
return tensor_to_reduce.tolist()
|
||||
@@ -89,11 +128,13 @@ def poll_and_all_reduce_attn_cp_tp_group(
|
||||
|
||||
|
||||
def poll_and_all_reduce_with_staging(
|
||||
decode_reqs, staging_handler, gloo_group: dist.ProcessGroup
|
||||
decode_reqs,
|
||||
staging_handler,
|
||||
gloo_group: dist.ProcessGroup,
|
||||
metadata_buffers: Optional[MetadataBuffers] = None,
|
||||
server_args: Optional[ServerArgs] = None,
|
||||
):
|
||||
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
|
||||
from sglang.srt.disaggregation.base import KVPoll
|
||||
|
||||
for decode_req in decode_reqs:
|
||||
if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
|
||||
decode_req
|
||||
@@ -107,6 +148,9 @@ def poll_and_all_reduce_with_staging(
|
||||
decode_req
|
||||
):
|
||||
raw_polls[i] = int(KVPoll.Transferring)
|
||||
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||
if metadata_buffers is not None and server_args is not None:
|
||||
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
|
||||
poll_tensor = torch.tensor(raw_polls, dtype=torch.uint8, device="cpu")
|
||||
dist.all_reduce(poll_tensor, op=dist.ReduceOp.MIN, group=gloo_group)
|
||||
return poll_tensor.tolist()
|
||||
|
||||
Reference in New Issue
Block a user