[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:
Shangming Cai
2026-05-26 21:59:42 +08:00
committed by GitHub
parent 98eb84497d
commit c8c1aed5e9
2 changed files with 104 additions and 51 deletions
+54 -45
View File
@@ -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,
+50 -6
View File
@@ -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()