fix(disagg): support pipeline-parallel hybrid-linear transfer (#32270)
This commit is contained in:
@@ -40,6 +40,7 @@ class KVArgs:
|
|||||||
kv_data_ptrs: List[int]
|
kv_data_ptrs: List[int]
|
||||||
kv_data_lens: List[int]
|
kv_data_lens: List[int]
|
||||||
kv_item_lens: List[int]
|
kv_item_lens: List[int]
|
||||||
|
kv_layer_ids: List[int]
|
||||||
aux_data_ptrs: List[int]
|
aux_data_ptrs: List[int]
|
||||||
aux_data_lens: List[int]
|
aux_data_lens: List[int]
|
||||||
aux_item_lens: List[int]
|
aux_item_lens: List[int]
|
||||||
@@ -47,6 +48,7 @@ class KVArgs:
|
|||||||
state_data_ptrs: List[List[int]]
|
state_data_ptrs: List[List[int]]
|
||||||
state_data_lens: List[List[int]]
|
state_data_lens: List[List[int]]
|
||||||
state_item_lens: List[List[int]]
|
state_item_lens: List[List[int]]
|
||||||
|
state_layer_ids: List[List[int]]
|
||||||
# Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
|
# Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
|
||||||
state_dim_per_tensor: List[List[int]]
|
state_dim_per_tensor: List[List[int]]
|
||||||
# Number of rows before the slice axis in each per-slot state tensor.
|
# Number of rows before the slice axis in each per-slot state tensor.
|
||||||
|
|||||||
@@ -451,6 +451,12 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
kv_args.kv_data_ptrs = kv_data_ptrs
|
kv_args.kv_data_ptrs = kv_data_ptrs
|
||||||
kv_args.kv_data_lens = kv_data_lens
|
kv_args.kv_data_lens = kv_data_lens
|
||||||
kv_args.kv_item_lens = kv_item_lens
|
kv_args.kv_item_lens = kv_item_lens
|
||||||
|
kv_args.kv_layer_ids = (
|
||||||
|
self.token_to_kv_pool.get_kv_layer_ids()
|
||||||
|
if self.draft_token_to_kv_pool is None
|
||||||
|
and hasattr(self.token_to_kv_pool, "get_kv_layer_ids")
|
||||||
|
else []
|
||||||
|
)
|
||||||
if self.transfer_backend == TransferBackend.NIXL:
|
if self.transfer_backend == TransferBackend.NIXL:
|
||||||
kv_args.kv_data_mem_kinds = kv_data_mem_kinds
|
kv_args.kv_data_mem_kinds = kv_data_mem_kinds
|
||||||
kv_args.page_size = self.token_to_kv_pool.page_size
|
kv_args.page_size = self.token_to_kv_pool.page_size
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ from sglang.srt.disaggregation.mooncake.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import (
|
from sglang.srt.disaggregation.utils import (
|
||||||
DisaggregationMode,
|
DisaggregationMode,
|
||||||
|
build_transfer_entry_pairs,
|
||||||
compute_mamba_state_slice_byte_blocks,
|
compute_mamba_state_slice_byte_blocks,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
|
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
|
||||||
@@ -128,6 +129,8 @@ class KVArgsRegisterInfo:
|
|||||||
# for mamba state different tp slice transfer
|
# for mamba state different tp slice transfer
|
||||||
dst_state_item_lens: List[List[int]]
|
dst_state_item_lens: List[List[int]]
|
||||||
dst_state_dim_per_tensor: List[List[int]]
|
dst_state_dim_per_tensor: List[List[int]]
|
||||||
|
dst_kv_layer_ids: List[int]
|
||||||
|
dst_state_layer_ids: List[List[int]]
|
||||||
# Note: always put the staging field at the final (since the staging field is optional and contains multiple inputs)
|
# Note: always put the staging field at the final (since the staging field is optional and contains multiple inputs)
|
||||||
staging: Optional[StagingRegisterInfo] = None
|
staging: Optional[StagingRegisterInfo] = None
|
||||||
|
|
||||||
@@ -150,8 +153,18 @@ class KVArgsRegisterInfo:
|
|||||||
dst_state_dim_per_tensor=(
|
dst_state_dim_per_tensor=(
|
||||||
unpack_int_lists(msg[11], "I") if len(msg) > 11 else []
|
unpack_int_lists(msg[11], "I") if len(msg) > 11 else []
|
||||||
),
|
),
|
||||||
|
dst_kv_layer_ids=(
|
||||||
|
list(struct.unpack(f"{len(msg[12]) // 4}I", msg[12]))
|
||||||
|
if len(msg) > 12 and msg[12] != b""
|
||||||
|
else []
|
||||||
|
),
|
||||||
|
dst_state_layer_ids=(
|
||||||
|
unpack_int_lists(msg[13], "I")
|
||||||
|
if len(msg) > 13 and msg[13] != b""
|
||||||
|
else []
|
||||||
|
),
|
||||||
# Note: always put the staging field at the final
|
# Note: always put the staging field at the final
|
||||||
staging=StagingRegisterInfo.from_zmq_fields(msg, 12),
|
staging=StagingRegisterInfo.from_zmq_fields(msg, 14),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -594,6 +607,8 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
executor: concurrent.futures.ThreadPoolExecutor,
|
executor: concurrent.futures.ThreadPoolExecutor,
|
||||||
state_type: Optional[StateType] = None,
|
state_type: Optional[StateType] = None,
|
||||||
force_flat: bool = False,
|
force_flat: bool = False,
|
||||||
|
src_layer_ids: Optional[List[int]] = None,
|
||||||
|
dst_layer_ids: Optional[List[int]] = None,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""
|
"""
|
||||||
Generic KV cache transfer supporting both MHA and MLA architectures.
|
Generic KV cache transfer supporting both MHA and MLA architectures.
|
||||||
@@ -612,17 +627,33 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
|
|
||||||
# Decode pp size should be equal to prefill pp size or 1
|
# Decode pp size should be equal to prefill pp size or 1
|
||||||
if self.is_mla_backend or self.is_hybrid_mla_backend or force_flat:
|
if self.is_mla_backend or self.is_hybrid_mla_backend or force_flat:
|
||||||
src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = (
|
# Layer IDs map PP-local buffers to global decode entries.
|
||||||
self.get_mla_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs, state_type)
|
# Registrations without them retain the existing PP mapping.
|
||||||
)
|
if src_layer_ids or dst_layer_ids:
|
||||||
layers_params = [
|
pairs = build_transfer_entry_pairs(
|
||||||
(
|
src_layer_ids,
|
||||||
src_kv_ptrs[layer_id],
|
dst_layer_ids,
|
||||||
dst_kv_ptrs[layer_id],
|
len(src_data_ptrs),
|
||||||
item_lens[layer_id],
|
len(dst_data_ptrs),
|
||||||
|
allow_positional_fallback=self.pp_size == 1,
|
||||||
)
|
)
|
||||||
for layer_id in range(layers_current_pp_stage)
|
layers_params = [
|
||||||
]
|
(src_data_ptrs[i], dst_data_ptrs[j], item_lens[i]) for i, j in pairs
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = (
|
||||||
|
self.get_mla_kv_ptrs_with_pp(
|
||||||
|
src_data_ptrs, dst_data_ptrs, state_type
|
||||||
|
)
|
||||||
|
)
|
||||||
|
layers_params = [
|
||||||
|
(
|
||||||
|
src_kv_ptrs[layer_id],
|
||||||
|
dst_kv_ptrs[layer_id],
|
||||||
|
item_lens[layer_id],
|
||||||
|
)
|
||||||
|
for layer_id in range(layers_current_pp_stage)
|
||||||
|
]
|
||||||
else:
|
else:
|
||||||
src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = (
|
src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = (
|
||||||
self.get_mha_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs)
|
self.get_mha_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs)
|
||||||
@@ -704,6 +735,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
dst_kv_ptrs: list[int],
|
dst_kv_ptrs: list[int],
|
||||||
dst_kv_indices: npt.NDArray[np.int32],
|
dst_kv_indices: npt.NDArray[np.int32],
|
||||||
executor: concurrent.futures.ThreadPoolExecutor,
|
executor: concurrent.futures.ThreadPoolExecutor,
|
||||||
|
dst_layer_ids: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
return self._send_kvcache_generic(
|
return self._send_kvcache_generic(
|
||||||
mooncake_session_id=mooncake_session_id,
|
mooncake_session_id=mooncake_session_id,
|
||||||
@@ -713,6 +745,8 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
prefill_data_indices=prefill_kv_indices,
|
prefill_data_indices=prefill_kv_indices,
|
||||||
dst_data_indices=dst_kv_indices,
|
dst_data_indices=dst_kv_indices,
|
||||||
executor=executor,
|
executor=executor,
|
||||||
|
src_layer_ids=self.kv_args.kv_layer_ids,
|
||||||
|
dst_layer_ids=dst_layer_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
def send_kvcache_slice(
|
def send_kvcache_slice(
|
||||||
@@ -993,6 +1027,10 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
src_slice_outer_counts = (
|
src_slice_outer_counts = (
|
||||||
src_slice_outer_counts[i] if i < len(src_slice_outer_counts) else []
|
src_slice_outer_counts[i] if i < len(src_slice_outer_counts) else []
|
||||||
)
|
)
|
||||||
|
src_state_layer_ids = self.kv_args.state_layer_ids
|
||||||
|
src_state_layer_ids = (
|
||||||
|
src_state_layer_ids[i] if i < len(src_state_layer_ids) else []
|
||||||
|
)
|
||||||
if target_rank_registration_info is not None:
|
if target_rank_registration_info is not None:
|
||||||
dst_data_ptrs = (
|
dst_data_ptrs = (
|
||||||
target_rank_registration_info.dst_state_data_ptrs[i]
|
target_rank_registration_info.dst_state_data_ptrs[i]
|
||||||
@@ -1009,8 +1047,14 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
if i < len(target_rank_registration_info.dst_state_dim_per_tensor)
|
if i < len(target_rank_registration_info.dst_state_dim_per_tensor)
|
||||||
else []
|
else []
|
||||||
)
|
)
|
||||||
|
dst_state_layer_ids = (
|
||||||
|
target_rank_registration_info.dst_state_layer_ids[i]
|
||||||
|
if i < len(target_rank_registration_info.dst_state_layer_ids)
|
||||||
|
else []
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
dst_data_ptrs, dst_item_lens, dst_dim_per_tensor = [], [], []
|
dst_data_ptrs, dst_item_lens, dst_dim_per_tensor = [], [], []
|
||||||
|
dst_state_layer_ids = []
|
||||||
dst_indices = (
|
dst_indices = (
|
||||||
req.dst_state_indices[i] if i < len(req.dst_state_indices) else []
|
req.dst_state_indices[i] if i < len(req.dst_state_indices) else []
|
||||||
)
|
)
|
||||||
@@ -1036,6 +1080,8 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
target_rank_registration_info.dst_attn_tp_size,
|
target_rank_registration_info.dst_attn_tp_size,
|
||||||
src_conv_shard_groups,
|
src_conv_shard_groups,
|
||||||
src_slice_outer_counts,
|
src_slice_outer_counts,
|
||||||
|
src_state_layer_ids,
|
||||||
|
dst_state_layer_ids,
|
||||||
)
|
)
|
||||||
or rc
|
or rc
|
||||||
)
|
)
|
||||||
@@ -1048,6 +1094,8 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
src_item_lens,
|
src_item_lens,
|
||||||
dst_data_ptrs,
|
dst_data_ptrs,
|
||||||
dst_indices,
|
dst_indices,
|
||||||
|
src_state_layer_ids,
|
||||||
|
dst_state_layer_ids,
|
||||||
)
|
)
|
||||||
or rc
|
or rc
|
||||||
)
|
)
|
||||||
@@ -1149,11 +1197,21 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
src_state_item_lens: list[int],
|
src_state_item_lens: list[int],
|
||||||
dst_state_data_ptrs: list[int],
|
dst_state_data_ptrs: list[int],
|
||||||
dst_mamba_index: list,
|
dst_mamba_index: list,
|
||||||
|
src_layer_ids: Optional[List[int]] = None,
|
||||||
|
dst_layer_ids: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
assert len(prefill_mamba_index) == 1, "Mamba should have single state index"
|
assert len(prefill_mamba_index) == 1, "Mamba should have single state index"
|
||||||
|
|
||||||
transfer_blocks = []
|
transfer_blocks = []
|
||||||
for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
|
pairs = build_transfer_entry_pairs(
|
||||||
|
src_layer_ids or [],
|
||||||
|
dst_layer_ids or [],
|
||||||
|
len(src_state_data_ptrs),
|
||||||
|
len(dst_state_data_ptrs),
|
||||||
|
allow_positional_fallback=self.pp_size == 1,
|
||||||
|
)
|
||||||
|
for i, j in pairs:
|
||||||
|
dst_state_ptr = dst_state_data_ptrs[j]
|
||||||
length = src_state_item_lens[i]
|
length = src_state_item_lens[i]
|
||||||
src_addr = src_state_data_ptrs[i] + length * int(prefill_mamba_index[0])
|
src_addr = src_state_data_ptrs[i] + length * int(prefill_mamba_index[0])
|
||||||
dst_addr = dst_state_ptr + length * int(dst_mamba_index[0])
|
dst_addr = dst_state_ptr + length * int(dst_mamba_index[0])
|
||||||
@@ -1176,6 +1234,8 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
dst_attn_tp_size: int,
|
dst_attn_tp_size: int,
|
||||||
src_state_conv_shard_groups: list = None,
|
src_state_conv_shard_groups: list = None,
|
||||||
src_state_slice_outer_counts: list[int] = None,
|
src_state_slice_outer_counts: list[int] = None,
|
||||||
|
src_layer_ids: Optional[List[int]] = None,
|
||||||
|
dst_layer_ids: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
"""Transfer Mamba states with TP slice support.
|
"""Transfer Mamba states with TP slice support.
|
||||||
|
|
||||||
@@ -1205,17 +1265,27 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
src_state_item_lens,
|
src_state_item_lens,
|
||||||
dst_state_data_ptrs,
|
dst_state_data_ptrs,
|
||||||
dst_mamba_index,
|
dst_mamba_index,
|
||||||
|
src_layer_ids,
|
||||||
|
dst_layer_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size
|
local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size
|
||||||
dst_tp_rank_in_group = dst_tp_rank % dst_attn_tp_size
|
dst_tp_rank_in_group = dst_tp_rank % dst_attn_tp_size
|
||||||
|
|
||||||
transfer_blocks = []
|
transfer_blocks = []
|
||||||
for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
|
pairs = build_transfer_entry_pairs(
|
||||||
|
src_layer_ids or [],
|
||||||
|
dst_layer_ids or [],
|
||||||
|
len(src_state_data_ptrs),
|
||||||
|
len(dst_state_data_ptrs),
|
||||||
|
allow_positional_fallback=self.pp_size == 1,
|
||||||
|
)
|
||||||
|
for i, j in pairs:
|
||||||
|
dst_state_ptr = dst_state_data_ptrs[j]
|
||||||
src_item_len = src_state_item_lens[i]
|
src_item_len = src_state_item_lens[i]
|
||||||
dst_item_len = dst_state_item_lens[i]
|
dst_item_len = dst_state_item_lens[j]
|
||||||
src_dim = src_state_dim_per_tensor[i]
|
src_dim = src_state_dim_per_tensor[i]
|
||||||
dst_dim = dst_state_dim_per_tensor[i]
|
dst_dim = dst_state_dim_per_tensor[j]
|
||||||
|
|
||||||
conv_shard_groups = (
|
conv_shard_groups = (
|
||||||
src_state_conv_shard_groups[i]
|
src_state_conv_shard_groups[i]
|
||||||
@@ -1370,7 +1440,11 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags(
|
skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags(
|
||||||
target_rank_registration_info
|
target_rank_registration_info
|
||||||
)
|
)
|
||||||
if len(kv_chunk.prefill_kv_indices) == 0 or skip_kv:
|
if (
|
||||||
|
len(kv_chunk.prefill_kv_indices) == 0
|
||||||
|
or not self.kv_args.kv_data_ptrs
|
||||||
|
or skip_kv
|
||||||
|
):
|
||||||
ret = 0
|
ret = 0
|
||||||
elif (
|
elif (
|
||||||
self.is_mla_backend
|
self.is_mla_backend
|
||||||
@@ -1384,6 +1458,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
target_rank_registration_info.dst_kv_ptrs,
|
target_rank_registration_info.dst_kv_ptrs,
|
||||||
chunked_dst_kv_indice,
|
chunked_dst_kv_indice,
|
||||||
executor,
|
executor,
|
||||||
|
target_rank_registration_info.dst_kv_layer_ids,
|
||||||
)
|
)
|
||||||
elif (
|
elif (
|
||||||
self.enable_staging
|
self.enable_staging
|
||||||
@@ -1919,9 +1994,20 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
packed_state_dim_per_tensor = pack_int_lists(
|
packed_state_dim_per_tensor = pack_int_lists(
|
||||||
getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I"
|
getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I"
|
||||||
)
|
)
|
||||||
|
packed_state_layer_ids = pack_int_lists(
|
||||||
|
self.kv_mgr.kv_args.state_layer_ids, "I"
|
||||||
|
)
|
||||||
|
packed_kv_layer_ids = b"".join(
|
||||||
|
struct.pack("I", layer_id)
|
||||||
|
for layer_id in self.kv_mgr.kv_args.kv_layer_ids
|
||||||
|
)
|
||||||
# Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1
|
# Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1
|
||||||
tp_rank = self.kv_mgr.kv_args.engine_rank
|
tp_rank = self.kv_mgr.kv_args.engine_rank
|
||||||
kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0]
|
kv_item_len = (
|
||||||
|
self.kv_mgr.kv_args.kv_item_lens[0]
|
||||||
|
if self.kv_mgr.kv_args.kv_item_lens
|
||||||
|
else 0
|
||||||
|
)
|
||||||
dst_tp_rank = str(tp_rank).encode("ascii")
|
dst_tp_rank = str(tp_rank).encode("ascii")
|
||||||
dst_attn_tp_size = str(self.kv_mgr.attn_tp_size).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_kv_item_len = str(kv_item_len).encode("ascii")
|
||||||
@@ -1953,6 +2039,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
dst_kv_item_len,
|
dst_kv_item_len,
|
||||||
packed_state_item_lens,
|
packed_state_item_lens,
|
||||||
packed_state_dim_per_tensor,
|
packed_state_dim_per_tensor,
|
||||||
|
packed_kv_layer_ids,
|
||||||
|
packed_state_layer_ids,
|
||||||
packed_staging_base_ptr,
|
packed_staging_base_ptr,
|
||||||
staging_total_size_str,
|
staging_total_size_str,
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ from sglang.srt.disaggregation.common.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import (
|
from sglang.srt.disaggregation.utils import (
|
||||||
DisaggregationMode,
|
DisaggregationMode,
|
||||||
|
build_transfer_entry_pairs,
|
||||||
compute_mamba_state_slice_byte_blocks,
|
compute_mamba_state_slice_byte_blocks,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -209,9 +210,11 @@ class KVArgsRegisterInfo:
|
|||||||
decode_tp_rank: int
|
decode_tp_rank: int
|
||||||
dst_kv_item_len: int
|
dst_kv_item_len: int
|
||||||
dst_kv_item_lens: list[int]
|
dst_kv_item_lens: list[int]
|
||||||
|
dst_kv_layer_ids: list[int] = dataclasses.field(default_factory=list)
|
||||||
dst_num_slots: Optional[int] = None
|
dst_num_slots: Optional[int] = None
|
||||||
dst_state_item_lens: List[List[int]] = dataclasses.field(default_factory=list)
|
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)
|
dst_state_dim_per_tensor: List[List[int]] = dataclasses.field(default_factory=list)
|
||||||
|
dst_state_layer_ids: List[List[int]] = dataclasses.field(default_factory=list)
|
||||||
dst_homogeneous_mem_kind: Optional[str] = None
|
dst_homogeneous_mem_kind: Optional[str] = None
|
||||||
kv_xfer_segments: Optional[List[_KVXferPreparedSegment]] = None
|
kv_xfer_segments: Optional[List[_KVXferPreparedSegment]] = None
|
||||||
# Keep last: optional, parsed from a variable-length tail of the ZMQ
|
# Keep last: optional, parsed from a variable-length tail of the ZMQ
|
||||||
@@ -249,6 +252,14 @@ class KVArgsRegisterInfo:
|
|||||||
dst_num_slots = (
|
dst_num_slots = (
|
||||||
int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else None
|
int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else None
|
||||||
)
|
)
|
||||||
|
dst_state_layer_ids = (
|
||||||
|
unpack_int_lists(msg[19], "I") if len(msg) > 19 and len(msg[19]) > 0 else []
|
||||||
|
)
|
||||||
|
dst_kv_layer_ids = (
|
||||||
|
list(struct.unpack(f"{len(msg[20]) // 4}I", msg[20]))
|
||||||
|
if len(msg) > 20 and msg[20] != b""
|
||||||
|
else []
|
||||||
|
)
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
room=str(msg[0].decode("ascii")),
|
room=str(msg[0].decode("ascii")),
|
||||||
@@ -265,9 +276,11 @@ class KVArgsRegisterInfo:
|
|||||||
decode_tp_rank=int(msg[10].decode("ascii")),
|
decode_tp_rank=int(msg[10].decode("ascii")),
|
||||||
dst_kv_item_len=dst_kv_item_len,
|
dst_kv_item_len=dst_kv_item_len,
|
||||||
dst_kv_item_lens=dst_kv_item_lens,
|
dst_kv_item_lens=dst_kv_item_lens,
|
||||||
|
dst_kv_layer_ids=dst_kv_layer_ids,
|
||||||
dst_num_slots=dst_num_slots,
|
dst_num_slots=dst_num_slots,
|
||||||
dst_state_item_lens=dst_state_item_lens,
|
dst_state_item_lens=dst_state_item_lens,
|
||||||
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
|
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
|
||||||
|
dst_state_layer_ids=dst_state_layer_ids,
|
||||||
staging=StagingRegisterInfo.from_zmq_fields(msg, 14),
|
staging=StagingRegisterInfo.from_zmq_fields(msg, 14),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -360,6 +373,9 @@ class NixlKVManager(CommonKVManager):
|
|||||||
is_mla_backend: Optional[bool] = False,
|
is_mla_backend: Optional[bool] = False,
|
||||||
):
|
):
|
||||||
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
||||||
|
self.transfer_source_rank = (
|
||||||
|
self.kv_args.pp_rank * self.server_args.tp_size + self.kv_args.engine_rank
|
||||||
|
)
|
||||||
self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds(
|
self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds(
|
||||||
getattr(self.kv_args, "kv_data_mem_kinds", None),
|
getattr(self.kv_args, "kv_data_mem_kinds", None),
|
||||||
len(self.kv_args.kv_data_ptrs),
|
len(self.kv_args.kv_data_ptrs),
|
||||||
@@ -367,6 +383,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
self.src_mem_kind = (
|
self.src_mem_kind = (
|
||||||
_homogeneous_kv_mem_kind(self.kv_args.kv_data_mem_kinds, "source")
|
_homogeneous_kv_mem_kind(self.kv_args.kv_data_mem_kinds, "source")
|
||||||
if disaggregation_mode == DisaggregationMode.PREFILL
|
if disaggregation_mode == DisaggregationMode.PREFILL
|
||||||
|
and self.kv_args.kv_data_mem_kinds
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
@@ -432,9 +449,10 @@ class NixlKVManager(CommonKVManager):
|
|||||||
self._num_slots_src: int = 0
|
self._num_slots_src: int = 0
|
||||||
|
|
||||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||||
self._num_slots_src = (
|
if self.kv_args.kv_item_lens:
|
||||||
self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0]
|
self._num_slots_src = (
|
||||||
)
|
self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0]
|
||||||
|
)
|
||||||
transfer_queue_size = envs.SGLANG_DISAGGREGATION_QUEUE_SIZE.get()
|
transfer_queue_size = envs.SGLANG_DISAGGREGATION_QUEUE_SIZE.get()
|
||||||
self.transfer_queues: List[FastQueue] = [
|
self.transfer_queues: List[FastQueue] = [
|
||||||
FastQueue() for _ in range(transfer_queue_size)
|
FastQueue() for _ in range(transfer_queue_size)
|
||||||
@@ -919,9 +937,6 @@ class NixlKVManager(CommonKVManager):
|
|||||||
peer_info.kv_xfer_segments = prepared_segments
|
peer_info.kv_xfer_segments = prepared_segments
|
||||||
|
|
||||||
def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
|
def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
|
||||||
assert self.src_mem_kind is not None
|
|
||||||
src_mem_kind = self.src_mem_kind
|
|
||||||
|
|
||||||
# If prefill does not run speculative decoding (the usual case),
|
# If prefill does not run speculative decoding (the usual case),
|
||||||
# decode with speculative decoding will have more kv items.
|
# decode with speculative decoding will have more kv items.
|
||||||
# Prefill having more kv items is impossible.
|
# Prefill having more kv items is impossible.
|
||||||
@@ -932,6 +947,10 @@ class NixlKVManager(CommonKVManager):
|
|||||||
"NIXL PD transfer: decode registered fewer KV regions "
|
"NIXL PD transfer: decode registered fewer KV regions "
|
||||||
f"({n_dst}) than prefill ({n_src}); unexpected geometry"
|
f"({n_dst}) than prefill ({n_src}); unexpected geometry"
|
||||||
)
|
)
|
||||||
|
if n_src == 0:
|
||||||
|
return
|
||||||
|
assert self.src_mem_kind is not None
|
||||||
|
src_mem_kind = self.src_mem_kind
|
||||||
decode_only_spec_dec = n_dst > n_src
|
decode_only_spec_dec = n_dst > n_src
|
||||||
if (
|
if (
|
||||||
self.is_mla_backend
|
self.is_mla_backend
|
||||||
@@ -979,8 +998,16 @@ class NixlKVManager(CommonKVManager):
|
|||||||
else self._num_slots_src
|
else self._num_slots_src
|
||||||
)
|
)
|
||||||
|
|
||||||
dst_kv_ptrs = peer_info.dst_kv_ptrs[:n_src]
|
pairs = build_transfer_entry_pairs(
|
||||||
dst_kv_item_lens = peer_info.dst_kv_item_lens[:n_src]
|
self.kv_args.kv_layer_ids,
|
||||||
|
peer_info.dst_kv_layer_ids,
|
||||||
|
n_src,
|
||||||
|
n_dst,
|
||||||
|
allow_positional_fallback=self.pp_size == 1,
|
||||||
|
)
|
||||||
|
dst_indices = [j for _, j in pairs]
|
||||||
|
dst_kv_ptrs = [peer_info.dst_kv_ptrs[j] for j in dst_indices]
|
||||||
|
dst_kv_item_lens = [peer_info.dst_kv_item_lens[j] for j in dst_indices]
|
||||||
dst_kv_data_lens = [
|
dst_kv_data_lens = [
|
||||||
item_len * dst_num_slots for item_len in dst_kv_item_lens
|
item_len * dst_num_slots for item_len in dst_kv_item_lens
|
||||||
]
|
]
|
||||||
@@ -1058,7 +1085,10 @@ 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_chunk=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
|
||||||
|
and self.kv_args.kv_data_ptrs
|
||||||
|
):
|
||||||
chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice]
|
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
|
# NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices
|
||||||
@@ -1077,7 +1107,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
|
|
||||||
notif = (
|
notif = (
|
||||||
f"{req.room}_kv_{kv_chunk.chunk_id}"
|
f"{req.room}_kv_{kv_chunk.chunk_id}"
|
||||||
f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
|
f"_{int(kv_chunk.is_last_chunk)}_{self.transfer_source_rank}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Decide which kv send path to use:
|
# Decide which kv send path to use:
|
||||||
@@ -1163,11 +1193,12 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_info.dst_state_data_ptrs,
|
dst_info.dst_state_data_ptrs,
|
||||||
req.dst_state_indices,
|
req.dst_state_indices,
|
||||||
dst_info.gpu_id,
|
dst_info.gpu_id,
|
||||||
f"{req.room}_state_{self.kv_args.engine_rank}",
|
f"{req.room}_state_{self.transfer_source_rank}",
|
||||||
decode_tp_size,
|
decode_tp_size,
|
||||||
decode_tp_rank=dst_info.decode_tp_rank,
|
decode_tp_rank=dst_info.decode_tp_rank,
|
||||||
dst_state_item_lens=dst_info.dst_state_item_lens,
|
dst_state_item_lens=dst_info.dst_state_item_lens,
|
||||||
dst_state_dim_per_tensor=dst_info.dst_state_dim_per_tensor,
|
dst_state_dim_per_tensor=dst_info.dst_state_dim_per_tensor,
|
||||||
|
dst_state_layer_ids=dst_info.dst_state_layer_ids,
|
||||||
)
|
)
|
||||||
handles.extend(
|
handles.extend(
|
||||||
h for h in state_xfer_handles if h is not None
|
h for h in state_xfer_handles if h is not None
|
||||||
@@ -1175,12 +1206,13 @@ class NixlKVManager(CommonKVManager):
|
|||||||
|
|
||||||
if kv_chunk.prefill_aux_index is None:
|
if kv_chunk.prefill_aux_index is None:
|
||||||
raise RuntimeError("Missing aux index for last chunk")
|
raise RuntimeError("Missing aux index for last chunk")
|
||||||
# When no KV pages were sent (decode-side cache hit),
|
# A no-KV notification still identifies its PP source.
|
||||||
# encode pp_rank in aux notif so receiver can mark
|
if (
|
||||||
# expected_kvs_per_pp[pp_rank] = 0.
|
len(kv_chunk.prefill_kv_indices) == 0
|
||||||
if len(kv_chunk.prefill_kv_indices) == 0:
|
or not self.kv_args.kv_data_ptrs
|
||||||
|
):
|
||||||
aux_notif = (
|
aux_notif = (
|
||||||
f"{req.room}_aux_nokv_{self.kv_args.engine_rank}"
|
f"{req.room}_aux_nokv_{self.transfer_source_rank}"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
aux_notif = f"{req.room}_aux"
|
aux_notif = f"{req.room}_aux"
|
||||||
@@ -1267,8 +1299,6 @@ class NixlKVManager(CommonKVManager):
|
|||||||
f"NIXL memory registration failed for {mem_kind} kv tensors"
|
f"NIXL memory registration failed for {mem_kind} kv tensors"
|
||||||
)
|
)
|
||||||
self.kv_descs.append(kv_descs)
|
self.kv_descs.append(kv_descs)
|
||||||
if not self.kv_descs:
|
|
||||||
raise Exception("NIXL memory registration failed for kv tensors")
|
|
||||||
aux_addrs = []
|
aux_addrs = []
|
||||||
for aux_data_ptr, aux_data_len in zip(
|
for aux_data_ptr, aux_data_len in zip(
|
||||||
self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens
|
self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens
|
||||||
@@ -1766,7 +1796,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
|
|
||||||
notif_tag = (
|
notif_tag = (
|
||||||
f"{req.room}_stg_{kv_chunk.chunk_id}_{int(kv_chunk.is_last_chunk)}"
|
f"{req.room}_stg_{kv_chunk.chunk_id}_{int(kv_chunk.is_last_chunk)}"
|
||||||
f"_{self.kv_args.engine_rank}_{chunk_idx}"
|
f"_{self.transfer_source_rank}_{chunk_idx}"
|
||||||
f"_{page_start}_{num_pages}_{req.agent_name}"
|
f"_{page_start}_{num_pages}_{req.agent_name}"
|
||||||
)
|
)
|
||||||
handle = self.send_kvcache_staged(
|
handle = self.send_kvcache_staged(
|
||||||
@@ -1842,6 +1872,8 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_state_indices: List[int],
|
dst_state_indices: List[int],
|
||||||
dst_gpu_id: int,
|
dst_gpu_id: int,
|
||||||
notif: str,
|
notif: str,
|
||||||
|
src_layer_ids: list[int] = None,
|
||||||
|
dst_layer_ids: list[int] = None,
|
||||||
):
|
):
|
||||||
"""Transfer Mamba states via RDMA."""
|
"""Transfer Mamba states via RDMA."""
|
||||||
assert len(prefill_state_indices) == 1, "Mamba should have single state index"
|
assert len(prefill_state_indices) == 1, "Mamba should have single state index"
|
||||||
@@ -1852,7 +1884,15 @@ class NixlKVManager(CommonKVManager):
|
|||||||
src_addrs = []
|
src_addrs = []
|
||||||
dst_addrs = []
|
dst_addrs = []
|
||||||
|
|
||||||
for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
|
pairs = build_transfer_entry_pairs(
|
||||||
|
src_layer_ids or [],
|
||||||
|
dst_layer_ids or [],
|
||||||
|
len(src_state_data_ptrs),
|
||||||
|
len(dst_state_data_ptrs),
|
||||||
|
allow_positional_fallback=self.pp_size == 1,
|
||||||
|
)
|
||||||
|
for i, j in pairs:
|
||||||
|
dst_state_ptr = dst_state_data_ptrs[j]
|
||||||
length = src_state_item_lens[i]
|
length = src_state_item_lens[i]
|
||||||
if length == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0:
|
if length == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0:
|
||||||
continue
|
continue
|
||||||
@@ -1895,6 +1935,8 @@ class NixlKVManager(CommonKVManager):
|
|||||||
decode_tp_rank: int,
|
decode_tp_rank: int,
|
||||||
src_state_conv_shard_groups: list = None,
|
src_state_conv_shard_groups: list = None,
|
||||||
src_state_slice_outer_counts: list[int] = None,
|
src_state_slice_outer_counts: list[int] = None,
|
||||||
|
src_layer_ids: list[int] = None,
|
||||||
|
dst_layer_ids: list[int] = None,
|
||||||
):
|
):
|
||||||
"""Transfer Mamba states with TP slice support via RDMA.
|
"""Transfer Mamba states with TP slice support via RDMA.
|
||||||
|
|
||||||
@@ -1922,6 +1964,8 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_state_indices,
|
dst_state_indices,
|
||||||
dst_gpu_id,
|
dst_gpu_id,
|
||||||
notif,
|
notif,
|
||||||
|
src_layer_ids=src_layer_ids,
|
||||||
|
dst_layer_ids=dst_layer_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size
|
local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size
|
||||||
@@ -1930,13 +1974,21 @@ class NixlKVManager(CommonKVManager):
|
|||||||
src_addrs = []
|
src_addrs = []
|
||||||
dst_addrs = []
|
dst_addrs = []
|
||||||
|
|
||||||
for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
|
pairs = build_transfer_entry_pairs(
|
||||||
|
src_layer_ids or [],
|
||||||
|
dst_layer_ids or [],
|
||||||
|
len(src_state_data_ptrs),
|
||||||
|
len(dst_state_data_ptrs),
|
||||||
|
allow_positional_fallback=self.pp_size == 1,
|
||||||
|
)
|
||||||
|
for i, j in pairs:
|
||||||
|
dst_state_ptr = dst_state_data_ptrs[j]
|
||||||
src_item_len = src_state_item_lens[i]
|
src_item_len = src_state_item_lens[i]
|
||||||
dst_item_len = dst_state_item_lens[i]
|
dst_item_len = dst_state_item_lens[j]
|
||||||
if src_item_len == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0:
|
if src_item_len == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0:
|
||||||
continue
|
continue
|
||||||
src_dim = src_state_dim_per_tensor[i]
|
src_dim = src_state_dim_per_tensor[i]
|
||||||
dst_dim = dst_state_dim_per_tensor[i]
|
dst_dim = dst_state_dim_per_tensor[j]
|
||||||
|
|
||||||
conv_shard_groups = (
|
conv_shard_groups = (
|
||||||
src_state_conv_shard_groups[i]
|
src_state_conv_shard_groups[i]
|
||||||
@@ -2007,6 +2059,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
decode_tp_rank: int = 0,
|
decode_tp_rank: int = 0,
|
||||||
dst_state_item_lens: List[List[int]] | None = None,
|
dst_state_item_lens: List[List[int]] | None = None,
|
||||||
dst_state_dim_per_tensor: List[List[int]] | None = None,
|
dst_state_dim_per_tensor: List[List[int]] | None = None,
|
||||||
|
dst_state_layer_ids: List[List[int]] | None = None,
|
||||||
):
|
):
|
||||||
"""Send state per hybrid component, dispatching by state_type[i]."""
|
"""Send state per hybrid component, dispatching by state_type[i]."""
|
||||||
state_types = getattr(self.kv_args, "state_types", []) or []
|
state_types = getattr(self.kv_args, "state_types", []) or []
|
||||||
@@ -2021,8 +2074,10 @@ class NixlKVManager(CommonKVManager):
|
|||||||
src_state_slice_outer_counts = (
|
src_state_slice_outer_counts = (
|
||||||
getattr(self.kv_args, "state_slice_outer_counts", []) or []
|
getattr(self.kv_args, "state_slice_outer_counts", []) or []
|
||||||
)
|
)
|
||||||
|
src_state_layer_ids = self.kv_args.state_layer_ids
|
||||||
dst_state_item_lens = dst_state_item_lens or []
|
dst_state_item_lens = dst_state_item_lens or []
|
||||||
dst_state_dim_per_tensor = dst_state_dim_per_tensor or []
|
dst_state_dim_per_tensor = dst_state_dim_per_tensor or []
|
||||||
|
dst_state_layer_ids = dst_state_layer_ids or []
|
||||||
|
|
||||||
handles = []
|
handles = []
|
||||||
for i, st in enumerate(state_types):
|
for i, st in enumerate(state_types):
|
||||||
@@ -2046,12 +2101,14 @@ class NixlKVManager(CommonKVManager):
|
|||||||
if i < len(src_state_slice_outer_counts)
|
if i < len(src_state_slice_outer_counts)
|
||||||
else []
|
else []
|
||||||
)
|
)
|
||||||
|
src_lids = src_state_layer_ids[i] if i < len(src_state_layer_ids) else []
|
||||||
dst_ptrs = dst_state_data_ptrs[i] if i < len(dst_state_data_ptrs) else []
|
dst_ptrs = dst_state_data_ptrs[i] if i < len(dst_state_data_ptrs) else []
|
||||||
dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else []
|
dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else []
|
||||||
dst_lens = dst_state_item_lens[i] if i < len(dst_state_item_lens) else []
|
dst_lens = dst_state_item_lens[i] if i < len(dst_state_item_lens) else []
|
||||||
dst_dims = (
|
dst_dims = (
|
||||||
dst_state_dim_per_tensor[i] if i < len(dst_state_dim_per_tensor) else []
|
dst_state_dim_per_tensor[i] if i < len(dst_state_dim_per_tensor) else []
|
||||||
)
|
)
|
||||||
|
dst_lids = dst_state_layer_ids[i] if i < len(dst_state_layer_ids) else []
|
||||||
comp_notif = f"{notif}_{i}"
|
comp_notif = f"{notif}_{i}"
|
||||||
|
|
||||||
if st == StateType.MAMBA:
|
if st == StateType.MAMBA:
|
||||||
@@ -2072,6 +2129,8 @@ class NixlKVManager(CommonKVManager):
|
|||||||
decode_tp_rank,
|
decode_tp_rank,
|
||||||
src_conv,
|
src_conv,
|
||||||
src_outer_counts,
|
src_outer_counts,
|
||||||
|
src_layer_ids=src_lids,
|
||||||
|
dst_layer_ids=dst_lids,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
h = self._send_mamba_state(
|
h = self._send_mamba_state(
|
||||||
@@ -2083,6 +2142,8 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_indices,
|
dst_indices,
|
||||||
dst_gpu_id,
|
dst_gpu_id,
|
||||||
comp_notif,
|
comp_notif,
|
||||||
|
src_layer_ids=src_lids,
|
||||||
|
dst_layer_ids=dst_lids,
|
||||||
)
|
)
|
||||||
elif st in (
|
elif st in (
|
||||||
StateType.SWA,
|
StateType.SWA,
|
||||||
@@ -2709,6 +2770,10 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
struct.pack("Q", item_len)
|
struct.pack("Q", item_len)
|
||||||
for item_len in self.kv_mgr.kv_args.kv_item_lens
|
for item_len in self.kv_mgr.kv_args.kv_item_lens
|
||||||
)
|
)
|
||||||
|
packed_kv_layer_ids = b"".join(
|
||||||
|
struct.pack("I", layer_id)
|
||||||
|
for layer_id in self.kv_mgr.kv_args.kv_layer_ids
|
||||||
|
)
|
||||||
packed_aux_data_ptrs = b"".join(
|
packed_aux_data_ptrs = b"".join(
|
||||||
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.aux_data_ptrs
|
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.aux_data_ptrs
|
||||||
)
|
)
|
||||||
@@ -2721,6 +2786,9 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
packed_state_dim_per_tensor = pack_int_lists(
|
packed_state_dim_per_tensor = pack_int_lists(
|
||||||
getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I"
|
getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I"
|
||||||
)
|
)
|
||||||
|
packed_state_layer_ids = pack_int_lists(
|
||||||
|
self.kv_mgr.kv_args.state_layer_ids, "I"
|
||||||
|
)
|
||||||
|
|
||||||
# Include staging allocator metadata if available
|
# Include staging allocator metadata if available
|
||||||
if (
|
if (
|
||||||
@@ -2733,10 +2801,12 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
else:
|
else:
|
||||||
packed_staging_base_ptr = b""
|
packed_staging_base_ptr = b""
|
||||||
staging_total_size_str = b""
|
staging_total_size_str = b""
|
||||||
dst_num_slots = (
|
if self.kv_mgr.kv_args.kv_item_lens:
|
||||||
self.kv_mgr.kv_args.kv_data_lens[0]
|
dst_kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0]
|
||||||
// self.kv_mgr.kv_args.kv_item_lens[0]
|
dst_num_slots = self.kv_mgr.kv_args.kv_data_lens[0] // dst_kv_item_len
|
||||||
)
|
else:
|
||||||
|
dst_kv_item_len = 0
|
||||||
|
dst_num_slots = 0
|
||||||
|
|
||||||
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
||||||
try:
|
try:
|
||||||
@@ -2755,7 +2825,7 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
|
str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
|
||||||
str(self.kv_mgr.attn_tp_size).encode("ascii"),
|
str(self.kv_mgr.attn_tp_size).encode("ascii"),
|
||||||
str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
|
str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
|
||||||
str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"),
|
str(dst_kv_item_len).encode("ascii"),
|
||||||
packed_state_item_lens,
|
packed_state_item_lens,
|
||||||
packed_state_dim_per_tensor,
|
packed_state_dim_per_tensor,
|
||||||
packed_staging_base_ptr,
|
packed_staging_base_ptr,
|
||||||
@@ -2763,6 +2833,8 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
str(dst_num_slots).encode("ascii"),
|
str(dst_num_slots).encode("ascii"),
|
||||||
packed_kv_data_mem_kinds,
|
packed_kv_data_mem_kinds,
|
||||||
packed_kv_item_lens,
|
packed_kv_item_lens,
|
||||||
|
packed_state_layer_ids,
|
||||||
|
packed_kv_layer_ids,
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
except zmq.ZMQError:
|
except zmq.ZMQError:
|
||||||
|
|||||||
@@ -219,6 +219,12 @@ class PrefillBootstrapQueue:
|
|||||||
kv_args.kv_data_ptrs = kv_data_ptrs
|
kv_args.kv_data_ptrs = kv_data_ptrs
|
||||||
kv_args.kv_data_lens = kv_data_lens
|
kv_args.kv_data_lens = kv_data_lens
|
||||||
kv_args.kv_item_lens = kv_item_lens
|
kv_args.kv_item_lens = kv_item_lens
|
||||||
|
kv_args.kv_layer_ids = (
|
||||||
|
self.token_to_kv_pool.get_kv_layer_ids()
|
||||||
|
if self.draft_token_to_kv_pool is None
|
||||||
|
and hasattr(self.token_to_kv_pool, "get_kv_layer_ids")
|
||||||
|
else []
|
||||||
|
)
|
||||||
if not self.is_mla_backend:
|
if not self.is_mla_backend:
|
||||||
kv_args.kv_head_num = self.token_to_kv_pool.head_num
|
kv_args.kv_head_num = self.token_to_kv_pool.head_num
|
||||||
kv_args.total_kv_head_num = (
|
kv_args.total_kv_head_num = (
|
||||||
|
|||||||
@@ -855,6 +855,53 @@ def compute_mamba_state_slice_byte_blocks(
|
|||||||
return blocks
|
return blocks
|
||||||
|
|
||||||
|
|
||||||
|
def build_transfer_entry_pairs(
|
||||||
|
src_layer_ids: List[int],
|
||||||
|
dst_layer_ids: List[int],
|
||||||
|
n_src: int,
|
||||||
|
n_dst: int,
|
||||||
|
allow_positional_fallback: bool = False,
|
||||||
|
) -> List[Tuple[int, int]]:
|
||||||
|
"""Pair prefill-local transfer entries with decode entries by layer id."""
|
||||||
|
if n_src == 0:
|
||||||
|
return []
|
||||||
|
if bool(src_layer_ids) != bool(dst_layer_ids):
|
||||||
|
if not allow_positional_fallback:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Layer metadata must be provided by both PD peers or neither"
|
||||||
|
)
|
||||||
|
src_layer_ids = []
|
||||||
|
dst_layer_ids = []
|
||||||
|
if src_layer_ids:
|
||||||
|
if len(src_layer_ids) != n_src or len(dst_layer_ids) != n_dst:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Layer metadata length must match transfer entries: "
|
||||||
|
f"src metadata={len(src_layer_ids)} entries={n_src}, "
|
||||||
|
f"dst metadata={len(dst_layer_ids)} entries={n_dst}"
|
||||||
|
)
|
||||||
|
# Layer ids can repeat across tensor groups (for example K/V or multiple
|
||||||
|
# state tensors), so pair occurrences in order rather than by plain lookup.
|
||||||
|
dst_pos = {}
|
||||||
|
for j, lid in enumerate(dst_layer_ids):
|
||||||
|
dst_pos.setdefault(lid, deque()).append(j)
|
||||||
|
pairs = []
|
||||||
|
for i, lid in enumerate(src_layer_ids):
|
||||||
|
if not dst_pos.get(lid):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Decode peer is missing a transfer entry for model layer {lid}"
|
||||||
|
)
|
||||||
|
pairs.append((i, dst_pos[lid].popleft()))
|
||||||
|
return pairs
|
||||||
|
if n_dst < n_src or (n_src != n_dst and not allow_positional_fallback):
|
||||||
|
# Without layer ids a positional pairing would silently transfer the
|
||||||
|
# wrong layers (e.g. PP prefill peered with a stale decode server).
|
||||||
|
raise RuntimeError(
|
||||||
|
"PP-heterogeneous transfer requires layer ids on "
|
||||||
|
f"both peers; got src={n_src} dst={n_dst} entries"
|
||||||
|
)
|
||||||
|
return [(i, i) for i in range(n_src)]
|
||||||
|
|
||||||
|
|
||||||
def append_state_component(
|
def append_state_component(
|
||||||
kv_args: KVArgs,
|
kv_args: KVArgs,
|
||||||
state_type: StateType,
|
state_type: StateType,
|
||||||
@@ -864,6 +911,7 @@ def append_state_component(
|
|||||||
dim_per_tensor: Optional[List[int]] = None,
|
dim_per_tensor: Optional[List[int]] = None,
|
||||||
conv_shard_groups: Optional[List[Optional[List[int]]]] = None,
|
conv_shard_groups: Optional[List[Optional[List[int]]]] = None,
|
||||||
slice_outer_counts: Optional[List[int]] = None,
|
slice_outer_counts: Optional[List[int]] = None,
|
||||||
|
layer_ids: Optional[List[int]] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Append one state component. Caller orders state_types consistently
|
"""Append one state component. Caller orders state_types consistently
|
||||||
on prefill and decode sides."""
|
on prefill and decode sides."""
|
||||||
@@ -874,6 +922,7 @@ def append_state_component(
|
|||||||
kv_args.state_dim_per_tensor.append(dim_per_tensor or [])
|
kv_args.state_dim_per_tensor.append(dim_per_tensor or [])
|
||||||
kv_args.state_conv_shard_groups.append(conv_shard_groups or [])
|
kv_args.state_conv_shard_groups.append(conv_shard_groups or [])
|
||||||
kv_args.state_slice_outer_counts.append(slice_outer_counts or [])
|
kv_args.state_slice_outer_counts.append(slice_outer_counts or [])
|
||||||
|
kv_args.state_layer_ids.append(layer_ids or [])
|
||||||
|
|
||||||
|
|
||||||
def setup_state_kv_args(
|
def setup_state_kv_args(
|
||||||
@@ -903,6 +952,7 @@ def setup_state_kv_args(
|
|||||||
kv_args.state_item_lens = []
|
kv_args.state_item_lens = []
|
||||||
kv_args.state_dim_per_tensor = []
|
kv_args.state_dim_per_tensor = []
|
||||||
kv_args.state_slice_outer_counts = []
|
kv_args.state_slice_outer_counts = []
|
||||||
|
kv_args.state_layer_ids = []
|
||||||
kv_args.is_hybrid_mla_backend = False
|
kv_args.is_hybrid_mla_backend = False
|
||||||
kv_args.state_conv_shard_groups = []
|
kv_args.state_conv_shard_groups = []
|
||||||
|
|
||||||
@@ -971,6 +1021,9 @@ def setup_state_kv_args(
|
|||||||
if hasattr(token_to_kv_pool, "get_state_slice_outer_counts")
|
if hasattr(token_to_kv_pool, "get_state_slice_outer_counts")
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
# Global layer ids let the sender pair src/dst entries when the
|
||||||
|
# prefill PP stage registers only its own subset of mamba layers.
|
||||||
|
layer_ids = token_to_kv_pool.get_state_layer_ids()
|
||||||
append_state_component(
|
append_state_component(
|
||||||
kv_args,
|
kv_args,
|
||||||
StateType.MAMBA,
|
StateType.MAMBA,
|
||||||
@@ -980,6 +1033,7 @@ def setup_state_kv_args(
|
|||||||
dim,
|
dim,
|
||||||
conv_shard_groups,
|
conv_shard_groups,
|
||||||
slice_outer_counts,
|
slice_outer_counts,
|
||||||
|
layer_ids,
|
||||||
)
|
)
|
||||||
elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)):
|
elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)):
|
||||||
if draft_token_to_kv_pool is not None and isinstance(
|
if draft_token_to_kv_pool is not None and isinstance(
|
||||||
|
|||||||
@@ -473,6 +473,7 @@ class MambaPool:
|
|||||||
enable=enable_memory_saver
|
enable=enable_memory_saver
|
||||||
)
|
)
|
||||||
num_mamba_layers = len(mamba_layer_ids)
|
num_mamba_layers = len(mamba_layer_ids)
|
||||||
|
self.mamba_layer_ids = list(mamba_layer_ids)
|
||||||
|
|
||||||
self.size = size
|
self.size = size
|
||||||
self.device = device
|
self.device = device
|
||||||
@@ -874,6 +875,7 @@ class MambaPool:
|
|||||||
return (
|
return (
|
||||||
not _is_npu
|
not _is_npu
|
||||||
and len(convs) > 0
|
and len(convs) > 0
|
||||||
|
and convs[0].shape[0] > 0
|
||||||
and convs[0].is_cuda
|
and convs[0].is_cuda
|
||||||
and all(c.dtype == torch.bfloat16 and c.is_contiguous() for c in convs)
|
and all(c.dtype == torch.bfloat16 and c.is_contiguous() for c in convs)
|
||||||
)
|
)
|
||||||
@@ -1070,6 +1072,16 @@ class MambaPool:
|
|||||||
dim_per_tensor += [sliceable_dim] * self.num_mamba_layers
|
dim_per_tensor += [sliceable_dim] * self.num_mamba_layers
|
||||||
return dim_per_tensor
|
return dim_per_tensor
|
||||||
|
|
||||||
|
def get_state_layer_ids(self):
|
||||||
|
"""Global model-layer id for each RDMA state entry.
|
||||||
|
|
||||||
|
Aligned element-wise with get_contiguous_buf_infos(), which flattens
|
||||||
|
the state list tensor-major x layer. Lets PD transfer match entries
|
||||||
|
by layer id when prefill (PP stage) holds a subset of the mamba layers.
|
||||||
|
"""
|
||||||
|
state_tensor_count = sum(1 for _ in self._iter_transfer_state_tensors())
|
||||||
|
return list(self.mamba_layer_ids) * state_tensor_count
|
||||||
|
|
||||||
def get_state_slice_outer_counts(self):
|
def get_state_slice_outer_counts(self):
|
||||||
"""Get the number of rows preceding each tensor's TP slice axis."""
|
"""Get the number of rows preceding each tensor's TP slice axis."""
|
||||||
outer_counts = []
|
outer_counts = []
|
||||||
@@ -3625,6 +3637,11 @@ class HybridLinearKVPool(KVCache):
|
|||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
return self.full_kv_pool.get_contiguous_buf_infos()
|
return self.full_kv_pool.get_contiguous_buf_infos()
|
||||||
|
|
||||||
|
def get_kv_layer_ids(self):
|
||||||
|
"""Global layer ids aligned with the full-attention KV buffers."""
|
||||||
|
layer_ids = list(self.full_attention_layer_id_mapping)
|
||||||
|
return layer_ids if self.use_mla else layer_ids * 2
|
||||||
|
|
||||||
def get_state_buf_infos(self):
|
def get_state_buf_infos(self):
|
||||||
mamba_data_ptrs, mamba_data_lens, mamba_item_lens = (
|
mamba_data_ptrs, mamba_data_lens, mamba_item_lens = (
|
||||||
self.mamba_pool.get_contiguous_buf_infos()
|
self.mamba_pool.get_contiguous_buf_infos()
|
||||||
@@ -3635,6 +3652,10 @@ class HybridLinearKVPool(KVCache):
|
|||||||
"""Get the sliceable dimension size for each mamba state tensor."""
|
"""Get the sliceable dimension size for each mamba state tensor."""
|
||||||
return self.mamba_pool.get_state_dim_per_tensor()
|
return self.mamba_pool.get_state_dim_per_tensor()
|
||||||
|
|
||||||
|
def get_state_layer_ids(self):
|
||||||
|
"""Global layer id per mamba state entry, aligned with get_state_buf_infos()."""
|
||||||
|
return self.mamba_pool.get_state_layer_ids()
|
||||||
|
|
||||||
def get_state_slice_outer_counts(self):
|
def get_state_slice_outer_counts(self):
|
||||||
"""Get the row count preceding each mamba state slice axis."""
|
"""Get the row count preceding each mamba state slice axis."""
|
||||||
return self.mamba_pool.get_state_slice_outer_counts()
|
return self.mamba_pool.get_state_slice_outer_counts()
|
||||||
|
|||||||
@@ -517,6 +517,8 @@ class UnifiedMambaPool(MambaPool):
|
|||||||
):
|
):
|
||||||
spec = unified_buffer.mamba_spec(sub_pool_name)
|
spec = unified_buffer.mamba_spec(sub_pool_name)
|
||||||
assert spec.layer_num == len(mamba_layer_ids)
|
assert spec.layer_num == len(mamba_layer_ids)
|
||||||
|
# PP disagg state transfer maps entries by global layer id.
|
||||||
|
self.mamba_layer_ids = list(mamba_layer_ids)
|
||||||
conv_views, temporal_view = unified_buffer.mamba_views_for(sub_pool_name)
|
conv_views, temporal_view = unified_buffer.mamba_views_for(sub_pool_name)
|
||||||
max_slots = unified_buffer.max_slots(sub_pool_name)
|
max_slots = unified_buffer.max_slots(sub_pool_name)
|
||||||
|
|
||||||
|
|||||||
@@ -131,6 +131,17 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
num_layers = kvc.layer_info.num_effective_layers
|
num_layers = kvc.layer_info.num_effective_layers
|
||||||
|
|
||||||
self._cell_size = self._compute_cell_size(kvc, num_layers)
|
self._cell_size = self._compute_cell_size(kvc, num_layers)
|
||||||
|
has_kv_on_another_pp_stage = (
|
||||||
|
self._cell_size == 0
|
||||||
|
and mambaish is not None
|
||||||
|
and bool(mambaish.full_attention_layer_ids)
|
||||||
|
and kvc.ps.pp_size > 1
|
||||||
|
)
|
||||||
|
self._zero_kv_max_tokens = (
|
||||||
|
torch.iinfo(torch.int64).max
|
||||||
|
if has_kv_on_another_pp_stage
|
||||||
|
else kvc.server_args.max_total_tokens or kvc.model_config.context_len
|
||||||
|
)
|
||||||
|
|
||||||
# EAGLE/STANDALONE: scale cell_size to account for draft model KV cache.
|
# EAGLE/STANDALONE: scale cell_size to account for draft model KV cache.
|
||||||
# Assumes draft and target share the same per-layer KV size (head_dim,
|
# Assumes draft and target share the same per-layer KV size (head_dim,
|
||||||
@@ -287,7 +298,11 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
def calculate_pool_sizes(
|
def calculate_pool_sizes(
|
||||||
self, available_bytes: int, page_size: int
|
self, available_bytes: int, page_size: int
|
||||||
) -> MemoryPoolConfig:
|
) -> MemoryPoolConfig:
|
||||||
max_total_num_tokens = available_bytes // self._cell_size
|
max_total_num_tokens = (
|
||||||
|
available_bytes // self._cell_size
|
||||||
|
if self._cell_size
|
||||||
|
else self._zero_kv_max_tokens
|
||||||
|
)
|
||||||
max_total_num_tokens = max_total_num_tokens // page_size * page_size
|
max_total_num_tokens = max_total_num_tokens // page_size * page_size
|
||||||
return MemoryPoolConfig(max_total_num_tokens=max_total_num_tokens)
|
return MemoryPoolConfig(max_total_num_tokens=max_total_num_tokens)
|
||||||
|
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
|||||||
from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat
|
from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
|
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
|
||||||
from sglang.srt.layers.utils import PPMissingLayer
|
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import (
|
from sglang.srt.layers.vocab_parallel_embedding import (
|
||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
@@ -628,6 +628,8 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
self.model = KimiLinearModel(
|
self.model = KimiLinearModel(
|
||||||
config, quant_config, prefix=maybe_prefix(prefix, "model")
|
config, quant_config, prefix=maybe_prefix(prefix, "model")
|
||||||
)
|
)
|
||||||
|
self.start_layer = self.model.start_layer
|
||||||
|
self.end_layer = self.model.end_layer
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
if self.pp_group.is_last_rank:
|
if self.pp_group.is_last_rank:
|
||||||
self.lm_head = ParallelLMHead(
|
self.lm_head = ParallelLMHead(
|
||||||
@@ -664,6 +666,19 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
else:
|
else:
|
||||||
return hidden_states
|
return hidden_states
|
||||||
|
|
||||||
|
def _is_non_local_pp_weight(self, name: str) -> bool:
|
||||||
|
if self.pp_group.world_size == 1:
|
||||||
|
return False
|
||||||
|
|
||||||
|
layer_id = get_layer_id(name)
|
||||||
|
if layer_id is not None:
|
||||||
|
return not (self.model.start_layer <= layer_id < self.model.end_layer)
|
||||||
|
if name.startswith("model.embed_tokens."):
|
||||||
|
return not self.pp_group.is_first_rank
|
||||||
|
if name.startswith(("model.norm.", "lm_head.")):
|
||||||
|
return not self.pp_group.is_last_rank
|
||||||
|
return False
|
||||||
|
|
||||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
|
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
|
||||||
stacked_params_mapping = [
|
stacked_params_mapping = [
|
||||||
# (param_name, shard_name, shard_id)
|
# (param_name, shard_name, shard_id)
|
||||||
@@ -703,6 +718,8 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
for args in weights:
|
for args in weights:
|
||||||
name, loaded_weight = args[:2]
|
name, loaded_weight = args[:2]
|
||||||
kwargs = args[2] if len(args) > 2 else {}
|
kwargs = args[2] if len(args) > 2 else {}
|
||||||
|
if self._is_non_local_pp_weight(name):
|
||||||
|
continue
|
||||||
if "rotary_emb.inv_freq" in name:
|
if "rotary_emb.inv_freq" in name:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -738,8 +755,6 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
# Skip loading extra bias for GPTQ models.
|
# Skip loading extra bias for GPTQ models.
|
||||||
if name.endswith(".bias") and name not in params_dict:
|
if name.endswith(".bias") and name not in params_dict:
|
||||||
continue
|
continue
|
||||||
# if is_pp_missing_parameter(name, self):
|
|
||||||
# continue
|
|
||||||
param = params_dict[name]
|
param = params_dict[name]
|
||||||
weight_loader = param.weight_loader
|
weight_loader = param.weight_loader
|
||||||
weight_loader(param, loaded_weight, shard_id)
|
weight_loader(param, loaded_weight, shard_id)
|
||||||
@@ -751,8 +766,6 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
if weight_name not in name:
|
if weight_name not in name:
|
||||||
continue
|
continue
|
||||||
name = name.replace(weight_name, param_name)
|
name = name.replace(weight_name, param_name)
|
||||||
# if is_pp_missing_parameter(name, self):
|
|
||||||
# continue
|
|
||||||
param = params_dict[name]
|
param = params_dict[name]
|
||||||
weight_loader = param.weight_loader
|
weight_loader = param.weight_loader
|
||||||
weight_loader(
|
weight_loader(
|
||||||
@@ -775,9 +788,6 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
name = maybe_remap_kv_scale_name(name, params_dict)
|
name = maybe_remap_kv_scale_name(name, params_dict)
|
||||||
if name is None:
|
if name is None:
|
||||||
continue
|
continue
|
||||||
# if is_pp_missing_parameter(name, self):
|
|
||||||
# continue
|
|
||||||
|
|
||||||
param = params_dict[name]
|
param = params_dict[name]
|
||||||
weight_loader = getattr(
|
weight_loader = getattr(
|
||||||
param, "weight_loader", default_weight_loader
|
param, "weight_loader", default_weight_loader
|
||||||
@@ -786,6 +796,8 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
loaded_params.add(name)
|
loaded_params.add(name)
|
||||||
|
|
||||||
for layer_id in self.config.full_attention_layer_ids:
|
for layer_id in self.config.full_attention_layer_ids:
|
||||||
|
if not self.model.start_layer <= layer_id < self.model.end_layer:
|
||||||
|
continue
|
||||||
self_attn = self.model.layers[layer_id].self_attn
|
self_attn = self.model.layers[layer_id].self_attn
|
||||||
w_kc, w_vc = self_attn.kv_b_proj.weight.unflatten(
|
w_kc, w_vc = self_attn.kv_b_proj.weight.unflatten(
|
||||||
0, (-1, self_attn.qk_nope_head_dim + self_attn.v_head_dim)
|
0, (-1, self_attn.qk_nope_head_dim + self_attn.v_head_dim)
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=240, stage="base-c", runner_config="4-gpu-h100")
|
register_cuda_ci(est_time=480, stage="base-c", runner_config="4-gpu-h100")
|
||||||
|
|
||||||
KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random"
|
KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random"
|
||||||
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
|
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
|
||||||
@@ -38,6 +38,7 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
|||||||
prefill_tp_size = 2
|
prefill_tp_size = 2
|
||||||
decode_tp_size = 1
|
decode_tp_size = 1
|
||||||
decode_base_gpu_id = 2
|
decode_base_gpu_id = 2
|
||||||
|
reference_parallel_args = ["--tp-size", "2"]
|
||||||
extra_prefill_args = SERVER_ARGS
|
extra_prefill_args = SERVER_ARGS
|
||||||
extra_decode_args = SERVER_ARGS
|
extra_decode_args = SERVER_ARGS
|
||||||
extra_prefill_env = SERVER_ENV
|
extra_prefill_env = SERVER_ENV
|
||||||
@@ -72,7 +73,9 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
|||||||
self.model,
|
self.model,
|
||||||
self.lb_url,
|
self.lb_url,
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
other_args=["--tp-size", "2", "--trust-remote-code"] + SERVER_ARGS,
|
other_args=self.reference_parallel_args
|
||||||
|
+ ["--trust-remote-code"]
|
||||||
|
+ SERVER_ARGS,
|
||||||
env=SERVER_ENV,
|
env=SERVER_ENV,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
@@ -101,5 +104,13 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
|||||||
assert_process_healthy(self, "decode", self.process_decode, self.decode_url)
|
assert_process_healthy(self, "decode", self.process_decode, self.decode_url)
|
||||||
|
|
||||||
|
|
||||||
|
class TestKimiLinearPipelineDisaggregation(TestKimiLinearHeterogeneousTPDisaggregation):
|
||||||
|
prefill_tp_size = 1
|
||||||
|
decode_tp_size = 1
|
||||||
|
decode_base_gpu_id = 2
|
||||||
|
reference_parallel_args = ["--tp-size", "1", "--pp-size", "2"]
|
||||||
|
extra_prefill_args = SERVER_ARGS + ["--pp-size", "2"]
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -451,8 +451,10 @@ class TestNixlTransferWorker(CustomTestCase):
|
|||||||
mgr.enable_staging = False
|
mgr.enable_staging = False
|
||||||
mgr._staging_ctx = None
|
mgr._staging_ctx = None
|
||||||
mgr.is_mla_backend = False
|
mgr.is_mla_backend = False
|
||||||
|
mgr.is_hybrid_mla_backend = False
|
||||||
mgr.attn_tp_size = 1
|
mgr.attn_tp_size = 1
|
||||||
mgr.kv_args = SimpleNamespace(engine_rank=0)
|
mgr.transfer_source_rank = 0
|
||||||
|
mgr.kv_args = SimpleNamespace(engine_rank=0, kv_data_ptrs=[0])
|
||||||
mgr.exceptions = {}
|
mgr.exceptions = {}
|
||||||
mgr.failure_lock = threading.Lock()
|
mgr.failure_lock = threading.Lock()
|
||||||
mgr.failure_records = {}
|
mgr.failure_records = {}
|
||||||
@@ -493,6 +495,7 @@ class TestNixlTransferWorker(CustomTestCase):
|
|||||||
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
||||||
self.assertNotIn(room, mgr.transfer_infos)
|
self.assertNotIn(room, mgr.transfer_infos)
|
||||||
self.assertNotIn(room, mgr.req_to_decode_prefix_len)
|
self.assertNotIn(room, mgr.req_to_decode_prefix_len)
|
||||||
|
mgr.send_aux.assert_called_once()
|
||||||
|
|
||||||
def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
|
def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
|
||||||
self,
|
self,
|
||||||
@@ -507,6 +510,7 @@ class TestNixlTransferWorker(CustomTestCase):
|
|||||||
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
||||||
self.assertIn(room, mgr.transfer_infos)
|
self.assertIn(room, mgr.transfer_infos)
|
||||||
self.assertIn(room, mgr.req_to_decode_prefix_len)
|
self.assertIn(room, mgr.req_to_decode_prefix_len)
|
||||||
|
mgr.send_kvcache.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
class TestNixlNotifications(CustomTestCase):
|
class TestNixlNotifications(CustomTestCase):
|
||||||
@@ -699,6 +703,7 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
mgr.agent = agent or StagingFakeAgent()
|
mgr.agent = agent or StagingFakeAgent()
|
||||||
mgr.attn_tp_size = 2
|
mgr.attn_tp_size = 2
|
||||||
mgr.is_mla_backend = False
|
mgr.is_mla_backend = False
|
||||||
|
mgr.transfer_source_rank = 1
|
||||||
mgr.kv_args = SimpleNamespace(
|
mgr.kv_args = SimpleNamespace(
|
||||||
gpu_id=1,
|
gpu_id=1,
|
||||||
engine_rank=1,
|
engine_rank=1,
|
||||||
|
|||||||
Reference in New Issue
Block a user