diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 19f2dd627..0fbdebde9 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -8,7 +8,7 @@ import threading import time import uuid from collections import defaultdict -from typing import TYPE_CHECKING, Dict, List, Optional, Set +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple import numpy as np import numpy.typing as npt @@ -118,6 +118,7 @@ class KVArgsRegisterInfo: decode_tp_size: int decode_tp_rank: int dst_kv_item_len: 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) # Keep last: optional, parsed from a variable-length tail of the ZMQ @@ -135,6 +136,9 @@ class KVArgsRegisterInfo: dst_state_dim_per_tensor = ( unpack_int_lists(msg[13], "I") if len(msg) > 13 and len(msg[13]) > 0 else [] ) + dst_num_slots = ( + int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else None + ) return cls( room=str(msg[0].decode("ascii")), @@ -149,12 +153,50 @@ class KVArgsRegisterInfo: 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_num_slots=dst_num_slots, dst_state_item_lens=dst_state_item_lens, dst_state_dim_per_tensor=dst_state_dim_per_tensor, staging=StagingRegisterInfo.from_zmq_fields(msg, 14), ) +def expand_page_indices_for_slice( + page_indices: npt.NDArray[np.int32], + num_ptr_pairs: int, + num_slots: int, + page_size: int, + num_groups: int = 1, + head_group_idx: int = 0, +) -> npt.NDArray[np.int32]: + """Map page slot indices to flat dlist indices for the slice prepped path. + + Dlist layout: num_ptr_pairs blocks of (num_slots * page_size * num_groups), + with [slot, token, group] interleaving. head_group_idx selects one group (0 for dst). + """ + token_offsets = np.arange(page_size, dtype=np.int32) + pair_stride = num_slots * page_size * num_groups + within_pair = ( + page_indices[:, None] * (page_size * num_groups) + + token_offsets[None, :] * num_groups + + head_group_idx + ).ravel() + pair_offsets = np.arange(num_ptr_pairs, dtype=np.int64) * pair_stride + return (pair_offsets[:, None] + within_pair[None, :]).ravel().astype(np.int32) + + +def repeat_indices_over_layers( + indices: npt.NDArray[np.int32], num_layers: int, layer_length: int +) -> npt.NDArray[np.int32]: + """Map per-slot token indices to flat indices in a pre-built descriptor list. + + Each of ``num_layers`` blocks has ``layer_length`` slots; block i is offset by + ``i * layer_length``. Works uniformly for both MLA (one ptr/layer) and MHA + (K+V ptrs, 2×N entries). + """ + offsets = np.arange(num_layers, dtype=np.int32) * layer_length + return (offsets[:, None] + indices[None, :]).ravel().astype(np.int32) + + @dataclasses.dataclass class TransferStatus: """Used by KV Receiver to know when a transfer is done.""" @@ -249,8 +291,18 @@ class NixlKVManager(CommonKVManager): self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() self.kv_buffer_tensors = None + self.prep_handles: Dict[str, Any] = {} + self.prep_handle_slice_src: Optional[Tuple[Any, int, int, int]] = ( + None # (handle, num_groups, num_ptr_pairs, num_slots) + ) + self.prep_handles_slice_dst: Dict[str, Tuple[Any, int, int]] = {} + # peer_name -> (handle, num_slots, head_group_idx) + self._num_slots_src: int = 0 if self.disaggregation_mode == DisaggregationMode.PREFILL: + self._num_slots_src = ( + self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0] + ) transfer_queue_size = envs.SGLANG_DISAGGREGATION_QUEUE_SIZE.get() self.transfer_queues: List[FastQueue] = [ FastQueue() for _ in range(transfer_queue_size) @@ -454,6 +506,187 @@ 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( + self, + peer_name: str, + kv_ptrs: list[int], + gpu_id: int, + num_slots: Optional[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). + """ + arrays = [] + for base_ptr, item_len, data_len in zip( + kv_ptrs, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens + ): + n = num_slots if num_slots is not None else (data_len // item_len) + addrs = np.arange(n, dtype=np.int64) * item_len + base_ptr + arrays.append( + np.column_stack( + [ + addrs, + np.full(n, item_len, dtype=np.int64), + np.full(n, gpu_id, dtype=np.int64), + ] + ) + ) + + self.prep_handles[peer_name] = self.agent.prep_xfer_dlist( + peer_name, np.vstack(arrays), "VRAM" + ) + assert ( + self.prep_handles[peer_name] is not None + ), f"prep_xfer_dlist returned None for peer '{peer_name}'" + + def _init_hetero_tp_prep_handle( + self, peer_name: str, decode_kv_args: KVArgsRegisterInfo + ): + """Pre-build NIXL dlists for TP-heterogeneous slice transfers. + + Src dlist shared across decode peers (same TP size). prefill_tp < decode_tp: + interleave num_groups per token, peers select via head_group_idx. + prefill_tp > decode_tp: num_groups=1. Dst dlist is per-peer. + """ + decode_tp_size = decode_kv_args.decode_tp_size + dst_kv_item_len = decode_kv_args.dst_kv_item_len + prefill_tp_size = self.attn_tp_size + + page_size = self.kv_args.page_size + + total_kv_heads = getattr(self.kv_args, "total_kv_head_num", 0) + if total_kv_heads <= 0: + total_kv_heads = self.kv_args.kv_head_num * prefill_tp_size + + src_heads_per_rank = max(1, total_kv_heads // prefill_tp_size) + dst_heads_per_rank = max(1, total_kv_heads // decode_tp_size) + bytes_per_head_slice = dst_kv_item_len // page_size // dst_heads_per_rank + + if prefill_tp_size > decode_tp_size: + # Multiple prefill ranks feed one decode rank: each prefill rank sends + # all its src heads to a specific head-range in the decode rank. + src_replication = max(1, prefill_tp_size // total_kv_heads) + local_tp_rank_in_group = self.kv_args.engine_rank % prefill_tp_size + num_groups = 1 + num_heads_to_send = src_heads_per_rank + head_group_idx = 0 + unique_head_idx = local_tp_rank_in_group // src_replication + dst_head_start = (unique_head_idx * src_heads_per_rank) % dst_heads_per_rank + dst_head_offset = dst_head_start * bytes_per_head_slice + else: + # One prefill rank feeds multiple decode ranks: interleave num_groups + # head-groups in the src dlist so each decode rank picks its slice. + dst_tp_rank_in_group = decode_kv_args.decode_tp_rank % decode_tp_size + num_groups = decode_tp_size // prefill_tp_size + num_heads_to_send = dst_heads_per_rank + src_head_start = ( + dst_tp_rank_in_group * dst_heads_per_rank + ) % src_heads_per_rank + head_group_idx = src_head_start // dst_heads_per_rank + dst_head_offset = 0 + + src_kv_item_len = self.kv_args.kv_item_lens[0] + bytes_per_token_to_send = num_heads_to_send * bytes_per_head_slice + bytes_per_token_src = src_kv_item_len // page_size + bytes_per_token_dst = dst_kv_item_len // page_size + + src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_pp = ( + self.get_mha_kv_ptrs_with_pp( + self.kv_args.kv_data_ptrs, decode_kv_args.dst_kv_ptrs + ) + ) + src_ptrs = list(src_k_ptrs[:layers_pp]) + list(src_v_ptrs[:layers_pp]) + dst_ptrs = list(dst_k_ptrs[:layers_pp]) + list(dst_v_ptrs[:layers_pp]) + num_ptr_pairs = len(src_ptrs) + + num_slots = self.kv_args.kv_data_lens[0] // src_kv_item_len + slots = np.arange(num_slots, dtype=np.int64) + tokens = np.arange(page_size, dtype=np.int64) # reused in dst dlist below + groups = np.arange(num_groups, dtype=np.int64) + + # Src dlist built once and shared. + if self.prep_handle_slice_src is None: + # (ptr, slot, token, group) → ravel; groups interleaved per token. + src_ptrs_arr = np.array(src_ptrs, dtype=np.int64) + addrs = ( + src_ptrs_arr[:, None, None, None] + + slots[None, :, None, None] * src_kv_item_len + + tokens[None, None, :, None] * bytes_per_token_src + + groups[None, None, None, :] * bytes_per_token_to_send + ).ravel() + src_array = np.column_stack( + [ + addrs, + np.full(len(addrs), bytes_per_token_to_send, dtype=np.int64), + np.full(len(addrs), self.kv_args.gpu_id, dtype=np.int64), + ] + ) + src_handle = self.agent.prep_xfer_dlist("", src_array, "VRAM") + assert ( + src_handle is not None + ), f"prep_xfer_dlist returned None for slice src (decode_tp_size={decode_tp_size})" + self.prep_handle_slice_src = ( + src_handle, + num_groups, + num_ptr_pairs, + num_slots, + ) + + # Dst dlist per-peer; use decode's slot count (may exceed prefill's). + num_slots_dst = ( + decode_kv_args.dst_num_slots + if decode_kv_args.dst_num_slots is not None + else num_slots + ) + dst_slots = np.arange(num_slots_dst, dtype=np.int64) + # (ptr, slot, token) → ravel. + dst_ptrs_arr = np.array(dst_ptrs, dtype=np.int64) + addrs = ( + dst_ptrs_arr[:, None, None] + + dst_slots[None, :, None] * dst_kv_item_len + + tokens[None, None, :] * bytes_per_token_dst + + dst_head_offset + ).ravel() + dst_array = np.column_stack( + [ + addrs, + np.full(len(addrs), bytes_per_token_to_send, dtype=np.int64), + np.full(len(addrs), decode_kv_args.gpu_id, dtype=np.int64), + ] + ) + dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, "VRAM") + assert ( + dst_handle is not None + ), f"prep_xfer_dlist returned None for slice dst for peer '{peer_name}'" + self.prep_handles_slice_dst[peer_name] = ( + dst_handle, + num_slots_dst, + head_group_idx, + ) + + def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo): + 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. + # 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._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, + ) + else: + self._init_hetero_tp_prep_handle(peer_info.agent_name, peer_info) + def transfer_worker(self, queue: FastQueue, staging_buffer=None): # Per-worker staging strategy: lazy-created on first chunk so we # see kv_buffer_tensors (set by ModelRunner after engine init). @@ -463,6 +696,7 @@ class NixlKVManager(CommonKVManager): while True: kv_chunk: TransferKVChunk = queue.get() room = kv_chunk.room + handles: List[Any] = [] try: if self.check_status(room) == KVPoll.Failed: continue @@ -481,7 +715,6 @@ class NixlKVManager(CommonKVManager): self.update_status(room, KVPoll.Transferring) reqs_to_be_processed = list(self.transfer_infos[room].values()) - handles: List = [] # Set when staging allocation/watermark is not yet ready and # the chunk has been re-enqueued. We then break out of the @@ -516,6 +749,11 @@ class NixlKVManager(CommonKVManager): : len(chunked_dst_kv_indice) ] + notif = ( + f"{req.room}_kv_{kv_chunk.chunk_id}" + f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}" + ) + # Decide which kv send path to use: # 1. Staging (heterogeneous TP, both sides have # registered staging, watermark/alloc ready) @@ -551,10 +789,6 @@ class NixlKVManager(CommonKVManager): # the slice path below. if kv_xfer_handle is None: - notif = ( - f"{req.room}_kv_{kv_chunk.chunk_id}" - f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}" - ) if self.is_mla_backend or ( decode_tp_size == self.attn_tp_size ): @@ -570,14 +804,8 @@ class NixlKVManager(CommonKVManager): kv_xfer_handle = self.send_kvcache_slice( req.agent_name, kv_chunk.prefill_kv_indices, - dst_info.dst_kv_ptrs, chunked_dst_kv_indice, - dst_info.gpu_id, notif, - prefill_tp_size=self.attn_tp_size, - decode_tp_size=decode_tp_size, - decode_tp_rank=dst_info.decode_tp_rank, - dst_kv_item_len=dst_info.dst_kv_item_len, ) handles.append(kv_xfer_handle) @@ -707,6 +935,8 @@ class NixlKVManager(CommonKVManager): return self.decode_kv_args_table[agent_name] = decode_kv_args self.agent.add_remote_agent(decode_kv_args.agent_metadata) + if self.disaggregation_mode == DisaggregationMode.PREFILL: + self._prepare_payload_xfer(decode_kv_args) def _send_kvcache_generic( self, @@ -721,6 +951,43 @@ class NixlKVManager(CommonKVManager): ): """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 + and "" in self.prep_handles + and peer_name in self.prep_handles + ): + src_prep = self.prep_handles[""] + dst_prep = self.prep_handles[peer_name] + info = self.decode_kv_args_table[peer_name] + num_slots_dst = ( + info.dst_num_slots + if info.dst_num_slots is not None + else self._num_slots_src + ) + num_layers = len(item_lens) + src_indices = repeat_indices_over_layers( + prefill_data_indices, num_layers, self._num_slots_src + ) + dst_indices = repeat_indices_over_layers( + dst_data_indices, num_layers, num_slots_dst + ) + xfer_handle = self.agent.make_prepped_xfer( + "WRITE", + src_prep, + src_indices, + dst_prep, + dst_indices, + notif.encode("ascii"), + ) + if not xfer_handle: + raise Exception("KVSender failed to create prepped transfer") + state = self.agent.transfer(xfer_handle) + if state == "ERR": + raise Exception("KVSender failed to post prepped transfer") + return xfer_handle + + # Non-prepped path: used for state transfers (SWA/NSA) via maybe_send_extra. # Convert pointer lists to np.uint64 arrays up front. # torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set, e.g. # 0xffff81ab54e01000). Casting here prevents overflow when these values @@ -850,139 +1117,46 @@ class NixlKVManager(CommonKVManager): self, peer_name: str, prefill_kv_indices: npt.NDArray[np.int32], - dst_kv_ptrs: list[int], dst_kv_indices: npt.NDArray[np.int32], - dst_gpu_id: int, notif: str, - prefill_tp_size: int, - decode_tp_size: int, - decode_tp_rank: int, - dst_kv_item_len: int, ): - # Get configuration from kv_args - local_tp_rank_in_group = self.kv_args.engine_rank % prefill_tp_size - dst_tp_rank_in_group = decode_tp_rank % decode_tp_size - - src_kv_item_len = self.kv_args.kv_item_lens[0] - page_size = self.kv_args.page_size - - # Use total KV head count (not per-rank) for correct head distribution. - # Per-rank kv_head_num is max(1, total//tp) which loses info when total < tp. - total_kv_heads = getattr(self.kv_args, "total_kv_head_num", 0) - if total_kv_heads <= 0: - total_kv_heads = self.kv_args.kv_head_num * prefill_tp_size - - src_heads_per_rank = max(1, total_kv_heads // prefill_tp_size) - dst_heads_per_rank = max(1, total_kv_heads // decode_tp_size) - - bytes_per_head_slice_to_send = ( - dst_kv_item_len // page_size // dst_heads_per_rank + # Prepped path: src dlist is shared per decode_tp_size; dst is per peer. + assert self.prep_handle_slice_src is not None + assert peer_name in self.prep_handles_slice_dst + src_handle, num_groups, num_ptr_pairs, num_slots_src = ( + self.prep_handle_slice_src ) - - # GQA replication: how many prefill ranks share the same KV head - src_replication = max(1, prefill_tp_size // total_kv_heads) - - # Determine which heads to send - if prefill_tp_size > decode_tp_size: - # Multiple prefill ranks to one decode rank - src_head_start_offset = 0 - num_heads_to_send = src_heads_per_rank - unique_head_idx = local_tp_rank_in_group // src_replication - dst_head_start_offset = ( - unique_head_idx * src_heads_per_rank - ) % dst_heads_per_rank - else: - # Send KVCache from 1 prefill instance to multiple decode instances - src_head_start_offset = ( - dst_tp_rank_in_group * dst_heads_per_rank - ) % src_heads_per_rank - num_heads_to_send = dst_heads_per_rank - dst_head_start_offset = 0 - - # torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set, e.g. - # 0xffff81ab54e01000). Use np.uint64 to prevent overflow on XPU. - kv_data_ptrs = np.array(self.kv_args.kv_data_ptrs, dtype=np.uint64) - dst_kv_ptrs = np.array(dst_kv_ptrs, dtype=np.uint64) - src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = ( - self.get_mha_kv_ptrs_with_pp(kv_data_ptrs, dst_kv_ptrs) - ) - # Calculate precise byte offset and length for the sub-slice within the token - src_head_slice_offset = src_head_start_offset * bytes_per_head_slice_to_send - dst_head_slice_offset = dst_head_start_offset * bytes_per_head_slice_to_send - heads_bytes_per_token_to_send = num_heads_to_send * bytes_per_head_slice_to_send - - src_dst_ptr_pairs = [ - ( - src_k_ptrs[layer_id], - dst_k_ptrs[layer_id], - ) - for layer_id in range(layers_current_pp_stage) - ] + [ - ( - src_v_ptrs[layer_id], - dst_v_ptrs[layer_id], - ) - for layer_id in range(layers_current_pp_stage) + dst_handle, num_slots_dst, head_group_idx = self.prep_handles_slice_dst[ + peer_name ] - - prefill_indices = np.asarray(prefill_kv_indices, dtype=np.uint64) - dst_indices = np.asarray(dst_kv_indices, dtype=np.uint64) - bytes_per_token_prefill = src_kv_item_len // page_size - bytes_per_token_decode = dst_kv_item_len // page_size - token_offsets = np.arange(page_size, dtype=np.uint64) - - src_addrs = [] - dst_addrs = [] - - for src_ptr, dst_ptr in src_dst_ptr_pairs: - src_page_bases = src_ptr + prefill_indices * src_kv_item_len - dst_page_bases = dst_ptr + dst_indices * dst_kv_item_len - - src_all = ( - src_page_bases[:, None] - + token_offsets[None, :] * bytes_per_token_prefill - + src_head_slice_offset - ).ravel() - dst_all = ( - dst_page_bases[:, None] - + token_offsets[None, :] * bytes_per_token_decode - + dst_head_slice_offset - ).ravel() - - src_addrs.append(src_all) - dst_addrs.append(dst_all) - - def make_req_array(addr_chunks, size, gpu): - if not addr_chunks: - return np.empty((0, 3), dtype=np.uint64) - flat_addrs = np.concatenate(addr_chunks).astype(np.uint64, copy=False) - return np.column_stack( - ( - flat_addrs, - np.full_like(flat_addrs, size, dtype=np.uint64), - np.full_like(flat_addrs, gpu, dtype=np.uint64), - ) - ) - - src_reqs = make_req_array( - src_addrs, heads_bytes_per_token_to_send, self.kv_args.gpu_id + page_size = self.kv_args.page_size + src_indices = expand_page_indices_for_slice( + np.asarray(prefill_kv_indices, dtype=np.int32), + num_ptr_pairs, + num_slots_src, + page_size, + num_groups=num_groups, + head_group_idx=head_group_idx, ) - dst_reqs = make_req_array(dst_addrs, heads_bytes_per_token_to_send, dst_gpu_id) - - # Use NIXL agent for transfer - src_descs = self.agent.get_xfer_descs(src_reqs, "VRAM") - dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM") - - xfer_handle = self.agent.initialize_xfer( - "WRITE", src_descs, dst_descs, peer_name, notif.encode("ascii") + dst_indices = expand_page_indices_for_slice( + np.asarray(dst_kv_indices, dtype=np.int32), + num_ptr_pairs, + num_slots_dst, + page_size, + ) + xfer_handle = self.agent.make_prepped_xfer( + "WRITE", + src_handle, + src_indices, + dst_handle, + dst_indices, + notif.encode("ascii"), ) if not xfer_handle: - raise Exception("Failed to create sliced KV transfer") - + raise Exception("KVSender failed to create prepped slice transfer") state = self.agent.transfer(xfer_handle) if state == "ERR": - raise Exception("Failed to post sliced KV transfer") - + raise Exception("KVSender failed to post prepped slice transfer") return xfer_handle def send_kvcache_staged( @@ -1910,6 +2084,10 @@ class NixlKVReceiver(CommonKVReceiver): else: packed_staging_base_ptr = b"" staging_total_size_str = b"" + dst_num_slots = ( + self.kv_mgr.kv_args.kv_data_lens[0] + // self.kv_mgr.kv_args.kv_item_lens[0] + ) with lock: sock.send_multipart( @@ -1931,6 +2109,7 @@ class NixlKVReceiver(CommonKVReceiver): packed_state_dim_per_tensor, packed_staging_base_ptr, staging_total_size_str, + str(dst_num_slots).encode("ascii"), ] )