From b030b1a5f3636ba929f59b74be4261a9c09c6364 Mon Sep 17 00:00:00 2001 From: ishandhanani <82981111+ishandhanani@users.noreply.github.com> Date: Sat, 27 Jun 2026 16:32:29 +0200 Subject: [PATCH] hisparse: support NIXL DRAM KV destinations for HiSparse (#27563) Co-authored-by: Zhangheng Co-authored-by: Shangming Cai --- .../docs/advanced_features/hisparse_guide.mdx | 9 +- python/sglang/srt/disaggregation/base/conn.py | 2 - python/sglang/srt/disaggregation/decode.py | 9 + python/sglang/srt/disaggregation/nixl/conn.py | 537 ++++++++++++++++-- python/sglang/srt/disaggregation/prefill.py | 3 +- .../srt/managers/hisparse_coordinator.py | 6 +- .../sparsity/core/sparse_coordinator.py | 1 + .../sglang/srt/mem_cache/sparsity/factory.py | 12 +- .../sglang/srt/model_executor/model_runner.py | 1 + .../test_disaggregation_dsv4.py | 103 ---- .../test_disaggregation_hisparse.py | 165 ++++++ .../disaggregation/test_nixl_backend_basic.py | 45 +- 12 files changed, 726 insertions(+), 167 deletions(-) create mode 100644 test/registered/disaggregation/test_disaggregation_hisparse.py 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(),