feat: support Kimi Linear PD disaggregation with DCP (#32837)
Co-authored-by: Yangmin Li <yangminl@nvidia.com>
This commit is contained in:
@@ -26,6 +26,31 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
"with MC_FORCE_TCP=1 (TCP transport, no RDMA)"
|
||||
)
|
||||
|
||||
if server_args.disaggregation_mode == "prefill" and server_args.dcp_size > 1:
|
||||
logger.warning(
|
||||
"DCP on a PD prefill server is supported when prefill and decode "
|
||||
"use the same DCP layout, but it usually adds communication "
|
||||
"overhead without improving prefill performance."
|
||||
)
|
||||
|
||||
if server_args.disaggregation_mode == "decode" and server_args.dcp_size > 1:
|
||||
if server_args.disaggregation_transfer_backend not in ("mooncake", "nixl"):
|
||||
raise ValueError(
|
||||
"PD decode DCP requires --disaggregation-transfer-backend "
|
||||
"mooncake or nixl, got "
|
||||
f"{server_args.disaggregation_transfer_backend!r}."
|
||||
)
|
||||
if server_args.disaggregation_decode_enable_radix_cache:
|
||||
raise ValueError(
|
||||
"PD decode DCP currently requires chunk cache; "
|
||||
"--disaggregation-decode-enable-radix-cache is not supported."
|
||||
)
|
||||
if server_args.enable_hierarchical_cache:
|
||||
raise ValueError(
|
||||
"PD decode DCP currently requires chunk cache; "
|
||||
"--enable-hierarchical-cache is not supported."
|
||||
)
|
||||
|
||||
if server_args.disaggregation_mode == "decode":
|
||||
if server_args.disaggregation_decode_enable_radix_cache:
|
||||
if server_args.enable_hisparse:
|
||||
|
||||
@@ -135,6 +135,7 @@ class BaseKVSender(ABC):
|
||||
self,
|
||||
kv_indices: npt.NDArray[np.int32],
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
"""
|
||||
Send the kv cache at the given kv indices and the extra cache/state at the given indices to the decoder server.
|
||||
|
||||
@@ -163,10 +163,13 @@ class CommonKVManager(BaseKVManager):
|
||||
self.bootstrap_host = server_args.host
|
||||
self.bootstrap_port = server_args.disaggregation_bootstrap_port
|
||||
self.dist_init_addr = server_args.dist_init_addr
|
||||
self.attn_tp_size = get_parallel().attn_tp_size
|
||||
self.attn_tp_rank = get_parallel().attn_tp_rank
|
||||
self.attn_cp_size = get_parallel().attn_cp_size
|
||||
self.attn_cp_rank = get_parallel().attn_cp_rank
|
||||
parallel = get_parallel()
|
||||
self.attn_tp_size = parallel.attn_tp_size
|
||||
self.attn_tp_rank = parallel.attn_tp_rank
|
||||
self.attn_cp_size = parallel.attn_cp_size
|
||||
self.attn_cp_rank = parallel.attn_cp_rank
|
||||
self.dcp_size = server_args.dcp_size
|
||||
self.dcp_rank = parallel.dcp_rank if self.dcp_size > 1 else 0
|
||||
self.attn_dp_size = get_attention_dp_size()
|
||||
self.attn_dp_rank = get_attention_dp_rank()
|
||||
self.system_dp_size = (
|
||||
@@ -268,6 +271,39 @@ class CommonKVManager(BaseKVManager):
|
||||
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
||||
)
|
||||
|
||||
def requires_dcp_relayout(self, dst_dcp_size: int, dst_dcp_rank: int) -> bool:
|
||||
if self.dcp_size == dst_dcp_size:
|
||||
if self.dcp_rank != dst_dcp_rank:
|
||||
raise RuntimeError(
|
||||
"PD peers must connect matching DCP ranks, got "
|
||||
f"prefill={self.dcp_rank}, decode={dst_dcp_rank}"
|
||||
)
|
||||
return False
|
||||
|
||||
if (
|
||||
self.dcp_size == 1
|
||||
and dst_dcp_size > 1
|
||||
and (self.is_mla_backend or self.is_hybrid_mla_backend)
|
||||
):
|
||||
return True
|
||||
|
||||
raise RuntimeError(
|
||||
f"Unsupported PD DCP topology: {self.dcp_size} -> {dst_dcp_size}"
|
||||
)
|
||||
|
||||
def prepare_dcp_token_item_lens(self, dst_page_item_lens: List[int]) -> List[int]:
|
||||
page_size = self.kv_args.page_size
|
||||
src_token_lens = [
|
||||
item_len // page_size for item_len in self.kv_args.kv_item_lens
|
||||
]
|
||||
dst_token_lens = [item_len // page_size for item_len in dst_page_item_lens]
|
||||
if src_token_lens != dst_token_lens:
|
||||
raise RuntimeError(
|
||||
"PD DCP source/destination KV geometry differs: "
|
||||
f"src={src_token_lens}, dst={dst_token_lens}"
|
||||
)
|
||||
return src_token_lens
|
||||
|
||||
def check_status(self, bootstrap_room: int) -> KVPoll:
|
||||
return self.request_status[bootstrap_room]
|
||||
|
||||
@@ -469,7 +505,11 @@ class CommonKVManager(BaseKVManager):
|
||||
|
||||
info: PrefillServerInfo = None
|
||||
try:
|
||||
url = f"http://{bootstrap_addr}/route?prefill_dp_rank={-1}&prefill_cp_rank={-1}&target_tp_rank={-1}&target_pp_rank={-1}"
|
||||
url = (
|
||||
f"http://{bootstrap_addr}/route?"
|
||||
f"prefill_dp_rank={-1}&prefill_cp_rank={-1}&"
|
||||
f"target_tp_rank={-1}&target_pp_rank={-1}"
|
||||
)
|
||||
response = requests.get(url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
@@ -501,6 +541,17 @@ class CommonKVManager(BaseKVManager):
|
||||
f"Both servers must use the same --kv-cache-dtype value."
|
||||
)
|
||||
|
||||
if self.dcp_size > 1:
|
||||
if not (self.is_mla_backend or self.is_hybrid_mla_backend):
|
||||
raise RuntimeError(
|
||||
"PD decode DCP requires an MLA or hybrid-MLA KV pool."
|
||||
)
|
||||
if info.attn_cp_size != 1:
|
||||
raise RuntimeError(
|
||||
"PD decode DCP currently requires prefill attention CP=1, "
|
||||
f"got {info.attn_cp_size}."
|
||||
)
|
||||
|
||||
self._resolve_rank_mapping(info)
|
||||
self.prefill_info_table[bootstrap_addr] = info
|
||||
logger.debug(f"Prefill parallel info for [{bootstrap_addr}]: {info}")
|
||||
@@ -1137,6 +1188,7 @@ class CommonKVSender(BaseKVSender):
|
||||
self,
|
||||
kv_indices: npt.NDArray[np.int32],
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
pass
|
||||
|
||||
|
||||
@@ -25,6 +25,7 @@ class TransferKVChunk:
|
||||
prefill_aux_index: Optional[int]
|
||||
state_indices: Optional[List]
|
||||
chunk_id: Optional[int] = None
|
||||
num_kv_tokens: Optional[int] = None
|
||||
trace_ctx: Union[TraceReqContext, TraceNullContext] = dataclasses.field(
|
||||
default_factory=TraceNullContext
|
||||
)
|
||||
@@ -127,3 +128,71 @@ def group_concurrent_contiguous(
|
||||
dst_groups = [g.tolist() for g in dst_groups]
|
||||
|
||||
return src_groups, dst_groups
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class DCPTokenTransferPlan:
|
||||
src_token_indices: npt.NDArray[np.int64]
|
||||
dst_token_indices: npt.NDArray[np.int64]
|
||||
|
||||
|
||||
def build_dcp_token_transfer_plan(
|
||||
src_page_indices: npt.NDArray[np.int32],
|
||||
dst_page_indices: npt.NDArray[np.int32],
|
||||
*,
|
||||
physical_page_size: int,
|
||||
dcp_size: int,
|
||||
dcp_rank: int,
|
||||
src_page_offset: int = 0,
|
||||
decode_prefix_len: int = 0,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
) -> DCPTokenTransferPlan:
|
||||
virtual_page_size = physical_page_size * dcp_size
|
||||
if decode_prefix_len % virtual_page_size != 0:
|
||||
raise ValueError(
|
||||
"PD DCP transfer requires decode_prefix_len to align to the virtual "
|
||||
f"DCP page size ({virtual_page_size}), got {decode_prefix_len}"
|
||||
)
|
||||
|
||||
src_pages = np.asarray(src_page_indices, dtype=np.int64)
|
||||
dst_pages = np.asarray(dst_page_indices, dtype=np.int64)
|
||||
source_capacity = src_pages.size * physical_page_size
|
||||
if num_kv_tokens is None:
|
||||
num_kv_tokens = source_capacity
|
||||
if not 0 <= num_kv_tokens <= source_capacity:
|
||||
raise ValueError(
|
||||
"num_kv_tokens must fit in the provided source pages, "
|
||||
f"got tokens={num_kv_tokens}, capacity={source_capacity}"
|
||||
)
|
||||
if src_pages.size == 0:
|
||||
empty = np.empty((0,), dtype=np.int64)
|
||||
return DCPTokenTransferPlan(empty, empty.copy())
|
||||
|
||||
chunk_start = decode_prefix_len + src_page_offset * physical_page_size
|
||||
first_owned_offset = (dcp_rank - chunk_start) % dcp_size
|
||||
owned_offsets = np.arange(
|
||||
first_owned_offset, num_kv_tokens, dcp_size, dtype=np.int64
|
||||
)
|
||||
src_token_indices = (
|
||||
src_pages[owned_offsets // physical_page_size] * physical_page_size
|
||||
+ owned_offsets % physical_page_size
|
||||
)
|
||||
|
||||
relative_positions = src_page_offset * physical_page_size + owned_offsets
|
||||
dst_local_offsets = relative_positions // dcp_size
|
||||
dst_page_ordinals = dst_local_offsets // physical_page_size
|
||||
if dst_page_ordinals.size and (
|
||||
dst_pages.size == 0 or int(dst_page_ordinals.max()) >= dst_pages.size
|
||||
):
|
||||
required_pages = int(dst_page_ordinals.max()) + 1
|
||||
raise ValueError(
|
||||
"Insufficient destination DCP pages: "
|
||||
f"required={required_pages}, provided={dst_pages.size}, "
|
||||
f"src_page_offset={src_page_offset}, dcp_rank={dcp_rank}"
|
||||
)
|
||||
|
||||
dst_token_indices = (
|
||||
dst_pages[dst_page_ordinals] * physical_page_size
|
||||
+ dst_local_offsets % physical_page_size
|
||||
)
|
||||
return DCPTokenTransferPlan(src_token_indices, dst_token_indices)
|
||||
|
||||
@@ -78,6 +78,7 @@ class FakeKVSender(BaseKVSender):
|
||||
self,
|
||||
kv_indices: npt.NDArray[np.int32],
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
self.has_sent = True
|
||||
logger.debug(
|
||||
|
||||
@@ -34,6 +34,7 @@ from sglang.srt.disaggregation.common.utils import (
|
||||
AuxDataCodec,
|
||||
FastQueue,
|
||||
TransferKVChunk,
|
||||
build_dcp_token_transfer_plan,
|
||||
group_concurrent_contiguous,
|
||||
pack_int_lists,
|
||||
unpack_int_lists,
|
||||
@@ -45,6 +46,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
DisaggregationMode,
|
||||
build_transfer_entry_pairs,
|
||||
compute_mamba_state_slice_byte_blocks,
|
||||
resolve_dcp_dst_entry_indices,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
|
||||
from sglang.srt.environ import envs
|
||||
@@ -131,6 +133,10 @@ class KVArgsRegisterInfo:
|
||||
dst_state_dim_per_tensor: List[List[int]]
|
||||
dst_kv_layer_ids: List[int]
|
||||
dst_state_layer_ids: List[List[int]]
|
||||
dst_dcp_size: int = 1
|
||||
dst_dcp_rank: int = 0
|
||||
requires_dcp_relayout: bool = False
|
||||
dcp_token_item_lens: Optional[List[int]] = None
|
||||
# Note: always put the staging field at the final (since the staging field is optional and contains multiple inputs)
|
||||
staging: Optional[StagingRegisterInfo] = None
|
||||
|
||||
@@ -163,6 +169,13 @@ class KVArgsRegisterInfo:
|
||||
if len(msg) > 13 and msg[13] != b""
|
||||
else []
|
||||
),
|
||||
# msg[14:16] belong to the staging field below; DCP trails it.
|
||||
dst_dcp_size=(
|
||||
int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else 1
|
||||
),
|
||||
dst_dcp_rank=(
|
||||
int(msg[17].decode("ascii")) if len(msg) > 17 and msg[17] != b"" else 0
|
||||
),
|
||||
# Note: always put the staging field at the final
|
||||
staging=StagingRegisterInfo.from_zmq_fields(msg, 14),
|
||||
)
|
||||
@@ -749,6 +762,106 @@ class MooncakeKVManager(CommonKVManager):
|
||||
dst_layer_ids=dst_layer_ids,
|
||||
)
|
||||
|
||||
def send_kvcache_dcp(
|
||||
self,
|
||||
mooncake_session_id: str,
|
||||
prefill_kv_indices: npt.NDArray[np.int32],
|
||||
dst_kv_ptrs: list[int],
|
||||
dst_kv_indices: npt.NDArray[np.int32],
|
||||
*,
|
||||
dcp_token_item_lens: List[int],
|
||||
dst_dcp_size: int,
|
||||
dst_dcp_rank: int,
|
||||
src_page_offset: int,
|
||||
decode_prefix_len: int,
|
||||
num_kv_tokens: int,
|
||||
executor: concurrent.futures.ThreadPoolExecutor,
|
||||
dst_layer_ids: List[int],
|
||||
) -> int:
|
||||
if num_kv_tokens is None:
|
||||
raise ValueError("PD DCP transfer requires num_kv_tokens")
|
||||
physical_page_size = self.kv_args.page_size
|
||||
plan = build_dcp_token_transfer_plan(
|
||||
prefill_kv_indices,
|
||||
dst_kv_indices,
|
||||
physical_page_size=physical_page_size,
|
||||
dcp_size=dst_dcp_size,
|
||||
dcp_rank=dst_dcp_rank,
|
||||
src_page_offset=src_page_offset,
|
||||
decode_prefix_len=decode_prefix_len,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
)
|
||||
if plan.src_token_indices.size == 0:
|
||||
return 0
|
||||
|
||||
src_layer_ids = self.kv_args.kv_layer_ids
|
||||
if src_layer_ids or dst_layer_ids:
|
||||
dst_indices = resolve_dcp_dst_entry_indices(
|
||||
src_layer_ids,
|
||||
dst_layer_ids,
|
||||
len(self.kv_args.kv_data_ptrs),
|
||||
len(dst_kv_ptrs),
|
||||
)
|
||||
src_kv_ptrs = self.kv_args.kv_data_ptrs
|
||||
dst_kv_ptrs = [dst_kv_ptrs[j] for j in dst_indices]
|
||||
else:
|
||||
src_kv_ptrs, dst_kv_ptrs, _ = self.get_mla_kv_ptrs_with_pp(
|
||||
self.kv_args.kv_data_ptrs,
|
||||
dst_kv_ptrs,
|
||||
)
|
||||
layers_current_pp_stage = len(src_kv_ptrs)
|
||||
src_groups, dst_groups = group_concurrent_contiguous(
|
||||
plan.src_token_indices,
|
||||
plan.dst_token_indices,
|
||||
)
|
||||
|
||||
layers_params = [
|
||||
(
|
||||
src_kv_ptrs[layer_id],
|
||||
dst_kv_ptrs[layer_id],
|
||||
dcp_token_item_lens[layer_id],
|
||||
)
|
||||
for layer_id in range(layers_current_pp_stage)
|
||||
]
|
||||
|
||||
def set_transfer_blocks(
|
||||
src_ptr: int, dst_ptr: int, token_item_len: int
|
||||
) -> List[Tuple[int, int, int]]:
|
||||
return [
|
||||
(
|
||||
src_ptr + int(src_group[0]) * token_item_len,
|
||||
dst_ptr + int(dst_group[0]) * token_item_len,
|
||||
len(src_group) * token_item_len,
|
||||
)
|
||||
for src_group, dst_group in zip(src_groups, dst_groups)
|
||||
]
|
||||
|
||||
def process_layer(src_ptr: int, dst_ptr: int, token_item_len: int) -> int:
|
||||
return self._transfer_data(
|
||||
mooncake_session_id,
|
||||
set_transfer_blocks(src_ptr, dst_ptr, token_item_len),
|
||||
)
|
||||
|
||||
if self.enable_custom_mem_pool:
|
||||
futures = [
|
||||
executor.submit(process_layer, src_ptr, dst_ptr, token_item_len)
|
||||
for src_ptr, dst_ptr, token_item_len in layers_params
|
||||
]
|
||||
for future in concurrent.futures.as_completed(futures):
|
||||
status = future.result()
|
||||
if status != 0:
|
||||
for pending in futures:
|
||||
pending.cancel()
|
||||
return status
|
||||
return 0
|
||||
|
||||
transfer_blocks = []
|
||||
for src_ptr, dst_ptr, token_item_len in layers_params:
|
||||
transfer_blocks.extend(
|
||||
set_transfer_blocks(src_ptr, dst_ptr, token_item_len)
|
||||
)
|
||||
return self._transfer_data(mooncake_session_id, transfer_blocks)
|
||||
|
||||
def send_kvcache_slice(
|
||||
self,
|
||||
mooncake_session_id: str,
|
||||
@@ -1437,7 +1550,18 @@ class MooncakeKVManager(CommonKVManager):
|
||||
)
|
||||
break
|
||||
|
||||
chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice]
|
||||
target_rank_registration_info: KVArgsRegisterInfo = (
|
||||
self.decode_kv_args_table[req.mooncake_session_id]
|
||||
)
|
||||
is_dcp_transfer = (
|
||||
target_rank_registration_info.requires_dcp_relayout
|
||||
)
|
||||
if is_dcp_transfer:
|
||||
chunked_dst_kv_indice = req.dst_kv_indices
|
||||
else:
|
||||
chunked_dst_kv_indice = req.dst_kv_indices[
|
||||
kv_chunk.index_slice
|
||||
]
|
||||
|
||||
# NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices
|
||||
# is mismatched with the dst_kv_indices when page size > 1, this should never happen.
|
||||
@@ -1447,13 +1571,12 @@ class MooncakeKVManager(CommonKVManager):
|
||||
logger.warning(
|
||||
f"len(chunked_dst_kv_indice) = {len(chunked_dst_kv_indice)}, len(kv_chunk.prefill_kv_indices) = {len(kv_chunk.prefill_kv_indices)}"
|
||||
)
|
||||
kv_chunk.prefill_kv_indices = kv_chunk.prefill_kv_indices[
|
||||
kv_chunk.prefill_kv_indices = (
|
||||
kv_chunk.prefill_kv_indices[
|
||||
: len(chunked_dst_kv_indice)
|
||||
]
|
||||
|
||||
target_rank_registration_info: KVArgsRegisterInfo = (
|
||||
self.decode_kv_args_table[req.mooncake_session_id]
|
||||
)
|
||||
|
||||
skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags(
|
||||
target_rank_registration_info
|
||||
)
|
||||
@@ -1463,6 +1586,27 @@ class MooncakeKVManager(CommonKVManager):
|
||||
or skip_kv
|
||||
):
|
||||
ret = 0
|
||||
elif is_dcp_transfer:
|
||||
dcp_token_item_lens = (
|
||||
target_rank_registration_info.dcp_token_item_lens
|
||||
)
|
||||
assert dcp_token_item_lens is not None
|
||||
ret = self.send_kvcache_dcp(
|
||||
req.mooncake_session_id,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
target_rank_registration_info.dst_kv_ptrs,
|
||||
chunked_dst_kv_indice,
|
||||
dcp_token_item_lens=dcp_token_item_lens,
|
||||
dst_dcp_size=target_rank_registration_info.dst_dcp_size,
|
||||
dst_dcp_rank=target_rank_registration_info.dst_dcp_rank,
|
||||
src_page_offset=kv_chunk.index_slice.start or 0,
|
||||
decode_prefix_len=req.decode_prefix_len or 0,
|
||||
num_kv_tokens=kv_chunk.num_kv_tokens,
|
||||
executor=executor,
|
||||
dst_layer_ids=(
|
||||
target_rank_registration_info.dst_kv_layer_ids
|
||||
),
|
||||
)
|
||||
elif (
|
||||
self.is_mla_backend
|
||||
or self.is_hybrid_mla_backend
|
||||
@@ -1670,9 +1814,19 @@ class MooncakeKVManager(CommonKVManager):
|
||||
continue
|
||||
mooncake_session_id = waiting_req_bytes[3].decode("ascii")
|
||||
if room == "None":
|
||||
self.decode_kv_args_table[mooncake_session_id] = (
|
||||
KVArgsRegisterInfo.from_zmq(waiting_req_bytes)
|
||||
decode_kv_args = KVArgsRegisterInfo.from_zmq(waiting_req_bytes)
|
||||
decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout(
|
||||
decode_kv_args.dst_dcp_size,
|
||||
decode_kv_args.dst_dcp_rank,
|
||||
)
|
||||
if decode_kv_args.requires_dcp_relayout:
|
||||
decode_kv_args.dcp_token_item_lens = (
|
||||
self.prepare_dcp_token_item_lens(
|
||||
[decode_kv_args.dst_kv_item_len]
|
||||
* len(self.kv_args.kv_item_lens)
|
||||
)
|
||||
)
|
||||
self.decode_kv_args_table[mooncake_session_id] = decode_kv_args
|
||||
with self.session_lock:
|
||||
if mooncake_session_id in self.failed_sessions:
|
||||
self.failed_sessions.remove(mooncake_session_id)
|
||||
@@ -1786,6 +1940,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
is_last_chunk: bool,
|
||||
aux_index: Optional[int] = None,
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
trace_ctx: Optional[Union[TraceReqContext, TraceNullContext]] = None,
|
||||
):
|
||||
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
||||
@@ -1824,6 +1979,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
is_last_chunk=is_last_chunk,
|
||||
prefill_aux_index=aux_index,
|
||||
state_indices=state_indices,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
trace_ctx=trace_ctx,
|
||||
)
|
||||
)
|
||||
@@ -1904,6 +2060,7 @@ class MooncakeKVSender(CommonKVSender):
|
||||
self,
|
||||
kv_indices: npt.NDArray[np.int32],
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
kv_indices, index_slice, is_last_chunk, should_skip = (
|
||||
self._prepare_send_indices(kv_indices, state_indices)
|
||||
@@ -1917,6 +2074,7 @@ class MooncakeKVSender(CommonKVSender):
|
||||
kv_indices,
|
||||
index_slice,
|
||||
False,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
trace_ctx=self.trace_ctx.copy_for_thread(),
|
||||
)
|
||||
else:
|
||||
@@ -1927,6 +2085,7 @@ class MooncakeKVSender(CommonKVSender):
|
||||
True,
|
||||
aux_index=self.aux_index,
|
||||
state_indices=state_indices,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
trace_ctx=self.trace_ctx.copy_for_thread(),
|
||||
)
|
||||
self._record_transfer_indices(kv_indices, state_indices)
|
||||
@@ -2030,6 +2189,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
||||
dst_tp_rank = str(tp_rank).encode("ascii")
|
||||
dst_attn_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
|
||||
dst_kv_item_len = str(kv_item_len).encode("ascii")
|
||||
dst_dcp_size = str(self.kv_mgr.dcp_size).encode("ascii")
|
||||
dst_dcp_rank = str(self.kv_mgr.dcp_rank).encode("ascii")
|
||||
if (
|
||||
self.kv_mgr.enable_staging
|
||||
and self.kv_mgr._staging_ctx.allocator is not None
|
||||
@@ -2062,6 +2223,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
||||
packed_state_layer_ids,
|
||||
packed_staging_base_ptr,
|
||||
staging_total_size_str,
|
||||
dst_dcp_size,
|
||||
dst_dcp_rank,
|
||||
]
|
||||
)
|
||||
except zmq.ZMQError:
|
||||
|
||||
@@ -1417,6 +1417,7 @@ class MoriKVSender(CommonKVSender):
|
||||
self,
|
||||
kv_indices: npt.NDArray[np.int32],
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
kv_indices, index_slice, is_last_chunk, should_skip = (
|
||||
self._prepare_send_indices(kv_indices, state_indices)
|
||||
|
||||
@@ -32,6 +32,7 @@ from sglang.srt.disaggregation.common.staging_handler import (
|
||||
from sglang.srt.disaggregation.common.utils import (
|
||||
FastQueue,
|
||||
TransferKVChunk,
|
||||
build_dcp_token_transfer_plan,
|
||||
group_concurrent_contiguous,
|
||||
pack_int_lists,
|
||||
unpack_int_lists,
|
||||
@@ -40,6 +41,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
DisaggregationMode,
|
||||
build_transfer_entry_pairs,
|
||||
compute_mamba_state_slice_byte_blocks,
|
||||
resolve_dcp_dst_entry_indices,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -158,6 +160,7 @@ class TransferInfo:
|
||||
required_dst_info_num: int
|
||||
dst_state_indices: List[List[int]]
|
||||
decode_prefix_len: Optional[int] = None # for decode radix cache
|
||||
is_dummy_rank: Optional[bool] = None
|
||||
# NOTE: optional staging field; populated via STAGING_RSP. Keep at the
|
||||
# end so positional construction in from_zmq() continues to work.
|
||||
staging: Optional[StagingTransferInfo] = None
|
||||
@@ -167,6 +170,8 @@ class TransferInfo:
|
||||
# When dst_kv_indices is empty due to a decode-side radix cache
|
||||
# full hit (decode_prefix_len > 0), the transfer is NOT dummy --
|
||||
# aux/state data still needs to be sent.
|
||||
if self.is_dummy_rank is not None:
|
||||
return self.is_dummy_rank
|
||||
if self.dst_kv_indices.size == 0 and self.decode_prefix_len:
|
||||
return False
|
||||
return self.dst_kv_indices.size == 0
|
||||
@@ -189,6 +194,11 @@ class TransferInfo:
|
||||
decode_prefix_len=(
|
||||
int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None
|
||||
), # hacky just add it into the message that will be sent
|
||||
is_dummy_rank=(
|
||||
bool(int(msg[9].decode("ascii")))
|
||||
if len(msg) > 9 and msg[9] != b""
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -211,6 +221,11 @@ class KVArgsRegisterInfo:
|
||||
dst_kv_item_len: int
|
||||
dst_kv_item_lens: list[int]
|
||||
dst_kv_layer_ids: list[int] = dataclasses.field(default_factory=list)
|
||||
dst_dcp_size: int = 1
|
||||
dst_dcp_rank: int = 0
|
||||
requires_dcp_relayout: bool = False
|
||||
dcp_token_item_lens: Optional[List[int]] = None
|
||||
dcp_dst_region_indices: Optional[List[int]] = None
|
||||
dst_num_slots: Optional[int] = None
|
||||
dst_state_item_lens: List[List[int]] = dataclasses.field(default_factory=list)
|
||||
dst_state_dim_per_tensor: List[List[int]] = dataclasses.field(default_factory=list)
|
||||
@@ -277,6 +292,13 @@ class KVArgsRegisterInfo:
|
||||
dst_kv_item_len=dst_kv_item_len,
|
||||
dst_kv_item_lens=dst_kv_item_lens,
|
||||
dst_kv_layer_ids=dst_kv_layer_ids,
|
||||
# msg[19:21] are the layer-id frames above; DCP trails them.
|
||||
dst_dcp_size=(
|
||||
int(msg[21].decode("ascii")) if len(msg) > 21 and msg[21] != b"" else 1
|
||||
),
|
||||
dst_dcp_rank=(
|
||||
int(msg[22].decode("ascii")) if len(msg) > 22 and msg[22] != b"" else 0
|
||||
),
|
||||
dst_num_slots=dst_num_slots,
|
||||
dst_state_item_lens=dst_state_item_lens,
|
||||
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
|
||||
@@ -973,6 +995,27 @@ class NixlKVManager(CommonKVManager):
|
||||
assert self.src_mem_kind is not None
|
||||
src_mem_kind = self.src_mem_kind
|
||||
decode_only_spec_dec = n_dst > n_src
|
||||
|
||||
if peer_info.requires_dcp_relayout:
|
||||
dst_indices = resolve_dcp_dst_entry_indices(
|
||||
self.kv_args.kv_layer_ids,
|
||||
peer_info.dst_kv_layer_ids,
|
||||
n_src,
|
||||
n_dst,
|
||||
)
|
||||
peer_info.dcp_dst_region_indices = dst_indices
|
||||
dst_kv_mem_kinds = [peer_info.dst_kv_mem_kinds[j] for j in dst_indices]
|
||||
dst_kv_item_lens = [peer_info.dst_kv_item_lens[j] for j in dst_indices]
|
||||
dst_mem_kind = _homogeneous_kv_mem_kind(
|
||||
dst_kv_mem_kinds,
|
||||
"PD DCP destination",
|
||||
)
|
||||
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
|
||||
peer_info.dcp_token_item_lens = self.prepare_dcp_token_item_lens(
|
||||
dst_kv_item_lens
|
||||
)
|
||||
return
|
||||
|
||||
if (
|
||||
self.is_mla_backend
|
||||
or self.is_hybrid_mla_backend
|
||||
@@ -1110,7 +1153,13 @@ class NixlKVManager(CommonKVManager):
|
||||
len(kv_chunk.prefill_kv_indices) > 0
|
||||
and self.kv_args.kv_data_ptrs
|
||||
):
|
||||
chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice]
|
||||
is_dcp_transfer = dst_info.requires_dcp_relayout
|
||||
if is_dcp_transfer:
|
||||
chunked_dst_kv_indice = req.dst_kv_indices
|
||||
else:
|
||||
chunked_dst_kv_indice = req.dst_kv_indices[
|
||||
kv_chunk.index_slice
|
||||
]
|
||||
|
||||
# NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices
|
||||
# is mismatched with the dst_kv_indices when page size > 1, this should never happen.
|
||||
@@ -1120,9 +1169,11 @@ class NixlKVManager(CommonKVManager):
|
||||
logger.warning(
|
||||
f"len(chunked_dst_kv_indice) = {len(chunked_dst_kv_indice)}, len(kv_chunk.prefill_kv_indices) = {len(kv_chunk.prefill_kv_indices)}"
|
||||
)
|
||||
kv_chunk.prefill_kv_indices = kv_chunk.prefill_kv_indices[
|
||||
kv_chunk.prefill_kv_indices = (
|
||||
kv_chunk.prefill_kv_indices[
|
||||
: len(chunked_dst_kv_indice)
|
||||
]
|
||||
)
|
||||
|
||||
src_prefill_kv_indices = kv_chunk.prefill_kv_indices
|
||||
|
||||
@@ -1164,7 +1215,18 @@ class NixlKVManager(CommonKVManager):
|
||||
break
|
||||
|
||||
if kv_xfer_handle is None:
|
||||
if (
|
||||
if is_dcp_transfer:
|
||||
kv_xfer_handle = self.send_kvcache_dcp(
|
||||
req.agent_name,
|
||||
src_prefill_kv_indices,
|
||||
dst_info,
|
||||
chunked_dst_kv_indice,
|
||||
src_page_offset=kv_chunk.index_slice.start or 0,
|
||||
decode_prefix_len=req.decode_prefix_len or 0,
|
||||
num_kv_tokens=kv_chunk.num_kv_tokens,
|
||||
notif=notif,
|
||||
)
|
||||
elif (
|
||||
self.is_mla_backend
|
||||
or self.is_hybrid_mla_backend
|
||||
or decode_tp_size == self.attn_tp_size
|
||||
@@ -1228,15 +1290,17 @@ class NixlKVManager(CommonKVManager):
|
||||
if kv_chunk.prefill_aux_index is None:
|
||||
raise RuntimeError("Missing aux index for last chunk")
|
||||
# A no-KV notification still identifies its PP source.
|
||||
# Empty non-final chunks do not consume chunk IDs, so a
|
||||
# final no-KV chunk_id equals the prior KV chunk count.
|
||||
aux_notif = f"{req.room}_aux"
|
||||
if (
|
||||
len(kv_chunk.prefill_kv_indices) == 0
|
||||
or not self.kv_args.kv_data_ptrs
|
||||
):
|
||||
aux_notif = (
|
||||
f"{req.room}_aux_nokv_{self.transfer_source_rank}"
|
||||
aux_notif += (
|
||||
f"_nokv_{self.transfer_source_rank}"
|
||||
f"_{kv_chunk.chunk_id}"
|
||||
)
|
||||
else:
|
||||
aux_notif = f"{req.room}_aux"
|
||||
aux_xfer_handle = self.send_aux(
|
||||
req.agent_name,
|
||||
kv_chunk.prefill_aux_index,
|
||||
@@ -1354,6 +1418,9 @@ class NixlKVManager(CommonKVManager):
|
||||
if agent_name in self.decode_kv_args_table:
|
||||
logger.info(f"Peer {agent_name} was already registered, ignoring.")
|
||||
return
|
||||
decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout(
|
||||
decode_kv_args.dst_dcp_size, decode_kv_args.dst_dcp_rank
|
||||
)
|
||||
self.decode_kv_args_table[agent_name] = decode_kv_args
|
||||
self.agent.add_remote_agent(decode_kv_args.agent_metadata)
|
||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
@@ -1373,6 +1440,7 @@ class NixlKVManager(CommonKVManager):
|
||||
src_mem_kind: str = "VRAM",
|
||||
dst_mem_kind: str = "VRAM",
|
||||
force_flat: bool = False,
|
||||
bypass_prepped: bool = False,
|
||||
):
|
||||
"""Generic KV cache transfer supporting both MHA and MLA architectures.
|
||||
Used by both send_kvcache and maybe_send_extra.
|
||||
@@ -1382,7 +1450,8 @@ class NixlKVManager(CommonKVManager):
|
||||
index) whose per-layer list must not be half-split into K/V."""
|
||||
# Prepped path (KV only; state transfers use the non-prepped path below).
|
||||
if (
|
||||
src_data_ptrs is self.kv_args.kv_data_ptrs
|
||||
not bypass_prepped
|
||||
and src_data_ptrs is self.kv_args.kv_data_ptrs
|
||||
and "" in self.prep_handles
|
||||
and peer_name in self.prep_handles
|
||||
):
|
||||
@@ -1550,6 +1619,63 @@ class NixlKVManager(CommonKVManager):
|
||||
dst_mem_kind=dst_mem_kind,
|
||||
)
|
||||
|
||||
def send_kvcache_dcp(
|
||||
self,
|
||||
peer_name: str,
|
||||
prefill_kv_indices: npt.NDArray[np.int32],
|
||||
dst_info: KVArgsRegisterInfo,
|
||||
dst_kv_indices: npt.NDArray[np.int32],
|
||||
*,
|
||||
src_page_offset: int,
|
||||
decode_prefix_len: int,
|
||||
num_kv_tokens: int,
|
||||
notif: str,
|
||||
):
|
||||
if self.src_mem_kind is None:
|
||||
raise RuntimeError("Missing NIXL source KV memory kind")
|
||||
if dst_info.dst_homogeneous_mem_kind is None:
|
||||
raise RuntimeError("Missing NIXL destination KV memory kind")
|
||||
if num_kv_tokens is None:
|
||||
raise ValueError("PD DCP transfer requires num_kv_tokens")
|
||||
|
||||
physical_page_size = self.kv_args.page_size
|
||||
plan = build_dcp_token_transfer_plan(
|
||||
prefill_kv_indices,
|
||||
dst_kv_indices,
|
||||
physical_page_size=physical_page_size,
|
||||
dcp_size=dst_info.dst_dcp_size,
|
||||
dcp_rank=dst_info.dst_dcp_rank,
|
||||
src_page_offset=src_page_offset,
|
||||
decode_prefix_len=decode_prefix_len,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
)
|
||||
if plan.src_token_indices.size == 0:
|
||||
self.agent.send_notif(peer_name, notif.encode("ascii"))
|
||||
return None
|
||||
|
||||
token_item_lens = dst_info.dcp_token_item_lens
|
||||
assert token_item_lens is not None
|
||||
dst_kv_ptrs = [
|
||||
dst_info.dst_kv_ptrs[dst_idx] for dst_idx in dst_info.dcp_dst_region_indices
|
||||
]
|
||||
|
||||
# Prepared handles encode page-level offsets, while DCP relayout needs
|
||||
# flat descriptors for the selected token rows.
|
||||
return self._send_kvcache_generic(
|
||||
peer_name=peer_name,
|
||||
src_data_ptrs=self.kv_args.kv_data_ptrs,
|
||||
dst_data_ptrs=dst_kv_ptrs,
|
||||
item_lens=token_item_lens,
|
||||
prefill_data_indices=plan.src_token_indices,
|
||||
dst_data_indices=plan.dst_token_indices,
|
||||
dst_gpu_id=dst_info.gpu_id,
|
||||
notif=notif,
|
||||
src_mem_kind=self.src_mem_kind,
|
||||
dst_mem_kind=dst_info.dst_homogeneous_mem_kind,
|
||||
force_flat=True,
|
||||
bypass_prepped=True,
|
||||
)
|
||||
|
||||
def send_kvcache_mixed(
|
||||
self,
|
||||
peer_name: str,
|
||||
@@ -2243,6 +2369,7 @@ class NixlKVManager(CommonKVManager):
|
||||
chunk_id: int,
|
||||
aux_index: Optional[int] = None,
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
||||
assert not is_last_chunk or (is_last_chunk and aux_index is not None)
|
||||
@@ -2273,6 +2400,7 @@ class NixlKVManager(CommonKVManager):
|
||||
chunk_id=chunk_id,
|
||||
prefill_aux_index=aux_index,
|
||||
state_indices=state_indices,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
)
|
||||
)
|
||||
return None
|
||||
@@ -2313,8 +2441,8 @@ class NixlKVManager(CommonKVManager):
|
||||
elif tag == "stg":
|
||||
self._handle_stg_notification(components, room)
|
||||
elif tag == "aux":
|
||||
# main's "nokv" marker (decode-side radix cache hit):
|
||||
# mark expected_kvs_per_pp[pp_rank] = 0 for this rank.
|
||||
# Main's "nokv" marker carries the number of earlier KV
|
||||
# chunks expected from this PP rank.
|
||||
self._handle_aux_notification(room, components)
|
||||
elif tag == "state":
|
||||
pp_rank = int(components[2]) if len(components) > 2 else 0
|
||||
@@ -2342,15 +2470,16 @@ class NixlKVManager(CommonKVManager):
|
||||
|
||||
Notification tag layouts:
|
||||
aux: {room}_aux -> 2 fields
|
||||
aux (nokv): {room}_aux_nokv_{pp_rank} -> 4 fields
|
||||
(decode-side radix cache hit; this pp_rank sent
|
||||
no KV pages, so expected_kvs_per_pp[pp_rank] = 0)
|
||||
aux (nokv): {room}_aux_nokv_{pp_rank}_{expected} -> 5 fields
|
||||
(the last chunk had no KV pages for this rank;
|
||||
`expected` is the number of prior KV chunks)
|
||||
"""
|
||||
self.transfer_statuses[room].received_aux = True
|
||||
# main's "nokv" marker (decode-side radix cache hit, see #19746).
|
||||
if len(components) > 3 and components[2] == "nokv":
|
||||
pp_rank = int(components[3])
|
||||
self.transfer_statuses[room].expected_kvs_per_pp[pp_rank] = 0
|
||||
expected = int(components[4]) if len(components) > 4 else 0
|
||||
self.transfer_statuses[room].expected_kvs_per_pp[pp_rank] = expected
|
||||
if self.transfer_statuses[room].num_pp_ranks_expected is None:
|
||||
self.transfer_statuses[room].num_pp_ranks_expected = (
|
||||
self.required_prefill_response_num_table.get(room, 1)
|
||||
@@ -2597,6 +2726,7 @@ class NixlKVSender(CommonKVSender):
|
||||
self,
|
||||
kv_indices: npt.NDArray[np.int32],
|
||||
state_indices: Optional[List] = None,
|
||||
num_kv_tokens: Optional[int] = None,
|
||||
):
|
||||
if self._send_failed:
|
||||
return
|
||||
@@ -2620,6 +2750,7 @@ class NixlKVSender(CommonKVSender):
|
||||
self.chunk_id,
|
||||
self.aux_index,
|
||||
state_indices,
|
||||
num_kv_tokens,
|
||||
)
|
||||
self._record_transfer_indices(kv_indices, state_indices)
|
||||
self.chunk_id += 1
|
||||
@@ -2737,6 +2868,7 @@ class NixlKVReceiver(CommonKVReceiver):
|
||||
str(self.required_dst_info_num).encode("ascii"),
|
||||
packed_state_indices,
|
||||
str(decode_prefix_len or 0).encode("ascii"),
|
||||
str(int(is_dummy)).encode("ascii"),
|
||||
]
|
||||
)
|
||||
except zmq.ZMQError:
|
||||
@@ -2865,6 +2997,8 @@ class NixlKVReceiver(CommonKVReceiver):
|
||||
packed_kv_item_lens,
|
||||
packed_state_layer_ids,
|
||||
packed_kv_layer_ids,
|
||||
str(self.kv_mgr.dcp_size).encode("ascii"),
|
||||
str(self.kv_mgr.dcp_rank).encode("ascii"),
|
||||
]
|
||||
)
|
||||
except zmq.ZMQError:
|
||||
|
||||
@@ -332,7 +332,8 @@ class PrefillBootstrapQueue:
|
||||
req.start_send_idx = decode_prefix_len
|
||||
num_kv_indices_to_send = num_kv_indices - decode_prefix_len
|
||||
num_pages = kv_to_page_num(
|
||||
num_kv_indices_to_send, self.token_to_kv_pool.page_size
|
||||
num_kv_indices_to_send,
|
||||
self.scheduler.token_to_kv_pool_allocator.page_size,
|
||||
)
|
||||
req.disagg_kv_sender.init(num_pages, req.metadata_buffer_index)
|
||||
req.pending_bootstrap = False
|
||||
@@ -1082,7 +1083,11 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
cached_end = len(req.prefix_indices) - req.host_hit_length
|
||||
if cached_end <= req.start_send_idx:
|
||||
return
|
||||
assert cached_end % self.token_to_kv_pool_allocator.page_size == 0
|
||||
if cached_end % self.token_to_kv_pool_allocator.page_size != 0:
|
||||
# DCP radix hits can end on a logical cache-page boundary that is
|
||||
# not a complete physical DCP page. The regular final send covers
|
||||
# the full range; only skip this optional early-send optimization.
|
||||
return
|
||||
# Early-send issues the KV read before this step's forward is enqueued,
|
||||
# but under overlap scheduling the PRIOR step's prefill forward may still
|
||||
# be writing these prefix pages on forward_stream. Record a completion
|
||||
@@ -1235,7 +1240,11 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
page_indices = kv_to_page_indices(kv_indices, page_size)
|
||||
if not req.disagg_kv_sender.should_send_kv_chunk(len(page_indices), last_chunk):
|
||||
return
|
||||
req.disagg_kv_sender.send(page_indices, state_indices)
|
||||
req.disagg_kv_sender.send(
|
||||
page_indices,
|
||||
state_indices,
|
||||
num_kv_tokens=end_idx - start_idx,
|
||||
)
|
||||
req.start_send_idx = end_idx
|
||||
|
||||
def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None:
|
||||
|
||||
@@ -934,6 +934,31 @@ def build_transfer_entry_pairs(
|
||||
return [(i, i) for i in range(n_src)]
|
||||
|
||||
|
||||
def resolve_dcp_dst_entry_indices(
|
||||
src_layer_ids: List[int],
|
||||
dst_layer_ids: List[int],
|
||||
n_src: int,
|
||||
n_dst: int,
|
||||
) -> List[int]:
|
||||
"""Destination entry index for each local KV entry, for a DCP relayout.
|
||||
|
||||
DCP re-splits the KV by context while PP re-splits it by layer, so the two
|
||||
index spaces only line up when neither peer is pipelined. Both backends
|
||||
need the same resolution, hence the shared helper.
|
||||
"""
|
||||
if not src_layer_ids and not dst_layer_ids:
|
||||
# Legacy/non-PP layout. n_dst may exceed n_src when the decode side
|
||||
# runs speculative decoding and the prefill side does not.
|
||||
return list(range(n_src))
|
||||
# A one-sided mapping is rejected by build_transfer_entry_pairs itself.
|
||||
return [
|
||||
j
|
||||
for _, j in build_transfer_entry_pairs(
|
||||
src_layer_ids, dst_layer_ids, n_src, n_dst
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def append_state_component(
|
||||
kv_args: KVArgs,
|
||||
state_type: StateType,
|
||||
|
||||
@@ -0,0 +1,488 @@
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
import unittest
|
||||
|
||||
import requests
|
||||
import torch
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||
PDDisaggregationServerBase,
|
||||
)
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
popen_launch_pd_server,
|
||||
popen_launch_server,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=1200, suite="nightly-8-gpu-b200", nightly=True)
|
||||
|
||||
KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct"
|
||||
PHYSICAL_PAGE_SIZE = 64
|
||||
CHUNKED_PREFILL_SIZE = 8192
|
||||
PARITY_PROMPT_LENGTHS = (63, 64, 65, 255, 256, 257, 8191, 8192, 8193)
|
||||
LONG_CONTEXT_TOKENS = int(os.environ.get("SGLANG_TEST_PD_LONG_CONTEXT_TOKENS", "32768"))
|
||||
LONG_CONTEXT_DEPTHS = tuple(
|
||||
float(value)
|
||||
for value in os.environ.get("SGLANG_TEST_PD_NIAH_DEPTHS", "0.1,0.5,0.9").split(",")
|
||||
)
|
||||
LOGPROB_ATOL = 0.20
|
||||
NIAH_KEY = "739391"
|
||||
KIMI_LINEAR_MAX_CONTEXT = 1_048_576
|
||||
|
||||
|
||||
def _has_eight_blackwell_gpus() -> bool:
|
||||
if not torch.cuda.is_available() or torch.cuda.device_count() < 8:
|
||||
return False
|
||||
return all(
|
||||
torch.cuda.get_device_capability(device_index) >= (10, 0)
|
||||
for device_index in range(8)
|
||||
)
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
_has_eight_blackwell_gpus(),
|
||||
"Kimi-Linear PD+DCP acceptance requires eight Blackwell GPUs",
|
||||
)
|
||||
class TestKimiLinearPDDCP4(GSM8KMixin, PDDisaggregationServerBase):
|
||||
model = KIMI_LINEAR_MODEL
|
||||
gsm8k_score_threshold = 0.88
|
||||
gsm8k_num_examples = 200
|
||||
gsm8k_num_threads = 4
|
||||
gsm8k_num_shots = 5
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
cls._collect_monolithic_references()
|
||||
cls.launch_all()
|
||||
|
||||
@classmethod
|
||||
def _monolithic_reference_args(cls):
|
||||
return [
|
||||
"--tp-size",
|
||||
"4",
|
||||
"--ep-size",
|
||||
"4",
|
||||
"--attention-backend",
|
||||
"tokenspeed_mla",
|
||||
"--kv-cache-dtype",
|
||||
"fp8_e4m3",
|
||||
"--trust-remote-code",
|
||||
"--random-seed",
|
||||
"0",
|
||||
"--dtype",
|
||||
"bfloat16",
|
||||
"--page-size",
|
||||
str(PHYSICAL_PAGE_SIZE),
|
||||
"--chunked-prefill-size",
|
||||
str(CHUNKED_PREFILL_SIZE),
|
||||
"--cuda-graph-max-bs-decode",
|
||||
"64",
|
||||
"--cuda-graph-backend-prefill",
|
||||
"disabled",
|
||||
"--mem-fraction-static",
|
||||
"0.80",
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def _tokenize(cls, text: str, *, add_special_tokens: bool):
|
||||
return cls.tokenizer.encode(text, add_special_tokens=add_special_tokens)
|
||||
|
||||
@classmethod
|
||||
def _repeat_to_length(cls, token_ids, length: int):
|
||||
if length == 0:
|
||||
return []
|
||||
assert token_ids, "filler must tokenize to at least one token"
|
||||
return (token_ids * math.ceil(length / len(token_ids)))[:length]
|
||||
|
||||
@classmethod
|
||||
def _chat_template_parts(cls):
|
||||
marker = "SGLANGPDCONTENTSENTINEL739391"
|
||||
marker_ids = cls._tokenize(marker, add_special_tokens=False)
|
||||
templated = cls.tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": marker}],
|
||||
tokenize=True,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
full_ids = templated if isinstance(templated, list) else templated["input_ids"]
|
||||
marker_starts = [
|
||||
index
|
||||
for index in range(len(full_ids) - len(marker_ids) + 1)
|
||||
if full_ids[index : index + len(marker_ids)] == marker_ids
|
||||
]
|
||||
assert len(marker_starts) == 1, (
|
||||
"Kimi chat template must contain the user-content marker exactly once, "
|
||||
f"found offsets={marker_starts}"
|
||||
)
|
||||
marker_start = marker_starts[0]
|
||||
return (
|
||||
full_ids[:marker_start],
|
||||
full_ids[marker_start + len(marker_ids) :],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _build_boundary_prompt(cls, target_length: int):
|
||||
prefix = cls._tokenize(
|
||||
"Read the repeated facts carefully.\n",
|
||||
add_special_tokens=True,
|
||||
)
|
||||
filler = cls._tokenize(
|
||||
"The sky is blue and the grass is green. ",
|
||||
add_special_tokens=False,
|
||||
)
|
||||
suffix = cls._tokenize(
|
||||
"\nAnswer with one word. The capital of France is",
|
||||
add_special_tokens=False,
|
||||
)
|
||||
middle_length = target_length - len(prefix) - len(suffix)
|
||||
assert middle_length >= 0, (
|
||||
f"target prompt length {target_length} is shorter than fixed prompt "
|
||||
f"content ({len(prefix) + len(suffix)})"
|
||||
)
|
||||
prompt_ids = prefix + cls._repeat_to_length(filler, middle_length) + suffix
|
||||
assert len(prompt_ids) == target_length
|
||||
assert target_length < KIMI_LINEAR_MAX_CONTEXT
|
||||
return prompt_ids
|
||||
|
||||
@classmethod
|
||||
def _build_niah_prompt(cls, target_length: int, depth: float):
|
||||
assert 0 <= depth <= 1
|
||||
chat_prefix, chat_suffix = cls._chat_template_parts()
|
||||
prefix = cls._tokenize(
|
||||
(
|
||||
"You will read a long collection of mundane records. One record "
|
||||
"contains an access code. Remember that code exactly.\n"
|
||||
),
|
||||
add_special_tokens=False,
|
||||
)
|
||||
filler = cls._tokenize(
|
||||
(
|
||||
"Archive note: the weather was mild, the office lights were on, "
|
||||
"and no unusual event was reported.\n"
|
||||
),
|
||||
add_special_tokens=False,
|
||||
)
|
||||
needle = cls._tokenize(
|
||||
f"IMPORTANT RECORD: The access code is {NIAH_KEY}.\n",
|
||||
add_special_tokens=False,
|
||||
)
|
||||
query = cls._tokenize(
|
||||
"\nWhat is the access code? Reply with the digits only.",
|
||||
add_special_tokens=False,
|
||||
)
|
||||
filler_length = (
|
||||
target_length
|
||||
- len(chat_prefix)
|
||||
- len(prefix)
|
||||
- len(needle)
|
||||
- len(query)
|
||||
- len(chat_suffix)
|
||||
)
|
||||
assert filler_length >= 0
|
||||
before_length = int(filler_length * depth)
|
||||
after_length = filler_length - before_length
|
||||
prompt_ids = (
|
||||
chat_prefix
|
||||
+ prefix
|
||||
+ cls._repeat_to_length(filler, before_length)
|
||||
+ needle
|
||||
+ cls._repeat_to_length(filler, after_length)
|
||||
+ query
|
||||
+ chat_suffix
|
||||
)
|
||||
assert len(prompt_ids) == target_length
|
||||
assert target_length < KIMI_LINEAR_MAX_CONTEXT
|
||||
return prompt_ids
|
||||
|
||||
@classmethod
|
||||
def _generate(
|
||||
cls,
|
||||
base_url: str,
|
||||
input_ids,
|
||||
*,
|
||||
max_new_tokens: int,
|
||||
ignore_eos: bool = True,
|
||||
):
|
||||
response = requests.post(
|
||||
base_url + "/generate",
|
||||
json={
|
||||
"input_ids": input_ids,
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
"ignore_eos": ignore_eos,
|
||||
},
|
||||
"return_logprob": True,
|
||||
},
|
||||
timeout=900,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
@staticmethod
|
||||
def _flush_cache(base_url: str):
|
||||
response = requests.post(
|
||||
base_url + "/flush_cache",
|
||||
params={"timeout": 30},
|
||||
timeout=120,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
@classmethod
|
||||
def _flush_pd_caches(cls):
|
||||
cls._flush_cache(cls.prefill_url)
|
||||
cls._flush_cache(cls.decode_url)
|
||||
|
||||
@classmethod
|
||||
def _collect_monolithic_references(cls):
|
||||
reference_process = popen_launch_server(
|
||||
cls.model,
|
||||
cls.base_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5,
|
||||
other_args=cls._monolithic_reference_args(),
|
||||
)
|
||||
try:
|
||||
cls.tokenizer = AutoTokenizer.from_pretrained(
|
||||
cls.model, trust_remote_code=True
|
||||
)
|
||||
cls.parity_prompts = [
|
||||
cls._build_boundary_prompt(target_length)
|
||||
for target_length in PARITY_PROMPT_LENGTHS
|
||||
]
|
||||
cls._flush_cache(cls.base_url)
|
||||
cls.parity_references = [
|
||||
cls._generate(cls.base_url, prompt, max_new_tokens=1)
|
||||
for prompt in cls.parity_prompts
|
||||
]
|
||||
cls._flush_cache(cls.base_url)
|
||||
cls.parity_batch_references = cls._generate(
|
||||
cls.base_url, cls.parity_prompts, max_new_tokens=1
|
||||
)
|
||||
|
||||
cls.niah_prompts = [
|
||||
cls._build_niah_prompt(LONG_CONTEXT_TOKENS, needle_depth)
|
||||
for needle_depth in LONG_CONTEXT_DEPTHS
|
||||
]
|
||||
cls._flush_cache(cls.base_url)
|
||||
cls.niah_references = [
|
||||
cls._generate(
|
||||
cls.base_url,
|
||||
prompt,
|
||||
max_new_tokens=16,
|
||||
ignore_eos=False,
|
||||
)
|
||||
for prompt in cls.niah_prompts
|
||||
]
|
||||
finally:
|
||||
kill_process_tree(reference_process.pid, wait_timeout=60)
|
||||
time.sleep(5)
|
||||
|
||||
@classmethod
|
||||
def start_prefill(cls):
|
||||
args = [
|
||||
"--trust-remote-code",
|
||||
"--disaggregation-mode",
|
||||
"prefill",
|
||||
"--disaggregation-bootstrap-port",
|
||||
cls.bootstrap_port,
|
||||
"--tp-size",
|
||||
"4",
|
||||
"--ep-size",
|
||||
"4",
|
||||
"--attention-backend",
|
||||
"tokenspeed_mla",
|
||||
"--kv-cache-dtype",
|
||||
"fp8_e4m3",
|
||||
"--dtype",
|
||||
"bfloat16",
|
||||
"--random-seed",
|
||||
"0",
|
||||
"--page-size",
|
||||
str(PHYSICAL_PAGE_SIZE),
|
||||
"--chunked-prefill-size",
|
||||
str(CHUNKED_PREFILL_SIZE),
|
||||
"--mem-fraction-static",
|
||||
"0.80",
|
||||
]
|
||||
args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_prefill = popen_launch_pd_server(
|
||||
cls.model,
|
||||
cls.prefill_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5,
|
||||
other_args=args,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def start_decode(cls):
|
||||
args = [
|
||||
"--trust-remote-code",
|
||||
"--disaggregation-mode",
|
||||
"decode",
|
||||
"--disaggregation-bootstrap-port",
|
||||
cls.bootstrap_port,
|
||||
"--tp-size",
|
||||
"4",
|
||||
"--dcp-size",
|
||||
"4",
|
||||
"--base-gpu-id",
|
||||
"4",
|
||||
"--attention-backend",
|
||||
"tokenspeed_mla",
|
||||
"--kv-cache-dtype",
|
||||
"fp8_e4m3",
|
||||
"--dcp-comm-backend",
|
||||
"a2a",
|
||||
"--dcp-replicate-q-proj",
|
||||
"--dtype",
|
||||
"bfloat16",
|
||||
"--random-seed",
|
||||
"0",
|
||||
"--page-size",
|
||||
str(PHYSICAL_PAGE_SIZE),
|
||||
"--cuda-graph-max-bs-decode",
|
||||
"64",
|
||||
"--cuda-graph-backend-prefill",
|
||||
"disabled",
|
||||
"--mem-fraction-static",
|
||||
"0.80",
|
||||
]
|
||||
args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_decode = popen_launch_pd_server(
|
||||
cls.model,
|
||||
cls.decode_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5,
|
||||
other_args=args,
|
||||
)
|
||||
|
||||
def _assert_output_parity(
|
||||
self, reference, actual, *, label: str, check_logprobs: bool = True
|
||||
):
|
||||
self.assertEqual(
|
||||
actual["output_ids"],
|
||||
reference["output_ids"],
|
||||
f"{label}: generated token IDs diverged",
|
||||
)
|
||||
if not check_logprobs:
|
||||
return
|
||||
|
||||
reference_logprobs = reference["meta_info"]["output_token_logprobs"]
|
||||
actual_logprobs = actual["meta_info"]["output_token_logprobs"]
|
||||
self.assertEqual(
|
||||
len(actual_logprobs),
|
||||
len(reference_logprobs),
|
||||
f"{label}: output logprob lengths diverged",
|
||||
)
|
||||
|
||||
max_delta = 0.0
|
||||
for token_index, (expected, observed) in enumerate(
|
||||
zip(reference_logprobs, actual_logprobs)
|
||||
):
|
||||
self.assertEqual(
|
||||
int(observed[1]),
|
||||
int(expected[1]),
|
||||
f"{label}: logprob token ID diverged at output {token_index}",
|
||||
)
|
||||
max_delta = max(max_delta, abs(float(observed[0]) - float(expected[0])))
|
||||
self.assertLessEqual(
|
||||
max_delta,
|
||||
LOGPROB_ATOL,
|
||||
f"{label}: max output logprob delta {max_delta:.6f}",
|
||||
)
|
||||
|
||||
def test_monolithic_pd_token_and_logprob_parity(self):
|
||||
self._flush_pd_caches()
|
||||
for target_length, prompt, reference in zip(
|
||||
PARITY_PROMPT_LENGTHS, self.parity_prompts, self.parity_references
|
||||
):
|
||||
with self.subTest(prompt_tokens=target_length):
|
||||
actual = self._generate(self.base_url, prompt, max_new_tokens=1)
|
||||
self.assertEqual(actual["meta_info"]["prompt_tokens"], target_length)
|
||||
self._assert_output_parity(
|
||||
reference,
|
||||
actual,
|
||||
label=f"sequential prompt_tokens={target_length}",
|
||||
)
|
||||
|
||||
def test_chunk_and_virtual_page_boundary_batch_parity(self):
|
||||
self._flush_pd_caches()
|
||||
actual_outputs = self._generate(
|
||||
self.base_url, self.parity_prompts, max_new_tokens=1
|
||||
)
|
||||
self.assertEqual(len(actual_outputs), len(self.parity_batch_references))
|
||||
for target_length, reference, actual in zip(
|
||||
PARITY_PROMPT_LENGTHS, self.parity_batch_references, actual_outputs
|
||||
):
|
||||
with self.subTest(prompt_tokens=target_length):
|
||||
self.assertEqual(actual["meta_info"]["prompt_tokens"], target_length)
|
||||
self._assert_output_parity(
|
||||
reference,
|
||||
actual,
|
||||
label=f"batched prompt_tokens={target_length}",
|
||||
check_logprobs=False,
|
||||
)
|
||||
|
||||
def test_long_context_needle_parity(self):
|
||||
self._flush_pd_caches()
|
||||
for needle_depth, prompt, reference in zip(
|
||||
LONG_CONTEXT_DEPTHS, self.niah_prompts, self.niah_references
|
||||
):
|
||||
with self.subTest(
|
||||
prompt_tokens=LONG_CONTEXT_TOKENS, needle_depth=needle_depth
|
||||
):
|
||||
self.assertIn(NIAH_KEY, reference["text"])
|
||||
actual = self._generate(
|
||||
self.base_url,
|
||||
prompt,
|
||||
max_new_tokens=16,
|
||||
ignore_eos=False,
|
||||
)
|
||||
self.assertIn(NIAH_KEY, actual["text"])
|
||||
self._assert_output_parity(
|
||||
reference,
|
||||
actual,
|
||||
label=(
|
||||
f"niah prompt_tokens={LONG_CONTEXT_TOKENS} "
|
||||
f"depth={needle_depth}"
|
||||
),
|
||||
)
|
||||
|
||||
def _assert_batch_completes(self, batch_size: int):
|
||||
response = requests.post(
|
||||
self.base_url + "/generate",
|
||||
json={
|
||||
"text": [
|
||||
f"Reply with one short word for request {index}: the sky is"
|
||||
for index in range(batch_size)
|
||||
],
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": 8,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
},
|
||||
timeout=300,
|
||||
)
|
||||
response.raise_for_status()
|
||||
outputs = response.json()
|
||||
self.assertIsInstance(outputs, list)
|
||||
self.assertEqual(len(outputs), batch_size)
|
||||
self.assertTrue(all(output["text"].strip() for output in outputs))
|
||||
|
||||
def test_decode_cuda_graph_and_eager_batch(self):
|
||||
self._assert_batch_completes(2)
|
||||
self._assert_batch_completes(2)
|
||||
self._assert_batch_completes(65)
|
||||
|
||||
def test_decode_physical_capacity_sanity(self):
|
||||
response = requests.get(self.decode_url + "/server_info", timeout=30)
|
||||
response.raise_for_status()
|
||||
self.assertGreater(response.json()["max_total_num_tokens"], 0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -215,6 +215,10 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
||||
b"64",
|
||||
b"DRAM,DRAM",
|
||||
b"".join(struct.pack("Q", item_len) for item_len in [1024, 2048]),
|
||||
pack_int_lists([[4], [4, 5]], "I"),
|
||||
b"".join(struct.pack("I", layer_id) for layer_id in [2, 7]),
|
||||
b"4",
|
||||
b"3",
|
||||
]
|
||||
|
||||
info = KVArgsRegisterInfo.from_zmq(msg)
|
||||
@@ -236,6 +240,10 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
||||
self.assertEqual(info.dst_kv_mem_kinds, ["DRAM", "DRAM"])
|
||||
self.assertEqual(info.dst_state_item_lens, state_item_lens)
|
||||
self.assertEqual(info.dst_state_dim_per_tensor, state_dims)
|
||||
self.assertEqual(info.dst_dcp_size, 4)
|
||||
self.assertEqual(info.dst_dcp_rank, 3)
|
||||
self.assertEqual(info.dst_state_layer_ids, [[4], [4, 5]])
|
||||
self.assertEqual(info.dst_kv_layer_ids, [2, 7])
|
||||
self.assertIsNotNone(info.staging)
|
||||
self.assertEqual(info.staging.base_ptr, staging_ptr)
|
||||
self.assertEqual(info.staging.total_size, 1048576)
|
||||
@@ -262,6 +270,8 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
||||
self.assertEqual(info.dst_state_item_lens, [])
|
||||
self.assertEqual(info.dst_state_dim_per_tensor, [])
|
||||
self.assertEqual(info.dst_kv_item_lens, [256])
|
||||
self.assertEqual(info.dst_dcp_size, 1)
|
||||
self.assertEqual(info.dst_dcp_rank, 0)
|
||||
self.assertIsNone(info.staging)
|
||||
|
||||
|
||||
@@ -445,6 +455,12 @@ class TestNixlTransferWorker(CustomTestCase):
|
||||
staging=None,
|
||||
kv_xfer_segments=None,
|
||||
dst_homogeneous_mem_kind="VRAM",
|
||||
# Non-DCP peer. Without this the worker raises AttributeError
|
||||
# and lands in the same Failed status the assertions expect,
|
||||
# so the transfer path would go unexercised.
|
||||
requires_dcp_relayout=False,
|
||||
dcp_dst_region_indices=None,
|
||||
dcp_token_item_lens=None,
|
||||
)
|
||||
}
|
||||
mgr.req_to_decode_prefix_len = {room: 4}
|
||||
@@ -496,6 +512,7 @@ class TestNixlTransferWorker(CustomTestCase):
|
||||
self.assertNotIn(room, mgr.transfer_infos)
|
||||
self.assertNotIn(room, mgr.req_to_decode_prefix_len)
|
||||
mgr.send_aux.assert_called_once()
|
||||
self.assertEqual(mgr.send_aux.call_args.args[-1], "21_aux_nokv_0_0")
|
||||
|
||||
def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
|
||||
self,
|
||||
|
||||
@@ -8,6 +8,7 @@ from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import sglang.srt.server_args as server_args_module
|
||||
from sglang.srt.arg_groups import pd_disaggregation_hook
|
||||
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
|
||||
from sglang.srt.entrypoints.sidecar import (
|
||||
SGLANG_GRPC_ENDPOINT_ENV,
|
||||
@@ -219,6 +220,56 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
||||
server_args = self._load_balance_args(disaggregation_mode="decode")
|
||||
self.assertEqual(server_args.load_balance_method, "round_robin")
|
||||
|
||||
def test_pd_prefill_dcp_warns_about_performance(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
disaggregation_mode="prefill",
|
||||
dcp_size=4,
|
||||
)
|
||||
with self.assertLogs(pd_disaggregation_hook.logger, level="WARNING") as logs:
|
||||
server_args._handle_pd_disaggregation()
|
||||
self.assertIn("without improving prefill performance", "\n".join(logs.output))
|
||||
|
||||
def test_pd_decode_dcp_forces_chunk_cache(self):
|
||||
server_args = self._load_balance_args(
|
||||
disaggregation_mode="decode",
|
||||
disaggregation_transfer_backend="mooncake",
|
||||
dcp_size=4,
|
||||
)
|
||||
self.assertTrue(server_args.disable_radix_cache)
|
||||
|
||||
def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
disaggregation_mode="decode",
|
||||
disaggregation_transfer_backend="fake",
|
||||
dcp_size=4,
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "mooncake or nixl"):
|
||||
server_args._handle_pd_disaggregation()
|
||||
|
||||
def test_pd_decode_dcp_rejects_radix_cache(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
disaggregation_mode="decode",
|
||||
disaggregation_transfer_backend="nixl",
|
||||
disaggregation_decode_enable_radix_cache=True,
|
||||
dcp_size=4,
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "currently requires chunk cache"):
|
||||
server_args._handle_pd_disaggregation()
|
||||
|
||||
def test_pd_decode_dcp_rejects_hierarchical_cache(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
disaggregation_mode="decode",
|
||||
disaggregation_transfer_backend="nixl",
|
||||
enable_hierarchical_cache=True,
|
||||
dcp_size=4,
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "--enable-hierarchical-cache"):
|
||||
server_args._handle_pd_disaggregation()
|
||||
|
||||
def test_pd_decode_radix_cache_rejects_hisparse(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
|
||||
Reference in New Issue
Block a user