diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 9d93721db..afcc0ac47 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -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 = ( diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 208fcadb2..2eda6fb81 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -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) ) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 157c88356..f4f49c3ba 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -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.