[PD][MoRI] Align hybrid state transfer with per-component schema (#26539)

This commit is contained in:
Niko Ma
2026-05-29 00:54:46 -07:00
committed by GitHub
parent a42a7654a2
commit 4d1163e6a9
2 changed files with 184 additions and 88 deletions
+159 -81
View File
@@ -35,6 +35,8 @@ from sglang.srt.disaggregation.common.conn import (
from sglang.srt.disaggregation.common.utils import (
AuxDataCodec,
group_concurrent_contiguous,
pack_int_lists,
unpack_int_lists,
)
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.server_args import ServerArgs
@@ -45,12 +47,33 @@ logger = logging.getLogger(__name__)
MORI_GUARD = b"MoriMsgGuard"
def _normalize_state_indices(
state_indices,
) -> Optional[npt.NDArray[np.int32]]:
def _normalize_state_indices_per_component(
state_indices: Optional[List],
) -> Optional[List[Optional[npt.NDArray[np.int32]]]]:
if state_indices is None:
return None
return np.asarray(state_indices, dtype=np.int32)
out: List[Optional[npt.NDArray[np.int32]]] = []
for entry in state_indices:
if entry is None:
out.append(None)
else:
out.append(np.asarray(entry, dtype=np.int32).ravel())
return out
def _pack_state_indices(
state_indices: Optional[List[Optional[npt.NDArray[np.int32]]]],
) -> bytes:
if not state_indices:
return b""
lists = [(arr.tolist() if arr is not None else []) for arr in state_indices]
return pack_int_lists(lists, "i")
def _unpack_state_indices(buf: bytes) -> List[npt.NDArray[np.int32]]:
if not buf:
return []
return [np.asarray(lst, dtype=np.int32) for lst in unpack_int_lists(buf, "i")]
def _pack_mem_desc_list(mems: List[MemoryDesc]) -> bytes:
@@ -67,6 +90,21 @@ def _unpack_mem_desc_list(blob: bytes) -> List[MemoryDesc]:
return [MemoryDesc.unpack(b) for b in desc_blobs]
def _pack_mem_desc_lists(mems_per_comp: List[List[MemoryDesc]]) -> bytes:
if not mems_per_comp:
return b""
return msgspec.msgpack.encode(
[[mem.pack() for mem in comp] for comp in mems_per_comp]
)
def _unpack_mem_desc_lists(blob: bytes) -> List[List[MemoryDesc]]:
if not blob:
return []
nested = msgspec.msgpack.decode(blob)
return [[MemoryDesc.unpack(b) for b in comp] for comp in nested]
@dataclasses.dataclass
class TransferInfo:
room: int
@@ -75,7 +113,7 @@ class TransferInfo:
engine_key: str
dst_kv_indices: npt.NDArray[np.int32]
dst_aux_index: int
dst_state_indices: npt.NDArray[np.int32]
dst_state_indices: List[npt.NDArray[np.int32]]
required_dst_info_num: int
is_dummy: bool
@@ -97,9 +135,9 @@ class TransferInfo:
dst_aux_index = -1
if len(payload) > 6 and payload[6]:
dst_state_indices = np.frombuffer(payload[6], dtype=np.int32)
dst_state_indices = _unpack_state_indices(payload[6])
else:
dst_state_indices = np.array([], dtype=np.int32)
dst_state_indices = []
required_dst_info_num = (
int(payload[7].decode("ascii")) if len(payload) > 7 else 1
@@ -125,13 +163,13 @@ class KVArgsRegisterInfo:
engine_desc: EngineDesc
dst_kv_mem_descs: List[MemoryDesc]
dst_aux_mem_descs: List[MemoryDesc]
dst_state_mem_descs: List[MemoryDesc]
dst_state_mem_descs: List[List[MemoryDesc]]
gpu_id: int
decode_tp_size: int
decode_tp_rank: int
dst_kv_item_len: int
dst_state_item_lens: List[int]
dst_state_dim_per_tensor: List[int]
dst_state_item_lens: List[List[int]]
dst_state_dim_per_tensor: List[List[int]]
@property
def engine_key(self) -> str:
@@ -144,19 +182,19 @@ class KVArgsRegisterInfo:
engine_desc = EngineDesc.unpack(payload[3])
dst_kv_mem_descs = _unpack_mem_desc_list(payload[4])
dst_aux_mem_descs = _unpack_mem_desc_list(payload[5])
dst_state_mem_descs = _unpack_mem_desc_list(payload[6])
dst_state_mem_descs = _unpack_mem_desc_lists(payload[6])
gpu_id = int(payload[7].decode("ascii"))
decode_tp_size = int(payload[8].decode("ascii"))
decode_tp_rank = int(payload[9].decode("ascii"))
dst_kv_item_len = int(payload[10].decode("ascii"))
dst_state_item_lens = (
list(struct.unpack(f"{len(payload[11]) // 4}I", payload[11]))
if len(payload) > 11 and len(payload[11]) > 0
unpack_int_lists(payload[11], "I")
if len(payload) > 11 and payload[11]
else []
)
dst_state_dim_per_tensor = (
list(struct.unpack(f"{len(payload[12]) // 4}I", payload[12]))
if len(payload) > 12 and len(payload[12]) > 0
unpack_int_lists(payload[12], "I")
if len(payload) > 12 and payload[12]
else []
)
return cls(
@@ -244,7 +282,7 @@ class MoriKVManager(CommonKVManager):
self.engine_desc = self.engine.get_engine_desc()
self.kv_mem_descs: List[MemoryDesc] = []
self.aux_mem_descs: List[MemoryDesc] = []
self.state_mem_descs: List[MemoryDesc] = []
self.state_mem_descs: List[List[MemoryDesc]] = []
self.transfer_lock = threading.Lock()
self._zmq_ctx = zmq.Context()
self._socket_local = threading.local()
@@ -338,6 +376,7 @@ class MoriKVManager(CommonKVManager):
self.kv_args.state_data_ptrs,
getattr(self.kv_args, "state_data_lens", []),
):
component_descs: List[MemoryDesc] = []
for ptr, length in zip(component_ptrs, component_lens):
desc = self.engine.register_memory(
ptr,
@@ -345,7 +384,8 @@ class MoriKVManager(CommonKVManager):
self.kv_args.gpu_id,
MemoryLocationType.GPU,
)
self.state_mem_descs.append(desc)
component_descs.append(desc)
self.state_mem_descs.append(component_descs)
def update_status(self, bootstrap_room: int, status: KVPoll):
current = self.request_status.get(bootstrap_room)
@@ -904,69 +944,106 @@ class MoriKVManager(CommonKVManager):
def send_state(
self,
peer_info: KVArgsRegisterInfo,
src_state_indices: npt.NDArray[np.int32],
dst_state_indices: npt.NDArray[np.int32],
src_state_indices: List[npt.NDArray[np.int32]],
dst_state_indices: List[npt.NDArray[np.int32]],
) -> List[TransferStatus]:
# Guard: no local state tensors -> no-op (e.g. SWA layers=0 on this PP rank)
if not self.state_mem_descs:
return []
state_type = getattr(self.kv_args, "state_type", "none")
if state_type == "none":
state_types = self.kv_args.state_types
if not state_types:
raise RuntimeError(
"PD state transfer failed: state_type is 'none' but state_indices were provided"
)
if not peer_info.dst_state_mem_descs:
raise RuntimeError(
f"PD state transfer failed: remote peer has no state descriptors "
f"(state_type={state_type}, prefill_tp_size={self.attn_tp_size}, "
f"decode_tp_size={peer_info.decode_tp_size})"
"PD state transfer failed: kv_args.state_types is empty but "
"state_indices were provided"
)
if len(peer_info.dst_state_mem_descs) != len(self.state_mem_descs):
raise RuntimeError(
f"PD state transfer failed: state descriptor count mismatch "
f"(local={len(self.state_mem_descs)}, remote={len(peer_info.dst_state_mem_descs)}), "
f"likely PP configuration mismatch (state_type={state_type})"
f"PD state transfer failed: state component count mismatch "
f"(local={len(self.state_mem_descs)}, "
f"remote={len(peer_info.dst_state_mem_descs)})"
)
if len(self.kv_args.state_item_lens) != len(self.state_mem_descs):
raise RuntimeError(
f"PD state transfer failed: local state_item_lens count "
f"({len(self.kv_args.state_item_lens)}) does not match state descriptor "
f"count ({len(self.state_mem_descs)}) (state_type={state_type})"
src_state_item_lens = self.kv_args.state_item_lens
src_state_dim_per_tensor = self.kv_args.state_dim_per_tensor
statuses: List[TransferStatus] = []
for i, st in enumerate(state_types):
src_indices = src_state_indices[i] if i < len(src_state_indices) else None
dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else None
if src_indices is None or src_indices.size == 0:
continue
if dst_indices is None or dst_indices.size == 0:
continue
src_descs = self.state_mem_descs[i]
dst_descs = peer_info.dst_state_mem_descs[i]
src_lens = src_state_item_lens[i] if i < len(src_state_item_lens) else []
dst_lens = (
peer_info.dst_state_item_lens[i]
if i < len(peer_info.dst_state_item_lens)
else []
)
src_dims = (
src_state_dim_per_tensor[i] if i < len(src_state_dim_per_tensor) else []
)
dst_dims = (
peer_info.dst_state_dim_per_tensor[i]
if i < len(peer_info.dst_state_dim_per_tensor)
else []
)
if state_type == "mamba":
return self._send_mamba_state(
peer_info, src_state_indices, dst_state_indices
)
elif state_type in ("swa", "dsa"):
return self._send_swa_dsa_state(
peer_info, src_state_indices, dst_state_indices, state_type
)
else:
raise RuntimeError(
f"PD state transfer failed: unknown state_type={state_type}"
)
if st == "mamba":
statuses.extend(
self._send_mamba_state(
peer_info,
src_indices,
dst_indices,
src_descs,
dst_descs,
src_lens,
dst_lens,
src_dims,
dst_dims,
)
)
elif st in ("swa", "dsa"):
statuses.extend(
self._send_swa_dsa_state(
peer_info,
src_indices,
dst_indices,
src_descs,
src_lens,
dst_descs,
st,
)
)
else:
raise RuntimeError(f"PD state transfer failed: unknown state_type={st}")
return statuses
def _send_mamba_state(
self,
peer_info: KVArgsRegisterInfo,
src_state_indices: npt.NDArray[np.int32],
dst_state_indices: npt.NDArray[np.int32],
src_state_mem_descs: List[MemoryDesc],
dst_state_mem_descs: List[MemoryDesc],
src_state_item_lens: List[int],
dst_state_item_lens: List[int],
src_state_dim_per_tensor: List[int],
dst_state_dim_per_tensor: List[int],
) -> List[TransferStatus]:
if len(src_state_indices) != 1 or len(dst_state_indices) != 1:
if src_state_indices.size != 1 or dst_state_indices.size != 1:
raise RuntimeError(
f"PD state transfer failed: mamba requires single state index, "
f"got src={len(src_state_indices)}, dst={len(dst_state_indices)}"
f"got src={src_state_indices.size}, dst={dst_state_indices.size}"
)
tp_mismatch = peer_info.decode_tp_size != self.attn_tp_size
src_state_dim_per_tensor = getattr(self.kv_args, "state_dim_per_tensor", [])
dst_state_dim_per_tensor = peer_info.dst_state_dim_per_tensor
# If dim info missing, silently degrade to whole-item copy (Mooncake compat)
if tp_mismatch and (
@@ -983,15 +1060,14 @@ class MoriKVManager(CommonKVManager):
src_idx = int(src_state_indices[0])
dst_idx = int(dst_state_indices[0])
statuses = []
statuses: List[TransferStatus] = []
local_tp_rank = self.kv_args.engine_rank % self.attn_tp_size
dst_tp_rank = peer_info.decode_tp_rank % peer_info.decode_tp_size
for i in range(len(self.state_mem_descs)):
src_desc = self.state_mem_descs[i]
dst_desc = peer_info.dst_state_mem_descs[i]
src_item_len = self.kv_args.state_item_lens[i]
for i, src_desc in enumerate(src_state_mem_descs):
dst_desc = dst_state_mem_descs[i]
src_item_len = src_state_item_lens[i]
if not tp_mismatch:
# same-TP: whole item copy
@@ -1000,7 +1076,7 @@ class MoriKVManager(CommonKVManager):
size = src_item_len
else:
# TP mismatch slice copy
dst_item_len = peer_info.dst_state_item_lens[i]
dst_item_len = dst_state_item_lens[i]
src_dim = src_state_dim_per_tensor[i]
dst_dim = dst_state_dim_per_tensor[i]
@@ -1044,6 +1120,9 @@ class MoriKVManager(CommonKVManager):
peer_info: KVArgsRegisterInfo,
src_state_indices: npt.NDArray[np.int32],
dst_state_indices: npt.NDArray[np.int32],
src_state_mem_descs: List[MemoryDesc],
src_state_item_lens: List[int],
dst_state_mem_descs: List[MemoryDesc],
state_type: str,
) -> List[TransferStatus]:
# TP mismatch check for non-MLA SWA
@@ -1057,17 +1136,17 @@ class MoriKVManager(CommonKVManager):
f"(prefill_tp_size={self.attn_tp_size}, decode_tp_size={peer_info.decode_tp_size})"
)
common_len = min(len(src_state_indices), len(dst_state_indices))
if common_len == 0 and max(len(src_state_indices), len(dst_state_indices)) > 0:
common_len = min(src_state_indices.size, dst_state_indices.size)
if common_len == 0 and max(src_state_indices.size, dst_state_indices.size) > 0:
raise RuntimeError(
f"No overlapping state indices for state_type={state_type}"
)
if len(src_state_indices) != len(dst_state_indices):
if src_state_indices.size != dst_state_indices.size:
logger.warning(
"State index length mismatch for %s: src=%d dst=%d; truncating to common prefix=%d",
state_type,
len(src_state_indices),
len(dst_state_indices),
src_state_indices.size,
dst_state_indices.size,
common_len,
)
src_state_indices = src_state_indices[:common_len]
@@ -1078,11 +1157,10 @@ class MoriKVManager(CommonKVManager):
*group_concurrent_contiguous(src_state_indices, dst_state_indices)
)
statuses = []
for i in range(len(self.state_mem_descs)):
src_desc = self.state_mem_descs[i]
dst_desc = peer_info.dst_state_mem_descs[i]
state_item_len = self.kv_args.state_item_lens[i]
statuses: List[TransferStatus] = []
for i, src_desc in enumerate(src_state_mem_descs):
dst_desc = dst_state_mem_descs[i]
state_item_len = src_state_item_lens[i]
statuses.extend(
self._submit_batch_transfer_plan(
@@ -1117,7 +1195,7 @@ class MoriKVManager(CommonKVManager):
index_slice: slice,
is_last_chunk: bool,
aux_index: Optional[int] = None,
state_indices: Optional[npt.NDArray[np.int32]] = None,
state_indices: Optional[List[npt.NDArray[np.int32]]] = None,
) -> Tuple[List[TransferStatus], Optional[List[TransferInfo]]]:
assert self.disaggregation_mode == DisaggregationMode.PREFILL
@@ -1233,7 +1311,9 @@ class MoriKVSender(CommonKVSender):
return
normalized_state = (
_normalize_state_indices(state_indices) if is_last_chunk else None
_normalize_state_indices_per_component(state_indices)
if is_last_chunk
else None
)
statuses, infos = self.kv_mgr.add_transfer_request(
self.bootstrap_room,
@@ -1382,18 +1462,16 @@ class MoriKVReceiver(CommonKVReceiver):
engine_desc_blob = self.kv_mgr.engine_desc.pack()
packed_kv_descs = _pack_mem_desc_list(self.kv_mgr.kv_mem_descs)
packed_aux_descs = _pack_mem_desc_list(self.kv_mgr.aux_mem_descs)
packed_state_descs = _pack_mem_desc_list(self.kv_mgr.state_mem_descs)
packed_state_descs = _pack_mem_desc_lists(self.kv_mgr.state_mem_descs)
gpu_id = str(self.kv_mgr.kv_args.gpu_id).encode("ascii")
decode_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
decode_tp_rank = str(self.kv_mgr.kv_args.engine_rank).encode("ascii")
kv_item_len = str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii")
packed_state_item_lens = b"".join(
struct.pack("I", item_len)
for item_len in self.kv_mgr.kv_args.state_item_lens
packed_state_item_lens = pack_int_lists(
self.kv_mgr.kv_args.state_item_lens, "I"
)
state_dim_per_tensor = getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", [])
packed_state_dim_per_tensor = b"".join(
struct.pack("I", dim) for dim in state_dim_per_tensor
packed_state_dim_per_tensor = pack_int_lists(
self.kv_mgr.kv_args.state_dim_per_tensor, "I"
)
for bootstrap_info in self.bootstrap_infos:
@@ -1432,13 +1510,13 @@ class MoriKVReceiver(CommonKVReceiver):
np.asarray(kv_indices, dtype=np.int32).tobytes() if kv_indices.size else b""
)
aux_bytes = str(aux_index).encode("ascii") if aux_index is not None else b""
normalized_state = _normalize_state_indices(state_indices)
normalized_state = _normalize_state_indices_per_component(state_indices)
for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
is_dummy = bootstrap_info.get("is_dummy", False)
if not is_dummy and normalized_state is not None:
state_bytes = normalized_state.tobytes()
state_bytes = _pack_state_indices(normalized_state)
else:
state_bytes = b""
with lock: