[PD] Refactor Disagg Conn and Fix Hang with total_request/total_tokens Balancing (#21299)

Co-authored-by: Weiliangl User <weiliangl@login-node.hosted.internal>
This commit is contained in:
weireweire
2026-03-31 18:01:50 +08:00
committed by GitHub
co-authored by Weiliangl User
parent acd37d8701
commit 4455d17619
8 changed files with 161 additions and 79 deletions
+11 -1
View File
@@ -122,13 +122,23 @@ class BaseKVReceiver(ABC):
@abstractmethod @abstractmethod
def init( def init(
self,
prefill_dp_rank: int,
):
"""
Resolve bootstrap metadata and mark the receiver ready for transfer metadata.
"""
...
@abstractmethod
def send_metadata(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List[int]] = None,
): ):
""" """
Set req's index metadata locally or notify the prefill server about the kv indices, aux index, and state_indices. Notify the prefill server about the kv indices, aux index, and state_indices.
""" """
... ...
@@ -489,20 +489,31 @@ class CommonKVReceiver(BaseKVReceiver):
mgr: CommonKVManager, mgr: CommonKVManager,
bootstrap_addr: str, bootstrap_addr: str,
bootstrap_room: Optional[int] = None, bootstrap_room: Optional[int] = None,
prefill_dp_rank: Optional[int] = None,
): ):
self.bootstrap_room = bootstrap_room self.bootstrap_room = bootstrap_room
self.bootstrap_addr = bootstrap_addr self.bootstrap_addr = bootstrap_addr
self.kv_mgr = mgr self.kv_mgr = mgr
self.conclude_state: Optional[KVPoll] = None
self.bootstrap_infos = None
self.prefill_info = None
self.prefill_dp_rank = None
self.target_tp_rank = None
self.target_tp_ranks = None
self.target_cp_ranks = None
self.target_pp_ranks = None
self.required_dst_info_num = None
self.required_prefill_response_num = None
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping) self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping)
def init(self, prefill_dp_rank: int):
if self.bootstrap_addr not in self.kv_mgr.prefill_info_table: if self.bootstrap_addr not in self.kv_mgr.prefill_info_table:
self.kv_mgr.record_failure( self.kv_mgr.record_failure(
self.bootstrap_room, self.bootstrap_room,
f"Prefill server with bootstrap_addr: {self.bootstrap_addr} is healthy before, but now it is down. Request (bootstrap_room: {self.bootstrap_room}) has been marked as failed.", f"Prefill server with bootstrap_addr: {self.bootstrap_addr} is healthy before, but now it is down. Request (bootstrap_room: {self.bootstrap_room}) has been marked as failed.",
) )
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed) self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
self.bootstrap_infos = None
return return
# Read pre-computed rank mapping from prefill_info (computed in try_ensure_parallel_info) # Read pre-computed rank mapping from prefill_info (computed in try_ensure_parallel_info)
@@ -520,11 +531,9 @@ class CommonKVReceiver(BaseKVReceiver):
self.required_prefill_response_num self.required_prefill_response_num
) )
assert (
prefill_dp_rank is not None
), "prefill_dp_rank must be resolved before creating receiver"
self.prefill_dp_rank = prefill_dp_rank self.prefill_dp_rank = prefill_dp_rank
self._setup_bootstrap_infos() self._setup_bootstrap_infos()
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput)
def _setup_bootstrap_infos(self): def _setup_bootstrap_infos(self):
all_bootstrap_infos = [] all_bootstrap_infos = []
@@ -562,6 +571,7 @@ class CommonKVReceiver(BaseKVReceiver):
self.bootstrap_room, self.bootstrap_room,
f"Could not fetch bootstrap info for: prefill_dp_rank: {self.prefill_dp_rank} prefill_cp_rank: {target_cp_rank} target_tp_rank: {target_tp_rank} and target_pp_rank {target_pp_rank}", f"Could not fetch bootstrap info for: prefill_dp_rank: {self.prefill_dp_rank} prefill_cp_rank: {target_cp_rank} target_tp_rank: {target_tp_rank} and target_pp_rank {target_pp_rank}",
) )
self.conclude_state = KVPoll.Failed
self.kv_mgr.update_status( self.kv_mgr.update_status(
self.bootstrap_room, KVPoll.Failed self.bootstrap_room, KVPoll.Failed
) )
@@ -645,6 +655,14 @@ class CommonKVReceiver(BaseKVReceiver):
def _register_kv_args(self): def _register_kv_args(self):
pass pass
def send_metadata(
self,
kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None,
):
raise NotImplementedError
def failure_exception(self): def failure_exception(self):
raise Exception("Fake KVReceiver Exception") raise Exception("Fake KVReceiver Exception")
+46 -38
View File
@@ -276,7 +276,7 @@ class DecodePreallocQueue:
# Queue for requests pending pre-allocation # Queue for requests pending pre-allocation
self.queue: List[DecodeRequest] = [] self.queue: List[DecodeRequest] = []
self.retracted_queue: List[Req] = [] self.retracted_queue: List[Req] = []
self.pending_reqs: List[Req] = [] self.pending_reqs: List[DecodeRequest] = []
self._ensure_retry_count: Dict[str, int] = {} self._ensure_retry_count: Dict[str, int] = {}
self._max_ensure_retries: int = 20 # scheduling cycles self._max_ensure_retries: int = 20 # scheduling cycles
self._ensure_last_attempt_time: Dict[str, float] = {} self._ensure_last_attempt_time: Dict[str, float] = {}
@@ -368,17 +368,20 @@ class DecodePreallocQueue:
req.retraction_mb_id = None req.retraction_mb_id = None
self.retracted_queue.append(req) self.retracted_queue.append(req)
else: else:
decode_req = self._create_receiver_and_enqueue(req)
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue # NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
if _is_fake_transfer(req, self.scheduler.server_args): if _is_fake_transfer(req, self.scheduler.server_args):
self._create_receiver_and_enqueue(req, 0) decode_req.kv_receiver.init(0)
return return
# Fast path: cache-only lookup, no network calls # Fast path: cache-only lookup, no network calls
prefill_dp_rank = self._resolve_prefill_dp_rank(req) prefill_dp_rank = self._resolve_prefill_dp_rank(req)
if prefill_dp_rank is not None: if prefill_dp_rank is not None:
self._create_receiver_and_enqueue(req, prefill_dp_rank) decode_req.kv_receiver.init(prefill_dp_rank)
else: return
self.pending_reqs.append(req)
self.pending_reqs.append(decode_req)
def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]: def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]:
if req.disagg_prefill_dp_rank is not None: if req.disagg_prefill_dp_rank is not None:
@@ -396,7 +399,7 @@ class DecodePreallocQueue:
return None return None
def _create_receiver_and_enqueue(self, req: Req, prefill_dp_rank: int) -> None: def _create_receiver_and_enqueue(self, req: Req) -> DecodeRequest:
backend = ( backend = (
TransferBackend.FAKE TransferBackend.FAKE
if _is_fake_transfer(req, self.scheduler.server_args) if _is_fake_transfer(req, self.scheduler.server_args)
@@ -408,12 +411,11 @@ class DecodePreallocQueue:
mgr=self.kv_manager, mgr=self.kv_manager,
bootstrap_addr=_bootstrap_addr(req), bootstrap_addr=_bootstrap_addr(req),
bootstrap_room=req.bootstrap_room, bootstrap_room=req.bootstrap_room,
prefill_dp_rank=prefill_dp_rank,
) )
self.queue.append( decode_req = DecodeRequest(req=req, kv_receiver=kv_receiver)
DecodeRequest(req=req, kv_receiver=kv_receiver, waiting_for_input=False) self.queue.append(decode_req)
) return decode_req
def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool: def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool:
if len(req.origin_input_ids) > self.max_total_num_tokens: if len(req.origin_input_ids) > self.max_total_num_tokens:
@@ -511,12 +513,12 @@ class DecodePreallocQueue:
raise ValueError(f"Unexpected poll case: {poll}") raise ValueError(f"Unexpected poll case: {poll}")
def _ensure_prefill_info( def _ensure_prefill_info(
self, addr_to_reqs: Dict[str, List[Req]] self, addr_to_reqs: Dict[str, List[DecodeRequest]]
) -> Tuple[Dict[str, List[Req]], List[Req]]: ) -> Tuple[Dict[str, List[DecodeRequest]], List[DecodeRequest]]:
"""Non-blocking ensure parallel info for each addr. """Non-blocking ensure parallel info for each addr.
Returns (ready_addrs, remaining_reqs).""" Returns (ready_addrs, remaining_reqs)."""
ready: Dict[str, List[Req]] = {} ready: Dict[str, List[DecodeRequest]] = {}
remaining: List[Req] = [] remaining: List[DecodeRequest] = []
now = time.monotonic() now = time.monotonic()
for bootstrap_addr, reqs in addr_to_reqs.items(): for bootstrap_addr, reqs in addr_to_reqs.items():
@@ -543,13 +545,17 @@ class DecodePreallocQueue:
if count >= self._max_ensure_retries: if count >= self._max_ensure_retries:
error_msg = f"Could not fetch prefill parallel info from {bootstrap_addr} after {count} attempts" error_msg = f"Could not fetch prefill parallel info from {bootstrap_addr} after {count} attempts"
logger.error(error_msg) logger.error(error_msg)
for req in reqs: for decode_req in reqs:
prepare_abort( prepare_abort(
req, error_msg, status_code=HTTPStatus.INTERNAL_SERVER_ERROR decode_req.req,
error_msg,
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
) )
if self.scheduler.enable_metrics: if self.scheduler.enable_metrics:
self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() self.scheduler.metrics_collector.increment_bootstrap_failed_reqs()
self.scheduler.stream_output([req], req.return_logprob) self.scheduler.stream_output(
[decode_req.req], decode_req.req.return_logprob
)
del self._ensure_retry_count[bootstrap_addr] del self._ensure_retry_count[bootstrap_addr]
del self._ensure_last_attempt_time[bootstrap_addr] del self._ensure_last_attempt_time[bootstrap_addr]
else: else:
@@ -558,46 +564,48 @@ class DecodePreallocQueue:
return ready, remaining return ready, remaining
def _resolve_pending_reqs(self) -> None: def _resolve_pending_reqs(self) -> None:
"""Batch-resolve prefill_dp_ranks for pending requests and create receivers.""" """Batch-resolve prefill_dp_ranks for pending requests and initialize receivers."""
if not self.pending_reqs: if not self.pending_reqs:
return return
# Group pending requests by bootstrap_addr # Group pending requests by bootstrap_addr
addr_to_reqs: Dict[str, List[Req]] = {} addr_to_reqs: Dict[str, List[DecodeRequest]] = {}
for req in self.pending_reqs: for decode_req in self.pending_reqs:
addr = _bootstrap_addr(req) addr = _bootstrap_addr(decode_req.req)
addr_to_reqs.setdefault(addr, []).append(req) addr_to_reqs.setdefault(addr, []).append(decode_req)
# Pass 1: ensure parallel info for each addr # Pass 1: ensure parallel info for each addr
ready_addrs, remaining = self._ensure_prefill_info(addr_to_reqs) ready_addrs, remaining = self._ensure_prefill_info(addr_to_reqs)
# Pass 2: resolve dp rank for addrs whose info is available resolved: List[Tuple[DecodeRequest, int]] = []
resolved = [] for bootstrap_addr, decode_reqs in ready_addrs.items():
for bootstrap_addr, reqs in ready_addrs.items(): need_query: List[DecodeRequest] = []
need_query: List[Req] = [] for decode_req in decode_reqs:
for req in reqs: prefill_dp_rank = self._resolve_prefill_dp_rank(decode_req.req)
prefill_dp_rank = self._resolve_prefill_dp_rank(req)
if prefill_dp_rank is not None: if prefill_dp_rank is not None:
resolved.append((req, prefill_dp_rank)) resolved.append((decode_req, prefill_dp_rank))
else: else:
need_query.append(req) need_query.append(decode_req)
# Pass 2: resolve dp rank for addrs whose info is available
if need_query: if need_query:
rooms = [req.bootstrap_room for req in need_query] rooms = [decode_req.req.bootstrap_room for decode_req in need_query]
room_to_rank = CommonKVReceiver.query_prefill_dp_ranks( room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
bootstrap_addr, rooms bootstrap_addr, rooms
) )
for req in need_query: for decode_req in need_query:
prefill_dp_rank = room_to_rank.get(str(req.bootstrap_room)) prefill_dp_rank = room_to_rank.get(
str(decode_req.req.bootstrap_room)
)
if prefill_dp_rank is not None: if prefill_dp_rank is not None:
resolved.append((req, int(prefill_dp_rank))) resolved.append((decode_req, int(prefill_dp_rank)))
else: else:
remaining.append(req) remaining.append(decode_req)
self.pending_reqs = remaining self.pending_reqs = remaining
for req, prefill_dp_rank in resolved: for decode_req, prefill_dp_rank in resolved:
self._create_receiver_and_enqueue(req, prefill_dp_rank) decode_req.kv_receiver.init(prefill_dp_rank)
def pop_preallocated( def pop_preallocated(
self, rids_to_check: Optional[List[str]] = None self, rids_to_check: Optional[List[str]] = None
@@ -726,7 +734,7 @@ class DecodePreallocQueue:
) )
assert decode_req.metadata_buffer_index is not None assert decode_req.metadata_buffer_index is not None
page_indices = kv_to_page_indices(kv_indices, page_size) page_indices = kv_to_page_indices(kv_indices, page_size)
decode_req.kv_receiver.init( decode_req.kv_receiver.send_metadata(
page_indices, decode_req.metadata_buffer_index, state_indices page_indices, decode_req.metadata_buffer_index, state_indices
) )
preallocated_reqs.append(decode_req) preallocated_reqs.append(decode_req)
+15 -10
View File
@@ -82,28 +82,33 @@ class FakeKVReceiver(BaseKVReceiver):
mgr: BaseKVManager, mgr: BaseKVManager,
bootstrap_addr: str, bootstrap_addr: str,
bootstrap_room: Optional[int] = None, bootstrap_room: Optional[int] = None,
prefill_dp_rank: Optional[int] = None,
): ):
self.has_init = False self.bootstrap_done = False
self.has_sent_metadata = False
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
if self.has_init is False: if not self.bootstrap_done:
# Assume handshake completed instantly return KVPoll.Bootstrapping
if not self.has_sent_metadata:
return KVPoll.WaitingForInput return KVPoll.WaitingForInput
else: logger.debug("FakeKVReceiver poll success")
# Assume transfer completed instantly return KVPoll.Success
logger.debug("FakeKVReceiver poll success")
return KVPoll.Success
def init( def init(
self,
prefill_dp_rank: int,
):
self.bootstrap_done = True
def send_metadata(
self, self,
kv_indices: list[int], kv_indices: list[int],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List[int]] = None,
): ):
self.has_init = True self.has_sent_metadata = True
logger.debug( logger.debug(
f"FakeKVReceiver init with kv_indices: {kv_indices}, aux_index: {aux_index}, state_indices: {state_indices}" f"FakeKVReceiver send_metadata with kv_indices: {kv_indices}, aux_index: {aux_index}, state_indices: {state_indices}"
) )
def failure_exception(self): def failure_exception(self):
@@ -1238,15 +1238,10 @@ class MooncakeKVReceiver(CommonKVReceiver):
mgr: MooncakeKVManager, mgr: MooncakeKVManager,
bootstrap_addr: str, bootstrap_addr: str,
bootstrap_room: Optional[int] = None, bootstrap_room: Optional[int] = None,
prefill_dp_rank: Optional[int] = None,
): ):
self.session_id = mgr.get_session_id() self.session_id = mgr.get_session_id()
self.conclude_state = None
self.init_time = None self.init_time = None
super().__init__(mgr, bootstrap_addr, bootstrap_room, prefill_dp_rank) super().__init__(mgr, bootstrap_addr, bootstrap_room)
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput)
def _register_kv_args(self): def _register_kv_args(self):
for bootstrap_info in self.bootstrap_infos: for bootstrap_info in self.bootstrap_infos:
@@ -1297,6 +1292,12 @@ class MooncakeKVReceiver(CommonKVReceiver):
) )
def init( def init(
self,
prefill_dp_rank: int,
):
super().init(prefill_dp_rank)
def send_metadata(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
@@ -985,17 +985,18 @@ class MoriKVReceiver(CommonKVReceiver):
mgr: MoriKVManager, mgr: MoriKVManager,
bootstrap_addr: str, bootstrap_addr: str,
bootstrap_room: Optional[int] = None, bootstrap_room: Optional[int] = None,
prefill_dp_rank: Optional[int] = None,
): ):
super().__init__(mgr, bootstrap_addr, bootstrap_room, prefill_dp_rank) super().__init__(mgr, bootstrap_addr, bootstrap_room)
self.conclude_state: Optional[KVPoll] = None
self.init_time: Optional[float] = None self.init_time: Optional[float] = None
if self.bootstrap_room is None or self.bootstrap_infos is None:
def init(
self,
prefill_dp_rank: int,
):
super().init(prefill_dp_rank)
if self.bootstrap_room is None:
return return
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room)
self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput)
self._register_kv_args()
def _register_kv_args(self): def _register_kv_args(self):
if self.bootstrap_infos is None: if self.bootstrap_infos is None:
@@ -1029,7 +1030,7 @@ class MoriKVReceiver(CommonKVReceiver):
] ]
) )
def init( def send_metadata(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
+8 -10
View File
@@ -957,20 +957,18 @@ class NixlKVReceiver(CommonKVReceiver):
mgr: NixlKVManager, mgr: NixlKVManager,
bootstrap_addr: str, bootstrap_addr: str,
bootstrap_room: Optional[int] = None, bootstrap_room: Optional[int] = None,
prefill_dp_rank: Optional[int] = None,
): ):
self.started_transfer = False self.started_transfer = False
self.conclude_state = None super().__init__(mgr, bootstrap_addr, bootstrap_room)
super().__init__(mgr, bootstrap_addr, bootstrap_room, prefill_dp_rank)
# Track this room with its bootstrap address for heartbeat monitoring
if hasattr(self.kv_mgr, "addr_to_rooms_tracker"):
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(
self.bootstrap_room
)
self.init_time = None self.init_time = None
def init( def init(
self,
prefill_dp_rank: int,
):
super().init(prefill_dp_rank)
def send_metadata(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
@@ -1026,7 +1024,7 @@ class NixlKVReceiver(CommonKVReceiver):
self.conclude_state = status self.conclude_state = status
return status return status
if not self.started_transfer: if not self.started_transfer:
return KVPoll.WaitingForInput # type: ignore return status
now = time.time() now = time.time()
elapsed = now - self.init_time elapsed = now - self.init_time
@@ -110,7 +110,6 @@ class TestDisaggregationDPAttention(PDDisaggregationServerBase):
class TestDisaggregationDPAttentionRoundRobin(TestDisaggregationDPAttention): class TestDisaggregationDPAttentionRoundRobin(TestDisaggregationDPAttention):
LOAD_BALANCE_METHOD = "round_robin" LOAD_BALANCE_METHOD = "round_robin"
# TODO: add test for other load balance methods
# TODO: add a balancedness metric # TODO: add a balancedness metric
def test_bench_serving(self): def test_bench_serving(self):
@@ -130,6 +129,48 @@ class TestDisaggregationDPAttentionRoundRobin(TestDisaggregationDPAttention):
self.assertEqual(result["completed"], 1000) self.assertEqual(result["completed"], 1000)
class TestDisaggregationDPAttentionTotalRequests(TestDisaggregationDPAttention):
LOAD_BALANCE_METHOD = "total_requests"
test_gsm8k = unittest.skip(
"Covered by base class; this class targets total_requests path."
)(TestDisaggregationDPAttention.test_gsm8k)
def test_bench_serving(self):
args = get_benchmark_args(
base_url=f"http://{self.base_host}:{self.lb_port}",
dataset_name="random",
tokenizer=self.model,
num_prompts=256,
random_input_len=2048,
random_output_len=512,
request_rate=float("inf"),
max_concurrency=128,
)
result = run_benchmark(args)
self.assertEqual(result["completed"], 256)
class TestDisaggregationDPAttentionTotalTokens(TestDisaggregationDPAttention):
LOAD_BALANCE_METHOD = "total_tokens"
test_gsm8k = unittest.skip(
"Covered by base class; this class targets total_tokens path."
)(TestDisaggregationDPAttention.test_gsm8k)
def test_bench_serving(self):
args = get_benchmark_args(
base_url=f"http://{self.base_host}:{self.lb_port}",
dataset_name="random",
tokenizer=self.model,
num_prompts=256,
random_input_len=2048,
random_output_len=512,
request_rate=float("inf"),
max_concurrency=128,
)
result = run_benchmark(args)
self.assertEqual(result["completed"], 256)
@unittest.skip( @unittest.skip(
"Skip this test until new testing logic in mini-lb has been updated in docker image." "Skip this test until new testing logic in mini-lb has been updated in docker image."
) )