From c8c1aed5e908a4dfec88781bff3ae053b589e867 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Tue, 26 May 2026 21:59:42 +0800 Subject: [PATCH] [PD] Fix cross-rank queue divergence by gating metadata readiness before all-reduce (#26394) Signed-off-by: Shangming Cai --- python/sglang/srt/disaggregation/decode.py | 99 ++++++++++++---------- python/sglang/srt/disaggregation/utils.py | 56 ++++++++++-- 2 files changed, 104 insertions(+), 51 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 65f2e4cb8..93208639f 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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, diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 030b0cb7a..d64fd0298 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -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()