Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions python/sglang/srt/disaggregation/mooncake/conn.py
Original file line number Diff line number Diff line change
Expand Up @@ -215,13 +215,13 @@ def __init__(
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
self.enable_trace = server_args.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 = (
Expand Down
33 changes: 31 additions & 2 deletions python/sglang/srt/disaggregation/mori/conn.py
Original file line number Diff line number Diff line change
Expand Up @@ -334,6 +334,7 @@ def __init__(
elif self.disaggregation_mode == DisaggregationMode.DECODE:
self.room_to_bootstrap_addr: Dict[int, str] = {}
self._start_decode_thread()
self._start_heartbeat_checker_thread()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Will this depend on #36160?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

No. #36160 only touches the PREFILL branch of Mori. This change is in DECODE branch and just wires in the existing common heartbeat checker so decode can detect a dead prefill node and fail the affected rooms instead of waiting out timeout.


def _init_engine(self) -> IOEngine:
if self.kv_args.ib_device:
Expand Down Expand Up @@ -902,6 +903,20 @@ def _get_mha_mem_desc_slices(
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
Expand All @@ -910,8 +925,18 @@ def _get_mha_mem_desc_slices(
"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

Expand All @@ -920,6 +945,10 @@ def _get_mla_mem_desc_slices(
) -> 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):
Expand Down Expand Up @@ -1083,7 +1112,7 @@ def send_kvcache(
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)
)
Expand Down
20 changes: 18 additions & 2 deletions python/sglang/srt/disaggregation/nixl/conn.py
Original file line number Diff line number Diff line change
Expand Up @@ -1089,6 +1089,14 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)
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
Expand All @@ -1104,7 +1112,15 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)
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).
Expand All @@ -1117,7 +1133,7 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)

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.
Expand Down
Loading