[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.base.conn import StateType
from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST,
DisaggregationMode, DisaggregationMode,
KVClassType, KVClassType,
MetadataBuffers, MetadataBuffers,
ReqToMetadataIdxAllocator, ReqToMetadataIdxAllocator,
TransferBackend, TransferBackend,
_is_fake_transfer,
get_kv_class, get_kv_class,
is_mla_backend, is_mla_backend,
poll_and_all_reduce, poll_and_all_reduce,
@@ -82,18 +82,10 @@ logger = logging.getLogger(__name__)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import Scheduler 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() 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: def _bootstrap_addr(req: Req) -> str:
# FIXME: make a property of a req # FIXME: make a property of a req
return NetworkAddress(req.bootstrap_host, req.bootstrap_port).to_host_port_str() 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) self.staging_handler.register_decode_req(dr.req.bootstrap_room, dr)
def _commit_transfer_to_req(self, decode_req: DecodeRequest) -> bool: def _commit_transfer_to_req(self, decode_req: DecodeRequest):
"""
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)
"""
idx = decode_req.metadata_buffer_index idx = decode_req.metadata_buffer_index
( (
output_id, output_id,
@@ -1409,11 +1396,25 @@ class DecodeTransferQueue:
if _is_fake_transfer(decode_req.req, self.scheduler.server_args): if _is_fake_transfer(decode_req.req, self.scheduler.server_args):
pass pass
elif actual_room == 0: elif actual_room == 0:
# Case 1: Metadata not ready yet (actual_room == 0) # Should never happen: _poll_with_metadata_gate already confirmed
# Keep request in queue and wait for next poll # readiness on all TP ranks. Abort deterministically to avoid
return False # 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: elif actual_room != expected_room:
# Case 2: Real corruption detected (mismatch) # Real corruption detected (mismatch)
# Abort the request and remove from the queue # Abort the request and remove from the queue
error_msg = ( error_msg = (
f"Context corruption detected: Request {decode_req.req.rid} " f"Context corruption detected: Request {decode_req.req.rid} "
@@ -1430,9 +1431,9 @@ class DecodeTransferQueue:
) )
decode_req.kv_receiver.clear() decode_req.kv_receiver.clear()
decode_req.kv_receiver = None 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.output_ids.append(output_id[0].item())
decode_req.req.cached_tokens = cached_tokens[0].item() decode_req.req.cached_tokens = cached_tokens[0].item()
decode_req.req.cached_tokens_device = cached_tokens[1].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.clear()
decode_req.kv_receiver = None decode_req.kv_receiver = None
decode_req.req.time_stats.set_wait_queue_entry_time() 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: def _poll_with_staging(self) -> list:
return poll_and_all_reduce_with_staging( 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): def _init_staging_handler(self, kv_manager):
@@ -1489,9 +1503,7 @@ class DecodeTransferQueue:
if self.enable_staging: if self.enable_staging:
polls = self._poll_with_staging() polls = self._poll_with_staging()
else: else:
polls = poll_and_all_reduce( polls = self._poll_with_metadata_gate()
[dr.kv_receiver for dr in self.queue], self.gloo_group
)
transferred_reqs = [] transferred_reqs = []
indices_to_remove = set() indices_to_remove = set()
@@ -1524,26 +1536,23 @@ class DecodeTransferQueue:
self.scheduler.metrics_collector.increment_transfer_failed_reqs() self.scheduler.metrics_collector.increment_transfer_failed_reqs()
continue continue
elif poll == KVPoll.Success: elif poll == KVPoll.Success:
should_remove = self._commit_transfer_to_req(decode_req) self._commit_transfer_to_req(decode_req)
if should_remove: indices_to_remove.add(i)
indices_to_remove.add(i) # Check if request was aborted due to corruption
# Check if request was aborted due to corruption if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
if isinstance(decode_req.req.finished_reason, FINISH_ABORT): self.scheduler.output_streamer.stream_output(
self.scheduler.output_streamer.stream_output( [decode_req.req],
[decode_req.req], decode_req.req.return_logprob,
decode_req.req.return_logprob, )
if self.scheduler.enable_hisparse:
self.scheduler.hisparse_coordinator.request_finished(
decode_req.req
) )
if self.scheduler.enable_hisparse: release_kv_cache(decode_req.req, self.tree_cache, is_insert=False)
self.scheduler.hisparse_coordinator.request_finished( if self.scheduler.metrics_reporter.enable_metrics:
decode_req.req self.scheduler.metrics_collector.increment_transfer_failed_reqs()
) else:
release_kv_cache( transferred_reqs.append(decode_req.req)
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 [ elif poll in [
KVPoll.Bootstrapping, KVPoll.Bootstrapping,
KVPoll.WaitingForInput, KVPoll.WaitingForInput,
+50 -6
View File
@@ -11,6 +11,7 @@ import numpy as np
import torch import torch
import torch.distributed as dist import torch.distributed as dist
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.utils import is_npu from sglang.srt.utils import is_npu
@@ -23,6 +24,7 @@ if TYPE_CHECKING:
CommonKVSender, CommonKVSender,
) )
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.server_args import ServerArgs
######################### #########################
# Constants & Enums # Constants & Enums
@@ -52,17 +54,54 @@ class DisaggregationMode(Enum):
FAILURE_PROB = float(os.getenv("DISAGGREGATION_TEST_FAILURE_PROB", 0)) 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 # at a certain prob, the poll is failed to simulate failure
if FAILURE_PROB > 0: if FAILURE_PROB > 0:
from sglang.srt.disaggregation.base import KVPoll
polls = [ polls = [
int(KVPoll.Failed) if random.random() < FAILURE_PROB else int(poller.poll()) int(KVPoll.Failed) if random.random() < FAILURE_PROB else int(poller.poll())
for poller in pollers for poller in pollers
] ]
else: else:
polls = [int(poller.poll()) for poller in pollers] 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") tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=gloo_group) dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=gloo_group)
return tensor_to_reduce.tolist() 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( 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.""" """Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
from sglang.srt.disaggregation.base import KVPoll
for decode_req in decode_reqs: for decode_req in decode_reqs:
if decode_req.kv_receiver.require_staging and not staging_handler.is_done( if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
decode_req decode_req
@@ -107,6 +148,9 @@ def poll_and_all_reduce_with_staging(
decode_req decode_req
): ):
raw_polls[i] = int(KVPoll.Transferring) 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") poll_tensor = torch.tensor(raw_polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(poll_tensor, op=dist.ReduceOp.MIN, group=gloo_group) dist.all_reduce(poll_tensor, op=dist.ReduceOp.MIN, group=gloo_group)
return poll_tensor.tolist() return poll_tensor.tolist()