[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_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
self.enable_trace = get_observability().enable_trace
|
self.enable_trace = get_observability().enable_trace
|
||||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||||
self.start_prefill_thread()
|
|
||||||
self.session_failures = defaultdict(int)
|
self.session_failures = defaultdict(int)
|
||||||
self.failed_sessions = set()
|
self.failed_sessions = set()
|
||||||
|
self.session_lock = threading.Lock()
|
||||||
|
self.start_prefill_thread()
|
||||||
# Per-room count of chunks not yet transferred; teardown waits for
|
# Per-room count of chunks not yet transferred; teardown waits for
|
||||||
# zero so a deferred chunk is not dropped by an early conclude.
|
# zero so a deferred chunk is not dropped by an early conclude.
|
||||||
self._staging_outstanding = defaultdict(int)
|
self._staging_outstanding = defaultdict(int)
|
||||||
self.session_lock = threading.Lock()
|
|
||||||
# Determine the number of threads to use for kv sender
|
# Determine the number of threads to use for kv sender
|
||||||
cpu_count = os.cpu_count()
|
cpu_count = os.cpu_count()
|
||||||
transfer_thread_pool_size = (
|
transfer_thread_pool_size = (
|
||||||
|
|||||||
@@ -334,6 +334,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
self.room_to_bootstrap_addr: Dict[int, str] = {}
|
self.room_to_bootstrap_addr: Dict[int, str] = {}
|
||||||
self._start_decode_thread()
|
self._start_decode_thread()
|
||||||
|
self._start_heartbeat_checker_thread()
|
||||||
|
|
||||||
def _init_engine(self) -> IOEngine:
|
def _init_engine(self) -> IOEngine:
|
||||||
if self.kv_args.ib_device:
|
if self.kv_args.ib_device:
|
||||||
@@ -902,6 +903,20 @@ class MoriKVManager(CommonKVManager):
|
|||||||
src_k_descs = src_descs[:num_local_layers]
|
src_k_descs = src_descs[:num_local_layers]
|
||||||
src_v_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
|
start_layer = self.kv_args.prefill_start_layer
|
||||||
end_layer = start_layer + num_local_layers
|
end_layer = start_layer + num_local_layers
|
||||||
dst_total_layers = len(dst_mem_descs) // 2
|
dst_total_layers = len(dst_mem_descs) // 2
|
||||||
@@ -910,8 +925,18 @@ class MoriKVManager(CommonKVManager):
|
|||||||
"Destination KV descriptors do not match prefill pp configuration"
|
"Destination KV descriptors do not match prefill pp configuration"
|
||||||
)
|
)
|
||||||
dst_k_descs = dst_mem_descs[start_layer:end_layer]
|
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_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
|
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]:
|
) -> tuple[List[MemoryDesc], List[MemoryDesc], int]:
|
||||||
src_descs = self.kv_mem_descs
|
src_descs = self.kv_mem_descs
|
||||||
num_local_layers = len(src_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
|
start_layer = self.kv_args.prefill_start_layer
|
||||||
end_layer = start_layer + num_local_layers
|
end_layer = start_layer + num_local_layers
|
||||||
if end_layer > len(dst_mem_descs):
|
if end_layer > len(dst_mem_descs):
|
||||||
@@ -1083,7 +1112,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
statuses: List[TransferStatus] = []
|
statuses: List[TransferStatus] = []
|
||||||
kv_item_len = self.kv_args.kv_item_lens[0]
|
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 = (
|
src_descs, dst_descs, layers_current_pp_stage = (
|
||||||
self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs)
|
self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1089,6 +1089,14 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
room = kv_chunk.room
|
room = kv_chunk.room
|
||||||
handles: List[Any] = []
|
handles: List[Any] = []
|
||||||
try:
|
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
|
# Counted at dequeue, before the status check, so
|
||||||
# `outstanding == 0` means nothing is dequeued or in flight --
|
# `outstanding == 0` means nothing is dequeued or in flight --
|
||||||
# the predicate the abort ack relies on. The flag survives
|
# the predicate the abort ack relies on. The flag survives
|
||||||
@@ -1104,7 +1112,15 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
self._maybe_ack_drained_abort(room)
|
self._maybe_ack_drained_abort(room)
|
||||||
continue
|
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
|
# Lazily build a per-worker staging strategy bound to this
|
||||||
# worker's private staging buffer (matches mooncake).
|
# worker's private staging buffer (matches mooncake).
|
||||||
@@ -1117,7 +1133,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
|
|
||||||
self.update_status(room, KVPoll.Transferring)
|
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.
|
# Note(kpham-sgl): Pack each DCP rank once into its fixed region.
|
||||||
# NIXL reads regions asynchronously; the chunk barrier prevents
|
# NIXL reads regions asynchronously; the chunk barrier prevents
|
||||||
# reuse until every transfer completes.
|
# reuse until every transfer completes.
|
||||||
|
|||||||
Reference in New Issue
Block a user