[PD] Consolidate shared logic into common backend (#25979)

Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
Shangming Cai
2026-05-23 10:41:20 +08:00
committed by GitHub
parent c8cea6d4aa
commit a241659d18
5 changed files with 284 additions and 403 deletions
+161 -3
View File
@@ -25,7 +25,10 @@ from sglang.srt.disaggregation.base.conn import (
KVPoll, KVPoll,
KVTransferMetric, KVTransferMetric,
) )
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import (
DisaggregationMode,
filter_kv_indices_for_cp_rank,
)
from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed import get_pp_group, get_world_group
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
@@ -594,6 +597,92 @@ class CommonKVManager(BaseKVManager):
return src_kv_ptrs, sliced_dst return src_kv_ptrs, sliced_dst
def _start_heartbeat_checker_thread(self):
"""Start the heartbeat checker thread for Decode worker."""
def heartbeat_checker():
while True:
time.sleep(self.heartbeat_interval)
with self.connection_lock:
addresses = list(self.prefill_info_table.keys())
for bootstrap_addr in addresses:
session = None
try:
with self.session_pool_lock:
session = self.session_pool[bootstrap_addr]
response = session.get(
f"http://{bootstrap_addr}/health",
timeout=(2, 3),
headers={"Connection": "keep-alive"},
)
if response.status_code == 200:
self.heartbeat_failures[bootstrap_addr] = 0
self._on_heartbeat_success(bootstrap_addr)
else:
logger.info(
f"Attempting to reconnect to {bootstrap_addr}..."
)
self.heartbeat_failures[bootstrap_addr] = (
self.heartbeat_failures.get(bootstrap_addr, 0) + 1
)
with self.session_pool_lock:
if bootstrap_addr in self.session_pool:
del self.session_pool[bootstrap_addr]
except Exception:
logger.info(f"Attempting to reconnect to {bootstrap_addr}...")
self.heartbeat_failures[bootstrap_addr] = (
self.heartbeat_failures.get(bootstrap_addr, 0) + 1
)
if (
self.heartbeat_failures.get(bootstrap_addr, 0)
>= self.max_failures
):
self._handle_node_failure(bootstrap_addr)
with self.session_pool_lock:
if bootstrap_addr in self.session_pool:
del self.session_pool[bootstrap_addr]
threading.Thread(target=heartbeat_checker, daemon=True).start()
def _on_heartbeat_success(self, bootstrap_addr: str):
"""Hook called on successful heartbeat. Override for backend-specific cleanup."""
pass
def _handle_node_failure(self, failed_bootstrap_addr: str):
"""Handle failure of a prefill node."""
with self.connection_lock:
keys_to_remove = [
k for k in self.connection_pool if k.startswith(failed_bootstrap_addr)
]
for k in keys_to_remove:
del self.connection_pool[k]
self.prefill_info_table.pop(failed_bootstrap_addr, None)
possible_affected_rooms = self.addr_to_rooms_tracker.get(
failed_bootstrap_addr, []
)
self.addr_to_rooms_tracker.pop(failed_bootstrap_addr, None)
affected_rooms = []
for room in possible_affected_rooms:
if (
room in self.request_status
and self.check_status(room) != KVPoll.Success
):
self.record_failure(
room,
f"Lost connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr})",
)
self.update_status(room, KVPoll.Failed)
affected_rooms.append(room)
logger.error(
f"Lost connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr}), "
f"{len(affected_rooms)} requests affected"
)
class CommonKVSender(BaseKVSender): class CommonKVSender(BaseKVSender):
def __init__( def __init__(
@@ -614,6 +703,7 @@ class CommonKVSender(BaseKVSender):
self._transfer_num_state_indices = 0 self._transfer_num_state_indices = 0
# inner state # inner state
self.curr_idx = 0 self.curr_idx = 0
self.init_time: Optional[float] = None
if self.kv_mgr.is_dummy_cp_rank: if self.kv_mgr.is_dummy_cp_rank:
# Non-authoritative CP ranks are dummy participants. # Non-authoritative CP ranks are dummy participants.
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput) self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput)
@@ -667,10 +757,10 @@ class CommonKVSender(BaseKVSender):
) )
def pop_decode_prefix_len(self) -> int: def pop_decode_prefix_len(self) -> int:
return 0 return self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, 0)
def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool: def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool:
return num_pages > 0 return num_pages > 0 or last_chunk
def get_transfer_metric(self) -> KVTransferMetric: def get_transfer_metric(self) -> KVTransferMetric:
total_bytes = self._transfer_num_kv_indices * self.kv_mgr.kv_item_lens_sum total_bytes = self._transfer_num_kv_indices * self.kv_mgr.kv_item_lens_sum
@@ -691,6 +781,36 @@ class CommonKVSender(BaseKVSender):
if component_indices is not None: if component_indices is not None:
self._transfer_num_state_indices += len(component_indices) self._transfer_num_state_indices += len(component_indices)
def _prepare_send_indices(
self,
kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List] = None,
) -> Tuple[npt.NDArray[np.int32], slice, bool, bool]:
"""Common pre-processing for send(): index tracking and CP-rank handling.
Returns:
(kv_indices, index_slice, is_last_chunk, should_skip)
If should_skip is True, the caller should return immediately.
"""
index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices))
self.curr_idx += len(kv_indices)
is_last_chunk = self.curr_idx == self.num_kv_indices
if self.kv_mgr.enable_all_cp_ranks_for_transfer:
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
self.kv_mgr,
kv_indices,
index_slice,
)
elif self.kv_mgr.is_dummy_cp_rank:
if not is_last_chunk:
return kv_indices, index_slice, is_last_chunk, True
else:
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success)
return kv_indices, index_slice, is_last_chunk, True
return kv_indices, index_slice, is_last_chunk, False
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
@@ -698,6 +818,25 @@ class CommonKVSender(BaseKVSender):
): ):
pass pass
def _check_bootstrap_timeout(self) -> Optional[KVPoll]:
if self.init_time is None:
return None
elapsed = time.time() - self.init_time
if elapsed < self.kv_mgr.bootstrap_timeout:
return None
logger.warning_once(
"Some requests timed out when bootstrapping, "
"which means prefill instances fail to receive the KV indices from the decode instance of this request. "
"If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT=600' (10 minutes) to relax the timeout condition. "
)
self.kv_mgr.record_failure(
self.bootstrap_room,
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s "
f"in KVPoll.Bootstrapping",
)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return KVPoll.Failed
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
pass pass
@@ -737,6 +876,7 @@ class CommonKVReceiver(BaseKVReceiver):
self.kv_mgr = mgr self.kv_mgr = mgr
self.conclude_state: Optional[KVPoll] = None self.conclude_state: Optional[KVPoll] = None
self.require_staging: bool = False self.require_staging: bool = False
self.init_time: Optional[float] = None
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room) self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping) self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping)
@@ -906,6 +1046,24 @@ class CommonKVReceiver(BaseKVReceiver):
): ):
raise NotImplementedError raise NotImplementedError
def _check_waiting_timeout(self) -> Optional[KVPoll]:
if self.init_time is None:
return None
elapsed = time.time() - self.init_time
if elapsed < self.kv_mgr.waiting_timeout:
return None
logger.warning_once(
"Some requests fail to receive KV Cache transfer done signal after bootstrapping. "
"If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_WAITING_TIMEOUT=600' (10 minutes) to relax the timeout condition. "
)
self.kv_mgr.record_failure(
self.bootstrap_room,
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s "
f"in KVPoll.WaitingForInput",
)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
return KVPoll.Failed
def failure_exception(self): def failure_exception(self):
raise Exception("Fake KVReceiver Exception") raise Exception("Fake KVReceiver Exception")
@@ -1,12 +1,27 @@
import ctypes
import dataclasses
import struct import struct
import threading import threading
from collections import deque from collections import deque
from typing import List, Tuple from typing import List, Optional, Tuple
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
@dataclasses.dataclass
class TransferKVChunk:
"""Work unit for KV cache transfer from prefill to decode."""
room: int
prefill_kv_indices: npt.NDArray[np.int32]
index_slice: slice
is_last_chunk: bool
prefill_aux_index: Optional[int]
state_indices: Optional[List]
chunk_id: Optional[int] = None
def pack_list_of_buffers(buffers: List[bytes]) -> bytes: def pack_list_of_buffers(buffers: List[bytes]) -> bytes:
if not buffers: if not buffers:
return b"" return b""
@@ -59,6 +74,26 @@ class FastQueue:
return self._buf.popleft() return self._buf.popleft()
class AuxDataCodec:
"""Handles serialization and deserialization of auxiliary data buffers."""
@staticmethod
def serialize_data_from_buffer(src_addr, data_length):
"""Serialize data from memory buffer to bytes."""
buffer = (ctypes.c_byte * data_length).from_address(src_addr)
return bytes(buffer)
@staticmethod
def deserialize_data_to_buffer(kv_args, buffer_index, aux_index, data):
"""Deserialize bytes into target memory buffer."""
dst_aux_ptr = kv_args.aux_data_ptrs[buffer_index]
item_len = kv_args.aux_item_lens[buffer_index]
dst_addr = dst_aux_ptr + item_len * aux_index
buffer = (ctypes.c_byte * len(data)).from_address(dst_addr)
buffer[:] = data
return
def group_concurrent_contiguous( def group_concurrent_contiguous(
src_indices: npt.NDArray[np.int32], dst_indices: npt.NDArray[np.int32] src_indices: npt.NDArray[np.int32], dst_indices: npt.NDArray[np.int32]
) -> Tuple[List[npt.NDArray[np.int32]], List[npt.NDArray[np.int32]]]: ) -> Tuple[List[npt.NDArray[np.int32]], List[npt.NDArray[np.int32]]]:
+30 -159
View File
@@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
import concurrent.futures import concurrent.futures
import ctypes
import dataclasses import dataclasses
import logging import logging
import os import os
@@ -29,7 +28,9 @@ from sglang.srt.disaggregation.common.staging_handler import (
StagingTransferInfo, StagingTransferInfo,
) )
from sglang.srt.disaggregation.common.utils import ( from sglang.srt.disaggregation.common.utils import (
AuxDataCodec,
FastQueue, FastQueue,
TransferKVChunk,
group_concurrent_contiguous, group_concurrent_contiguous,
pack_int_lists, pack_int_lists,
unpack_int_lists, unpack_int_lists,
@@ -37,10 +38,7 @@ from sglang.srt.disaggregation.common.utils import (
from sglang.srt.disaggregation.mooncake.utils import ( from sglang.srt.disaggregation.mooncake.utils import (
check_mooncake_custom_mem_pool_enabled, check_mooncake_custom_mem_pool_enabled,
) )
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import DisaggregationMode
DisaggregationMode,
filter_kv_indices_for_cp_rank,
)
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -64,17 +62,6 @@ class KVTransferError(Exception):
return f"KVTransferError(bootstrap_room={self.bootstrap_room}): {self.failure_reason}" return f"KVTransferError(bootstrap_room={self.bootstrap_room}): {self.failure_reason}"
# prefill
@dataclasses.dataclass
class TransferKVChunk:
room: int
prefill_kv_indices: npt.NDArray[np.int32]
index_slice: slice
is_last_chunk: bool
prefill_aux_index: Optional[int]
state_indices: Optional[List]
# decode # decode
@dataclasses.dataclass @dataclasses.dataclass
class TransferInfo: class TransferInfo:
@@ -162,26 +149,6 @@ class KVArgsRegisterInfo:
) )
class AuxDataCodec:
"""Handles serialization and deserialization of auxiliary data buffers"""
@staticmethod
def serialize_data_from_buffer(src_addr, data_length):
"""Serialize data from memory buffer to bytes"""
buffer = (ctypes.c_byte * data_length).from_address(src_addr)
return bytes(buffer)
@staticmethod
def deserialize_data_to_buffer(kv_args, buffer_index, aux_index, data):
"""Deserialize bytes into target memory buffer"""
dst_aux_ptr = kv_args.aux_data_ptrs[buffer_index]
item_len = kv_args.aux_item_lens[buffer_index]
dst_addr = dst_aux_ptr + item_len * aux_index
buffer = (ctypes.c_byte * len(data)).from_address(dst_addr)
buffer[:] = data
return
class MooncakeKVManager(CommonKVManager): class MooncakeKVManager(CommonKVManager):
AUX_DATA_HEADER = b"AUX_DATA" AUX_DATA_HEADER = b"AUX_DATA"
@@ -1478,62 +1445,8 @@ class MooncakeKVManager(CommonKVManager):
) )
self.update_status(bootstrap_room, status) self.update_status(bootstrap_room, status)
def heartbeat_checker():
while True:
time.sleep(self.heartbeat_interval)
with self.connection_lock:
addresses = list(self.prefill_info_table.keys())
for bootstrap_addr in addresses:
session = None
try:
with self.session_pool_lock:
session = self.session_pool[bootstrap_addr]
response = session.get(
f"http://{bootstrap_addr}/health",
timeout=(2, 3),
headers={"Connection": "keep-alive"},
)
if response.status_code == 200:
self.heartbeat_failures[bootstrap_addr] = 0
current_rooms = self.addr_to_rooms_tracker[
bootstrap_addr
].copy()
for bootstrap_room in current_rooms:
# Remove KVPoll.Success requests from the tracker
if bootstrap_room not in self.request_status:
self.addr_to_rooms_tracker[bootstrap_addr].discard(
bootstrap_room
)
else:
logger.info(
f"Attempting to reconnect to {bootstrap_addr}..."
)
self.heartbeat_failures[bootstrap_addr] = (
self.heartbeat_failures.get(bootstrap_addr, 0) + 1
)
with self.session_pool_lock:
if bootstrap_addr in self.session_pool:
del self.session_pool[bootstrap_addr]
except Exception:
logger.info(f"Attempting to reconnect to {bootstrap_addr}...")
self.heartbeat_failures[bootstrap_addr] = (
self.heartbeat_failures.get(bootstrap_addr, 0) + 1
)
if (
self.heartbeat_failures.get(bootstrap_addr, 0)
>= self.max_failures
):
self._handle_node_failure(bootstrap_addr)
with self.session_pool_lock:
if bootstrap_addr in self.session_pool:
del self.session_pool[bootstrap_addr]
threading.Thread(target=decode_thread).start() threading.Thread(target=decode_thread).start()
threading.Thread(target=heartbeat_checker).start() self._start_heartbeat_checker_thread()
def add_transfer_request( def add_transfer_request(
self, self,
@@ -1583,6 +1496,13 @@ class MooncakeKVManager(CommonKVManager):
def get_session_id(self): def get_session_id(self):
return self.engine.get_session_id() return self.engine.get_session_id()
def _on_heartbeat_success(self, bootstrap_addr: str):
current_rooms = self.addr_to_rooms_tracker[bootstrap_addr].copy()
for bootstrap_room in current_rooms:
# Remove KVPoll.Success requests from the tracker
if bootstrap_room not in self.request_status:
self.addr_to_rooms_tracker[bootstrap_addr].discard(bootstrap_room)
def _run_one_probe_pass(self) -> None: def _run_one_probe_pass(self) -> None:
with self.session_lock: with self.session_lock:
snapshot = list(self.failed_sessions) snapshot = list(self.failed_sessions)
@@ -1666,34 +1586,16 @@ class MooncakeKVSender(CommonKVSender):
self.conclude_state = None self.conclude_state = None
self.init_time = time.time() self.init_time = time.time()
def pop_decode_prefix_len(self) -> int:
return self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, 0)
def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool:
return num_pages > 0 or last_chunk
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List] = None, state_indices: Optional[List] = None,
): ):
index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) kv_indices, index_slice, is_last_chunk, should_skip = (
self.curr_idx += len(kv_indices) self._prepare_send_indices(kv_indices, state_indices)
is_last_chunk = self.curr_idx == self.num_kv_indices )
if should_skip:
# Special handling for cp return
if self.kv_mgr.enable_all_cp_ranks_for_transfer:
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
self.kv_mgr,
kv_indices,
index_slice,
)
elif self.kv_mgr.is_dummy_cp_rank:
if not is_last_chunk:
return
else:
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success)
return
if not is_last_chunk: if not is_last_chunk:
self.kv_mgr.add_transfer_request( self.kv_mgr.add_transfer_request(
@@ -1719,21 +1621,9 @@ class MooncakeKVSender(CommonKVSender):
if status in (KVPoll.Success, KVPoll.Failed): if status in (KVPoll.Success, KVPoll.Failed):
self.conclude_state = status self.conclude_state = status
elif status == KVPoll.Bootstrapping: elif status == KVPoll.Bootstrapping:
if self.init_time is not None: timeout_result = self._check_bootstrap_timeout()
now = time.time() if timeout_result is not None:
elapsed = now - self.init_time return timeout_result
if elapsed >= self.kv_mgr.bootstrap_timeout:
logger.warning_once(
"Some requests timed out when bootstrapping, "
"which means prefill instances fail to receive the KV indices from the decode instance of this request. "
"If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT=600' (10 minutes) to relax the timeout condition. "
)
self.kv_mgr.record_failure(
self.bootstrap_room,
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s in KVPoll.Bootstrapping",
)
self.conclude_state = KVPoll.Failed
return KVPoll.Failed
return status return status
else: else:
@@ -1819,12 +1709,6 @@ class MooncakeKVReceiver(CommonKVReceiver):
] ]
) )
def init(
self,
prefill_dp_rank: int,
):
super().init(prefill_dp_rank)
def send_metadata( def send_metadata(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
@@ -1874,33 +1758,20 @@ class MooncakeKVReceiver(CommonKVReceiver):
self.init_time = time.time() self.init_time = time.time()
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
if self.conclude_state is None: if self.conclude_state is not None:
status = self.kv_mgr.check_status(self.bootstrap_room)
if status in (KVPoll.Success, KVPoll.Failed):
self.conclude_state = status
elif status == KVPoll.WaitingForInput:
if self.init_time is not None:
now = time.time()
elapsed = now - self.init_time
if elapsed >= self.kv_mgr.waiting_timeout:
logger.warning_once(
"Some requests fail to receive KV Cache transfer done signal after bootstrapping. "
"If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_WAITING_TIMEOUT=600' (10 minutes) to relax the timeout condition. "
)
self.kv_mgr.record_failure(
self.bootstrap_room,
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s in KVPoll.WaitingForInput",
)
self.conclude_state = KVPoll.Failed
return KVPoll.Failed
return status
else:
return self.conclude_state return self.conclude_state
status = self.kv_mgr.check_status(self.bootstrap_room)
if status in (KVPoll.Success, KVPoll.Failed):
self.conclude_state = status
elif status == KVPoll.WaitingForInput:
timeout_result = self._check_waiting_timeout()
if timeout_result is not None:
return timeout_result
return status
def failure_exception(self): def failure_exception(self):
# Explicitly set the status to failure since this request has failed in another rank
if self.conclude_state is None: if self.conclude_state is None:
self.conclude_state = KVPoll.Failed self.conclude_state = KVPoll.Failed
+27 -64
View File
@@ -1,6 +1,5 @@
from __future__ import annotations from __future__ import annotations
import ctypes
import dataclasses import dataclasses
import logging import logging
import os import os
@@ -33,11 +32,11 @@ from sglang.srt.disaggregation.common.conn import (
CommonKVReceiver, CommonKVReceiver,
CommonKVSender, CommonKVSender,
) )
from sglang.srt.disaggregation.common.utils import group_concurrent_contiguous from sglang.srt.disaggregation.common.utils import (
from sglang.srt.disaggregation.utils import ( AuxDataCodec,
DisaggregationMode, group_concurrent_contiguous,
filter_kv_indices_for_cp_rank,
) )
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import get_int_env_var from sglang.srt.utils.common import get_int_env_var
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
@@ -176,22 +175,6 @@ class KVArgsRegisterInfo:
) )
class AuxDataCodec:
@staticmethod
def serialize_data_from_buffer(src_addr, data_length):
buffer = (ctypes.c_byte * data_length).from_address(src_addr)
return bytes(buffer)
@staticmethod
def deserialize_data_to_buffer(kv_args, buffer_index, aux_index, data):
dst_aux_ptr = kv_args.aux_data_ptrs[buffer_index]
item_len = kv_args.aux_item_lens[buffer_index]
dst_addr = dst_aux_ptr + item_len * aux_index
buffer = (ctypes.c_byte * len(data)).from_address(dst_addr)
buffer[:] = data
return
@dataclasses.dataclass @dataclasses.dataclass
class TPSliceConfig: class TPSliceConfig:
page_size: int page_size: int
@@ -1132,7 +1115,7 @@ class MoriKVManager(CommonKVManager):
bootstrap_room: int, bootstrap_room: int,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
index_slice: slice, index_slice: slice,
is_last: bool, is_last_chunk: bool,
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[npt.NDArray[np.int32]] = None, state_indices: Optional[npt.NDArray[np.int32]] = None,
) -> Tuple[List[TransferStatus], Optional[List[TransferInfo]]]: ) -> Tuple[List[TransferStatus], Optional[List[TransferInfo]]]:
@@ -1163,7 +1146,7 @@ class MoriKVManager(CommonKVManager):
self.update_status(bootstrap_room, KVPoll.Failed) self.update_status(bootstrap_room, KVPoll.Failed)
return [], list(transfer_infos.values()) return [], list(transfer_infos.values())
targets.append(TransferTarget(info=info, peer_info=peer_info)) targets.append(TransferTarget(info=info, peer_info=peer_info))
if is_last: if is_last_chunk:
target_infos_snapshot = list(transfer_infos.values()) target_infos_snapshot = list(transfer_infos.values())
result_statuses: List[TransferStatus] = [] result_statuses: List[TransferStatus] = []
@@ -1179,7 +1162,7 @@ class MoriKVManager(CommonKVManager):
) )
if ( if (
is_last is_last_chunk
and state_indices is not None and state_indices is not None
and not info.is_dummy and not info.is_dummy
and self.state_mem_descs and self.state_mem_descs
@@ -1191,7 +1174,7 @@ class MoriKVManager(CommonKVManager):
) )
if ( if (
is_last is_last_chunk
and aux_index is not None and aux_index is not None
and info.dst_aux_index >= 0 and info.dst_aux_index >= 0
and self.pp_group.is_last_rank and self.pp_group.is_last_rank
@@ -1212,7 +1195,7 @@ class MoriKVManager(CommonKVManager):
) )
return result_statuses, target_infos_snapshot return result_statuses, target_infos_snapshot
if is_last: if is_last_chunk:
with self.transfer_lock: with self.transfer_lock:
# Keep transfer_infos alive until sender.clear() so abort/failure # Keep transfer_infos alive until sender.clear() so abort/failure
# paths can still recover notification targets after posting. # paths can still recover notification targets after posting.
@@ -1243,38 +1226,28 @@ class MoriKVSender(CommonKVSender):
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List] = None, state_indices: Optional[List] = None,
): ):
index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) kv_indices, index_slice, is_last_chunk, should_skip = (
self.curr_idx += len(kv_indices) self._prepare_send_indices(kv_indices, state_indices)
is_last = self.curr_idx == self.num_kv_indices )
if should_skip:
return
# Special handling for cp normalized_state = (
if self.kv_mgr.enable_all_cp_ranks_for_transfer: _normalize_state_indices(state_indices) if is_last_chunk else None
kv_indices, index_slice = filter_kv_indices_for_cp_rank( )
self.kv_mgr,
kv_indices,
index_slice,
)
elif self.kv_mgr.is_dummy_cp_rank:
if not is_last:
return
else:
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success)
return
normalized_state = _normalize_state_indices(state_indices) if is_last else None
statuses, infos = self.kv_mgr.add_transfer_request( statuses, infos = self.kv_mgr.add_transfer_request(
self.bootstrap_room, self.bootstrap_room,
kv_indices, kv_indices,
index_slice, index_slice,
is_last, is_last_chunk,
aux_index=self.aux_index if is_last else None, aux_index=self.aux_index if is_last_chunk else None,
state_indices=normalized_state, state_indices=normalized_state,
) )
self.transfer_statuses.extend(statuses) self.transfer_statuses.extend(statuses)
self._record_transfer_indices(kv_indices, None) self._record_transfer_indices(kv_indices, None)
if infos is not None: if infos is not None:
self.pending_infos = infos self.pending_infos = infos
if is_last: if is_last_chunk:
self.sent_last_chunk = True self.sent_last_chunk = True
self._maybe_finalize_if_room_failed() self._maybe_finalize_if_room_failed()
@@ -1295,15 +1268,9 @@ class MoriKVSender(CommonKVSender):
status = self.kv_mgr.check_status(self.bootstrap_room) status = self.kv_mgr.check_status(self.bootstrap_room)
if status == KVPoll.Bootstrapping: if status == KVPoll.Bootstrapping:
elapsed = time.time() - self.init_time timeout_result = self._check_bootstrap_timeout()
if elapsed >= self.kv_mgr.bootstrap_timeout: if timeout_result is not None:
reason = ( self._finalize_failure()
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s "
"waiting for decode handshake"
)
self.kv_mgr.record_failure(self.bootstrap_room, reason)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
self._finalize_failure(reason)
return KVPoll.Failed return KVPoll.Failed
return status return status
@@ -1499,14 +1466,10 @@ class MoriKVReceiver(CommonKVReceiver):
self.conclude_state = status self.conclude_state = status
return status return status
if status == KVPoll.WaitingForInput and self.init_time is not None: if status == KVPoll.WaitingForInput:
elapsed = time.time() - self.init_time timeout_result = self._check_waiting_timeout()
if elapsed >= self.kv_mgr.waiting_timeout: if timeout_result is not None:
reason = f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s waiting for KV transfer" return timeout_result
self.kv_mgr.record_failure(self.bootstrap_room, reason)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
self.conclude_state = KVPoll.Failed
return KVPoll.Failed
return status return status
+30 -176
View File
@@ -26,14 +26,12 @@ from sglang.srt.disaggregation.common.conn import (
from sglang.srt.disaggregation.common.staging_handler import StagingRegisterInfo from sglang.srt.disaggregation.common.staging_handler import StagingRegisterInfo
from sglang.srt.disaggregation.common.utils import ( from sglang.srt.disaggregation.common.utils import (
FastQueue, FastQueue,
TransferKVChunk,
group_concurrent_contiguous, group_concurrent_contiguous,
pack_int_lists, pack_int_lists,
unpack_int_lists, unpack_int_lists,
) )
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import DisaggregationMode
DisaggregationMode,
filter_kv_indices_for_cp_rank,
)
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -104,17 +102,6 @@ class TransferInfo:
) )
@dataclasses.dataclass
class TransferKVChunk:
room: int
prefill_kv_indices: npt.NDArray[np.int32]
index_slice: slice
is_last: bool
chunk_id: int
prefill_aux_index: Optional[int]
state_indices: Optional[List]
@dataclasses.dataclass @dataclasses.dataclass
class KVArgsRegisterInfo: class KVArgsRegisterInfo:
"""Contains base pointers and other info which only needs to be sent once by KVReceiver. Received by prefill bootstrap thread.""" """Contains base pointers and other info which only needs to be sent once by KVReceiver. Received by prefill bootstrap thread."""
@@ -176,7 +163,7 @@ class TransferStatus:
received_kvs_per_pp: Dict[int, Set[int]] = dataclasses.field( received_kvs_per_pp: Dict[int, Set[int]] = dataclasses.field(
default_factory=lambda: defaultdict(set) default_factory=lambda: defaultdict(set)
) )
# Expected chunk count per pp_rank (set when is_last=True): {pp_rank: expected_count} # Expected chunk count per pp_rank (set when is_last_chunk=True): {pp_rank: expected_count}
expected_kvs_per_pp: Dict[int, int] = dataclasses.field(default_factory=dict) expected_kvs_per_pp: Dict[int, int] = dataclasses.field(default_factory=dict)
# Number of PP ranks expected to send data. # Number of PP ranks expected to send data.
num_pp_ranks_expected: Optional[int] = None num_pp_ranks_expected: Optional[int] = None
@@ -186,12 +173,8 @@ class TransferStatus:
received_state_per_pp: Set[int] = dataclasses.field(default_factory=set) received_state_per_pp: Set[int] = dataclasses.field(default_factory=set)
# Whether state data is expected (set based on state_type). # Whether state data is expected (set based on state_type).
expects_state: bool = False expects_state: bool = False
# Mark as failed
is_failure: bool = False
def is_done(self): def is_done(self):
if self.is_failure:
return True
if self.num_pp_ranks_expected is None or not self.received_aux: if self.num_pp_ranks_expected is None or not self.received_aux:
return False return False
# If state data is expected, check all PP ranks have sent it # If state data is expected, check all PP ranks have sent it
@@ -209,9 +192,6 @@ class TransferStatus:
return False return False
return True return True
def is_failed(self):
return self.is_failure
class NixlKVManager(CommonKVManager): class NixlKVManager(CommonKVManager):
def __init__( def __init__(
@@ -471,92 +451,6 @@ class NixlKVManager(CommonKVManager):
) )
self._staging_ctx.prefetched_rooms.add(room) self._staging_ctx.prefetched_rooms.add(room)
def _start_heartbeat_checker_thread(self):
"""
Start the heartbeat checker thread for Decode worker.
TODO (smor): unite nixl heartbeat checker with mooncake's.
"""
def heartbeat_checker():
while True:
time.sleep(self.heartbeat_interval)
with self.connection_lock:
addresses = list(self.prefill_info_table.keys())
for bootstrap_addr in addresses:
session = None
try:
with self.session_pool_lock:
session = self.session_pool[bootstrap_addr]
response = session.get(
f"http://{bootstrap_addr}/health",
timeout=(2, 3),
headers={"Connection": "keep-alive"},
)
if response.status_code == 200:
self.heartbeat_failures[bootstrap_addr] = 0
else:
logger.info(
f"Attempting to reconnect to {bootstrap_addr}..."
)
self.heartbeat_failures[bootstrap_addr] = (
self.heartbeat_failures.get(bootstrap_addr, 0) + 1
)
with self.session_pool_lock:
if bootstrap_addr in self.session_pool:
del self.session_pool[bootstrap_addr]
except Exception:
logger.info(f"Attempting to reconnect to {bootstrap_addr}...")
self.heartbeat_failures[bootstrap_addr] = (
self.heartbeat_failures.get(bootstrap_addr, 0) + 1
)
if (
self.heartbeat_failures.get(bootstrap_addr, 0)
>= self.max_failures
):
self._handle_node_failure(bootstrap_addr)
with self.session_pool_lock:
if bootstrap_addr in self.session_pool:
del self.session_pool[bootstrap_addr]
threading.Thread(target=heartbeat_checker, daemon=True).start()
def _handle_node_failure(self, failed_bootstrap_addr):
"""Handle failure of a prefill node."""
with self.connection_lock:
keys_to_remove = [
k for k in self.connection_pool if k.startswith(failed_bootstrap_addr)
]
for k in keys_to_remove:
del self.connection_pool[k]
self.prefill_info_table.pop(failed_bootstrap_addr, None)
possible_affected_rooms = self.addr_to_rooms_tracker.get(
failed_bootstrap_addr, []
)
self.addr_to_rooms_tracker.pop(failed_bootstrap_addr, None)
# Mark all pending transfers associated with the failed node as failed
affected_rooms = []
for room in possible_affected_rooms:
if (
room in self.transfer_statuses
and not self.transfer_statuses[room].is_done()
):
# Mark the transfer as failed
self.transfer_statuses[room].is_failure = True
affected_rooms.append(room)
logger.error(
f"Lost connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr}), "
f"{len(affected_rooms)} transfers affected"
)
for room in possible_affected_rooms:
logger.error(f"Let room {room} be failed due to prefill down")
self.update_status(room, KVPoll.Failed)
def check_status(self, bootstrap_room: int): def check_status(self, bootstrap_room: int):
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput) return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
@@ -606,7 +500,7 @@ class NixlKVManager(CommonKVManager):
# Skip KV RDMA transfer when there are no pages to send # Skip KV RDMA transfer when there are no pages to send
# (e.g., decode-side radix cache matched the entire prefix). # (e.g., decode-side radix cache matched the entire prefix).
# Aux data is still sent below when is_last=True. # Aux data is still sent below when is_last_chunk=True.
if len(kv_chunk.prefill_kv_indices) > 0: if len(kv_chunk.prefill_kv_indices) > 0:
chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice] chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice]
@@ -659,7 +553,7 @@ class NixlKVManager(CommonKVManager):
if kv_xfer_handle is None: if kv_xfer_handle is None:
notif = ( notif = (
f"{req.room}_kv_{kv_chunk.chunk_id}" f"{req.room}_kv_{kv_chunk.chunk_id}"
f"_{int(kv_chunk.is_last)}_{self.kv_args.engine_rank}" f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
) )
if self.is_mla_backend or ( if self.is_mla_backend or (
decode_tp_size == self.attn_tp_size decode_tp_size == self.attn_tp_size
@@ -688,7 +582,7 @@ class NixlKVManager(CommonKVManager):
handles.append(kv_xfer_handle) handles.append(kv_xfer_handle)
if kv_chunk.is_last: if kv_chunk.is_last_chunk:
dst_info = self.decode_kv_args_table[req.agent_name] dst_info = self.decode_kv_args_table[req.agent_name]
if kv_chunk.state_indices: if kv_chunk.state_indices:
state_xfer_handles = self.maybe_send_extra( state_xfer_handles = self.maybe_send_extra(
@@ -739,7 +633,7 @@ class NixlKVManager(CommonKVManager):
break break
time.sleep(0) time.sleep(0)
if kv_chunk.is_last: if kv_chunk.is_last_chunk:
self.update_status(room, KVPoll.Success) self.update_status(room, KVPoll.Success)
# Drop per-room state on Success (parity with mooncake # Drop per-room state on Success (parity with mooncake
# transfer_worker; staging prefetch sets are NIXL-only). # transfer_worker; staging prefetch sets are NIXL-only).
@@ -1265,7 +1159,7 @@ class NixlKVManager(CommonKVManager):
return (None, True) return (None, True)
notif_tag = ( notif_tag = (
f"{req.room}_stg_{kv_chunk.chunk_id}_{int(kv_chunk.is_last)}" f"{req.room}_stg_{kv_chunk.chunk_id}_{int(kv_chunk.is_last_chunk)}"
f"_{self.kv_args.engine_rank}_{chunk_idx}" f"_{self.kv_args.engine_rank}_{chunk_idx}"
f"_{page_start}_{num_pages}_{req.agent_name}" f"_{page_start}_{num_pages}_{req.agent_name}"
) )
@@ -1574,13 +1468,13 @@ class NixlKVManager(CommonKVManager):
bootstrap_room: int, bootstrap_room: int,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
index_slice: slice, index_slice: slice,
is_last: bool, is_last_chunk: bool,
chunk_id: int, chunk_id: int,
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List] = None, state_indices: Optional[List] = None,
): ):
assert self.disaggregation_mode == DisaggregationMode.PREFILL assert self.disaggregation_mode == DisaggregationMode.PREFILL
assert not is_last or (is_last and aux_index is not None) assert not is_last_chunk or (is_last_chunk and aux_index is not None)
# Prefetch STAGING_REQ to decode before enqueueing so decode has # Prefetch STAGING_REQ to decode before enqueueing so decode has
# already allocated staging by the time the worker picks up the # already allocated staging by the time the worker picks up the
@@ -1601,7 +1495,7 @@ class NixlKVManager(CommonKVManager):
room=bootstrap_room, room=bootstrap_room,
prefill_kv_indices=kv_indices, prefill_kv_indices=kv_indices,
index_slice=index_slice, index_slice=index_slice,
is_last=is_last, is_last_chunk=is_last_chunk,
chunk_id=chunk_id, chunk_id=chunk_id,
prefill_aux_index=aux_index, prefill_aux_index=aux_index,
state_indices=state_indices, state_indices=state_indices,
@@ -1628,9 +1522,9 @@ class NixlKVManager(CommonKVManager):
tag = components[1] tag = components[1]
if tag == "kv": if tag == "kv":
chunk_id = int(components[2]) chunk_id = int(components[2])
is_last = bool(int(components[3])) is_last_chunk = bool(int(components[3]))
pp_rank = int(components[4]) if len(components) > 4 else 0 pp_rank = int(components[4]) if len(components) > 4 else 0
self._track_kv_arrival(room, chunk_id, is_last, pp_rank) self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
elif tag == "stg": elif tag == "stg":
self._handle_stg_notification(components, room) self._handle_stg_notification(components, room)
elif tag == "aux": elif tag == "aux":
@@ -1647,13 +1541,13 @@ class NixlKVManager(CommonKVManager):
Format: {room}_stg_{chunk_id}_{is_last}_{pp_rank}_{chunk_idx}_{page_start}_{num_pages}_{agent_name} Format: {room}_stg_{chunk_id}_{is_last}_{pp_rank}_{chunk_idx}_{page_start}_{num_pages}_{agent_name}
""" """
chunk_id = int(components[2]) chunk_id = int(components[2])
is_last = bool(int(components[3])) is_last_chunk = bool(int(components[3]))
pp_rank = int(components[4]) pp_rank = int(components[4])
chunk_idx = int(components[5]) chunk_idx = int(components[5])
page_start = int(components[6]) page_start = int(components[6])
num_pages = int(components[7]) num_pages = int(components[7])
agent_name = components[8] if len(components) > 8 else "" agent_name = components[8] if len(components) > 8 else ""
self._track_kv_arrival(room, chunk_id, is_last, pp_rank) self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
self._handle_staging_chunk_arrived( self._handle_staging_chunk_arrived(
room, chunk_idx, page_start, num_pages, agent_name room, chunk_idx, page_start, num_pages, agent_name
) )
@@ -1683,10 +1577,12 @@ class NixlKVManager(CommonKVManager):
): ):
self._maybe_submit_last_scatter(room) self._maybe_submit_last_scatter(room)
def _track_kv_arrival(self, room: int, chunk_id: int, is_last: bool, pp_rank: int): def _track_kv_arrival(
self, room: int, chunk_id: int, is_last_chunk: bool, pp_rank: int
):
"""Update transfer status tracking for a kv chunk arrival.""" """Update transfer status tracking for a kv chunk arrival."""
self.transfer_statuses[room].received_kvs_per_pp[pp_rank].add(chunk_id) self.transfer_statuses[room].received_kvs_per_pp[pp_rank].add(chunk_id)
if is_last: if is_last_chunk:
self.transfer_statuses[room].expected_kvs_per_pp[pp_rank] = chunk_id + 1 self.transfer_statuses[room].expected_kvs_per_pp[pp_rank] = chunk_id + 1
if self.transfer_statuses[room].num_pp_ranks_expected is None: if self.transfer_statuses[room].num_pp_ranks_expected is None:
self.transfer_statuses[room].num_pp_ranks_expected = ( self.transfer_statuses[room].num_pp_ranks_expected = (
@@ -1827,12 +1723,6 @@ class NixlKVSender(CommonKVSender):
self._send_error: Optional[Exception] = None self._send_error: Optional[Exception] = None
self._transfer_start_time: Optional[float] = None self._transfer_start_time: Optional[float] = None
def pop_decode_prefix_len(self) -> int:
return self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, 0)
def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool:
return num_pages > 0 or last_chunk
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
@@ -1841,23 +1731,11 @@ class NixlKVSender(CommonKVSender):
if self._send_failed: if self._send_failed:
return return
index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) kv_indices, index_slice, is_last_chunk, should_skip = (
self.curr_idx += len(kv_indices) self._prepare_send_indices(kv_indices, state_indices)
is_last = self.curr_idx == self.num_kv_indices )
if should_skip:
# Special handling for cp return
if self.kv_mgr.enable_all_cp_ranks_for_transfer:
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
self.kv_mgr,
kv_indices,
index_slice,
)
elif self.kv_mgr.is_dummy_cp_rank:
if not is_last:
return
else:
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success)
return
if self._transfer_start_time is None and ( if self._transfer_start_time is None and (
len(kv_indices) > 0 or state_indices is not None len(kv_indices) > 0 or state_indices is not None
@@ -1868,14 +1746,14 @@ class NixlKVSender(CommonKVSender):
self.bootstrap_room, self.bootstrap_room,
kv_indices, kv_indices,
index_slice, index_slice,
is_last, is_last_chunk,
self.chunk_id, self.chunk_id,
self.aux_index, self.aux_index,
state_indices, state_indices,
) )
self._record_transfer_indices(kv_indices, state_indices) self._record_transfer_indices(kv_indices, state_indices)
self.chunk_id += 1 self.chunk_id += 1
if is_last: if is_last_chunk:
self.has_sent = True self.has_sent = True
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
@@ -1892,9 +1770,6 @@ class NixlKVSender(CommonKVSender):
) )
return status return status
def clear(self):
super().clear()
def failure_exception(self): def failure_exception(self):
if self._send_error is not None: if self._send_error is not None:
raise self._send_error raise self._send_error
@@ -1915,12 +1790,6 @@ class NixlKVReceiver(CommonKVReceiver):
super().__init__(mgr, bootstrap_addr, bootstrap_room) super().__init__(mgr, bootstrap_addr, bootstrap_room)
self.init_time = None self.init_time = None
def init(
self,
prefill_dp_rank: int,
):
super().init(prefill_dp_rank)
def send_metadata( def send_metadata(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
@@ -1997,31 +1866,16 @@ class NixlKVReceiver(CommonKVReceiver):
if not self.started_transfer: if not self.started_transfer:
return status return status
now = time.time() timeout_result = self._check_waiting_timeout()
elapsed = now - self.init_time if timeout_result is not None:
return timeout_result
if elapsed >= self.kv_mgr.waiting_timeout:
logger.error(f"Request {self.bootstrap_room} waiting_timeout")
self.kv_mgr.record_failure(
self.bootstrap_room,
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s in KVPoll.WaitingForInput",
)
self.conclude_state = KVPoll.Failed
return KVPoll.Failed
self.kv_mgr.update_transfer_status() self.kv_mgr.update_transfer_status()
if self.kv_mgr.check_transfer_done(self.bootstrap_room): # type: ignore if self.kv_mgr.check_transfer_done(self.bootstrap_room): # type: ignore
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].discard( self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].discard(
self.bootstrap_room self.bootstrap_room
) )
# Check if the transfer failed self.conclude_state = KVPoll.Success
if self.kv_mgr.transfer_statuses[self.bootstrap_room].is_failed():
self.conclude_state = KVPoll.Failed
logger.error(
f"Transfer for room {self.bootstrap_room} failed due to node failure"
)
else:
self.conclude_state = KVPoll.Success
del self.kv_mgr.transfer_statuses[self.bootstrap_room] del self.kv_mgr.transfer_statuses[self.bootstrap_room]
return self.conclude_state # type: ignore return self.conclude_state # type: ignore
return KVPoll.WaitingForInput # type: ignore return KVPoll.WaitingForInput # type: ignore