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