fix(disagg): support pipeline-parallel hybrid-linear transfer (#32270)

This commit is contained in:
YAMY
2026-07-25 13:34:38 -07:00
committed by GitHub
parent 659d349b61
commit 91f386a5b2
12 changed files with 351 additions and 57 deletions
@@ -40,6 +40,7 @@ class KVArgs:
kv_data_ptrs: List[int]
kv_data_lens: List[int]
kv_item_lens: List[int]
kv_layer_ids: List[int]
aux_data_ptrs: List[int]
aux_data_lens: List[int]
aux_item_lens: List[int]
@@ -47,6 +48,7 @@ class KVArgs:
state_data_ptrs: List[List[int]]
state_data_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.
state_dim_per_tensor: List[List[int]]
# 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_lens = kv_data_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:
kv_args.kv_data_mem_kinds = kv_data_mem_kinds
kv_args.page_size = self.token_to_kv_pool.page_size
+105 -17
View File
@@ -43,6 +43,7 @@ from sglang.srt.disaggregation.mooncake.utils import (
)
from sglang.srt.disaggregation.utils import (
DisaggregationMode,
build_transfer_entry_pairs,
compute_mamba_state_slice_byte_blocks,
)
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
@@ -128,6 +129,8 @@ class KVArgsRegisterInfo:
# for mamba state different tp slice transfer
dst_state_item_lens: 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)
staging: Optional[StagingRegisterInfo] = None
@@ -150,8 +153,18 @@ class KVArgsRegisterInfo:
dst_state_dim_per_tensor=(
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
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,
state_type: Optional[StateType] = None,
force_flat: bool = False,
src_layer_ids: Optional[List[int]] = None,
dst_layer_ids: Optional[List[int]] = None,
) -> int:
"""
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
if self.is_mla_backend or self.is_hybrid_mla_backend or force_flat:
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],
# Layer IDs map PP-local buffers to global decode entries.
# Registrations without them retain the existing PP mapping.
if src_layer_ids or dst_layer_ids:
pairs = build_transfer_entry_pairs(
src_layer_ids,
dst_layer_ids,
len(src_data_ptrs),
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:
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)
@@ -704,6 +735,7 @@ class MooncakeKVManager(CommonKVManager):
dst_kv_ptrs: list[int],
dst_kv_indices: npt.NDArray[np.int32],
executor: concurrent.futures.ThreadPoolExecutor,
dst_layer_ids: Optional[List[int]] = None,
):
return self._send_kvcache_generic(
mooncake_session_id=mooncake_session_id,
@@ -713,6 +745,8 @@ class MooncakeKVManager(CommonKVManager):
prefill_data_indices=prefill_kv_indices,
dst_data_indices=dst_kv_indices,
executor=executor,
src_layer_ids=self.kv_args.kv_layer_ids,
dst_layer_ids=dst_layer_ids,
)
def send_kvcache_slice(
@@ -993,6 +1027,10 @@ class MooncakeKVManager(CommonKVManager):
src_slice_outer_counts = (
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:
dst_data_ptrs = (
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)
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:
dst_data_ptrs, dst_item_lens, dst_dim_per_tensor = [], [], []
dst_state_layer_ids = []
dst_indices = (
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,
src_conv_shard_groups,
src_slice_outer_counts,
src_state_layer_ids,
dst_state_layer_ids,
)
or rc
)
@@ -1048,6 +1094,8 @@ class MooncakeKVManager(CommonKVManager):
src_item_lens,
dst_data_ptrs,
dst_indices,
src_state_layer_ids,
dst_state_layer_ids,
)
or rc
)
@@ -1149,11 +1197,21 @@ class MooncakeKVManager(CommonKVManager):
src_state_item_lens: list[int],
dst_state_data_ptrs: list[int],
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"
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]
src_addr = src_state_data_ptrs[i] + length * int(prefill_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,
src_state_conv_shard_groups: list = 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.
@@ -1205,17 +1265,27 @@ class MooncakeKVManager(CommonKVManager):
src_state_item_lens,
dst_state_data_ptrs,
dst_mamba_index,
src_layer_ids,
dst_layer_ids,
)
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
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]
dst_item_len = dst_state_item_lens[i]
dst_item_len = dst_state_item_lens[j]
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 = (
src_state_conv_shard_groups[i]
@@ -1370,7 +1440,11 @@ class MooncakeKVManager(CommonKVManager):
skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags(
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
elif (
self.is_mla_backend
@@ -1384,6 +1458,7 @@ class MooncakeKVManager(CommonKVManager):
target_rank_registration_info.dst_kv_ptrs,
chunked_dst_kv_indice,
executor,
target_rank_registration_info.dst_kv_layer_ids,
)
elif (
self.enable_staging
@@ -1919,9 +1994,20 @@ class MooncakeKVReceiver(CommonKVReceiver):
packed_state_dim_per_tensor = pack_int_lists(
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
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_attn_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
dst_kv_item_len = str(kv_item_len).encode("ascii")
@@ -1953,6 +2039,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
dst_kv_item_len,
packed_state_item_lens,
packed_state_dim_per_tensor,
packed_kv_layer_ids,
packed_state_layer_ids,
packed_staging_base_ptr,
staging_total_size_str,
]
+100 -28
View File
@@ -38,6 +38,7 @@ from sglang.srt.disaggregation.common.utils import (
)
from sglang.srt.disaggregation.utils import (
DisaggregationMode,
build_transfer_entry_pairs,
compute_mamba_state_slice_byte_blocks,
)
from sglang.srt.environ import envs
@@ -209,9 +210,11 @@ class KVArgsRegisterInfo:
decode_tp_rank: int
dst_kv_item_len: 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_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_layer_ids: List[List[int]] = dataclasses.field(default_factory=list)
dst_homogeneous_mem_kind: Optional[str] = None
kv_xfer_segments: Optional[List[_KVXferPreparedSegment]] = None
# Keep last: optional, parsed from a variable-length tail of the ZMQ
@@ -249,6 +252,14 @@ class KVArgsRegisterInfo:
dst_num_slots = (
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(
room=str(msg[0].decode("ascii")),
@@ -265,9 +276,11 @@ class KVArgsRegisterInfo:
decode_tp_rank=int(msg[10].decode("ascii")),
dst_kv_item_len=dst_kv_item_len,
dst_kv_item_lens=dst_kv_item_lens,
dst_kv_layer_ids=dst_kv_layer_ids,
dst_num_slots=dst_num_slots,
dst_state_item_lens=dst_state_item_lens,
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),
)
@@ -360,6 +373,9 @@ class NixlKVManager(CommonKVManager):
is_mla_backend: Optional[bool] = False,
):
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(
getattr(self.kv_args, "kv_data_mem_kinds", None),
len(self.kv_args.kv_data_ptrs),
@@ -367,6 +383,7 @@ class NixlKVManager(CommonKVManager):
self.src_mem_kind = (
_homogeneous_kv_mem_kind(self.kv_args.kv_data_mem_kinds, "source")
if disaggregation_mode == DisaggregationMode.PREFILL
and self.kv_args.kv_data_mem_kinds
else None
)
try:
@@ -432,9 +449,10 @@ class NixlKVManager(CommonKVManager):
self._num_slots_src: int = 0
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self._num_slots_src = (
self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0]
)
if self.kv_args.kv_item_lens:
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()
self.transfer_queues: List[FastQueue] = [
FastQueue() for _ in range(transfer_queue_size)
@@ -919,9 +937,6 @@ class NixlKVManager(CommonKVManager):
peer_info.kv_xfer_segments = prepared_segments
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),
# decode with speculative decoding will have more kv items.
# Prefill having more kv items is impossible.
@@ -932,6 +947,10 @@ class NixlKVManager(CommonKVManager):
"NIXL PD transfer: decode registered fewer KV regions "
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
if (
self.is_mla_backend
@@ -979,8 +998,16 @@ class NixlKVManager(CommonKVManager):
else self._num_slots_src
)
dst_kv_ptrs = peer_info.dst_kv_ptrs[:n_src]
dst_kv_item_lens = peer_info.dst_kv_item_lens[:n_src]
pairs = build_transfer_entry_pairs(
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 = [
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
# (e.g., decode-side radix cache matched the entire prefix).
# 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]
# NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices
@@ -1077,7 +1107,7 @@ class NixlKVManager(CommonKVManager):
notif = (
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:
@@ -1163,11 +1193,12 @@ class NixlKVManager(CommonKVManager):
dst_info.dst_state_data_ptrs,
req.dst_state_indices,
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_rank=dst_info.decode_tp_rank,
dst_state_item_lens=dst_info.dst_state_item_lens,
dst_state_dim_per_tensor=dst_info.dst_state_dim_per_tensor,
dst_state_layer_ids=dst_info.dst_state_layer_ids,
)
handles.extend(
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:
raise RuntimeError("Missing aux index for last chunk")
# When no KV pages were sent (decode-side cache hit),
# encode pp_rank in aux notif so receiver can mark
# expected_kvs_per_pp[pp_rank] = 0.
if len(kv_chunk.prefill_kv_indices) == 0:
# A no-KV notification still identifies its PP source.
if (
len(kv_chunk.prefill_kv_indices) == 0
or not self.kv_args.kv_data_ptrs
):
aux_notif = (
f"{req.room}_aux_nokv_{self.kv_args.engine_rank}"
f"{req.room}_aux_nokv_{self.transfer_source_rank}"
)
else:
aux_notif = f"{req.room}_aux"
@@ -1267,8 +1299,6 @@ class NixlKVManager(CommonKVManager):
f"NIXL memory registration failed for {mem_kind} kv tensors"
)
self.kv_descs.append(kv_descs)
if not self.kv_descs:
raise Exception("NIXL memory registration failed for kv tensors")
aux_addrs = []
for aux_data_ptr, aux_data_len in zip(
self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens
@@ -1766,7 +1796,7 @@ class NixlKVManager(CommonKVManager):
notif_tag = (
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}"
)
handle = self.send_kvcache_staged(
@@ -1842,6 +1872,8 @@ class NixlKVManager(CommonKVManager):
dst_state_indices: List[int],
dst_gpu_id: int,
notif: str,
src_layer_ids: list[int] = None,
dst_layer_ids: list[int] = None,
):
"""Transfer Mamba states via RDMA."""
assert len(prefill_state_indices) == 1, "Mamba should have single state index"
@@ -1852,7 +1884,15 @@ class NixlKVManager(CommonKVManager):
src_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]
if length == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0:
continue
@@ -1895,6 +1935,8 @@ class NixlKVManager(CommonKVManager):
decode_tp_rank: int,
src_state_conv_shard_groups: list = 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.
@@ -1922,6 +1964,8 @@ class NixlKVManager(CommonKVManager):
dst_state_indices,
dst_gpu_id,
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
@@ -1930,13 +1974,21 @@ class NixlKVManager(CommonKVManager):
src_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]
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:
continue
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 = (
src_state_conv_shard_groups[i]
@@ -2007,6 +2059,7 @@ class NixlKVManager(CommonKVManager):
decode_tp_rank: int = 0,
dst_state_item_lens: 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]."""
state_types = getattr(self.kv_args, "state_types", []) or []
@@ -2021,8 +2074,10 @@ class NixlKVManager(CommonKVManager):
src_state_slice_outer_counts = (
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_dim_per_tensor = dst_state_dim_per_tensor or []
dst_state_layer_ids = dst_state_layer_ids or []
handles = []
for i, st in enumerate(state_types):
@@ -2046,12 +2101,14 @@ class NixlKVManager(CommonKVManager):
if i < len(src_state_slice_outer_counts)
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_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_dims = (
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}"
if st == StateType.MAMBA:
@@ -2072,6 +2129,8 @@ class NixlKVManager(CommonKVManager):
decode_tp_rank,
src_conv,
src_outer_counts,
src_layer_ids=src_lids,
dst_layer_ids=dst_lids,
)
else:
h = self._send_mamba_state(
@@ -2083,6 +2142,8 @@ class NixlKVManager(CommonKVManager):
dst_indices,
dst_gpu_id,
comp_notif,
src_layer_ids=src_lids,
dst_layer_ids=dst_lids,
)
elif st in (
StateType.SWA,
@@ -2709,6 +2770,10 @@ class NixlKVReceiver(CommonKVReceiver):
struct.pack("Q", item_len)
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(
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(
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
if (
@@ -2733,10 +2801,12 @@ class NixlKVReceiver(CommonKVReceiver):
else:
packed_staging_base_ptr = b""
staging_total_size_str = b""
dst_num_slots = (
self.kv_mgr.kv_args.kv_data_lens[0]
// self.kv_mgr.kv_args.kv_item_lens[0]
)
if self.kv_mgr.kv_args.kv_item_lens:
dst_kv_item_len = 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)
try:
@@ -2755,7 +2825,7 @@ class NixlKVReceiver(CommonKVReceiver):
str(self.kv_mgr.kv_args.gpu_id).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.kv_item_lens[0]).encode("ascii"),
str(dst_kv_item_len).encode("ascii"),
packed_state_item_lens,
packed_state_dim_per_tensor,
packed_staging_base_ptr,
@@ -2763,6 +2833,8 @@ class NixlKVReceiver(CommonKVReceiver):
str(dst_num_slots).encode("ascii"),
packed_kv_data_mem_kinds,
packed_kv_item_lens,
packed_state_layer_ids,
packed_kv_layer_ids,
]
)
except zmq.ZMQError:
@@ -219,6 +219,12 @@ class PrefillBootstrapQueue:
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_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:
kv_args.kv_head_num = self.token_to_kv_pool.head_num
kv_args.total_kv_head_num = (
+54
View File
@@ -855,6 +855,53 @@ def compute_mamba_state_slice_byte_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(
kv_args: KVArgs,
state_type: StateType,
@@ -864,6 +911,7 @@ def append_state_component(
dim_per_tensor: Optional[List[int]] = None,
conv_shard_groups: Optional[List[Optional[List[int]]]] = None,
slice_outer_counts: Optional[List[int]] = None,
layer_ids: Optional[List[int]] = None,
) -> None:
"""Append one state component. Caller orders state_types consistently
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_conv_shard_groups.append(conv_shard_groups 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(
@@ -903,6 +952,7 @@ def setup_state_kv_args(
kv_args.state_item_lens = []
kv_args.state_dim_per_tensor = []
kv_args.state_slice_outer_counts = []
kv_args.state_layer_ids = []
kv_args.is_hybrid_mla_backend = False
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")
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(
kv_args,
StateType.MAMBA,
@@ -980,6 +1033,7 @@ def setup_state_kv_args(
dim,
conv_shard_groups,
slice_outer_counts,
layer_ids,
)
elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)):
if draft_token_to_kv_pool is not None and isinstance(
@@ -473,6 +473,7 @@ class MambaPool:
enable=enable_memory_saver
)
num_mamba_layers = len(mamba_layer_ids)
self.mamba_layer_ids = list(mamba_layer_ids)
self.size = size
self.device = device
@@ -874,6 +875,7 @@ class MambaPool:
return (
not _is_npu
and len(convs) > 0
and convs[0].shape[0] > 0
and convs[0].is_cuda
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
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):
"""Get the number of rows preceding each tensor's TP slice axis."""
outer_counts = []
@@ -3625,6 +3637,11 @@ class HybridLinearKVPool(KVCache):
def get_contiguous_buf_infos(self):
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):
mamba_data_ptrs, mamba_data_lens, mamba_item_lens = (
self.mamba_pool.get_contiguous_buf_infos()
@@ -3635,6 +3652,10 @@ class HybridLinearKVPool(KVCache):
"""Get the sliceable dimension size for each mamba state 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):
"""Get the row count preceding each mamba state slice axis."""
return self.mamba_pool.get_state_slice_outer_counts()
@@ -517,6 +517,8 @@ class UnifiedMambaPool(MambaPool):
):
spec = unified_buffer.mamba_spec(sub_pool_name)
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)
max_slots = unified_buffer.max_slots(sub_pool_name)
@@ -131,6 +131,17 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
num_layers = kvc.layer_info.num_effective_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.
# Assumes draft and target share the same per-layer KV size (head_dim,
@@ -287,7 +298,11 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
def calculate_pool_sizes(
self, available_bytes: int, page_size: int
) -> 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
return MemoryPoolConfig(max_total_num_tokens=max_total_num_tokens)
+20 -8
View File
@@ -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.quantization.base_config import QuantizationConfig
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 (
ParallelLMHead,
VocabParallelEmbedding,
@@ -628,6 +628,8 @@ class KimiLinearForCausalLM(nn.Module):
self.model = KimiLinearModel(
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()
if self.pp_group.is_last_rank:
self.lm_head = ParallelLMHead(
@@ -664,6 +666,19 @@ class KimiLinearForCausalLM(nn.Module):
else:
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]]):
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
@@ -703,6 +718,8 @@ class KimiLinearForCausalLM(nn.Module):
for args in weights:
name, loaded_weight = args[:2]
kwargs = args[2] if len(args) > 2 else {}
if self._is_non_local_pp_weight(name):
continue
if "rotary_emb.inv_freq" in name:
continue
@@ -738,8 +755,6 @@ class KimiLinearForCausalLM(nn.Module):
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
# if is_pp_missing_parameter(name, self):
# continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
@@ -751,8 +766,6 @@ class KimiLinearForCausalLM(nn.Module):
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
# if is_pp_missing_parameter(name, self):
# continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(
@@ -775,9 +788,6 @@ class KimiLinearForCausalLM(nn.Module):
name = maybe_remap_kv_scale_name(name, params_dict)
if name is None:
continue
# if is_pp_missing_parameter(name, self):
# continue
param = params_dict[name]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
@@ -786,6 +796,8 @@ class KimiLinearForCausalLM(nn.Module):
loaded_params.add(name)
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
w_kc, w_vc = self_attn.kv_b_proj.weight.unflatten(
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,
)
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"
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
@@ -38,6 +38,7 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
prefill_tp_size = 2
decode_tp_size = 1
decode_base_gpu_id = 2
reference_parallel_args = ["--tp-size", "2"]
extra_prefill_args = SERVER_ARGS
extra_decode_args = SERVER_ARGS
extra_prefill_env = SERVER_ENV
@@ -72,7 +73,9 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
self.model,
self.lb_url,
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,
)
try:
@@ -101,5 +104,13 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
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__":
unittest.main()
@@ -451,8 +451,10 @@ class TestNixlTransferWorker(CustomTestCase):
mgr.enable_staging = False
mgr._staging_ctx = None
mgr.is_mla_backend = False
mgr.is_hybrid_mla_backend = False
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.failure_lock = threading.Lock()
mgr.failure_records = {}
@@ -493,6 +495,7 @@ class TestNixlTransferWorker(CustomTestCase):
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
self.assertNotIn(room, mgr.transfer_infos)
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(
self,
@@ -507,6 +510,7 @@ class TestNixlTransferWorker(CustomTestCase):
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
self.assertIn(room, mgr.transfer_infos)
self.assertIn(room, mgr.req_to_decode_prefix_len)
mgr.send_kvcache.assert_called_once()
class TestNixlNotifications(CustomTestCase):
@@ -699,6 +703,7 @@ class TestNixlStaging(CustomTestCase):
mgr.agent = agent or StagingFakeAgent()
mgr.attn_tp_size = 2
mgr.is_mla_backend = False
mgr.transfer_source_rank = 1
mgr.kv_args = SimpleNamespace(
gpu_id=1,
engine_rank=1,