diff --git a/docs_new/docs/advanced_features/hisparse_guide.mdx b/docs_new/docs/advanced_features/hisparse_guide.mdx
index 78b71288a..f3dc2c321 100644
--- a/docs_new/docs/advanced_features/hisparse_guide.mdx
+++ b/docs_new/docs/advanced_features/hisparse_guide.mdx
@@ -108,10 +108,15 @@ Pass as a JSON string via `--hisparse-config`:
int |
Ratio of logical pool size to device pool size, determining host memory capacity |
+
+ swap_in_block_size |
+ int / 960 |
+ CUDA thread-block size for the HiSparse swap-in kernel |
+
-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.
diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py
index f1b49b2c8..074feae11 100644
--- a/python/sglang/srt/disaggregation/base/conn.py
+++ b/python/sglang/srt/disaggregation/base/conn.py
@@ -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,
diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py
index 83571ced1..83888afae 100644
--- a/python/sglang/srt/disaggregation/decode.py
+++ b/python/sglang/srt/disaggregation/decode.py
@@ -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 = (
diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py
index 2e6b79c38..30f0bc8ca 100644
--- a/python/sglang/srt/disaggregation/nixl/conn.py
+++ b/python/sglang/srt/disaggregation/nixl/conn.py
@@ -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,
]
)
diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py
index ee8c82aaa..acc8c46c6 100644
--- a/python/sglang/srt/disaggregation/prefill.py
+++ b/python/sglang/srt/disaggregation/prefill.py
@@ -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)
diff --git a/python/sglang/srt/managers/hisparse_coordinator.py b/python/sglang/srt/managers/hisparse_coordinator.py
index 2ea0482f8..5e8f1bf8c 100644
--- a/python/sglang/srt/managers/hisparse_coordinator.py
+++ b/python/sglang/srt/managers/hisparse_coordinator.py
@@ -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
diff --git a/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py b/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py
index cf7c1d06c..5d2e4849b 100644
--- a/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py
+++ b/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py
@@ -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
diff --git a/python/sglang/srt/mem_cache/sparsity/factory.py b/python/sglang/srt/mem_cache/sparsity/factory.py
index 86804d656..7bd141760 100644
--- a/python/sglang/srt/mem_cache/sparsity/factory.py
+++ b/python/sglang/srt/mem_cache/sparsity/factory.py
@@ -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,
diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py
index afcfafd36..b03125d29 100644
--- a/python/sglang/srt/model_executor/model_runner.py
+++ b/python/sglang/srt/model_executor/model_runner.py
@@ -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()
diff --git a/test/registered/disaggregation/test_disaggregation_dsv4.py b/test/registered/disaggregation/test_disaggregation_dsv4.py
index 5d54c4436..a60ef9fe2 100644
--- a/test/registered/disaggregation/test_disaggregation_dsv4.py
+++ b/test/registered/disaggregation/test_disaggregation_dsv4.py
@@ -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()
diff --git a/test/registered/disaggregation/test_disaggregation_hisparse.py b/test/registered/disaggregation/test_disaggregation_hisparse.py
new file mode 100644
index 000000000..a80c7e277
--- /dev/null
+++ b/test/registered/disaggregation/test_disaggregation_hisparse.py
@@ -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()
diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py
index 6ff816a7c..4ea3ebee3 100644
--- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py
+++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py
@@ -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(),