[PD] Align defensive protocol behavior across Mooncake, NIXL, and Mori (#35281)

Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
jambow0320
2026-08-31 10:57:56 +08:00
committed by GitHub
co-authored by Shangming Cai
parent df75ec5f77
commit 7700602278
3 changed files with 51 additions and 6 deletions
@@ -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 = (
+31 -2
View File
@@ -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)
)
+18 -2
View File
@@ -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.