[PD] Refactor hybrid state transfer (#24932)

This commit is contained in:
Ke Bao
2026-05-12 13:16:54 +08:00
committed by GitHub
parent 91907b7b93
commit d7f4761a48
13 changed files with 594 additions and 377 deletions
@@ -35,10 +35,11 @@ class AscendKVManager(MooncakeKVManager):
self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens
) )
# Batch register state/extra pool data buffers # Batch register state/extra pool data buffers
if self.kv_args.state_data_ptrs and self.kv_args.state_data_lens: for component_ptrs, component_lens in zip(
self.engine.batch_register( self.kv_args.state_data_ptrs or [],
self.kv_args.state_data_ptrs, self.kv_args.state_data_lens self.kv_args.state_data_lens or [],
) ):
self.engine.batch_register(component_ptrs, component_lens)
def send_kvcache( def send_kvcache(
self, self,
+15 -8
View File
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import dataclasses import dataclasses
import enum
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, List, Optional from typing import TYPE_CHECKING, List, Optional
@@ -13,6 +14,12 @@ if TYPE_CHECKING:
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
class StateType(str, enum.Enum):
MAMBA = "mamba"
SWA = "swa"
NSA = "nsa"
@dataclasses.dataclass @dataclasses.dataclass
class KVTransferMetric: class KVTransferMetric:
# Backends that cannot isolate transfer latency can leave this as None. # Backends that cannot isolate transfer latency can leave this as None.
@@ -28,12 +35,12 @@ class KVArgs:
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]
state_data_ptrs: List[int] state_types: List[StateType]
state_data_lens: List[int] state_data_ptrs: List[List[int]]
state_item_lens: List[int] state_data_lens: List[List[int]]
state_type: str # "none", "mamba", "swa", "nsa" state_item_lens: List[List[int]]
# for mamba state different tp slice transfer # Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
state_dim_per_tensor: List[int] # dimension to slice for each state tensor state_dim_per_tensor: List[List[int]]
ib_device: str ib_device: str
ib_traffic_class: str ib_traffic_class: str
gpu_id: int gpu_id: int
@@ -96,7 +103,7 @@ class BaseKVSender(ABC):
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
""" """
Send the kv cache at the given kv indices and the extra cache/state at the given indices to the decoder server. Send the kv cache at the given kv indices and the extra cache/state at the given indices to the decoder server.
@@ -154,7 +161,7 @@ class BaseKVReceiver(ABC):
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
decode_prefix_len: Optional[int] = None, decode_prefix_len: Optional[int] = None,
): ):
""" """
@@ -96,7 +96,7 @@ class CommonKVManager(BaseKVManager):
): ):
self.kv_args = args self.kv_args = args
self.kv_item_lens_sum = sum(args.kv_item_lens) self.kv_item_lens_sum = sum(args.kv_item_lens)
self.state_item_lens_sum = sum(args.state_item_lens) self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp)
self.is_mla_backend = is_mla_backend self.is_mla_backend = is_mla_backend
self.disaggregation_mode = disaggregation_mode self.disaggregation_mode = disaggregation_mode
self.server_args = server_args self.server_args = server_args
@@ -520,16 +520,18 @@ class CommonKVSender(BaseKVSender):
def _record_transfer_indices( def _record_transfer_indices(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]], state_indices: Optional[List],
): ):
self._transfer_num_kv_indices += len(kv_indices) self._transfer_num_kv_indices += len(kv_indices)
if state_indices is not None: if state_indices:
self._transfer_num_state_indices += len(state_indices) for component_indices in state_indices:
if component_indices is not None:
self._transfer_num_state_indices += len(component_indices)
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
pass pass
@@ -1,3 +1,4 @@
import struct
import threading import threading
from collections import deque from collections import deque
from typing import List, Tuple from typing import List, Tuple
@@ -6,6 +7,39 @@ import numpy as np
import numpy.typing as npt import numpy.typing as npt
def pack_list_of_buffers(buffers: List[bytes]) -> bytes:
if not buffers:
return b""
n = len(buffers)
header = struct.pack(f"<{n+1}I", n, *(len(b) for b in buffers))
return header + b"".join(buffers)
def unpack_list_of_buffers(buf: bytes) -> List[bytes]:
if buf == b"":
return []
(n,) = struct.unpack("<I", buf[:4])
lens = struct.unpack(f"<{n}I", buf[4 : 4 + 4 * n])
out = []
offset = 4 + 4 * n
for length in lens:
out.append(buf[offset : offset + length])
offset += length
return out
def pack_int_lists(lists, fmt: str) -> bytes:
return pack_list_of_buffers([struct.pack(f"<{len(a)}{fmt}", *a) for a in lists])
def unpack_int_lists(buf: bytes, fmt: str) -> List[List[int]]:
width = struct.calcsize(fmt)
return [
list(struct.unpack(f"<{len(b)//width}{fmt}", b))
for b in unpack_list_of_buffers(buf)
]
class FastQueue: class FastQueue:
def __init__(self): def __init__(self):
self._buf = deque() self._buf = deque()
+32 -20
View File
@@ -34,6 +34,7 @@ from torch.distributed import ProcessGroup
from sglang.srt.configs.mamba_utils import Mamba2CacheParams from sglang.srt.configs.mamba_utils import Mamba2CacheParams
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.disaggregation.base import KVPoll from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST, FAKE_BOOTSTRAP_HOST,
@@ -56,17 +57,14 @@ from sglang.srt.managers.schedule_policy import match_prefix_for_req
from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import (
kv_to_page_indices, kv_to_page_indices,
page_align_floor, page_align_floor,
release_kv_cache, release_kv_cache,
) )
from sglang.srt.mem_cache.memory_pool import ( from sglang.srt.mem_cache.memory_pool import (
HybridLinearKVPool,
HybridReqToTokenPool, HybridReqToTokenPool,
KVCache, KVCache,
NSATokenToKVPool,
ReqToTokenPool, ReqToTokenPool,
) )
from sglang.srt.observability.req_time_stats import ( from sglang.srt.observability.req_time_stats import (
@@ -366,7 +364,12 @@ class DecodePreallocQueue:
self.metadata_buffers.get_buf_infos() self.metadata_buffers.get_buf_infos()
) )
setup_state_kv_args(kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool) setup_state_kv_args(
kv_args,
self.token_to_kv_pool,
self.draft_token_to_kv_pool,
req_to_token_pool=getattr(self, "req_to_token_pool", None),
)
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
kv_args.gpu_id = self.scheduler.gpu_id kv_args.gpu_id = self.scheduler.gpu_id
@@ -809,45 +812,54 @@ class DecodePreallocQueue:
) )
page_size = self.token_to_kv_pool_allocator.page_size page_size = self.token_to_kv_pool_allocator.page_size
# Prepare extra pool indices for hybrid models seq_len = len(decode_req.req.origin_input_ids)
if isinstance(self.token_to_kv_pool, HybridLinearKVPool):
# Mamba hybrid model: single mamba state index def _mamba_payload():
state_indices = [ return [
self.req_to_token_pool.req_index_to_mamba_index_mapping[ self.req_to_token_pool.req_index_to_mamba_index_mapping[
decode_req.req.req_pool_idx decode_req.req.req_pool_idx
] ]
.cpu() .cpu()
.numpy() .numpy()
] ]
elif isinstance(self.token_to_kv_pool, BaseSWAKVPool):
seq_len = len(decode_req.req.origin_input_ids)
window_size = self.scheduler.sliding_window_size
def _swa_payload():
window_size = self.scheduler.sliding_window_size
window_start = max(0, seq_len - window_size) window_start = max(0, seq_len - window_size)
window_start = page_align_floor(window_start, page_size) window_start = page_align_floor(window_start, page_size)
window_kv_indices_full = self.req_to_token_pool.req_to_token[ window_kv_indices_full = self.req_to_token_pool.req_to_token[
decode_req.req.req_pool_idx, window_start:seq_len decode_req.req.req_pool_idx, window_start:seq_len
] ]
# Translate to SWA pool indices
window_kv_indices_swa = ( window_kv_indices_swa = (
self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa( self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
window_kv_indices_full window_kv_indices_full
) )
) )
state_indices = window_kv_indices_swa.cpu().numpy() return kv_to_page_indices(
state_indices = kv_to_page_indices(state_indices, page_size) window_kv_indices_swa.cpu().numpy(), page_size
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool): )
seq_len = len(decode_req.req.origin_input_ids)
def _nsa_payload():
kv_indices_full = self.req_to_token_pool.req_to_token[ kv_indices_full = self.req_to_token_pool.req_to_token[
decode_req.req.req_pool_idx, :seq_len decode_req.req.req_pool_idx, :seq_len
] ]
state_indices = kv_indices_full.cpu().numpy()
# Indexer lives on device pool; always use device page_size # Indexer lives on device pool; always use device page_size
device_page_size = self.token_to_kv_pool.page_size device_page_size = self.token_to_kv_pool.page_size
state_indices = kv_to_page_indices(state_indices, device_page_size) return kv_to_page_indices(
kv_indices_full.cpu().numpy(), device_page_size
)
state_types = self.kv_manager.kv_args.state_types
state_indices: Optional[List] = []
for st in state_types:
if st == StateType.MAMBA:
state_indices.append(_mamba_payload())
elif st == StateType.SWA:
state_indices.append(_swa_payload())
elif st == StateType.NSA:
state_indices.append(_nsa_payload())
else: else:
state_indices = None state_indices.append(None)
decode_req.metadata_buffer_index = ( decode_req.metadata_buffer_index = (
self.req_to_metadata_buffer_idx_allocator.alloc() self.req_to_metadata_buffer_idx_allocator.alloc()
@@ -71,7 +71,7 @@ class FakeKVSender(BaseKVSender):
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
self.has_sent = True self.has_sent = True
logger.debug( logger.debug(
@@ -111,7 +111,7 @@ class FakeKVReceiver(BaseKVReceiver):
self, self,
kv_indices: list[int], kv_indices: list[int],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
decode_prefix_len: Optional[int] = None, decode_prefix_len: Optional[int] = None,
): ):
self.has_sent_metadata = True self.has_sent_metadata = True
+137 -107
View File
@@ -14,7 +14,7 @@ from typing import List, Optional, Tuple
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType
from sglang.srt.disaggregation.common.conn import ( from sglang.srt.disaggregation.common.conn import (
CommonKVBootstrapServer, CommonKVBootstrapServer,
CommonKVManager, CommonKVManager,
@@ -30,6 +30,8 @@ from sglang.srt.disaggregation.common.staging_handler import (
from sglang.srt.disaggregation.common.utils import ( from sglang.srt.disaggregation.common.utils import (
FastQueue, FastQueue,
group_concurrent_contiguous, group_concurrent_contiguous,
pack_int_lists,
unpack_int_lists,
) )
from sglang.srt.disaggregation.mooncake.utils import ( from sglang.srt.disaggregation.mooncake.utils import (
check_mooncake_custom_mem_pool_enabled, check_mooncake_custom_mem_pool_enabled,
@@ -64,7 +66,7 @@ class TransferKVChunk:
index_slice: slice index_slice: slice
is_last_chunk: bool is_last_chunk: bool
prefill_aux_index: Optional[int] prefill_aux_index: Optional[int]
state_indices: Optional[List[int]] state_indices: Optional[List]
# decode # decode
@@ -76,7 +78,7 @@ class TransferInfo:
mooncake_session_id: str mooncake_session_id: str
dst_kv_indices: npt.NDArray[np.int32] dst_kv_indices: npt.NDArray[np.int32]
dst_aux_index: int dst_aux_index: int
dst_state_indices: List[int] dst_state_indices: List[List[int]] # parallel to receiver's state_types
required_dst_info_num: int required_dst_info_num: int
is_dummy: bool is_dummy: bool
decode_prefix_len: Optional[int] = None decode_prefix_len: Optional[int] = None
@@ -93,10 +95,7 @@ class TransferInfo:
else: else:
dst_kv_indices = np.frombuffer(msg[4], dtype=np.int32) dst_kv_indices = np.frombuffer(msg[4], dtype=np.int32)
dst_aux_index = int(msg[5].decode("ascii")) dst_aux_index = int(msg[5].decode("ascii"))
if msg[6] == b"": dst_state_indices = unpack_int_lists(msg[6], "i")
dst_state_indices = []
else:
dst_state_indices = list(np.frombuffer(msg[6], dtype=np.int32))
is_dummy = False is_dummy = False
return cls( return cls(
room=int(msg[0].decode("ascii")), room=int(msg[0].decode("ascii")),
@@ -123,13 +122,13 @@ class KVArgsRegisterInfo:
mooncake_session_id: str mooncake_session_id: str
dst_kv_ptrs: list[int] dst_kv_ptrs: list[int]
dst_aux_ptrs: list[int] dst_aux_ptrs: list[int]
dst_state_data_ptrs: list[int] dst_state_data_ptrs: List[List[int]] # parallel to state_types (same below)
dst_tp_rank: int dst_tp_rank: int
dst_attn_tp_size: int dst_attn_tp_size: int
dst_kv_item_len: int dst_kv_item_len: int
# for mamba state different tp slice transfer # for mamba state different tp slice transfer
dst_state_item_lens: list[int] dst_state_item_lens: List[List[int]]
dst_state_dim_per_tensor: list[int] dst_state_dim_per_tensor: List[List[int]]
# HiSparse: decode host pool stores KV at token granularity # HiSparse: decode host pool stores KV at token granularity
enable_hisparse: bool = False enable_hisparse: bool = False
# 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)
@@ -144,19 +143,15 @@ class KVArgsRegisterInfo:
mooncake_session_id=msg[3].decode("ascii"), mooncake_session_id=msg[3].decode("ascii"),
dst_kv_ptrs=list(struct.unpack(f"{len(msg[4])//8}Q", msg[4])), dst_kv_ptrs=list(struct.unpack(f"{len(msg[4])//8}Q", msg[4])),
dst_aux_ptrs=list(struct.unpack(f"{len(msg[5])//8}Q", msg[5])), dst_aux_ptrs=list(struct.unpack(f"{len(msg[5])//8}Q", msg[5])),
dst_state_data_ptrs=list(struct.unpack(f"{len(msg[6])//8}Q", msg[6])), dst_state_data_ptrs=unpack_int_lists(msg[6], "Q"),
dst_tp_rank=int(msg[7].decode("ascii")), dst_tp_rank=int(msg[7].decode("ascii")),
dst_attn_tp_size=int(msg[8].decode("ascii")), dst_attn_tp_size=int(msg[8].decode("ascii")),
dst_kv_item_len=int(msg[9].decode("ascii")), dst_kv_item_len=int(msg[9].decode("ascii")),
dst_state_item_lens=( dst_state_item_lens=(
list(struct.unpack(f"{len(msg[10])//4}I", msg[10])) unpack_int_lists(msg[10], "I") if len(msg) > 10 else []
if len(msg) > 10 and len(msg[10]) > 0
else []
), ),
dst_state_dim_per_tensor=( dst_state_dim_per_tensor=(
list(struct.unpack(f"{len(msg[11])//4}I", msg[11])) unpack_int_lists(msg[11], "I") if len(msg) > 11 else []
if len(msg) > 11 and len(msg[11]) > 0
else []
), ),
enable_hisparse=( enable_hisparse=(
msg[12].decode("ascii") == "1" if len(msg) > 12 else False msg[12].decode("ascii") == "1" if len(msg) > 12 else False
@@ -272,11 +267,11 @@ class MooncakeKVManager(CommonKVManager):
self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens
) )
# Batch register state/extra pool data buffers for ptrs, lens in zip(
if self.kv_args.state_data_ptrs and self.kv_args.state_data_lens:
self.engine.batch_register(
self.kv_args.state_data_ptrs, self.kv_args.state_data_lens self.kv_args.state_data_ptrs, self.kv_args.state_data_lens
) ):
if ptrs and lens:
self.engine.batch_register(ptrs, lens)
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Staging buffer methods (all delegate to staging_handler.py) # Staging buffer methods (all delegate to staging_handler.py)
@@ -966,88 +961,133 @@ class MooncakeKVManager(CommonKVManager):
def maybe_send_extra( def maybe_send_extra(
self, self,
req: TransferInfo, req: TransferInfo,
prefill_state_indices: list[int], prefill_state_indices: List,
dst_state_data_ptrs: list[int],
executor: concurrent.futures.ThreadPoolExecutor, executor: concurrent.futures.ThreadPoolExecutor,
target_rank_registration_info: Optional[KVArgsRegisterInfo] = None, target_rank_registration_info: Optional[KVArgsRegisterInfo] = None,
): ):
"""Send state or extra pool data with type-specific handling.""" rc = 0
state_type = getattr(self.kv_args, "state_type", "none") state_types = getattr(self.kv_args, "state_types", [])
for i, st in enumerate(state_types):
indices = (
prefill_state_indices[i] if i < len(prefill_state_indices) else None
)
if indices is None:
continue
src_data_ptrs = self.kv_args.state_data_ptrs[i]
src_item_lens = self.kv_args.state_item_lens[i]
src_dim_per_tensor = (
self.kv_args.state_dim_per_tensor[i]
if i < len(self.kv_args.state_dim_per_tensor)
else []
)
if target_rank_registration_info is not None:
dst_data_ptrs = (
target_rank_registration_info.dst_state_data_ptrs[i]
if i < len(target_rank_registration_info.dst_state_data_ptrs)
else []
)
dst_item_lens = (
target_rank_registration_info.dst_state_item_lens[i]
if i < len(target_rank_registration_info.dst_state_item_lens)
else []
)
dst_dim_per_tensor = (
target_rank_registration_info.dst_state_dim_per_tensor[i]
if i < len(target_rank_registration_info.dst_state_dim_per_tensor)
else []
)
else:
dst_data_ptrs, dst_item_lens, dst_dim_per_tensor = [], [], []
dst_indices = (
req.dst_state_indices[i] if i < len(req.dst_state_indices) else []
)
if state_type == "mamba": if st == StateType.MAMBA:
# Check if we need slice transfer for different TP sizes
if ( if (
target_rank_registration_info is not None target_rank_registration_info is not None
and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size and self.attn_tp_size
!= target_rank_registration_info.dst_attn_tp_size
): ):
return self._send_mamba_state_slice( rc = (
self._send_mamba_state_slice(
req, req,
prefill_state_indices, indices,
dst_state_data_ptrs, src_data_ptrs,
target_rank_registration_info.dst_state_item_lens, src_item_lens,
target_rank_registration_info.dst_state_dim_per_tensor, src_dim_per_tensor,
dst_data_ptrs,
dst_indices,
dst_item_lens,
dst_dim_per_tensor,
target_rank_registration_info.dst_tp_rank, target_rank_registration_info.dst_tp_rank,
target_rank_registration_info.dst_attn_tp_size, target_rank_registration_info.dst_attn_tp_size,
) )
else: or rc
return self._send_mamba_state(
req,
prefill_state_indices,
dst_state_data_ptrs,
) )
elif state_type in ["swa", "nsa"]: else:
# Non-MLA SWA / NSA hybrid models do not support different TP sizes yet. rc = (
self._send_mamba_state(
req,
indices,
src_data_ptrs,
src_item_lens,
dst_data_ptrs,
dst_indices,
)
or rc
)
elif st in (StateType.SWA, StateType.NSA):
if ( if (
target_rank_registration_info is not None target_rank_registration_info is not None
and not self.is_mla_backend and not self.is_mla_backend
and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size and self.attn_tp_size
!= target_rank_registration_info.dst_attn_tp_size
): ):
raise RuntimeError( raise RuntimeError(
f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {st.upper()} hybrid models yet."
) )
dst_state_indices = req.dst_state_indices src_indices = list(indices)
if len(prefill_state_indices) > len(dst_state_indices): dst_indices_local = list(dst_indices)
if len(src_indices) > len(dst_indices_local):
logger.warning( logger.warning(
f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(dst_state_indices)}" f"len(prefill_state_indices) = {len(src_indices)}, len(dst_state_indices) = {len(dst_indices_local)}"
) )
prefill_state_indices = prefill_state_indices[: len(dst_state_indices)] src_indices = src_indices[: len(dst_indices_local)]
elif len(prefill_state_indices) < len(dst_state_indices): elif len(src_indices) < len(dst_indices_local):
logger.warning( logger.warning(
f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(dst_state_indices)}" f"len(prefill_state_indices) = {len(src_indices)}, len(dst_state_indices) = {len(dst_indices_local)}"
) )
dst_state_indices = dst_state_indices[: len(prefill_state_indices)] dst_indices_local = dst_indices_local[: len(src_indices)]
# Reuse _send_kvcache_generic interface to send extra pool data rc = (
prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32) self._send_kvcache_generic(
dst_state_indices = np.array(dst_state_indices, dtype=np.int32)
return self._send_kvcache_generic(
mooncake_session_id=req.mooncake_session_id, mooncake_session_id=req.mooncake_session_id,
src_data_ptrs=self.kv_args.state_data_ptrs, src_data_ptrs=src_data_ptrs,
dst_data_ptrs=dst_state_data_ptrs, dst_data_ptrs=dst_data_ptrs,
item_lens=self.kv_args.state_item_lens, item_lens=src_item_lens,
prefill_data_indices=prefill_state_indices, prefill_data_indices=np.array(src_indices, dtype=np.int32),
dst_data_indices=dst_state_indices, dst_data_indices=np.array(dst_indices_local, dtype=np.int32),
executor=executor, executor=executor,
) )
else: or rc
return 0 )
return rc
def _send_mamba_state( def _send_mamba_state(
self, self,
req: TransferInfo, req: TransferInfo,
prefill_mamba_index: list[int], prefill_mamba_index: list,
src_state_data_ptrs: list[int],
src_state_item_lens: list[int],
dst_state_data_ptrs: list[int], dst_state_data_ptrs: list[int],
dst_mamba_index: list,
): ):
"""Transfer Mamba states."""
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 = []
prefill_state_data_ptrs = self.kv_args.state_data_ptrs
prefill_state_item_lens = self.kv_args.state_item_lens
for i, dst_state_ptr in enumerate(dst_state_data_ptrs): for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
length = prefill_state_item_lens[i] length = src_state_item_lens[i]
src_addr = prefill_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(req.dst_state_indices[0]) dst_addr = dst_state_ptr + length * int(dst_mamba_index[0])
transfer_blocks.append((src_addr, dst_addr, length)) transfer_blocks.append((src_addr, dst_addr, length))
return self._transfer_data(req.mooncake_session_id, transfer_blocks) return self._transfer_data(req.mooncake_session_id, transfer_blocks)
@@ -1055,8 +1095,12 @@ class MooncakeKVManager(CommonKVManager):
def _send_mamba_state_slice( def _send_mamba_state_slice(
self, self,
req: TransferInfo, req: TransferInfo,
prefill_mamba_index: list[int], prefill_mamba_index: list,
src_state_data_ptrs: list[int],
src_state_item_lens: list[int],
src_state_dim_per_tensor: list[int],
dst_state_data_ptrs: list[int], dst_state_data_ptrs: list[int],
dst_mamba_index: list,
dst_state_item_lens: list[int], dst_state_item_lens: list[int],
dst_state_dim_per_tensor: list[int], dst_state_dim_per_tensor: list[int],
dst_tp_rank: int, dst_tp_rank: int,
@@ -1078,33 +1122,33 @@ class MooncakeKVManager(CommonKVManager):
) )
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 = []
prefill_state_data_ptrs = self.kv_args.state_data_ptrs
prefill_state_item_lens = self.kv_args.state_item_lens
src_state_dim_per_tensor = getattr(self.kv_args, "state_dim_per_tensor", [])
# If no dimension info available, fall back to regular transfer # If no dimension info available, fall back to regular transfer
if not src_state_dim_per_tensor or not dst_state_dim_per_tensor: if not src_state_dim_per_tensor or not dst_state_dim_per_tensor:
return self._send_mamba_state(req, prefill_mamba_index, dst_state_data_ptrs) return self._send_mamba_state(
req,
prefill_mamba_index,
src_state_data_ptrs,
src_state_item_lens,
dst_state_data_ptrs,
dst_mamba_index,
)
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 = []
for i, dst_state_ptr in enumerate(dst_state_data_ptrs): for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
src_item_len = prefill_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[i]
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[i]
# Calculate bytes per dimension slice
# item_len = dim * trailing_dims_size, so trailing_dims_size = item_len / dim # item_len = dim * trailing_dims_size, so trailing_dims_size = item_len / dim
src_bytes_per_dim = src_item_len // src_dim src_bytes_per_dim = src_item_len // src_dim
dst_bytes_per_dim = dst_item_len // dst_dim dst_bytes_per_dim = dst_item_len // dst_dim
# Determine slicing parameters based on TP configuration
if self.attn_tp_size > dst_attn_tp_size: if self.attn_tp_size > dst_attn_tp_size:
# Multiple prefill ranks send to 1 decode rank # Multiple prefill ranks send to 1 decode rank
# Each prefill sends all its dims to the appropriate offset in decode
src_dim_start = 0 src_dim_start = 0
num_dims_to_send = src_dim num_dims_to_send = src_dim
writers_per_decode = self.attn_tp_size // dst_attn_tp_size writers_per_decode = self.attn_tp_size // dst_attn_tp_size
@@ -1112,26 +1156,21 @@ class MooncakeKVManager(CommonKVManager):
dst_dim_start = local_writer_idx * src_dim dst_dim_start = local_writer_idx * src_dim
else: else:
# 1 prefill rank sends to multiple decode ranks # 1 prefill rank sends to multiple decode ranks
# Prefill sends a slice of its dims to each decode rank
src_dim_start = (dst_tp_rank_in_group * dst_dim) % src_dim src_dim_start = (dst_tp_rank_in_group * dst_dim) % src_dim
num_dims_to_send = dst_dim num_dims_to_send = dst_dim
dst_dim_start = 0 dst_dim_start = 0
# Calculate byte offsets
src_dim_offset = src_dim_start * src_bytes_per_dim src_dim_offset = src_dim_start * src_bytes_per_dim
dst_dim_offset = dst_dim_start * dst_bytes_per_dim dst_dim_offset = dst_dim_start * dst_bytes_per_dim
bytes_to_send = num_dims_to_send * src_bytes_per_dim bytes_to_send = num_dims_to_send * src_bytes_per_dim
# Calculate addresses for this state tensor
src_addr = ( src_addr = (
prefill_state_data_ptrs[i] src_state_data_ptrs[i]
+ src_item_len * int(prefill_mamba_index[0]) + src_item_len * int(prefill_mamba_index[0])
+ src_dim_offset + src_dim_offset
) )
dst_addr = ( dst_addr = (
dst_state_ptr dst_state_ptr + dst_item_len * int(dst_mamba_index[0]) + dst_dim_offset
+ dst_item_len * int(req.dst_state_indices[0])
+ dst_dim_offset
) )
transfer_blocks.append((src_addr, dst_addr, bytes_to_send)) transfer_blocks.append((src_addr, dst_addr, bytes_to_send))
@@ -1297,11 +1336,10 @@ class MooncakeKVManager(CommonKVManager):
break break
if kv_chunk.is_last_chunk: if kv_chunk.is_last_chunk:
if kv_chunk.state_indices is not None: if kv_chunk.state_indices:
self.maybe_send_extra( self.maybe_send_extra(
req, req,
kv_chunk.state_indices, kv_chunk.state_indices,
target_rank_registration_info.dst_state_data_ptrs,
executor, executor,
target_rank_registration_info, target_rank_registration_info,
) )
@@ -1576,7 +1614,7 @@ class MooncakeKVManager(CommonKVManager):
index_slice: slice, index_slice: slice,
is_last_chunk: bool, is_last_chunk: bool,
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
assert self.disaggregation_mode == DisaggregationMode.PREFILL assert self.disaggregation_mode == DisaggregationMode.PREFILL
assert not is_last_chunk or (is_last_chunk and aux_index is not None) assert not is_last_chunk or (is_last_chunk and aux_index is not None)
@@ -1672,7 +1710,7 @@ class MooncakeKVSender(CommonKVSender):
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices))
self.curr_idx += len(kv_indices) self.curr_idx += len(kv_indices)
@@ -1769,19 +1807,14 @@ class MooncakeKVReceiver(CommonKVReceiver):
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
) )
packed_state_data_ptrs = b"".join( packed_state_data_ptrs = pack_int_lists(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs self.kv_mgr.kv_args.state_data_ptrs, "Q"
) )
# Pack state_item_lens and state_dim_per_tensor for mamba state slice transfer packed_state_item_lens = pack_int_lists(
packed_state_item_lens = b"".join( self.kv_mgr.kv_args.state_item_lens, "I"
struct.pack("I", item_len)
for item_len in self.kv_mgr.kv_args.state_item_lens
) )
state_dim_per_tensor = getattr( packed_state_dim_per_tensor = pack_int_lists(
self.kv_mgr.kv_args, "state_dim_per_tensor", [] getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I"
)
packed_state_dim_per_tensor = b"".join(
struct.pack("I", dim) for dim in state_dim_per_tensor
) )
# 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
@@ -1834,7 +1867,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
decode_prefix_len: Optional[int] = None, decode_prefix_len: Optional[int] = None,
): ):
if self.bootstrap_infos is None: if self.bootstrap_infos is None:
@@ -1868,11 +1901,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
kv_indices.tobytes() if not is_dummy else b"", kv_indices.tobytes() if not is_dummy else b"",
str(aux_index).encode("ascii") if not is_dummy else b"", str(aux_index).encode("ascii") if not is_dummy else b"",
( (
np.array( pack_int_lists(state_indices, "i")
state_indices, if not is_dummy and state_indices
dtype=np.int32,
).tobytes()
if not is_dummy and state_indices is not None
else b"" else b""
), ),
str(self.required_dst_info_num).encode("ascii"), str(self.required_dst_info_num).encode("ascii"),
@@ -351,9 +351,11 @@ class MoriKVManager(CommonKVManager):
MemoryLocationType.CPU, MemoryLocationType.CPU,
) )
self.aux_mem_descs.append(desc) self.aux_mem_descs.append(desc)
for ptr, length in zip( for component_ptrs, component_lens in zip(
self.kv_args.state_data_ptrs, getattr(self.kv_args, "state_data_lens", []) self.kv_args.state_data_ptrs,
getattr(self.kv_args, "state_data_lens", []),
): ):
for ptr, length in zip(component_ptrs, component_lens):
desc = self.engine.register_memory( desc = self.engine.register_memory(
ptr, ptr,
length, length,
@@ -1239,7 +1241,7 @@ class MoriKVSender(CommonKVSender):
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices))
self.curr_idx += len(kv_indices) self.curr_idx += len(kv_indices)
@@ -1453,7 +1455,7 @@ class MoriKVReceiver(CommonKVReceiver):
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
decode_prefix_len: Optional[int] = None, decode_prefix_len: Optional[int] = None,
): ):
if self.bootstrap_infos is None or self.bootstrap_room is None: if self.bootstrap_infos is None or self.bootstrap_room is None:
+132 -102
View File
@@ -13,7 +13,7 @@ from typing import Dict, List, Optional, Set
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType
from sglang.srt.disaggregation.common.conn import ( from sglang.srt.disaggregation.common.conn import (
CommonKVBootstrapServer, CommonKVBootstrapServer,
CommonKVManager, CommonKVManager,
@@ -23,6 +23,8 @@ from sglang.srt.disaggregation.common.conn import (
from sglang.srt.disaggregation.common.utils import ( from sglang.srt.disaggregation.common.utils import (
FastQueue, FastQueue,
group_concurrent_contiguous, group_concurrent_contiguous,
pack_int_lists,
unpack_int_lists,
) )
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
DisaggregationMode, DisaggregationMode,
@@ -62,7 +64,7 @@ class TransferInfo:
dst_kv_indices: npt.NDArray[np.int32] dst_kv_indices: npt.NDArray[np.int32]
dst_aux_index: int dst_aux_index: int
required_dst_info_num: int required_dst_info_num: int
dst_state_indices: List[int] dst_state_indices: List[List[int]]
decode_prefix_len: Optional[int] = None # for decode radix cache decode_prefix_len: Optional[int] = None # for decode radix cache
def is_dummy(self): def is_dummy(self):
@@ -76,11 +78,9 @@ class TransferInfo:
@classmethod @classmethod
def from_zmq(cls, msg: List[bytes]): def from_zmq(cls, msg: List[bytes]):
# Parse state_indices from msg[7] if present dst_state_indices = (
if len(msg) > 7 and msg[7] != b"": unpack_int_lists(msg[7], "i") if len(msg) > 7 and msg[7] != b"" else []
dst_state_indices = list(np.frombuffer(msg[7], dtype=np.int32)) )
else:
dst_state_indices = []
return cls( return cls(
room=int(msg[0].decode("ascii")), room=int(msg[0].decode("ascii")),
@@ -105,7 +105,7 @@ class TransferKVChunk:
is_last: bool is_last: bool
chunk_id: int chunk_id: int
prefill_aux_index: Optional[int] prefill_aux_index: Optional[int]
state_indices: Optional[List[int]] state_indices: Optional[List]
@dataclasses.dataclass @dataclasses.dataclass
@@ -119,29 +119,24 @@ class KVArgsRegisterInfo:
agent_metadata: bytes agent_metadata: bytes
dst_kv_ptrs: list[int] dst_kv_ptrs: list[int]
dst_aux_ptrs: list[int] dst_aux_ptrs: list[int]
dst_state_data_ptrs: list[int] dst_state_data_ptrs: List[List[int]]
gpu_id: int gpu_id: int
decode_tp_size: int decode_tp_size: int
decode_tp_rank: int decode_tp_rank: int
dst_kv_item_len: int dst_kv_item_len: int
dst_state_item_lens: 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[int] = dataclasses.field(default_factory=list) dst_state_dim_per_tensor: List[List[int]] = dataclasses.field(default_factory=list)
@classmethod @classmethod
def from_zmq(cls, msg: List[bytes]): def from_zmq(cls, msg: List[bytes]):
# Parse state_data_ptrs from msg[7] if present dst_state_data_ptrs = (
if len(msg) > 7 and msg[7] != b"": unpack_int_lists(msg[7], "Q") if len(msg) > 7 and msg[7] != b"" else []
dst_state_data_ptrs = list(struct.unpack(f"{len(msg[7]) // 8}Q", msg[7])) )
else: dst_state_item_lens = (
dst_state_data_ptrs = [] unpack_int_lists(msg[12], "I") if len(msg) > 12 and len(msg[12]) > 0 else []
)
dst_state_item_lens = [] dst_state_dim_per_tensor = (
dst_state_dim_per_tensor = [] unpack_int_lists(msg[13], "I") if len(msg) > 13 and len(msg[13]) > 0 else []
if len(msg) > 12 and len(msg[12]) > 0:
dst_state_item_lens = list(struct.unpack(f"{len(msg[12]) // 4}I", msg[12]))
if len(msg) > 13 and len(msg[13]) > 0:
dst_state_dim_per_tensor = list(
struct.unpack(f"{len(msg[13]) // 4}I", msg[13])
) )
return cls( return cls(
@@ -445,10 +440,9 @@ class NixlKVManager(CommonKVManager):
handles.append(kv_xfer_handle) handles.append(kv_xfer_handle)
if kv_chunk.is_last: if kv_chunk.is_last and kv_chunk.state_indices:
if kv_chunk.state_indices is not None:
dst_info = self.decode_kv_args_table[req.agent_name] dst_info = self.decode_kv_args_table[req.agent_name]
state_xfer_handle = self.maybe_send_extra( state_xfer_handles = self.maybe_send_extra(
req.agent_name, req.agent_name,
kv_chunk.state_indices, kv_chunk.state_indices,
dst_info.dst_state_data_ptrs, dst_info.dst_state_data_ptrs,
@@ -460,8 +454,7 @@ class NixlKVManager(CommonKVManager):
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,
) )
if state_xfer_handle is not None: handles.extend(h for h in state_xfer_handles if h is not None)
handles.append(state_xfer_handle)
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")
@@ -528,15 +521,18 @@ class NixlKVManager(CommonKVManager):
if not self.aux_descs: if not self.aux_descs:
raise Exception("NIXL memory registration failed for aux tensors") raise Exception("NIXL memory registration failed for aux tensors")
# Register state/extra pool data buffers if present
if self.kv_args.state_data_ptrs and self.kv_args.state_data_lens:
state_addrs = [] state_addrs = []
for state_data_ptr, state_data_len in zip( for comp_ptrs, comp_lens in zip(
self.kv_args.state_data_ptrs, self.kv_args.state_data_lens self.kv_args.state_data_ptrs or [],
self.kv_args.state_data_lens or [],
): ):
for state_data_ptr, state_data_len in zip(comp_ptrs, comp_lens):
if state_data_ptr == 0 or state_data_len == 0:
continue
state_addrs.append( state_addrs.append(
(state_data_ptr, state_data_len, self.kv_args.gpu_id, "") (state_data_ptr, state_data_len, self.kv_args.gpu_id, "")
) )
if state_addrs:
self.state_descs = self.agent.register_memory(state_addrs, "VRAM") self.state_descs = self.agent.register_memory(state_addrs, "VRAM")
logger.debug( logger.debug(
f"Register state tensors, len(state_addrs)= {len(state_addrs)}" f"Register state tensors, len(state_addrs)= {len(state_addrs)}"
@@ -871,6 +867,8 @@ class NixlKVManager(CommonKVManager):
self, self,
peer_name: str, peer_name: str,
prefill_state_indices: List[int], prefill_state_indices: List[int],
src_state_data_ptrs: list[int],
src_state_item_lens: list[int],
dst_state_data_ptrs: list[int], dst_state_data_ptrs: list[int],
dst_state_indices: List[int], dst_state_indices: List[int],
dst_gpu_id: int, dst_gpu_id: int,
@@ -885,14 +883,11 @@ class NixlKVManager(CommonKVManager):
src_addrs = [] src_addrs = []
dst_addrs = [] dst_addrs = []
prefill_state_data_ptrs = self.kv_args.state_data_ptrs
prefill_state_item_lens = self.kv_args.state_item_lens
for i, dst_state_ptr in enumerate(dst_state_data_ptrs): for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
length = prefill_state_item_lens[i] length = src_state_item_lens[i]
src_addr = prefill_state_data_ptrs[i] + length * int( if length == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0:
prefill_state_indices[0] continue
) src_addr = src_state_data_ptrs[i] + length * int(prefill_state_indices[0])
dst_addr = dst_state_ptr + length * int(dst_state_indices[0]) dst_addr = dst_state_ptr + length * int(dst_state_indices[0])
src_addrs.append((src_addr, length, self.kv_args.gpu_id)) src_addrs.append((src_addr, length, self.kv_args.gpu_id))
dst_addrs.append((dst_addr, length, dst_gpu_id)) dst_addrs.append((dst_addr, length, dst_gpu_id))
@@ -918,12 +913,15 @@ class NixlKVManager(CommonKVManager):
self, self,
peer_name: str, peer_name: str,
prefill_state_indices: List[int], prefill_state_indices: List[int],
src_state_data_ptrs: list[int],
src_state_item_lens: list[int],
src_state_dim_per_tensor: list[int],
dst_state_data_ptrs: list[int], dst_state_data_ptrs: list[int],
dst_state_indices: List[int], dst_state_indices: List[int],
dst_gpu_id: int,
notif: str,
dst_state_item_lens: list[int], dst_state_item_lens: list[int],
dst_state_dim_per_tensor: list[int], dst_state_dim_per_tensor: list[int],
dst_gpu_id: int,
notif: str,
decode_tp_size: int, decode_tp_size: int,
decode_tp_rank: int, decode_tp_rank: int,
): ):
@@ -940,14 +938,12 @@ class NixlKVManager(CommonKVManager):
) )
assert len(prefill_state_indices) == 1, "Mamba should have single state index" assert len(prefill_state_indices) == 1, "Mamba should have single state index"
prefill_state_data_ptrs = self.kv_args.state_data_ptrs
prefill_state_item_lens = self.kv_args.state_item_lens
src_state_dim_per_tensor = getattr(self.kv_args, "state_dim_per_tensor", [])
if not src_state_dim_per_tensor or not dst_state_dim_per_tensor: if not src_state_dim_per_tensor or not dst_state_dim_per_tensor:
return self._send_mamba_state( return self._send_mamba_state(
peer_name, peer_name,
prefill_state_indices, prefill_state_indices,
src_state_data_ptrs,
src_state_item_lens,
dst_state_data_ptrs, dst_state_data_ptrs,
dst_state_indices, dst_state_indices,
dst_gpu_id, dst_gpu_id,
@@ -961,8 +957,10 @@ class NixlKVManager(CommonKVManager):
dst_addrs = [] dst_addrs = []
for i, dst_state_ptr in enumerate(dst_state_data_ptrs): for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
src_item_len = prefill_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[i]
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] src_dim = src_state_dim_per_tensor[i]
dst_dim = dst_state_dim_per_tensor[i] dst_dim = dst_state_dim_per_tensor[i]
@@ -985,7 +983,7 @@ class NixlKVManager(CommonKVManager):
bytes_to_send = num_dims_to_send * src_bytes_per_dim bytes_to_send = num_dims_to_send * src_bytes_per_dim
src_addr = ( src_addr = (
prefill_state_data_ptrs[i] src_state_data_ptrs[i]
+ src_item_len * int(prefill_state_indices[0]) + src_item_len * int(prefill_state_indices[0])
+ src_dim_offset + src_dim_offset
) )
@@ -1017,67 +1015,101 @@ class NixlKVManager(CommonKVManager):
def maybe_send_extra( def maybe_send_extra(
self, self,
peer_name: str, peer_name: str,
prefill_state_indices: List[int], prefill_state_indices: List[List[int]],
dst_state_data_ptrs: list[int], dst_state_data_ptrs: List[List[int]],
dst_state_indices: List[int], dst_state_indices: List[List[int]],
dst_gpu_id: int, dst_gpu_id: int,
notif: str, notif: str,
decode_tp_size: int, decode_tp_size: int,
decode_tp_rank: int = 0, decode_tp_rank: int = 0,
dst_state_item_lens: list[int] | None = None, dst_state_item_lens: List[List[int]] | None = None,
dst_state_dim_per_tensor: list[int] | None = None, dst_state_dim_per_tensor: List[List[int]] | None = None,
): ):
"""Send state or extra pool data with type-specific handling.""" """Send state per hybrid component, dispatching by state_type[i]."""
state_type = getattr(self.kv_args, "state_type", "none") state_types = getattr(self.kv_args, "state_types", []) or []
src_state_data_ptrs = self.kv_args.state_data_ptrs or []
src_state_item_lens = self.kv_args.state_item_lens or []
src_state_dim_per_tensor = (
getattr(self.kv_args, "state_dim_per_tensor", []) or []
)
dst_state_item_lens = dst_state_item_lens or []
dst_state_dim_per_tensor = dst_state_dim_per_tensor or []
if state_type == "mamba": handles = []
for i, st in enumerate(state_types):
src_indices = (
prefill_state_indices[i] if i < len(prefill_state_indices) else None
)
if src_indices is None or len(src_indices) == 0:
continue
src_ptrs = src_state_data_ptrs[i] if i < len(src_state_data_ptrs) else []
src_lens = src_state_item_lens[i] if i < len(src_state_item_lens) else []
src_dims = (
src_state_dim_per_tensor[i] if i < len(src_state_dim_per_tensor) 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 []
)
comp_notif = f"{notif}_{i}"
if st == StateType.MAMBA:
if self.attn_tp_size != decode_tp_size: if self.attn_tp_size != decode_tp_size:
return self._send_mamba_state_slice( h = self._send_mamba_state_slice(
peer_name, peer_name,
prefill_state_indices, src_indices,
dst_state_data_ptrs, src_ptrs,
dst_state_indices, src_lens,
src_dims,
dst_ptrs,
dst_indices,
dst_lens,
dst_dims,
dst_gpu_id, dst_gpu_id,
notif, comp_notif,
dst_state_item_lens or [],
dst_state_dim_per_tensor or [],
decode_tp_size, decode_tp_size,
decode_tp_rank, decode_tp_rank,
) )
return self._send_mamba_state( else:
h = self._send_mamba_state(
peer_name, peer_name,
prefill_state_indices, src_indices,
dst_state_data_ptrs, src_ptrs,
dst_state_indices, src_lens,
dst_ptrs,
dst_indices,
dst_gpu_id, dst_gpu_id,
notif, comp_notif,
) )
elif state_type in ["swa", "nsa"]: elif st in (StateType.SWA, StateType.NSA):
if not self.is_mla_backend and self.attn_tp_size != decode_tp_size: if not self.is_mla_backend and self.attn_tp_size != decode_tp_size:
raise RuntimeError( raise RuntimeError(
f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {st.upper()} hybrid models yet."
) )
if len(prefill_state_indices) != len(dst_state_indices): if len(src_indices) != len(dst_indices):
raise RuntimeError( raise RuntimeError(
f"State index length mismatch: prefill={len(prefill_state_indices)}, " f"State index length mismatch at component {i}: "
f"dst={len(dst_state_indices)}" f"prefill={len(src_indices)}, dst={len(dst_indices)}"
) )
return self._send_kvcache_generic( h = self._send_kvcache_generic(
peer_name=peer_name, peer_name=peer_name,
src_data_ptrs=self.kv_args.state_data_ptrs, src_data_ptrs=src_ptrs,
dst_data_ptrs=dst_state_data_ptrs, dst_data_ptrs=dst_ptrs,
item_lens=self.kv_args.state_item_lens, item_lens=src_lens,
prefill_data_indices=np.array(prefill_state_indices, dtype=np.int32), prefill_data_indices=np.array(src_indices, dtype=np.int32),
dst_data_indices=np.array(dst_state_indices, dtype=np.int32), dst_data_indices=np.array(dst_indices, dtype=np.int32),
dst_gpu_id=dst_gpu_id, dst_gpu_id=dst_gpu_id,
notif=notif, notif=comp_notif,
) )
else: else:
if state_type != "none":
raise RuntimeError( raise RuntimeError(
f"PD Disaggregation via NIXL does NOT support {state_type} hybrid models yet." f"PD Disaggregation via NIXL does NOT support {st} hybrid models yet."
) )
return None if h is not None:
handles.append(h)
return handles
def add_transfer_request( def add_transfer_request(
self, self,
@@ -1087,7 +1119,7 @@ class NixlKVManager(CommonKVManager):
is_last: bool, is_last: bool,
chunk_id: int, chunk_id: int,
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
assert self.disaggregation_mode == DisaggregationMode.PREFILL assert self.disaggregation_mode == DisaggregationMode.PREFILL
assert not is_last or (is_last and aux_index is not None) assert not is_last or (is_last and aux_index is not None)
@@ -1225,7 +1257,7 @@ class NixlKVSender(CommonKVSender):
def send( def send(
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
): ):
if self._send_failed: if self._send_failed:
return return
@@ -1314,7 +1346,7 @@ class NixlKVReceiver(CommonKVReceiver):
self, self,
kv_indices: npt.NDArray[np.int32], kv_indices: npt.NDArray[np.int32],
aux_index: Optional[int] = None, aux_index: Optional[int] = None,
state_indices: Optional[List[int]] = None, state_indices: Optional[List] = None,
decode_prefix_len: Optional[int] = None, decode_prefix_len: Optional[int] = None,
): ):
if self.bootstrap_infos is None: if self.bootstrap_infos is None:
@@ -1333,6 +1365,13 @@ class NixlKVReceiver(CommonKVReceiver):
logger.debug( logger.debug(
f"Sending to prefill server with bootstrap room {self.bootstrap_room} {is_dummy=}" f"Sending to prefill server with bootstrap room {self.bootstrap_room} {is_dummy=}"
) )
packed_state_indices = (
pack_int_lists(
[(idx if idx is not None else []) for idx in state_indices], "i"
)
if not is_dummy and state_indices is not None
else b""
)
with lock: with lock:
sock.send_multipart( sock.send_multipart(
[ [
@@ -1344,11 +1383,7 @@ class NixlKVReceiver(CommonKVReceiver):
kv_indices.tobytes() if not is_dummy else b"", kv_indices.tobytes() if not is_dummy else b"",
str(aux_index).encode("ascii"), str(aux_index).encode("ascii"),
str(self.required_dst_info_num).encode("ascii"), str(self.required_dst_info_num).encode("ascii"),
( packed_state_indices,
np.array(state_indices, dtype=np.int32).tobytes()
if not is_dummy and state_indices is not None
else b""
),
str(decode_prefix_len or 0).encode("ascii"), str(decode_prefix_len or 0).encode("ascii"),
] ]
) )
@@ -1408,19 +1443,14 @@ class NixlKVReceiver(CommonKVReceiver):
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
) )
packed_state_data_ptrs = b"".join( packed_state_data_ptrs = pack_int_lists(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs self.kv_mgr.kv_args.state_data_ptrs or [], "Q"
) )
packed_state_item_lens = pack_int_lists(
packed_state_item_lens = b"".join( self.kv_mgr.kv_args.state_item_lens or [], "I"
struct.pack("I", item_len)
for item_len in self.kv_mgr.kv_args.state_item_lens
) )
state_dim_per_tensor = getattr( packed_state_dim_per_tensor = pack_int_lists(
self.kv_mgr.kv_args, "state_dim_per_tensor", [] getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I"
)
packed_state_dim_per_tensor = b"".join(
struct.pack("I", dim) for dim in state_dim_per_tensor
) )
with lock: with lock:
+35 -26
View File
@@ -27,6 +27,7 @@ from typing import TYPE_CHECKING, List, Optional
import torch import torch
from sglang.srt.disaggregation.base import KVPoll from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.disaggregation.common.conn import CommonKVManager from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST, FAKE_BOOTSTRAP_HOST,
@@ -48,14 +49,12 @@ from sglang.srt.managers.schedule_batch import (
Req, Req,
ScheduleBatch, ScheduleBatch,
) )
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import (
kv_to_page_indices, kv_to_page_indices,
kv_to_page_num, kv_to_page_num,
maybe_cache_unfinished_req, maybe_cache_unfinished_req,
release_kv_cache, release_kv_cache,
) )
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, NSATokenToKVPool
from sglang.srt.observability.req_time_stats import set_schedule_time_batch from sglang.srt.observability.req_time_stats import set_schedule_time_batch
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -177,7 +176,13 @@ class PrefillBootstrapQueue:
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
kv_args.gpu_id = self.scheduler.gpu_id kv_args.gpu_id = self.scheduler.gpu_id
setup_state_kv_args(kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool) req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
setup_state_kv_args(
kv_args,
self.token_to_kv_pool,
self.draft_token_to_kv_pool,
req_to_token_pool=req_to_token_pool,
)
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER) kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class( kv_manager = kv_manager_class(
@@ -771,52 +776,56 @@ class SchedulerDisaggregationPrefillMixin:
.cpu() .cpu()
.numpy() .numpy()
) )
state_indices = None state_indices: Optional[List] = None
if last_chunk: if last_chunk:
self.disagg_metadata_buffers.set_buf(req) self.disagg_metadata_buffers.set_buf(req)
# Prepare extra pool indices for hybrid models seq_len = len(req.fill_ids)
if isinstance(
self.token_to_kv_pool_allocator.get_kvcache(), HybridLinearKVPool def _mamba_payload():
): return [
# Mamba hybrid model: send single mamba state index
state_indices = [
self.req_to_token_pool.req_index_to_mamba_index_mapping[ self.req_to_token_pool.req_index_to_mamba_index_mapping[
req.req_pool_idx req.req_pool_idx
] ]
.cpu() .cpu()
.numpy() .numpy()
] ]
elif isinstance(
self.token_to_kv_pool_allocator.get_kvcache(), BaseSWAKVPool def _swa_payload():
):
# SWA hybrid model: send last window KV indices
seq_len = len(req.fill_ids)
window_size = self.sliding_window_size window_size = self.sliding_window_size
window_start = max(0, seq_len - window_size) window_start = max(0, seq_len - window_size)
window_start = (window_start // page_size) * page_size window_start = (window_start // page_size) * page_size
window_kv_indices_full = self.req_to_token_pool.req_to_token[ window_kv_indices_full = self.req_to_token_pool.req_to_token[
req.req_pool_idx, window_start:seq_len req.req_pool_idx, window_start:seq_len
] ]
# Translate to SWA pool indices
window_kv_indices_swa = ( window_kv_indices_swa = (
self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa( self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
window_kv_indices_full window_kv_indices_full
) )
) )
state_indices = window_kv_indices_swa.cpu().numpy() return kv_to_page_indices(
state_indices = kv_to_page_indices(state_indices, page_size) window_kv_indices_swa.cpu().numpy(), page_size
elif isinstance( )
self.token_to_kv_pool_allocator.get_kvcache(), NSATokenToKVPool
): def _nsa_payload():
seq_len = len(req.fill_ids)
kv_indices_full = self.req_to_token_pool.req_to_token[ kv_indices_full = self.req_to_token_pool.req_to_token[
req.req_pool_idx, :seq_len req.req_pool_idx, :seq_len
] ]
state_indices = kv_indices_full.cpu().numpy() return kv_to_page_indices(kv_indices_full.cpu().numpy(), page_size)
state_indices = kv_to_page_indices(state_indices, page_size)
state_types = (
self.disagg_prefill_bootstrap_queue.kv_manager.kv_args.state_types
)
state_indices = []
for st in state_types:
if st == StateType.MAMBA:
state_indices.append(_mamba_payload())
elif st == StateType.SWA:
state_indices.append(_swa_payload())
elif st == StateType.NSA:
state_indices.append(_nsa_payload())
else:
state_indices.append(None)
page_indices = kv_to_page_indices(kv_indices, page_size) page_indices = kv_to_page_indices(kv_indices, page_size)
if not req.disagg_kv_sender.should_send_kv_chunk(len(page_indices), last_chunk): if not req.disagg_kv_sender.should_send_kv_chunk(len(page_indices), last_chunk):
+61 -26
View File
@@ -5,7 +5,7 @@ import random
from collections import deque from collections import deque
from contextlib import nullcontext from contextlib import nullcontext
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Literal, Optional, Tuple, Type, overload from typing import TYPE_CHECKING, List, Literal, Optional, Tuple, Type, overload
import numpy as np import numpy as np
import torch import torch
@@ -15,7 +15,7 @@ from sglang.srt.environ import envs
from sglang.srt.utils import is_npu from sglang.srt.utils import is_npu
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.disaggregation.base.conn import KVArgs from sglang.srt.disaggregation.base.conn import KVArgs, StateType
from sglang.srt.disaggregation.common.conn import ( from sglang.srt.disaggregation.common.conn import (
CommonKVBootstrapServer, CommonKVBootstrapServer,
CommonKVManager, CommonKVManager,
@@ -532,57 +532,92 @@ def is_mla_backend(target_kv_pool) -> bool:
return isinstance(target_kv_pool, (MLATokenToKVPool, DeepSeekV4TokenToKVPool)) return isinstance(target_kv_pool, (MLATokenToKVPool, DeepSeekV4TokenToKVPool))
def append_state_component(
kv_args: KVArgs,
state_type: StateType,
data_ptrs: List[int],
data_lens: List[int],
item_lens: List[int],
dim_per_tensor: Optional[List[int]] = None,
) -> None:
"""Append one state component. Caller orders state_types consistently
on prefill and decode sides."""
kv_args.state_types.append(state_type)
kv_args.state_data_ptrs.append(data_ptrs)
kv_args.state_data_lens.append(data_lens)
kv_args.state_item_lens.append(item_lens)
kv_args.state_dim_per_tensor.append(dim_per_tensor or [])
def setup_state_kv_args( def setup_state_kv_args(
kv_args: KVArgs, kv_args: KVArgs,
token_to_kv_pool, token_to_kv_pool,
draft_token_to_kv_pool=None, draft_token_to_kv_pool=None,
req_to_token_pool=None,
) -> None: ) -> None:
"""Populate ``kv_args`` state-buffer fields from the given pool. """Populate ``kv_args`` state-buffer fields from the given pool.
Shared by prefill and decode bootstrap paths so the state_type dispatch Shared by prefill and decode bootstrap paths so the state_type dispatch
lives in one place. lives in one place.
""" """
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, NSATokenToKVPool from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, NSATokenToKVPool
if not hasattr(token_to_kv_pool, "get_state_buf_infos"): kv_args.state_types = []
kv_args.state_data_ptrs = [] kv_args.state_data_ptrs = []
kv_args.state_data_lens = [] kv_args.state_data_lens = []
kv_args.state_item_lens = [] kv_args.state_item_lens = []
kv_args.state_type = "none" kv_args.state_dim_per_tensor = []
return
state_data_ptrs, state_data_lens, state_item_lens = ( if hasattr(token_to_kv_pool, "get_state_buf_infos"):
token_to_kv_pool.get_state_buf_infos() data_ptrs, data_lens, item_lens = token_to_kv_pool.get_state_buf_infos()
)
kv_args.state_data_ptrs = state_data_ptrs
kv_args.state_data_lens = state_data_lens
kv_args.state_item_lens = state_item_lens
# DeepSeekV4TokenToKVPool inherits BaseSWAKVPool; its heterogeneous # DeepSeekV4TokenToKVPool inherits BaseSWAKVPool; its heterogeneous
# state list is described per-entry via get_state_buf_infos. # state list is described per-entry via get_state_buf_infos.
if isinstance(token_to_kv_pool, BaseSWAKVPool): if isinstance(token_to_kv_pool, BaseSWAKVPool):
kv_args.state_type = "swa" append_state_component(
kv_args, StateType.SWA, data_ptrs, data_lens, item_lens
)
elif isinstance(token_to_kv_pool, HybridLinearKVPool): elif isinstance(token_to_kv_pool, HybridLinearKVPool):
kv_args.state_type = "mamba" dim = (
# Get state dimension info for cross-TP slice transfer token_to_kv_pool.get_state_dim_per_tensor()
if hasattr(token_to_kv_pool, "get_state_dim_per_tensor"): if hasattr(token_to_kv_pool, "get_state_dim_per_tensor")
kv_args.state_dim_per_tensor = token_to_kv_pool.get_state_dim_per_tensor() else None
)
append_state_component(
kv_args, StateType.MAMBA, data_ptrs, data_lens, item_lens, dim
)
elif isinstance(token_to_kv_pool, NSATokenToKVPool): elif isinstance(token_to_kv_pool, NSATokenToKVPool):
kv_args.state_type = "nsa"
if draft_token_to_kv_pool is not None and isinstance( if draft_token_to_kv_pool is not None and isinstance(
draft_token_to_kv_pool, NSATokenToKVPool draft_token_to_kv_pool, NSATokenToKVPool
): ):
( (
draft_state_data_ptrs, draft_data_ptrs,
draft_state_data_lens, draft_data_lens,
draft_state_item_lens, draft_item_lens,
) = draft_token_to_kv_pool.get_state_buf_infos() ) = draft_token_to_kv_pool.get_state_buf_infos()
kv_args.state_data_ptrs += draft_state_data_ptrs data_ptrs = data_ptrs + draft_data_ptrs
kv_args.state_data_lens += draft_state_data_lens data_lens = data_lens + draft_data_lens
kv_args.state_item_lens += draft_state_item_lens item_lens = item_lens + draft_item_lens
else: append_state_component(
kv_args.state_type = "none" kv_args, StateType.NSA, data_ptrs, data_lens, item_lens
)
if (
StateType.MAMBA not in kv_args.state_types
and req_to_token_pool is not None
and hasattr(req_to_token_pool, "get_state_buf_infos")
):
data_ptrs, data_lens, item_lens = req_to_token_pool.get_state_buf_infos()
if data_ptrs:
dim = (
req_to_token_pool.get_state_dim_per_tensor()
if hasattr(req_to_token_pool, "get_state_dim_per_tensor")
else None
)
append_state_component(
kv_args, StateType.MAMBA, data_ptrs, data_lens, item_lens, dim
)
def prepare_abort(req: Req, error_message: str, status_code=None): def prepare_abort(req: Req, error_message: str, status_code=None):
@@ -624,6 +624,12 @@ class HybridReqToTokenPool(ReqToTokenPool):
def get_speculative_mamba2_params_all_layers(self) -> MambaPool.SpeculativeState: def get_speculative_mamba2_params_all_layers(self) -> MambaPool.SpeculativeState:
return self.mamba_pool.get_speculative_mamba2_params_all_layers() return self.mamba_pool.get_speculative_mamba2_params_all_layers()
def get_state_buf_infos(self):
return self.mamba_pool.get_contiguous_buf_infos()
def get_state_dim_per_tensor(self):
return self.mamba_pool.get_state_dim_per_tensor()
def get_mamba_ping_pong_other_idx(self, mamba_next_track_idx: int) -> int: def get_mamba_ping_pong_other_idx(self, mamba_next_track_idx: int) -> int:
if self.mamba_ping_pong_track_buffer_size == 2: if self.mamba_ping_pong_track_buffer_size == 2:
return 1 - mamba_next_track_idx return 1 - mamba_next_track_idx
@@ -0,0 +1,49 @@
import unittest
import numpy as np
from sglang.srt.disaggregation.common.utils import (
pack_int_lists,
pack_list_of_buffers,
unpack_int_lists,
unpack_list_of_buffers,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=2, suite="stage-a-test-cpu")
class TestDisaggregationWire(unittest.TestCase):
def test_int_lists_roundtrip(self):
cases = [
("Q", [[1, 2, 3], [4]]),
("I", [[10, 20], [30, 40, 50]]),
("i", [[-1, 2], [3, -4, 5]]),
]
for fmt, sample in cases:
packed = pack_int_lists(sample, fmt)
self.assertEqual(unpack_int_lists(packed, fmt), sample, msg=fmt)
def test_pack_accepts_ndarray(self):
arrs = [
np.array([1, 2, 3], dtype=np.int32),
np.array([4, 5], dtype=np.int32),
]
packed = pack_int_lists(arrs, "i")
self.assertEqual(unpack_int_lists(packed, "i"), [[1, 2, 3], [4, 5]])
def test_empty_outer_list(self):
self.assertEqual(pack_int_lists([], "Q"), b"")
self.assertEqual(unpack_int_lists(b"", "Q"), [])
def test_empty_inner_list(self):
packed = pack_int_lists([[]], "I")
self.assertEqual(unpack_int_lists(packed, "I"), [[]])
def test_list_of_buffers_roundtrip(self):
bufs = [b"abc", b"", b"de", b"x" * 17]
self.assertEqual(unpack_list_of_buffers(pack_list_of_buffers(bufs)), bufs)
if __name__ == "__main__":
unittest.main()