hisparse: support NIXL DRAM KV destinations for HiSparse (#27563)

Co-authored-by: Zhangheng <hzh0425@apache.org>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
ishandhanani
2026-06-27 22:32:29 +08:00
committed by GitHub
co-authored by Zhangheng Shangming Cai
parent c1b5c7e499
commit b030b1a5f3
12 changed files with 726 additions and 167 deletions
@@ -108,10 +108,15 @@ Pass as a JSON string via `--hisparse-config`:
<td>int</td>
<td>Ratio of logical pool size to device pool size, determining host memory capacity</td>
</tr>
<tr>
<td><code>swap_in_block_size</code></td>
<td>int / 960</td>
<td>CUDA thread-block size for the HiSparse swap-in kernel</td>
</tr>
</tbody>
</table>
Example: `--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10}'`
Example: `--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10, "swap_in_block_size": 960}'`
## Deployment
@@ -149,7 +154,7 @@ python3 -m sglang.launch_server \
--dist-init-addr 127.0.0.1:5757 \
--nnodes 1 --node-rank 0 \
--enable-hisparse \
--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10}'
--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10, "swap_in_block_size": 960}'
```
> **Note**: For DSA models, `--kv-cache-dtype` defaults to `auto`, which resolves to `fp8_e4m3` on SM100+ (Blackwell) and `bfloat16` on older architectures. The DSA decode backend is automatically selected based on KV dtype (`bfloat16` → `flashmla_sparse`, `fp8_e4m3` → `flashmla_kv`). DSA backend flags apply only to DSA models; DeepSeek V4 uses its own `dsv4` attention backend.
@@ -100,7 +100,6 @@ class BaseKVManager(ABC):
class BaseKVSender(ABC):
@abstractmethod
def __init__(
self,
@@ -156,7 +155,6 @@ class BaseKVSender(ABC):
class BaseKVReceiver(ABC):
@abstractmethod
def __init__(
self,
@@ -403,6 +403,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_data_ptrs, kv_data_lens, kv_item_lens = (
transfer_kv_pool.get_contiguous_buf_infos()
)
kv_data_mem_kinds = (
["DRAM"] * len(kv_data_ptrs)
if self.scheduler.enable_hisparse
else ["VRAM"] * len(kv_data_ptrs)
)
if self.scheduler.enable_hisparse and isinstance(
self.token_to_kv_pool, DeepSeekV4TokenToKVPool
):
@@ -413,6 +418,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_data_ptrs += device_kv_data_ptrs[c4_layer_num:]
kv_data_lens += device_kv_data_lens[c4_layer_num:]
kv_item_lens += device_kv_item_lens[c4_layer_num:]
kv_data_mem_kinds += ["VRAM"] * len(device_kv_data_ptrs[c4_layer_num:])
if self.draft_token_to_kv_pool is not None:
# We should also transfer draft model kv cache. The indices are
# always shared with a target model.
@@ -422,10 +428,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_data_ptrs += draft_kv_data_ptrs
kv_data_lens += draft_kv_data_lens
kv_item_lens += draft_kv_item_lens
kv_data_mem_kinds += ["VRAM"] * len(draft_kv_data_ptrs)
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_lens
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
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
+484 -53
View File
@@ -54,6 +54,87 @@ except ImportError:
logger = logging.getLogger(__name__)
GUARD = "NixlMsgGuard".encode("ascii")
KV_MEM_KINDS = {"VRAM", "DRAM"}
def _normalize_kv_mem_kinds(kinds: Optional[List[str]], expected_len: int) -> List[str]:
if kinds is None:
return ["VRAM"] * expected_len
kinds = [str(kind) for kind in kinds]
if len(kinds) != expected_len:
raise ValueError(
f"kv_data_mem_kinds length mismatch: got {len(kinds)}, "
f"expected {expected_len}"
)
invalid = sorted(set(kinds) - KV_MEM_KINDS)
if invalid:
raise ValueError(f"Unsupported NIXL KV memory kind(s): {invalid}")
return kinds
def _pack_kv_mem_kinds(kinds: List[str]) -> bytes:
return ",".join(kinds).encode("ascii")
def _unpack_kv_mem_kinds(buf: bytes, expected_len: int) -> List[str]:
if not buf:
return ["VRAM"] * expected_len
return _normalize_kv_mem_kinds(buf.decode("ascii").split(","), expected_len)
def _nixl_device_id(mem_kind: str, gpu_id: int) -> int:
return gpu_id if mem_kind == "VRAM" else 0
def _homogeneous_kv_mem_kind(kinds: List[str], context: str) -> str:
unique = set(kinds)
if len(unique) != 1:
raise NotImplementedError(
f"NIXL {context} mixed KV memory kinds are not implemented safely yet: "
f"{sorted(unique)}"
)
return next(iter(unique))
@dataclasses.dataclass(frozen=True)
class _KVXferMemSegment:
start: int
end: int
src_mem_kind: str
dst_mem_kind: str
def _kv_xfer_mem_segments(
src_kinds: List[str], dst_kinds: List[str]
) -> List[_KVXferMemSegment]:
if len(src_kinds) != len(dst_kinds):
raise ValueError(
f"KV source/destination memory kind length mismatch: "
f"src={len(src_kinds)}, dst={len(dst_kinds)}"
)
if not src_kinds:
return []
segments = []
start = 0
cur = (src_kinds[0], dst_kinds[0])
for i, pair in enumerate(zip(src_kinds, dst_kinds)):
if pair == cur:
continue
segments.append(_KVXferMemSegment(start, i, cur[0], cur[1]))
start = i
cur = pair
segments.append(_KVXferMemSegment(start, len(src_kinds), cur[0], cur[1]))
return segments
@dataclasses.dataclass
class _KVXferPreparedSegment:
start: int
end: int
src_handle: Any
dst_handle: Any
dst_num_slots: int
@dataclasses.dataclass
@@ -113,21 +194,42 @@ class KVArgsRegisterInfo:
agent_name: str
agent_metadata: bytes
dst_kv_ptrs: list[int]
dst_kv_mem_kinds: list[str]
dst_aux_ptrs: list[int]
dst_state_data_ptrs: List[List[int]]
gpu_id: int
decode_tp_size: int
decode_tp_rank: int
dst_kv_item_len: int
dst_kv_item_lens: list[int]
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_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
# frame in from_zmq() below, so positional construction stays stable.
staging: Optional[StagingRegisterInfo] = None
@classmethod
def from_zmq(cls, msg: List[bytes]):
dst_kv_ptrs = list(struct.unpack(f"{len(msg[5]) // 8}Q", msg[5]))
dst_kv_mem_kinds = (
_unpack_kv_mem_kinds(msg[17], len(dst_kv_ptrs))
if len(msg) > 17
else ["VRAM"] * len(dst_kv_ptrs)
)
dst_kv_item_len = int(msg[11].decode("ascii"))
dst_kv_item_lens = (
list(struct.unpack(f"{len(msg[18]) // 8}Q", msg[18]))
if len(msg) > 18 and msg[18] != b""
else [dst_kv_item_len] * len(dst_kv_ptrs)
)
if len(dst_kv_item_lens) != len(dst_kv_ptrs):
raise ValueError(
"dst_kv_item_lens length mismatch: "
f"got {len(dst_kv_item_lens)}, expected {len(dst_kv_ptrs)}"
)
dst_state_data_ptrs = (
unpack_int_lists(msg[7], "Q") if len(msg) > 7 and msg[7] != b"" else []
)
@@ -147,13 +249,15 @@ class KVArgsRegisterInfo:
dst_port=int(msg[2].decode("ascii")),
agent_name=msg[3].decode("ascii"),
agent_metadata=msg[4],
dst_kv_ptrs=list(struct.unpack(f"{len(msg[5]) // 8}Q", msg[5])),
dst_kv_ptrs=dst_kv_ptrs,
dst_kv_mem_kinds=dst_kv_mem_kinds,
dst_aux_ptrs=list(struct.unpack(f"{len(msg[6]) // 8}Q", msg[6])),
dst_state_data_ptrs=dst_state_data_ptrs,
gpu_id=int(msg[8].decode("ascii")),
decode_tp_size=int(msg[9].decode("ascii")),
decode_tp_rank=int(msg[10].decode("ascii")),
dst_kv_item_len=int(msg[11].decode("ascii")),
dst_kv_item_len=dst_kv_item_len,
dst_kv_item_lens=dst_kv_item_lens,
dst_num_slots=dst_num_slots,
dst_state_item_lens=dst_state_item_lens,
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
@@ -216,6 +320,10 @@ class TransferStatus:
received_state_per_pp: Set[int] = dataclasses.field(default_factory=set)
# Whether state data is expected (set based on state_type).
expects_state: bool = False
# KV part notifications for mixed-memory transfers. Keyed by
# (pp_rank, chunk_id); normal homogeneous transfers bypass this.
received_kv_parts_per_pp: Optional[Dict[Tuple[int, int], Set[int]]] = None
expected_kv_parts_per_pp: Optional[Dict[Tuple[int, int], int]] = None
def is_done(self):
if self.num_pp_ranks_expected is None or not self.received_aux:
@@ -245,6 +353,15 @@ class NixlKVManager(CommonKVManager):
is_mla_backend: Optional[bool] = False,
):
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
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),
)
self.src_mem_kind = (
_homogeneous_kv_mem_kind(self.kv_args.kv_data_mem_kinds, "source")
if disaggregation_mode == DisaggregationMode.PREFILL
else None
)
try:
from nixl._api import nixl_agent, nixl_agent_config, nixl_thread_sync_t
except ImportError as e:
@@ -304,6 +421,7 @@ class NixlKVManager(CommonKVManager):
)
self.prep_handles_slice_dst: Dict[str, Tuple[Any, int, int]] = {}
# peer_name -> (handle, num_slots, head_group_idx)
self.prep_handles_segment_src: Dict[Tuple[int, int, str], Any] = {}
self._num_slots_src: int = 0
if self.disaggregation_mode == DisaggregationMode.PREFILL:
@@ -513,47 +631,99 @@ class NixlKVManager(CommonKVManager):
def check_status(self, bootstrap_room: int):
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
def _init_equal_tp_prep_handle(
def _prep_equal_tp_dlist(
self,
peer_name: str,
kv_ptrs: list[int],
kv_item_lens: list[int],
kv_data_lens: list[int],
gpu_id: int,
num_slots: Optional[int] = None,
mem_kind: str = "VRAM",
kv_xfer_lens: Optional[list[int]] = None,
):
"""Pre-build NIXL dlist: all KV slots × all layers.
peer_name="" = src side; agent name = dst side. num_slots overrides the local
slot count — pass decode's count for the dst dlist (may differ from prefill).
Uses prefill's kv_item_lens as stride; requires equal per-slot byte size (equal-TP or MLA).
"""
if kv_xfer_lens is None:
kv_xfer_lens = kv_item_lens
if not (
len(kv_ptrs) == len(kv_item_lens) == len(kv_data_lens) == len(kv_xfer_lens)
):
raise ValueError(
"NIXL prepared dlist geometry length mismatch: "
f"ptrs={len(kv_ptrs)}, item_lens={len(kv_item_lens)}, "
f"data_lens={len(kv_data_lens)}, xfer_lens={len(kv_xfer_lens)}"
)
device_id = _nixl_device_id(mem_kind, gpu_id)
arrays = []
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set).
# Convert once at entry; all downstream arithmetic stays in uint64.
kv_ptrs_u64 = np.array(kv_ptrs, dtype=np.uint64)
for base_ptr, item_len, data_len in zip(
kv_ptrs_u64, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens
for base_ptr, item_len, data_len, xfer_len in zip(
kv_ptrs_u64, kv_item_lens, kv_data_lens, kv_xfer_lens
):
if xfer_len > item_len:
raise ValueError(
"NIXL prepared dlist transfer length exceeds item stride: "
f"xfer_len={xfer_len}, item_len={item_len}, mem_kind={mem_kind}"
)
n = num_slots if num_slots is not None else (data_len // item_len)
addrs = np.arange(n, dtype=np.uint64) * np.uint64(item_len) + base_ptr
arrays.append(
np.column_stack(
[
addrs,
np.full(n, item_len, dtype=np.uint64),
np.full(n, gpu_id, dtype=np.uint64),
np.full(n, xfer_len, dtype=np.uint64),
np.full(n, device_id, dtype=np.uint64),
]
)
)
self.prep_handles[peer_name] = self.agent.prep_xfer_dlist(
peer_name, np.vstack(arrays), "VRAM"
)
prep_handle = self.agent.prep_xfer_dlist(peer_name, np.vstack(arrays), mem_kind)
assert (
self.prep_handles[peer_name] is not None
prep_handle is not None
), f"prep_xfer_dlist returned None for peer '{peer_name}'"
return prep_handle
def _init_equal_tp_prep_handle(
self,
peer_name: str,
kv_ptrs: list[int],
gpu_id: int,
num_slots: Optional[int] = None,
mem_kind: str = "VRAM",
kv_item_lens: Optional[list[int]] = None,
kv_data_lens: Optional[list[int]] = None,
kv_xfer_lens: Optional[list[int]] = None,
):
"""Pre-build NIXL dlist: all KV slots × all layers.
peer_name="" = src side; agent name = dst side. num_slots overrides the local
slot count — pass decode's count for the dst dlist (may differ from prefill).
Uses prefill's kv_item_lens as stride; requires equal per-slot byte size (equal-TP or MLA).
Source dlists use prefill geometry; destination dlists must use decode
stride geometry but source transfer lengths, because HiSparse can transfer
directly into a host pool whose slot stride differs from prefill.
"""
if kv_item_lens is None:
kv_item_lens = self.kv_args.kv_item_lens
if kv_data_lens is None:
kv_data_lens = self.kv_args.kv_data_lens
self.prep_handles[peer_name] = self._prep_equal_tp_dlist(
peer_name,
kv_ptrs,
kv_item_lens,
kv_data_lens,
gpu_id,
num_slots=num_slots,
mem_kind=mem_kind,
kv_xfer_lens=kv_xfer_lens,
)
def _init_hetero_tp_prep_handle(
self, peer_name: str, decode_kv_args: KVArgsRegisterInfo
self,
peer_name: str,
decode_kv_args: KVArgsRegisterInfo,
src_mem_kind: str = "VRAM",
dst_mem_kind: str = "VRAM",
):
"""Pre-build NIXL dlists for TP-heterogeneous slice transfers.
@@ -630,10 +800,14 @@ class NixlKVManager(CommonKVManager):
[
addrs,
np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
np.full(len(addrs), self.kv_args.gpu_id, dtype=np.uint64),
np.full(
len(addrs),
_nixl_device_id(src_mem_kind, self.kv_args.gpu_id),
dtype=np.uint64,
),
]
)
src_handle = self.agent.prep_xfer_dlist("", src_array, "VRAM")
src_handle = self.agent.prep_xfer_dlist("", src_array, src_mem_kind)
assert (
src_handle is not None
), f"prep_xfer_dlist returned None for slice src (decode_tp_size={decode_tp_size})"
@@ -663,10 +837,14 @@ class NixlKVManager(CommonKVManager):
[
addrs,
np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
np.full(len(addrs), decode_kv_args.gpu_id, dtype=np.uint64),
np.full(
len(addrs),
_nixl_device_id(dst_mem_kind, decode_kv_args.gpu_id),
dtype=np.uint64,
),
]
)
dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, "VRAM")
dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, dst_mem_kind)
assert (
dst_handle is not None
), f"prep_xfer_dlist returned None for slice dst for peer '{peer_name}'"
@@ -676,25 +854,119 @@ class NixlKVManager(CommonKVManager):
head_group_idx,
)
def _init_mixed_equal_tp_prep_handles(
self,
peer_info: KVArgsRegisterInfo,
mem_segments: List[_KVXferMemSegment],
):
prepared_segments = []
for seg in mem_segments:
src_key = (seg.start, seg.end, seg.src_mem_kind)
src_handle = self.prep_handles_segment_src.get(src_key)
if src_handle is None:
src_handle = self._prep_equal_tp_dlist(
"",
self.kv_args.kv_data_ptrs[seg.start : seg.end],
self.kv_args.kv_item_lens[seg.start : seg.end],
self.kv_args.kv_data_lens[seg.start : seg.end],
self.kv_args.gpu_id,
mem_kind=seg.src_mem_kind,
)
self.prep_handles_segment_src[src_key] = src_handle
dst_num_slots = (
peer_info.dst_num_slots
if peer_info.dst_num_slots is not None
else self._num_slots_src
)
dst_kv_item_lens = peer_info.dst_kv_item_lens[seg.start : seg.end]
dst_kv_data_lens = [
item_len * dst_num_slots for item_len in dst_kv_item_lens
]
dst_handle = self._prep_equal_tp_dlist(
peer_info.agent_name,
peer_info.dst_kv_ptrs[seg.start : seg.end],
dst_kv_item_lens,
dst_kv_data_lens,
peer_info.gpu_id,
num_slots=peer_info.dst_num_slots,
mem_kind=seg.dst_mem_kind,
kv_xfer_lens=self.kv_args.kv_item_lens[seg.start : seg.end],
)
prepared_segments.append(
_KVXferPreparedSegment(
start=seg.start,
end=seg.end,
src_handle=src_handle,
dst_handle=dst_handle,
dst_num_slots=dst_num_slots,
)
)
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 self.is_mla_backend or peer_info.decode_tp_size == self.attn_tp_size:
# Safe to use prefill's kv_item_lens for the dst dlist stride:
# equal_tp guarantees identical heads-per-rank (same item_len);
# MLA latent shape is TP-invariant.
dst_mem_kind = None
try:
dst_mem_kind = _homogeneous_kv_mem_kind(
peer_info.dst_kv_mem_kinds, "destination"
)
except NotImplementedError:
mem_segments = _kv_xfer_mem_segments(
self.kv_args.kv_data_mem_kinds, peer_info.dst_kv_mem_kinds
)
if not mem_segments:
raise ValueError("NIXL KV transfer has no KV memory segments")
self._init_mixed_equal_tp_prep_handles(peer_info, mem_segments)
return
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
# Build the shared src dlist on the first equal-TP/MLA peer; later
# peers reuse it. Skipped entirely on heterogeneous-TP-only setups.
if "" not in self.prep_handles:
self._init_equal_tp_prep_handle(
"", self.kv_args.kv_data_ptrs, self.kv_args.gpu_id
"",
self.kv_args.kv_data_ptrs,
self.kv_args.gpu_id,
mem_kind=src_mem_kind,
)
dst_num_slots = (
peer_info.dst_num_slots
if peer_info.dst_num_slots is not None
else self._num_slots_src
)
dst_kv_item_lens = peer_info.dst_kv_item_lens
dst_kv_data_lens = [
item_len * dst_num_slots for item_len in dst_kv_item_lens
]
self._init_equal_tp_prep_handle(
peer_info.agent_name,
peer_info.dst_kv_ptrs,
peer_info.gpu_id,
num_slots=peer_info.dst_num_slots,
mem_kind=dst_mem_kind,
kv_item_lens=dst_kv_item_lens,
kv_data_lens=dst_kv_data_lens,
kv_xfer_lens=self.kv_args.kv_item_lens,
)
else:
self._init_hetero_tp_prep_handle(peer_info.agent_name, peer_info)
dst_mem_kind = _homogeneous_kv_mem_kind(
peer_info.dst_kv_mem_kinds, "destination"
)
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
if dst_mem_kind != "VRAM":
raise NotImplementedError(
"NIXL heterogeneous-TP direct-to-host KV transfer is not "
"implemented safely yet"
)
self._init_hetero_tp_prep_handle(
peer_info.agent_name,
peer_info,
src_mem_kind=src_mem_kind,
dst_mem_kind=dst_mem_kind,
)
def transfer_worker(self, queue: FastQueue, staging_buffer=None):
# Per-worker staging strategy: lazy-created on first chunk so we
@@ -758,6 +1030,8 @@ class NixlKVManager(CommonKVManager):
: len(chunked_dst_kv_indice)
]
src_prefill_kv_indices = kv_chunk.prefill_kv_indices
notif = (
f"{req.room}_kv_{kv_chunk.chunk_id}"
f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
@@ -782,6 +1056,7 @@ class NixlKVManager(CommonKVManager):
kv_xfer_handle, deferred = self._do_staging_transfer(
staging_strategy,
kv_chunk,
src_prefill_kv_indices,
req,
dst_info,
queue,
@@ -801,23 +1076,41 @@ class NixlKVManager(CommonKVManager):
if self.is_mla_backend or (
decode_tp_size == self.attn_tp_size
):
kv_xfer_handle = self.send_kvcache(
req.agent_name,
kv_chunk.prefill_kv_indices,
dst_info.dst_kv_ptrs,
chunked_dst_kv_indice,
dst_info.gpu_id,
notif,
)
if dst_info.kv_xfer_segments is None:
if dst_info.dst_homogeneous_mem_kind is None:
raise RuntimeError(
"Missing NIXL destination KV memory kind"
)
kv_xfer_handle = self.send_kvcache(
req.agent_name,
src_prefill_kv_indices,
dst_info.dst_kv_ptrs,
chunked_dst_kv_indice,
dst_info.gpu_id,
notif,
dst_mem_kind=(
dst_info.dst_homogeneous_mem_kind
),
)
else:
handles.extend(
self.send_kvcache_mixed(
req.agent_name,
src_prefill_kv_indices,
chunked_dst_kv_indice,
notif,
)
)
else:
kv_xfer_handle = self.send_kvcache_slice(
req.agent_name,
kv_chunk.prefill_kv_indices,
src_prefill_kv_indices,
chunked_dst_kv_indice,
notif,
)
handles.append(kv_xfer_handle)
if kv_xfer_handle is not None:
handles.append(kv_xfer_handle)
if kv_chunk.is_last_chunk:
dst_info = self.decode_kv_args_table[req.agent_name]
@@ -863,10 +1156,16 @@ class NixlKVManager(CommonKVManager):
continue
while handles:
states = [self.agent.check_xfer_state(h) for h in handles]
if any(s == "ERR" for s in states):
raise RuntimeError(f"NIXL transfer encountered ERR room={room}")
if all(s == "DONE" for s in states):
all_done = True
for handle in handles:
state = self.agent.check_xfer_state(handle)
if state == "ERR":
raise RuntimeError(
f"NIXL transfer encountered ERR room={room}"
)
if state != "DONE":
all_done = False
if all_done:
break
time.sleep(0)
@@ -899,13 +1198,34 @@ class NixlKVManager(CommonKVManager):
self.update_status(room, KVPoll.Failed)
def register_buffer_to_engine(self):
kv_addrs = []
for kv_data_ptr, kv_data_len in zip(
self.kv_args.kv_data_ptrs, self.kv_args.kv_data_lens
self.kv_descs = []
kv_addrs_by_mem_kind = {"VRAM": [], "DRAM": []}
for kv_data_ptr, kv_data_len, kv_mem_kind in zip(
self.kv_args.kv_data_ptrs,
self.kv_args.kv_data_lens,
self.kv_args.kv_data_mem_kinds,
):
kv_addrs.append((kv_data_ptr, kv_data_len, self.kv_args.gpu_id, ""))
self.kv_descs = self.agent.register_memory(kv_addrs, "VRAM")
logger.debug(f"Register kv tensors, len(kv_addr)= {len(kv_addrs)}")
kv_addrs_by_mem_kind[kv_mem_kind].append(
(
kv_data_ptr,
kv_data_len,
_nixl_device_id(kv_mem_kind, self.kv_args.gpu_id),
"",
)
)
for mem_kind in ("VRAM", "DRAM"):
kv_addrs = kv_addrs_by_mem_kind[mem_kind]
if not kv_addrs:
continue
kv_descs = self.agent.register_memory(kv_addrs, mem_kind)
logger.debug(
f"Register kv tensors, kind={mem_kind}, len(kv_addr)= {len(kv_addrs)}"
)
if not kv_descs:
raise Exception(
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 = []
@@ -957,9 +1277,12 @@ class NixlKVManager(CommonKVManager):
dst_data_indices: npt.NDArray[np.int32],
dst_gpu_id: int,
notif: str,
src_mem_kind: str = "VRAM",
dst_mem_kind: str = "VRAM",
):
"""Generic KV cache transfer supporting both MHA and MLA architectures.
Used by both send_kvcache and maybe_send_extra."""
# Prepped path (KV only; state transfers use the non-prepped path below).
if (
src_data_ptrs is self.kv_args.kv_data_ptrs
@@ -1079,14 +1402,18 @@ class NixlKVManager(CommonKVManager):
)
)
src_reqs = make_req_array(src_addrs, src_lens, self.kv_args.gpu_id)
dst_reqs = make_req_array(dst_addrs, dst_lens, dst_gpu_id)
src_reqs = make_req_array(
src_addrs, src_lens, _nixl_device_id(src_mem_kind, self.kv_args.gpu_id)
)
dst_reqs = make_req_array(
dst_addrs, dst_lens, _nixl_device_id(dst_mem_kind, dst_gpu_id)
)
logger.debug(
f"len(src_addrs): before group: {len(prefill_data_indices)}, after group: {len(src_addrs)}"
)
src_descs = self.agent.get_xfer_descs(src_reqs, "VRAM")
dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM")
src_descs = self.agent.get_xfer_descs(src_reqs, src_mem_kind)
dst_descs = self.agent.get_xfer_descs(dst_reqs, dst_mem_kind)
# Transfer data
xfer_handle = self.agent.initialize_xfer(
"WRITE",
@@ -1110,7 +1437,9 @@ class NixlKVManager(CommonKVManager):
dst_kv_indices: npt.NDArray[np.int32],
dst_gpu_id: int,
notif: str,
dst_mem_kind: str = "VRAM",
):
assert self.src_mem_kind is not None
return self._send_kvcache_generic(
peer_name=peer_name,
src_data_ptrs=self.kv_args.kv_data_ptrs,
@@ -1120,8 +1449,50 @@ class NixlKVManager(CommonKVManager):
dst_data_indices=dst_kv_indices,
dst_gpu_id=dst_gpu_id,
notif=notif,
src_mem_kind=self.src_mem_kind,
dst_mem_kind=dst_mem_kind,
)
def send_kvcache_mixed(
self,
peer_name: str,
prefill_kv_indices: npt.NDArray[np.int32],
dst_kv_indices: npt.NDArray[np.int32],
notif: str,
):
info = self.decode_kv_args_table[peer_name]
segments = info.kv_xfer_segments
assert segments is not None
if not segments:
raise RuntimeError(f"Missing NIXL mixed KV transfer plan for {peer_name}")
num_parts = len(segments)
handles = []
for part_idx, seg in enumerate(segments):
num_layers = seg.end - seg.start
src_indices = repeat_indices_over_layers(
prefill_kv_indices, num_layers, self._num_slots_src
)
dst_indices = repeat_indices_over_layers(
dst_kv_indices, num_layers, seg.dst_num_slots
)
part_notif = f"{notif}_part_{part_idx}_{num_parts}"
xfer_handle = self.agent.make_prepped_xfer(
"WRITE",
seg.src_handle,
src_indices,
seg.dst_handle,
dst_indices,
part_notif.encode("ascii"),
)
if not xfer_handle:
raise Exception("KVSender failed to create mixed prepped transfer")
state = self.agent.transfer(xfer_handle)
if state == "ERR":
raise Exception("KVSender failed to post mixed prepped transfer")
handles.append(xfer_handle)
return handles
def send_kvcache_slice(
self,
peer_name: str,
@@ -1300,6 +1671,7 @@ class NixlKVManager(CommonKVManager):
self,
staging_strategy,
kv_chunk: TransferKVChunk,
src_prefill_kv_indices: npt.NDArray[np.int32],
req: TransferInfo,
dst_info: KVArgsRegisterInfo,
queue: FastQueue,
@@ -1348,7 +1720,7 @@ class NixlKVManager(CommonKVManager):
)
handle = self.send_kvcache_staged(
req.agent_name,
kv_chunk.prefill_kv_indices,
src_prefill_kv_indices,
dst_info.staging.base_ptr + c_offset,
dst_info.staging.total_size - c_offset,
dst_info.gpu_id,
@@ -1693,6 +2065,7 @@ class NixlKVManager(CommonKVManager):
for msg in messages:
# Notification tag layouts (underscore-separated):
# kv: {room}_kv_{chunk_id}_{is_last}_{pp_rank} -> 5 fields
# kvpart:{room}_kv_{chunk_id}_{is_last}_{pp_rank}_part_{i}_{n}-> 8 fields
# stg: {room}_stg_{chunk_id}_{is_last}_{pp_rank}_{chunk_idx}
# _{page_start}_{num_pages}_{agent_name} -> 9 fields
# aux: {room}_aux -> 2 fields
@@ -1707,7 +2080,17 @@ class NixlKVManager(CommonKVManager):
chunk_id = int(components[2])
is_last_chunk = bool(int(components[3]))
pp_rank = int(components[4]) if len(components) > 4 else 0
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
if len(components) > 7 and components[5] == "part":
self._track_kv_part_arrival(
room,
chunk_id,
is_last_chunk,
pp_rank,
int(components[6]),
int(components[7]),
)
else:
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
elif tag == "stg":
self._handle_stg_notification(components, room)
elif tag == "aux":
@@ -1778,6 +2161,45 @@ class NixlKVManager(CommonKVManager):
):
self._maybe_submit_last_scatter(room)
def _track_kv_part_arrival(
self,
room: int,
chunk_id: int,
is_last_chunk: bool,
pp_rank: int,
part_idx: int,
num_parts: int,
):
"""Track one segment of a mixed-memory KV transfer."""
if num_parts <= 1:
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
return
if part_idx < 0 or part_idx >= num_parts:
raise RuntimeError(
f"NIXL KV part index out of range for room={room}, "
f"chunk={chunk_id}, pp_rank={pp_rank}: part={part_idx}, "
f"num_parts={num_parts}"
)
key = (pp_rank, chunk_id)
status = self.transfer_statuses[room]
if status.received_kv_parts_per_pp is None:
status.received_kv_parts_per_pp = defaultdict(set)
if status.expected_kv_parts_per_pp is None:
status.expected_kv_parts_per_pp = {}
expected = status.expected_kv_parts_per_pp.setdefault(key, num_parts)
if expected != num_parts:
raise RuntimeError(
f"NIXL KV part count mismatch for room={room}, chunk={chunk_id}, "
f"pp_rank={pp_rank}: got {num_parts}, expected {expected}"
)
parts = status.received_kv_parts_per_pp[key]
parts.add(part_idx)
if len(parts) == num_parts:
status.received_kv_parts_per_pp.pop(key, None)
status.expected_kv_parts_per_pp.pop(key, None)
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
def _handle_staging_chunk_arrived(
self,
room: int,
@@ -2095,6 +2517,13 @@ class NixlKVReceiver(CommonKVReceiver):
packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
)
packed_kv_data_mem_kinds = _pack_kv_mem_kinds(
self.kv_mgr.kv_args.kv_data_mem_kinds
)
packed_kv_item_lens = b"".join(
struct.pack("Q", item_len)
for item_len in self.kv_mgr.kv_args.kv_item_lens
)
packed_aux_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.aux_data_ptrs
)
@@ -2145,6 +2574,8 @@ class NixlKVReceiver(CommonKVReceiver):
packed_staging_base_ptr,
staging_total_size_str,
str(dst_num_slots).encode("ascii"),
packed_kv_data_mem_kinds,
packed_kv_item_lens,
]
)
+1 -2
View File
@@ -264,7 +264,7 @@ class PrefillBootstrapQueue:
def finalize_bootstrap(self, req: Req) -> bool:
"""Initialize the sender after bootstrap completes.
Returns False if no metadata buffer is available (non-terminal)."""
assert req.pending_bootstrap, f"finalize_bootstrap is not idempotent"
assert req.pending_bootstrap, "finalize_bootstrap is not idempotent"
if not self.ensure_metadata_buffer(req):
return False
@@ -737,7 +737,6 @@ class SchedulerDisaggregationPrefillMixin:
undone_reqs: List[Req] = []
# Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue
for req, poll in zip(self.disagg_prefill_inflight_queue, polls):
if rids_to_check is not None:
if req.rid not in rids_to_check:
undone_reqs.append(req)
@@ -57,12 +57,14 @@ class HiSparseCoordinator:
device: str,
tp_group,
host_to_device_ratio: int = 2,
swap_in_block_size: int = 960,
):
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
self.top_k = top_k
self.device_buffer_size = device_buffer_size
self.device = device
self.swap_in_block_size = swap_in_block_size
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
self.is_dsv4_hisparse = isinstance(
@@ -815,8 +817,6 @@ class HiSparseCoordinator:
top_k_indices = self.top_k_device_locs_buffer[:num_reqs]
top_k_indices.fill_(-1)
# todo, adjustable for performance
block_size = 1024
swap_in_fn = (
load_cache_to_device_buffer_dsv4_mla
if self.is_dsv4_hisparse
@@ -837,7 +837,7 @@ class HiSparseCoordinator:
num_top_k=self.top_k,
hot_buffer_size=self.device_buffer_size,
page_size=1,
block_size=block_size,
block_size=self.swap_in_block_size,
num_real_reqs=self.num_real_reqs,
)
return top_k_indices
@@ -58,6 +58,7 @@ class SparseConfig:
top_k: int = 2048
device_buffer_size: int = 4096
host_to_device_ratio: int = 2
swap_in_block_size: int = 960
algorithm: Optional[str] = None
backend: Optional[str] = None
page_size: Optional[int] = None
@@ -62,7 +62,7 @@ def _parse_sparse_config(server_args) -> SparseConfig:
"""Parse hierarchical sparse config from JSON string.
Required fields with defaults: top_k (2048), device_buffer_size (2*top_k),
host_to_device_ratio (2).
host_to_device_ratio (2), swap_in_block_size (960).
Optional fields (default None): algorithm, backend, min_sparse_prompt_len,
page_size. All remaining fields go to sparse_extra_config.
"""
@@ -78,11 +78,20 @@ def _parse_sparse_config(server_args) -> SparseConfig:
top_k = extra_config.pop("top_k", 2048)
device_buffer_size = extra_config.pop("device_buffer_size", 2 * top_k)
host_to_device_ratio = extra_config.pop("host_to_device_ratio", 2)
swap_in_block_size = extra_config.pop("swap_in_block_size", 960)
if device_buffer_size < top_k:
raise ValueError(
f"device_buffer_size ({device_buffer_size}) must be no smaller than top_k ({top_k})"
)
if not isinstance(swap_in_block_size, int) or isinstance(swap_in_block_size, bool):
raise ValueError(
f"swap_in_block_size must be an integer, got {swap_in_block_size!r}"
)
if swap_in_block_size <= 0 or swap_in_block_size > 1024:
raise ValueError(
f"swap_in_block_size ({swap_in_block_size}) must be in the range [1, 1024]"
)
algorithm = extra_config.pop("algorithm", None)
backend = extra_config.pop("backend", None)
@@ -93,6 +102,7 @@ def _parse_sparse_config(server_args) -> SparseConfig:
top_k=top_k,
device_buffer_size=device_buffer_size,
host_to_device_ratio=host_to_device_ratio,
swap_in_block_size=swap_in_block_size,
algorithm=algorithm,
backend=backend,
page_size=page_size,
@@ -856,6 +856,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
else self.tp_group.cpu_group
),
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
swap_in_block_size=hisparse_cfg.swap_in_block_size,
)
self.init_routed_experts_capturer()
@@ -17,10 +17,6 @@ register_cuda_ci(est_time=500, stage="base-c", runner_config="deepep-8-gpu-h200"
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}'
DSV4_FLASH_LOADER_CONFIG = '{"enable_multithread_load": true, "num_threads": 64}'
DSV4_HISPARSE_CONFIG = (
'{"top_k":512,"device_buffer_size":4096,"host_to_device_ratio":2}'
)
DSV4_FLASH_ENV = {
"SGLANG_DSV4_FP4_EXPERTS": "0",
@@ -129,104 +125,5 @@ class TestDisaggregationDSV4(SpecDecodingMixin, PDDisaggregationServerBase, GSM8
)
class TestDisaggregationDSV4HiSparseMooncake(PDDisaggregationServerBase, GSM8KMixin):
gsm8k_accuracy_thres = 0.93
gsm8k_num_questions = 200
gsm8k_num_shots = 20
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = try_cached_model(DSV4_FLASH_MODEL)
cls.start_prefill()
cls.start_decode()
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
cls.launch_lb()
@classmethod
def start_prefill(cls):
prefill_args = [
"--trust-remote-code",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--watchdog-timeout",
"900",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
env=DSV4_FLASH_ENV,
)
@classmethod
def start_decode(cls):
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--base-gpu-id",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--enable-hisparse",
"--hisparse-config",
DSV4_HISPARSE_CONFIG,
"--watchdog-timeout",
"900",
]
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
env=DSV4_FLASH_ENV,
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,165 @@
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
)
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
is_in_ci,
popen_launch_pd_server,
try_cached_model,
)
register_cuda_ci(est_time=1000, stage="extra-b", runner_config="deepep-8-gpu-h200")
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
DSV4_FLASH_LOADER_CONFIG = '{"enable_multithread_load": true, "num_threads": 64}'
DSV4_HISPARSE_CONFIG = (
'{"top_k":512,"device_buffer_size":4096,"host_to_device_ratio":2}'
)
DSV4_FLASH_ENV = {
"SGLANG_DSV4_FP4_EXPERTS": "0",
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256",
}
DSV4_NIXL_SERVER_LAUNCH_TIMEOUT = 1800
def _has_nixl():
try:
import nixl._api # noqa: F401
except Exception:
return False
return True
class TestDisaggregationDSV4HiSparseBase(PDDisaggregationServerBase, GSM8KMixin):
gsm8k_accuracy_thres = 0.93
gsm8k_num_questions = 200
gsm8k_num_shots = 20
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = try_cached_model(DSV4_FLASH_MODEL)
cls.start_prefill()
cls.start_decode()
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
cls.launch_lb()
@classmethod
def start_prefill(cls):
prefill_args = [
"--trust-remote-code",
"--disaggregation-mode",
"prefill",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--watchdog-timeout",
"900",
]
prefill_args += cls.transfer_backend + cls.rdma_devices
cls.process_prefill = popen_launch_pd_server(
cls.model,
cls.prefill_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=prefill_args,
env=DSV4_FLASH_ENV,
)
@classmethod
def start_decode(cls):
decode_args = [
"--trust-remote-code",
"--disaggregation-mode",
"decode",
"--disaggregation-bootstrap-port",
cls.bootstrap_port,
"--tp",
4,
"--base-gpu-id",
4,
"--page-size",
256,
"--chunked-prefill-size",
8192,
"--max-running-requests",
16,
"--mem-fraction-static",
0.9,
"--skip-server-warmup",
"--reasoning-parser",
"deepseek-v4",
"--tool-call-parser",
"deepseekv4",
"--model-loader-extra-config",
DSV4_FLASH_LOADER_CONFIG,
"--enable-hisparse",
"--hisparse-config",
DSV4_HISPARSE_CONFIG,
"--watchdog-timeout",
"900",
]
decode_args += cls.transfer_backend + cls.rdma_devices
cls.process_decode = popen_launch_pd_server(
cls.model,
cls.decode_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=decode_args,
env=DSV4_FLASH_ENV,
)
@unittest.skipUnless(
is_in_ci() or _has_nixl(),
"NIXL is required for DSV4 HiSparse disaggregation coverage.",
)
class TestDisaggregationDSV4HiSparseNixl(TestDisaggregationDSV4HiSparseBase):
@classmethod
def setUpClass(cls):
PDDisaggregationServerBase.setUpClass.__func__(cls)
cls.transfer_backend = ["--disaggregation-transfer-backend", "nixl"]
cls.model = try_cached_model(DSV4_FLASH_MODEL)
cls.start_prefill()
cls.start_decode()
cls.wait_server_ready(
cls.prefill_url + "/health",
timeout=DSV4_NIXL_SERVER_LAUNCH_TIMEOUT,
process=cls.process_prefill,
)
cls.wait_server_ready(
cls.decode_url + "/health",
timeout=DSV4_NIXL_SERVER_LAUNCH_TIMEOUT,
process=cls.process_decode,
)
cls.launch_lb()
if __name__ == "__main__":
unittest.main()
@@ -212,6 +212,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
pack_int_lists(state_dims, "I"),
struct.pack("Q", staging_ptr),
b"1048576",
b"64",
b"DRAM,DRAM",
b"".join(struct.pack("Q", item_len) for item_len in [1024, 2048]),
]
info = KVArgsRegisterInfo.from_zmq(msg)
@@ -228,6 +231,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
self.assertEqual(info.decode_tp_size, 4)
self.assertEqual(info.decode_tp_rank, 1)
self.assertEqual(info.dst_kv_item_len, 1024)
self.assertEqual(info.dst_kv_item_lens, [1024, 2048])
self.assertEqual(info.dst_num_slots, 64)
self.assertEqual(info.dst_kv_mem_kinds, ["DRAM", "DRAM"])
self.assertEqual(info.dst_state_item_lens, state_item_lens)
self.assertEqual(info.dst_state_dim_per_tensor, state_dims)
self.assertIsNotNone(info.staging)
@@ -255,6 +261,7 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
self.assertEqual(info.dst_state_data_ptrs, [])
self.assertEqual(info.dst_state_item_lens, [])
self.assertEqual(info.dst_state_dim_per_tensor, [])
self.assertEqual(info.dst_kv_item_lens, [256])
self.assertIsNone(info.staging)
@@ -523,6 +530,33 @@ class TestNixlStaging(CustomTestCase):
mgr.server_args = SimpleNamespace(chunked_prefill_size=4)
return mgr
def test_register_buffer_to_engine_groups_kv_memory_kinds_in_one_pass(self):
agent = StagingFakeAgent(register_result=["desc"])
mgr = self._make_manager(agent)
mgr.kv_args.kv_data_ptrs = [0x1000, 0x2000, 0x3000]
mgr.kv_args.kv_data_lens = [64, 128, 256]
mgr.kv_args.kv_data_mem_kinds = ["VRAM", "DRAM", "VRAM"]
mgr.kv_args.aux_data_ptrs = [0x4000]
mgr.kv_args.aux_data_lens = [32]
mgr.kv_args.state_data_ptrs = []
mgr.kv_args.state_data_lens = []
mgr.register_buffer_to_engine()
self.assertEqual(
agent.register_memory_calls,
[
(
[(0x1000, 64, 1, ""), (0x3000, 256, 1, "")],
"VRAM",
),
([(0x2000, 128, 0, "")], "DRAM"),
([(0x4000, 32, 0, "")], "DRAM"),
],
)
self.assertEqual(mgr.kv_descs, [["desc"], ["desc"]])
self.assertEqual(mgr.aux_descs, ["desc"])
def test_register_staging_memory_uses_vram_and_fails_on_empty_descs(self):
agent = StagingFakeAgent(register_result=["staging"])
mgr = self._make_manager(agent)
@@ -601,7 +635,12 @@ class TestNixlStaging(CustomTestCase):
},
):
handle, deferred = mgr._do_staging_transfer(
strategy, kv_chunk, req, SimpleNamespace(), queue
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
req,
SimpleNamespace(),
queue,
)
self.assertIsNone(handle)
@@ -640,6 +679,7 @@ class TestNixlStaging(CustomTestCase):
mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
SimpleNamespace(room=3, agent_name="decode_agent"),
SimpleNamespace(),
FakeQueue(),
@@ -666,12 +706,14 @@ class TestNixlStaging(CustomTestCase):
agent_name="decode_agent",
agent_metadata=b"",
dst_kv_ptrs=[],
dst_kv_mem_kinds=[],
dst_aux_ptrs=[],
dst_state_data_ptrs=[],
gpu_id=5,
decode_tp_size=1,
decode_tp_rank=0,
dst_kv_item_len=128,
dst_kv_item_lens=[],
staging=SimpleNamespace(base_ptr=0x8000, total_size=4096),
)
calls = []
@@ -682,6 +724,7 @@ class TestNixlStaging(CustomTestCase):
handle, deferred = mgr._do_staging_transfer(
strategy,
kv_chunk,
kv_chunk.prefill_kv_indices,
SimpleNamespace(room=3, agent_name="decode_agent"),
dst_info,
FakeQueue(),