[PD] Consolidate shared logic into common backend (#25979)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -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]]]:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user