[PD] Align defensive protocol behavior across Mooncake, NIXL, and Mori (#35281)
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Shangming Cai
parent
df75ec5f77
commit
7700602278
@@ -219,13 +219,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||
self.enable_trace = get_observability().enable_trace
|
||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
self.start_prefill_thread()
|
||||
self.session_failures = defaultdict(int)
|
||||
self.failed_sessions = set()
|
||||
self.session_lock = threading.Lock()
|
||||
self.start_prefill_thread()
|
||||
# Per-room count of chunks not yet transferred; teardown waits for
|
||||
# zero so a deferred chunk is not dropped by an early conclude.
|
||||
self._staging_outstanding = defaultdict(int)
|
||||
self.session_lock = threading.Lock()
|
||||
# Determine the number of threads to use for kv sender
|
||||
cpu_count = os.cpu_count()
|
||||
transfer_thread_pool_size = (
|
||||
|
||||
@@ -334,6 +334,7 @@ class MoriKVManager(CommonKVManager):
|
||||
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
self.room_to_bootstrap_addr: Dict[int, str] = {}
|
||||
self._start_decode_thread()
|
||||
self._start_heartbeat_checker_thread()
|
||||
|
||||
def _init_engine(self) -> IOEngine:
|
||||
if self.kv_args.ib_device:
|
||||
@@ -902,6 +903,20 @@ class MoriKVManager(CommonKVManager):
|
||||
src_k_descs = src_descs[:num_local_layers]
|
||||
src_v_descs = src_descs[num_local_layers:]
|
||||
|
||||
# Both peers expose the same PP-local layout. Their descriptor indices
|
||||
# are already aligned, so applying the Prefill rank's global layer
|
||||
# offset would incorrectly index into a local list.
|
||||
if len(src_descs) == len(dst_mem_descs):
|
||||
dst_k_descs = dst_mem_descs[:num_local_layers]
|
||||
dst_v_descs = dst_mem_descs[num_local_layers:]
|
||||
return (
|
||||
src_k_descs,
|
||||
src_v_descs,
|
||||
dst_k_descs,
|
||||
dst_v_descs,
|
||||
num_local_layers,
|
||||
)
|
||||
|
||||
start_layer = self.kv_args.prefill_start_layer
|
||||
end_layer = start_layer + num_local_layers
|
||||
dst_total_layers = len(dst_mem_descs) // 2
|
||||
@@ -910,8 +925,18 @@ class MoriKVManager(CommonKVManager):
|
||||
"Destination KV descriptors do not match prefill pp configuration"
|
||||
)
|
||||
dst_k_descs = dst_mem_descs[start_layer:end_layer]
|
||||
if (
|
||||
num_local_layers < dst_total_layers
|
||||
and dst_total_layers % num_local_layers != 0
|
||||
):
|
||||
# Decode has draft-model KV while Prefill has target-model KV only:
|
||||
# [K_main..., V_main..., draft_K..., draft_V...].
|
||||
multiplier_ratio = dst_total_layers // num_local_layers
|
||||
dst_v_offset = num_local_layers * multiplier_ratio
|
||||
else:
|
||||
dst_v_offset = dst_total_layers
|
||||
dst_v_descs = dst_mem_descs[
|
||||
dst_total_layers + start_layer : dst_total_layers + end_layer
|
||||
dst_v_offset + start_layer : dst_v_offset + end_layer
|
||||
]
|
||||
return src_k_descs, src_v_descs, dst_k_descs, dst_v_descs, num_local_layers
|
||||
|
||||
@@ -920,6 +945,10 @@ class MoriKVManager(CommonKVManager):
|
||||
) -> tuple[List[MemoryDesc], List[MemoryDesc], int]:
|
||||
src_descs = self.kv_mem_descs
|
||||
num_local_layers = len(src_descs)
|
||||
# Same-PP peers register matching local descriptor lists.
|
||||
if len(src_descs) == len(dst_mem_descs):
|
||||
return src_descs, dst_mem_descs, num_local_layers
|
||||
|
||||
start_layer = self.kv_args.prefill_start_layer
|
||||
end_layer = start_layer + num_local_layers
|
||||
if end_layer > len(dst_mem_descs):
|
||||
@@ -1083,7 +1112,7 @@ class MoriKVManager(CommonKVManager):
|
||||
statuses: List[TransferStatus] = []
|
||||
kv_item_len = self.kv_args.kv_item_lens[0]
|
||||
|
||||
if self.is_mla_backend:
|
||||
if self.is_mla_backend or self.is_hybrid_mla_backend:
|
||||
src_descs, dst_descs, layers_current_pp_stage = (
|
||||
self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs)
|
||||
)
|
||||
|
||||
@@ -1089,6 +1089,14 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
room = kv_chunk.room
|
||||
handles: List[Any] = []
|
||||
try:
|
||||
if room not in self.request_status:
|
||||
logger.debug(
|
||||
"Skipping chunk for room %s because it has been cleared",
|
||||
room,
|
||||
)
|
||||
self._staging_outstanding.pop(room, None)
|
||||
continue
|
||||
|
||||
# Counted at dequeue, before the status check, so
|
||||
# `outstanding == 0` means nothing is dequeued or in flight --
|
||||
# the predicate the abort ack relies on. The flag survives
|
||||
@@ -1104,7 +1112,15 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
self._maybe_ack_drained_abort(room)
|
||||
continue
|
||||
|
||||
assert room in self.transfer_infos
|
||||
room_transfer_infos = self.transfer_infos.get(room)
|
||||
if room_transfer_infos is None:
|
||||
logger.debug(
|
||||
"Skipping chunk for room %s because its transfer metadata "
|
||||
"has been cleared",
|
||||
room,
|
||||
)
|
||||
self._staging_outstanding.pop(room, None)
|
||||
continue
|
||||
|
||||
# Lazily build a per-worker staging strategy bound to this
|
||||
# worker's private staging buffer (matches mooncake).
|
||||
@@ -1117,7 +1133,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
|
||||
self.update_status(room, KVPoll.Transferring)
|
||||
|
||||
reqs_to_be_processed = list(self.transfer_infos[room].values())
|
||||
reqs_to_be_processed = list(room_transfer_infos.values())
|
||||
# Note(kpham-sgl): Pack each DCP rank once into its fixed region.
|
||||
# NIXL reads regions asynchronously; the chunk barrier prevents
|
||||
# reuse until every transfer completes.
|
||||
|
||||
Reference in New Issue
Block a user